Download code/tests/test_arope.py from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 15.7 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/test_arope.py
- Command line
-
hf download hf://teawhite/ActionRoPE/code/tests/test_arope.py
-
curl -L -o test_arope.py https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/test_arope.py
15.7 kB
| """ARoPE 模型补丁的自测(一张 H200 足够,DiT bf16 ~10 GB)。 | |
| CUDA_VISIBLE_DEVICES=0 /opt/dlami/nvme/zhiyangdeng/ActionRoPE/.venv/bin/python -m pytest \ | |
| /opt/dlami/nvme/zhiyangdeng/ActionRoPE/tests/test_arope.py -s -v | |
| 每个用例都能单独跑(-k),模型只在 module 级 fixture 里加载一次。 | |
| 实测数字追加写到 $AROPE_TEST_NUMBERS(默认 $AROPE_TEST_DIR=outputs/test_artifacts 下的 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)] | |
| # 测试产物只落在项目内(outputs/test_artifacts),不写任何机器相关的临时目录 | |
| 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, "arope_numbers.json")) | |
| # SPEC 的几何常量 | |
| C, F, LAT_H, LAT_W = 48, 21, 30, 52 | |
| TOK_H, TOK_W = 15, 26 | |
| L_TXT, D_TXT = 512, 4096 | |
| def record(**kv): | |
| """把实测数字追加进 json(多个用例各写各的键)。""" | |
| 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) | |
| # 原始键集:test_f 用它做对照,且必须在任何 install 之前抓 | |
| model._arope_test_baseline_keys = set(model.state_dict().keys()) | |
| return model | |
| 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 是 scheduler.timesteps[id].to(bf16),形状 [1] | |
| timestep = torch.tensor([500.0], device=device, dtype=torch.bfloat16) | |
| return latents, context, timestep | |
| def err_stats(out, ref): | |
| out, ref = out.float(), ref.float() | |
| diff = (out - ref).abs() | |
| return { | |
| "max_abs": diff.max().item(), | |
| "mean_abs": diff.mean().item(), | |
| "rel_l2": (diff.norm() / ref.norm().clamp_min(1e-12)).item(), | |
| "ref_absmax": ref.abs().max().item(), | |
| "ref_std": ref.std().item(), | |
| } | |
| # ---------------------------------------------------------------- (a) | |
| def test_a_equivalence_with_model_fn(dit): | |
| from diffsynth.pipelines.wan_video import model_fn_wan_video | |
| from actionrope.arope import arope_forward | |
| latents, context, timestep = make_inputs() | |
| ref = model_fn_wan_video(dit=dit, latents=latents, timestep=timestep, context=context, | |
| fuse_vae_embedding_in_latents=True) | |
| out = arope_forward(dit, latents, timestep, context, offset_px=None, mask_input=None) | |
| assert out.shape == (1, C, F, LAT_H, LAT_W) | |
| s = err_stats(out, ref) | |
| record(a_equiv_max_abs=s["max_abs"], a_equiv_rel_l2=s["rel_l2"], a_ref_absmax=s["ref_absmax"], | |
| a_ref_std=s["ref_std"], a_bitwise_equal=bool(torch.equal(out, ref))) | |
| assert torch.isfinite(out).all() | |
| assert s["rel_l2"] <= 1e-2 | |
| # ---------------------------------------------------------------- (b) | |
| def test_b_per_cell_context_equivalence(dit): | |
| from actionrope.arope import arope_forward | |
| latents, context, timestep = make_inputs() | |
| ref = arope_forward(dit, latents, timestep, context) | |
| ctx_cells = context.unsqueeze(1).expand(1, F, L_TXT, D_TXT).contiguous() | |
| # 按内容去重 | |
| out = arope_forward(dit, latents, timestep, ctx_cells) | |
| s = err_stats(out, ref) | |
| # 按 context_ids 去重(dataset 会给) | |
| ids = torch.zeros(1, F, dtype=torch.long) | |
| out2 = arope_forward(dit, latents, timestep, ctx_cells, context_ids=ids) | |
| s2 = err_stats(out2, ref) | |
| record(b_percell_max_abs=s["max_abs"], b_percell_rel_l2=s["rel_l2"], | |
| b_percell_ids_max_abs=s2["max_abs"], b_percell_ids_rel_l2=s2["rel_l2"], | |
| b_bitwise_equal=bool(torch.equal(out, ref))) | |
| assert s["rel_l2"] <= 1e-2 and s2["rel_l2"] <= 1e-2 | |
| # 两个不同的串真的会分到不同 cell:把后半段 cell 换成另一串,输出必须变 | |
| other = torch.randn_like(context) | |
| ctx_mix = ctx_cells.clone() | |
| ctx_mix[:, 11:] = other | |
| out3 = arope_forward(dit, latents, timestep, ctx_mix, context_ids=torch.tensor([[0] * 11 + [1] * 10])) | |
| s3 = err_stats(out3, ref) | |
| record(b_mixed_context_rel_l2=s3["rel_l2"]) | |
| assert s3["rel_l2"] > 1e-3 | |
| # ---------------------------------------------------------------- (c) | |
| def test_c_world_freqs_shift(dit): | |
| from actionrope.arope import build_world_freqs | |
| f, h, w = F, TOK_H, TOK_W | |
| table = build_world_freqs(dit, f, h, w, None, device="cuda") # 查表,[S,1,64] | |
| assert table.shape == (f * h * w, 1, 64) | |
| ref_table = torch.cat([ | |
| dit.freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1), | |
| dit.freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1), | |
| dit.freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1), | |
| ], dim=-1).reshape(f * h * w, 1, -1).to("cuda") | |
| assert torch.equal(table, ref_table) | |
| # 零偏移走现算路径,必须与查表逐位相同 | |
| zero = build_world_freqs(dit, f, h, w, torch.zeros(1, f, 2), device="cuda") | |
| assert zero.shape == (1, f * h * w, 1, 64) | |
| zero_max_diff = (zero[0] - table).abs().max().item() | |
| record(c_zero_offset_bitwise_equal=bool(torch.equal(zero[0], table)), c_zero_offset_max_diff=zero_max_diff) | |
| assert zero_max_diff <= 1e-6 | |
| # 每帧右移 1 token:w 轴分量 = 查表的 w 索引 +1(列 j 拿到 w_table[j+1]) | |
| shift = build_world_freqs(dit, f, h, w, torch.tensor([[[1.0, 0.0]]]).expand(1, f, 2), device="cuda") | |
| shift = shift.view(f, h, w, 64) | |
| manual = torch.cat([ | |
| dit.freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1), | |
| dit.freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1), | |
| dit.freqs[2][1:w + 1].view(1, 1, w, -1).expand(f, h, w, -1), | |
| ], dim=-1).to("cuda") | |
| d = (shift - manual).abs().max().item() | |
| record(c_shift_w_plus1_max_diff=d, c_shift_w_plus1_bitwise_equal=bool(torch.equal(shift, manual))) | |
| assert d <= 1e-6 | |
| # f 分量直接查表,逐位相同;h 分量是现算的,与 CPU 上算的表只有 cos/sin 库函数的 ULP 级差异 | |
| assert torch.equal(shift[..., :22], table.view(f, h, w, 64)[..., :22]) | |
| d_h = (shift[..., 22:43] - table.view(f, h, w, 64)[..., 22:43]).abs().max().item() | |
| record(c_recomputed_h_axis_max_diff=d_h) | |
| assert d_h <= 1e-6 | |
| # 每帧下移 1 token(dy=+1):h 轴分量 = 查表的 h 索引 +1 | |
| shift_h = build_world_freqs(dit, f, h, w, torch.tensor([[[0.0, 1.0]]]).expand(1, f, 2), device="cuda").view(f, h, w, 64) | |
| manual_h = dit.freqs[1][1:h + 1].view(1, h, 1, -1).expand(f, h, w, -1).to("cuda") | |
| assert (shift_h[..., 22:43] - manual_h).abs().max().item() <= 1e-6 | |
| # 逐帧不同的偏移:第 k 帧右移 k token,只有对应帧被移 | |
| off = torch.zeros(1, f, 2) | |
| off[0, :, 0] = torch.arange(f) | |
| var = build_world_freqs(dit, f, h, w, off, device="cuda").view(f, h, w, 64) | |
| for k in (0, 3, 20): | |
| manual_k = dit.freqs[2][k:k + w].view(1, w, -1).expand(h, w, -1).to("cuda") | |
| assert (var[k, ..., 43:] - manual_k).abs().max().item() <= 1e-6 | |
| # 小数偏移:相位是连续的(0.5 token 落在 0 与 1 之间) | |
| half = build_world_freqs(dit, 1, 1, 1, torch.tensor([[[0.5, 0.0]]]), device="cuda").view(64) | |
| ang = torch.angle(half[43:]) | |
| inv = 1.0 / (10000.0 ** (torch.arange(0, 42, 2).double() / 42)) | |
| assert torch.allclose(ang.cpu(), 0.5 * inv, atol=1e-9) | |
| # ---------------------------------------------------------------- (d) | |
| def test_d_known_mask_geometry(): | |
| from actionrope.arope import known_mask, loss_weight_map | |
| # (+64, 0) px = +2 token:每帧最右 2 列 new | |
| off = torch.tensor([[[64.0, 0.0]]]).expand(1, F, 2) / 32 | |
| m = known_mask(off, TOK_H, TOK_W) | |
| assert m.shape == (1, F, TOK_H, TOK_W) and m.dtype == torch.bool | |
| assert m[..., :, :TOK_W - 2].all() and (~m[..., :, TOK_W - 2:]).all() | |
| # (0, −32) px = −1 token:最上 1 行 new | |
| off = torch.tensor([[[0.0, -32.0]]]).expand(1, F, 2) / 32 | |
| m = known_mask(off, TOK_H, TOK_W) | |
| assert (~m[..., 0, :]).all() and m[..., 1:, :].all() | |
| # 零偏移:全 known | |
| assert known_mask(torch.zeros(1, F, 2), TOK_H, TOK_W).all() | |
| # 逐帧不同(帧 k 右移 k token):帧 k 最右 k 列 new | |
| off = torch.zeros(1, F, 2) | |
| off[0, :, 0] = torch.arange(F) | |
| m = known_mask(off, TOK_H, TOK_W) | |
| for k in range(F): | |
| n_new = (~m[0, k]).any(dim=0).sum().item() | |
| assert n_new == min(k, TOK_W), (k, n_new) | |
| # loss_weight_map:latent 分辨率,偏移除以 16 | |
| w = loss_weight_map(torch.tensor([[[64.0, 0.0]]]).expand(1, F, 2), new_weight=2.0) | |
| assert w.shape == (1, 1, F, LAT_H, LAT_W) and w.dtype == torch.float32 | |
| assert (w[..., :, :LAT_W - 4] == 1).all() and (w[..., :, LAT_W - 4:] == 2).all() | |
| w = loss_weight_map(torch.tensor([[[0.0, -32.0]]]).expand(1, F, 2), new_weight=3.0) | |
| assert (w[..., :2, :] == 3).all() and (w[..., 2:, :] == 1).all() | |
| assert (loss_weight_map(torch.zeros(2, F, 2)) == 1).all() | |
| # 半格边界取闭区间:+16 px = +0.5 token ⇒ 最右列世界坐标 25.5,仍 known | |
| assert known_mask(torch.tensor([[[16.0, 0.0]]]) / 32, TOK_H, TOK_W).all() | |
| assert not known_mask(torch.tensor([[[17.0, 0.0]]]) / 32, TOK_H, TOK_W)[..., -1].any() | |
| # ---------------------------------------------------------------- (e) | |
| def test_e_training_step(dit): | |
| from actionrope.arope import arope_forward, install_arope, loss_weight_map | |
| n_params_before = sum(p.numel() for p in dit.parameters()) | |
| install_arope(dit, mask_channel=True) | |
| n_params_after = sum(p.numel() for p in dit.parameters()) | |
| assert hasattr(dit, "arope_mask_embedding") | |
| assert dit.arope_mask_embedding.weight.dtype == torch.bfloat16 | |
| assert dit.arope_mask_embedding.weight.device.type == "cuda" | |
| latents, context, timestep = make_inputs(seed=1) | |
| off = torch.zeros(1, F, 2, device="cuda") | |
| off[0, :, 0] = torch.linspace(0, 96, F) # 全程向右走 3 token | |
| off[0, :, 1] = torch.linspace(0, -20, F) | |
| mask = loss_weight_map(off.cpu()).to("cuda").eq(1.0).float() # known=1 / new=0 | |
| ctx_cells = context.unsqueeze(1).expand(1, F, L_TXT, D_TXT).contiguous() | |
| # 零初始化 ⇒ mask 通道不改变输出(eval、无梯度) | |
| with torch.no_grad(): | |
| o_none = arope_forward(dit, latents, timestep, context, offset_px=None) | |
| o_ref = arope_forward(dit, latents, timestep, context, offset_px=off) | |
| o_mask = arope_forward(dit, latents, timestep, context, offset_px=off, mask_input=mask) | |
| assert torch.equal(o_ref, o_mask) | |
| # 世界 RoPE 真的在起作用:非零偏移必须改变输出,但量级不变 | |
| s_off = err_stats(o_ref, o_none) | |
| assert s_off["rel_l2"] > 1e-3 | |
| assert abs(o_ref.float().std().item() / o_none.float().std().item() - 1) < 0.2 | |
| record(e_zero_init_mask_no_change=True, e_params_added=n_params_after - n_params_before, | |
| e_params_total=n_params_after, e_offset_vs_plain_rel_l2=s_off["rel_l2"], | |
| e_out_std_plain=o_none.float().std().item(), e_out_std_offset=o_ref.float().std().item()) | |
| # 全参数训练一步:前向 + 反向,梯度检查点开 | |
| try: | |
| dit.train().requires_grad_(True) | |
| torch.cuda.synchronize() | |
| torch.cuda.reset_peak_memory_stats() | |
| t0 = time.time() | |
| out = arope_forward(dit, latents, timestep, ctx_cells, offset_px=off, mask_input=mask, | |
| use_gradient_checkpointing=True, | |
| context_ids=torch.zeros(1, F, dtype=torch.long)) | |
| target = torch.randn_like(out) | |
| wmap = loss_weight_map(off.cpu()).to(out.device) | |
| # 与 SPEC 的 loss 形式一致:Σ w·(pred−target)² / Σ w,w 在通道维广播,所以分母也要乘通道数 | |
| loss = ((out.float() - target.float()) ** 2 * wmap).sum() / (wmap.sum() * out.shape[1]) | |
| loss.backward() | |
| torch.cuda.synchronize() | |
| dt = time.time() - t0 | |
| peak_gb = torch.cuda.max_memory_allocated() / 1024 ** 3 | |
| assert torch.isfinite(loss) | |
| g_mask = dit.arope_mask_embedding.weight.grad | |
| assert g_mask is not None and torch.isfinite(g_mask).all() | |
| mask_grad_norm = g_mask.float().norm().item() | |
| assert mask_grad_norm > 0 | |
| # 抽查其它参数的梯度非零 | |
| 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, | |
| "text_embedding.0.weight": dit.text_embedding[0].weight, | |
| "time_embedding.0.weight": dit.time_embedding[0].weight, | |
| } | |
| grad_norms = {} | |
| for name, p in checks.items(): | |
| assert p.grad is not None, name | |
| grad_norms[name] = p.grad.float().norm().item() | |
| assert 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()) | |
| record(e_loss=loss.item(), e_peak_mem_gb=peak_gb, e_fwd_bwd_sec=dt, | |
| e_mask_conv_grad_norm=mask_grad_norm, e_grad_norms=grad_norms, | |
| e_params_with_nonzero_grad=f"{n_with_grad}/{n_total}") | |
| assert n_with_grad == n_total | |
| finally: | |
| # 还原:其它用例要在干净的 eval 模型上跑 | |
| for p in dit.parameters(): | |
| p.grad = None | |
| dit.eval().requires_grad_(False) | |
| install_arope(dit, mask_channel=False) | |
| torch.cuda.empty_cache() | |
| # ---------------------------------------------------------------- (f) | |
| def test_f_state_dict_keys_unchanged(dit): | |
| from actionrope.arope import install_arope | |
| baseline = dit._arope_test_baseline_keys | |
| install_arope(dit, mask_channel=False) | |
| assert set(dit.state_dict().keys()) == baseline | |
| assert not hasattr(dit, "arope_mask_embedding") | |
| install_arope(dit, mask_channel=True) | |
| keys = set(dit.state_dict().keys()) | |
| extra = keys - baseline | |
| assert extra == {"arope_mask_embedding.weight", "arope_mask_embedding.bias"} | |
| assert baseline <= keys | |
| install_arope(dit, mask_channel=False) | |
| assert set(dit.state_dict().keys()) == baseline | |
| record(f_keys_unchanged=True, f_extra_keys_with_mask=sorted(extra)) | |