Download code/baseline/linear.py from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 6.37 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/baseline/linear.py
- Command line
-
hf download hf://teawhite/ActionRoPE/code/baseline/linear.py
-
curl -L -o linear.py https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/baseline/linear.py
6.37 kB
| """`linear` 臂:ReactiveGWM 的逐块线性动作偏置(30 个 bias-free Linear,加在每个 DiT block 入口)。 | |
| 上游机制(vender/ReactiveGWM,只读) | |
| ----------------------------------- | |
| * 模型 `inference/models/dit.py::WanModelAction`(同一个 Wan2.2-TI2V-5B 基座,dim 3072 / 30 层): | |
| - L207-209 `self.action_embedders = nn.ModuleList([nn.Linear(num_buttons, dim, bias=False) for _ in range(num_layers)])` | |
| —— 每个 block 一个 **无偏置** 线性层,把动作向量投到 hidden 维;没有门控(docstring:"No gates")。 | |
| - L222-226 `_bin_action`:原始逐帧键盘 one-hot `[B, T_raw, K]` 用 `adaptive_max_pool1d` 压到 latent 帧数 `[B, f, K]` | |
| (每个 latent 帧 = 4 个视频帧的"按过即 1")。 | |
| - L228-234 `_inject_action`:`emb = action_embedders[i](action[B,f,K].to(x.dtype))` → `[B, f, C]`, | |
| 沿 h·w 展开成逐 token 偏置 `[B, f·h·w, C]`,`x = x + bias`。 | |
| - L284-286 前向:`for i, block: x = _inject_action(x, action, i, f, h, w); x = block(x, ctx, t_mod, freqs)` | |
| —— 偏置加在 **block 入口的残差流**(进 norm1/self-attn 之前),不是 adaLN、不是 cross-attn。 | |
| * 动作向量 `training/data/action_utils.py` / `inference/utils/actions.py`:parquet 的 10 个按键列 | |
| (`inference/constants.py::SF_BUTTON_COLS`,UP/DOWN/LEFT/RIGHT + 6 攻击键)0/1 值,`hold_last_upsample`(10 帧窗) | |
| 补成逐视频帧 dense one-hot,不做归一化。 | |
| * 初始化:`training/bidirectional/train.py` L17-18 / L134-137:"ActionModule keys stay at their default (zero / xavier) init", | |
| 其余权重从 Wan2.2-TI2V-5B 按形状拷贝(`_transfer_weights`)。训练时 action_embedders 与 DiT 全参数同 lr 训练; | |
| 推理有 `action_cfg_scale`(`inference/pipeline.py` L218-224,动作置零做无动作分支)。 | |
| * 上游参数量:30 × 10 × 3072 = 921,600(10 键)。 | |
| 本实现的对应(baseline/SPEC.md 的钩子) | |
| -------------------------------------- | |
| * `encode`:取 `action_inputs["offset_tok"]`(float32 `[b, 21, 2]`,逐 cell 累计屏幕位移 (dx, dy),单位 token, | |
| 与 ARoPE 臂**完全相同的控制信号**——旧报告口径"带动作的臂收到相同的逐 cell 累计偏移")。 | |
| 上游的"逐帧 one-hot → adaptive_max_pool 到 f 个 latent 帧"这一步,在我们这边由 sidecar 的 | |
| `frames_to_cells`(cell 内逐帧偏移取均值)已经做完,dataset 直接给逐 cell 的 `[b, 21, 2]`, | |
| 所以这里不再池化,只做形状校验。不归一化(上游也不归一化;零初始化下量纲只影响学习动态,不影响装上瞬间的等价性)。 | |
| cell 0 的偏移恒为 0 ⇒ 无偏置 Linear 在条件帧上的偏置恒为 0,与"首帧是干净条件帧"一致。 | |
| * `block_pre(i, x)`:`x + expand(action_embedders[i](offset_tok))`,与上游 `_inject_action` 同位置同算式。 | |
| 小矩阵乘在 fp32 里算再转回 x.dtype:offset_tok 最大约 ±11 token,bf16 在该量级的分辨率是 1/16 token = 2 px, | |
| fp32 保住亚像素信息;上游输入是 0/1 one-hot,bf16 本来就无损,所以它直接 `.to(x.dtype)`。 | |
| * 参数量:30 × 2 × 3072 = **184,320 ≈ 0.18M**(与旧报告 +0.18M 一致:旧报告就是 2 维偏移输入)。 | |
| * 初始化:**全零**(SPEC 要求"装上瞬间输出与原模型逐位相同";上游 docstring 也把 zero 列为默认之一)。 | |
| 零权重 ⇒ 偏置张量逐位为 0 ⇒ `x + 0 == x` 逐位相等。 | |
| * ckpt 键:`action_embedders.{i}.weight`(与上游键布局同名;integration 阶段导出时加 `arm.` 前缀,与 DiT 键无冲突)。 | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from baseline.base import ActionArm | |
| # 上游 SF_BUTTON_COLS 是 10 维 one-hot;这里换成 SPEC 的 (dx, dy) 累计偏移,2 维 | |
| ACTION_DIM = 2 | |
| class LinearArm(ActionArm): | |
| name = "linear" | |
| def __init__(self, dim: int = 3072, num_layers: int = 30, in_dim: int = ACTION_DIM): | |
| super().__init__() | |
| self.dim, self.num_layers, self.in_dim = dim, num_layers, in_dim | |
| self.action_embedders = self._build(dim, num_layers, in_dim) | |
| def _build(dim: int, num_layers: int, in_dim: int) -> nn.ModuleList: | |
| # 与上游 dit.py L207-209 同构:每层一个 bias-free Linear;权重零初始化(见模块 docstring) | |
| embedders = nn.ModuleList([nn.Linear(in_dim, dim, bias=False) for _ in range(num_layers)]) | |
| for lin in embedders: | |
| nn.init.zeros_(lin.weight) | |
| return embedders | |
| def install(self, dit) -> None: | |
| """按 dit 的 dim / 层数建参数,并放到 dit 的 device / dtype(与 DiT 一起以 bf16 训练)。""" | |
| dim, num_layers = int(dit.dim), len(dit.blocks) | |
| if (dim, num_layers) != (self.dim, self.num_layers): | |
| self.dim, self.num_layers = dim, num_layers | |
| self.action_embedders = self._build(dim, num_layers, self.in_dim) | |
| ref = dit.patch_embedding.weight | |
| self.to(device=ref.device, dtype=ref.dtype) | |
| # ---- 钩子 ---- | |
| def encode(self, action_inputs: dict | None, grid: dict): | |
| """offset_tok [b, f, 2] → 同一张量(float32)。None ⇒ 不注入(对应上游 keyboard_action=None 的分支)。""" | |
| if action_inputs is None: | |
| return None | |
| off = action_inputs["offset_tok"] | |
| b, f = grid["b"], grid["f"] | |
| assert off.shape == (b, f, self.in_dim), f"offset_tok 形状应为 [{b}, {f}, {self.in_dim}],得到 {tuple(off.shape)}" | |
| return off.to(torch.float32) | |
| def block_pre(self, idx: int, x: torch.Tensor, feats, grid: dict) -> torch.Tensor: | |
| """上游 `_inject_action`:逐帧线性偏置沿帧内 h·w 个 token 广播后加到残差流。""" | |
| if feats is None: | |
| return x | |
| b, f, n = grid["b"], grid["f"], grid["n_tok_per_frame"] | |
| weight = self.action_embedders[idx].weight | |
| emb = F.linear(feats.to(device=x.device), weight.to(torch.float32)).to(x.dtype) # [b, f, C] | |
| bias = emb.unsqueeze(2).expand(b, f, n, self.dim).reshape(b, f * n, self.dim) | |
| return x + bias | |
| # ---- 工具 ---- | |
| def zero_init_check(self) -> bool: | |
| return all(bool((lin.weight == 0).all()) for lin in self.action_embedders) | |