File size: 15,651 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
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
"""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))


@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)
    # 原始键集: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)

@torch.no_grad()
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)

@torch.no_grad()
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))