Download code/tests/test_arm_linear.py from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 10.1 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/test_arm_linear.py
- Command line
-
hf download hf://teawhite/ActionRoPE/code/tests/test_arm_linear.py
-
curl -L -o test_arm_linear.py https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/test_arm_linear.py
10.1 kB
| """`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)) | |
| 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 | |
| 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. 零初始化等价 | |
| 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) | |