"""`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) @staticmethod 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)