"""动作臂基类: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