File size: 1,930 Bytes
880dff9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
"""动作臂基类:baseline/SPEC.md 里钩子接口的默认实现(全部恒等)。

各 baseline(linear / xattn / prompt / adaln)继承它,只覆盖自己用到的钩子。
`arope_forward(arm=None)` 时一个钩子都不会被调用,所以基类本身不影响 ARoPE 臂。
"""
from __future__ import annotations

import torch
import torch.nn as nn


class ActionArm(nn.Module):
    name = "base"

    def install(self, dit) -> None:
        """需要按 dit 的 dim / num_heads / 层数建参数时在这里做(build_arm 会调用一次)。"""

    # ---- 钩子:签名固定,见 baseline/SPEC.md ----
    def encode(self, action_inputs: dict | None, grid: dict):
        """动作输入 → 特征(任意张量或 tuple),只算一次;None 表示这个臂不需要。"""
        return None

    def modify_t_mod(self, t_mod: torch.Tensor, feats, grid: dict) -> torch.Tensor:
        return t_mod

    def extra_context(self, context: torch.Tensor, feats, grid: dict) -> torch.Tensor:
        return context

    def after_patch(self, x: torch.Tensor, feats, grid: dict) -> torch.Tensor:
        return x

    def block_pre(self, idx: int, x: torch.Tensor, feats, grid: dict) -> torch.Tensor:
        return x

    def block_mid(self, idx: int, x: torch.Tensor, feats, grid: dict) -> torch.Tensor:
        return x

    def block_pre_ffn(self, idx: int, x: torch.Tensor, feats, grid: dict) -> torch.Tensor:
        """文本 cross-attn 之后、FFN 之前。"""
        return x

    def block_post(self, idx: int, x: torch.Tensor, feats, grid: dict) -> torch.Tensor:
        return x

    # ---- 工具 ----
    def n_new_params(self) -> int:
        return sum(p.numel() for p in self.parameters())

    def zero_init_check(self) -> bool:
        """"装上瞬间不改变输出"的粗检:至少有一条通向残差流的路径是零。子类按自己的结构覆盖。"""
        return True