File size: 6,365 Bytes
880dff9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 | """`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)
|