Download code/baseline/base.py from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 1.93 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/baseline/base.py
- Command line
-
hf download hf://teawhite/ActionRoPE/code/baseline/base.py
-
curl -L -o base.py https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/baseline/base.py
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 | |