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)