ActionRoPE / code /baseline /base.py
teawhite's picture
add docs+code
880dff9 verified
Raw History Blame Contribute Delete
1.93 kB
"""动作臂基类: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