File size: 10,108 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 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 | """`linear` 臂(ReactiveGWM 逐块线性偏置)的自测:baseline/SPEC.md 的测试 1–4。
CUDA_VISIBLE_DEVICES=4 /opt/dlami/nvme/zhiyangdeng/ActionRoPE/.venv/bin/python -m pytest \
/opt/dlami/nvme/zhiyangdeng/ActionRoPE/tests/test_arm_linear.py -s -v
模型只在 module 级 fixture 里加载一次(DiT bf16 ~10 GB);实测数字追加写到 $AROPE_TEST_NUMBERS
(默认 outputs/test_artifacts/arm_linear_numbers.json)。
"""
from __future__ import annotations
import json
import os
import time
import pytest
import torch
os.environ.setdefault("DIFFSYNTH_SKIP_DOWNLOAD", "True")
ROOT = "/opt/dlami/nvme/zhiyangdeng/ActionRoPE"
MODEL_DIR = f"{ROOT}/models/Wan2.2-TI2V-5B"
DIT_FILES = [f"{MODEL_DIR}/diffusion_pytorch_model-0000{i}-of-00003.safetensors" for i in (1, 2, 3)]
TEST_DIR = os.environ.get("AROPE_TEST_DIR", os.path.join(ROOT, "outputs", "test_artifacts"))
NUMBERS_PATH = os.environ.get("AROPE_TEST_NUMBERS", os.path.join(TEST_DIR, "arm_linear_numbers.json"))
C, F, LAT_H, LAT_W = 48, 21, 30, 52
TOK_H, TOK_W = 15, 26
L_TXT, D_TXT = 512, 4096
DIM, N_LAYERS, ACTION_DIM = 3072, 30, 2
N_NEW_PARAMS_EXPECTED = N_LAYERS * ACTION_DIM * DIM # 184,320 ≈ 旧报告的 +0.18M
def record(**kv):
data = {}
try:
with open(NUMBERS_PATH) as fp:
data = json.load(fp)
except (FileNotFoundError, json.JSONDecodeError):
pass
data.update({k: (float(v) if isinstance(v, (int, float)) and not isinstance(v, bool) else v) for k, v in kv.items()})
os.makedirs(os.path.dirname(NUMBERS_PATH), exist_ok=True)
with open(NUMBERS_PATH, "w") as fp:
json.dump(data, fp, indent=2, ensure_ascii=False)
print("\n[numbers]", json.dumps(kv, ensure_ascii=False))
@pytest.fixture(scope="module")
def dit():
from diffsynth.pipelines.wan_video import ModelConfig, WanVideoPipeline
pipe = WanVideoPipeline.from_pretrained(
torch_dtype=torch.bfloat16, device="cuda",
model_configs=[ModelConfig(path=DIT_FILES)],
tokenizer_config=None, redirect_common_files=False,
)
model = pipe.dit
model.eval().requires_grad_(False)
model._baseline_keys = set(model.state_dict().keys())
return model
@pytest.fixture(scope="module")
def arm(dit):
from baseline.linear import LinearArm
a = LinearArm()
a.install(dit)
a.eval().requires_grad_(False)
return a
def make_inputs(seed=0, batch=1, device="cuda"):
g = torch.Generator(device="cpu").manual_seed(seed)
latents = torch.randn(batch, C, F, LAT_H, LAT_W, generator=g).to(device=device, dtype=torch.bfloat16)
context = torch.randn(batch, L_TXT, D_TXT, generator=g).to(device=device, dtype=torch.bfloat16)
timestep = torch.tensor([500.0], device=device, dtype=torch.bfloat16)
return latents, context, timestep
def make_action_inputs(batch=1, device="cuda"):
"""SPEC 的 action_inputs:全程向右走 3 token、向上 20 px(与 test_arope 的 test_e 同一条轨迹)。"""
off_px = torch.zeros(batch, F, 2)
off_px[:, :, 0] = torch.linspace(0, 96, F)
off_px[:, :, 1] = torch.linspace(0, -20, F)
off_tok = off_px / 32
delta = torch.zeros_like(off_tok)
delta[:, 1:] = off_tok[:, 1:] - off_tok[:, :-1]
return {
"offset_px": off_px.to(device),
"offset_tok": off_tok.to(device),
"delta_tok": delta.to(device),
"action_idx": torch.full((batch, F), 7, dtype=torch.long, device=device), # 7 = moving right
}
def err_stats(out, ref):
out, ref = out.float(), ref.float()
diff = (out - ref).abs()
return {"max_abs": diff.max().item(), "rel_l2": (diff.norm() / ref.norm().clamp_min(1e-12)).item()}
# ---------------------------------------------------------------- 1. 零初始化等价
@torch.no_grad()
def test_1_zero_init_equivalence(dit, arm):
from actionrope.arope import arope_forward
assert arm.zero_init_check()
latents, context, timestep = make_inputs()
ai = make_action_inputs()
ref = arope_forward(dit, latents, timestep, context)
out = arope_forward(dit, latents, timestep, context, arm=arm, action_inputs=ai)
s = err_stats(out, ref)
bitwise = bool(torch.equal(out, ref))
# action_inputs=None ⇒ 不注入(上游 keyboard_action=None 分支),同样逐位相同
out_none = arope_forward(dit, latents, timestep, context, arm=arm, action_inputs=None)
record(equiv_max_abs=s["max_abs"], equiv_rel_l2=s["rel_l2"], equiv_bitwise_equal=bitwise,
equiv_none_bitwise_equal=bool(torch.equal(out_none, ref)))
assert torch.isfinite(out).all()
assert bitwise or s["rel_l2"] <= 1e-6
assert torch.equal(out_none, ref)
# ---------------------------------------------------------------- 2. 扰动后有变化 + 梯度
def test_2_perturbed_changes_and_grads(dit, arm):
from actionrope.arope import arope_forward
latents, context, timestep = make_inputs(seed=1)
ai = make_action_inputs()
g = torch.Generator(device="cpu").manual_seed(123)
saved = {k: v.clone() for k, v in arm.state_dict().items()}
try:
with torch.no_grad():
ref = arope_forward(dit, latents, timestep, context)
for lin in arm.action_embedders:
lin.weight.copy_(torch.randn(lin.weight.shape, generator=g) * 0.02)
out = arope_forward(dit, latents, timestep, context, arm=arm, action_inputs=ai)
s = err_stats(out, ref)
assert torch.isfinite(out).all()
assert s["rel_l2"] > 1e-3
# 动作为零(offset 全 0)时即使权重非零也不改变输出:bias-free Linear 的性质
with torch.no_grad():
zero_ai = {k: torch.zeros_like(v) for k, v in ai.items()}
out_zero_action = arope_forward(dit, latents, timestep, context, arm=arm, action_inputs=zero_ai)
assert torch.equal(out_zero_action, ref)
# 前向 + 反向,梯度检查点开,DiT 与 arm 都要有梯度
dit.train().requires_grad_(True)
arm.train().requires_grad_(True)
torch.cuda.synchronize()
torch.cuda.reset_peak_memory_stats()
t0 = time.time()
out = arope_forward(dit, latents, timestep, context, arm=arm, action_inputs=ai,
use_gradient_checkpointing=True)
target = torch.randn_like(out)
loss = ((out.float() - target.float()) ** 2).mean()
loss.backward()
torch.cuda.synchronize()
dt = time.time() - t0
peak_gb = torch.cuda.max_memory_allocated() / 1024 ** 3
assert torch.isfinite(loss)
arm_grad_norms = []
for i, lin in enumerate(arm.action_embedders):
assert lin.weight.grad is not None, f"action_embedders.{i} 无梯度"
assert torch.isfinite(lin.weight.grad).all()
arm_grad_norms.append(lin.weight.grad.float().norm().item())
assert all(n > 0 for n in arm_grad_norms)
checks = {
"patch_embedding.weight": dit.patch_embedding.weight,
"blocks.0.self_attn.q.weight": dit.blocks[0].self_attn.q.weight,
"blocks.15.cross_attn.k.weight": dit.blocks[15].cross_attn.k.weight,
"blocks.29.ffn.2.weight": dit.blocks[29].ffn[2].weight,
"head.head.weight": dit.head.head.weight,
"time_embedding.0.weight": dit.time_embedding[0].weight,
}
dit_grad_norms = {}
for name, p in checks.items():
assert p.grad is not None, name
dit_grad_norms[name] = p.grad.float().norm().item()
assert dit_grad_norms[name] > 0, name
n_with_grad = sum(1 for p in dit.parameters() if p.grad is not None and p.grad.abs().sum() > 0)
n_total = sum(1 for p in dit.parameters())
grad_ok = (n_with_grad == n_total) and all(n > 0 for n in arm_grad_norms)
record(perturbed_rel_l2=s["rel_l2"], perturbed_max_abs=s["max_abs"], perturbed_out_std=out.float().std().item(),
ref_out_std=ref.float().std().item(), loss=loss.item(), peak_mem_gb=peak_gb, fwd_bwd_sec=dt,
arm_grad_norm_min=min(arm_grad_norms), arm_grad_norm_max=max(arm_grad_norms),
dit_grad_norms=dit_grad_norms, dit_params_with_nonzero_grad=f"{n_with_grad}/{n_total}", grad_ok=grad_ok)
assert grad_ok
finally:
for p in list(dit.parameters()) + list(arm.parameters()):
p.grad = None
dit.eval().requires_grad_(False)
with torch.no_grad():
arm.load_state_dict(saved)
arm.eval().requires_grad_(False)
torch.cuda.empty_cache()
# ---------------------------------------------------------------- 3. 参数量
def test_3_param_count(dit, arm):
n_new = arm.n_new_params()
n_dit = sum(p.numel() for p in dit.parameters())
print(f"\n[linear] 新增参数 {n_new:,} ({n_new / 1e6:.3f}M);DiT {n_dit:,};上游 10 键版 = {N_LAYERS * 10 * DIM:,}")
record(n_new_params=n_new, n_new_params_M=n_new / 1e6, n_dit_params=n_dit,
n_upstream_10button_params=N_LAYERS * 10 * DIM)
assert n_new == N_NEW_PARAMS_EXPECTED
assert all(p.dtype == torch.bfloat16 and p.device.type == "cuda" for p in arm.parameters())
# ---------------------------------------------------------------- 4. state_dict 键集
def test_4_state_dict_keys(dit, arm):
from baseline.linear import LinearArm
keys = set(arm.state_dict().keys())
assert keys == {f"action_embedders.{i}.weight" for i in range(N_LAYERS)}
# 装 arm 不改 DiT 的键集;加 `arm.` 前缀后与 DiT 键无冲突
assert set(dit.state_dict().keys()) == dit._baseline_keys
assert not ({f"arm.{k}" for k in keys} & dit._baseline_keys)
# strict 加载到新建的同结构臂,并保持等价(形状、值)
fresh = LinearArm()
fresh.install(dit)
missing, unexpected = fresh.load_state_dict(arm.state_dict(), strict=True)
assert not missing and not unexpected
for k, v in arm.state_dict().items():
assert torch.equal(fresh.state_dict()[k], v)
record(state_dict_keys=sorted(keys)[:3] + ["..."], state_dict_strict_load_ok=True)
|