File size: 6,987 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
"""baseline 臂接入训练 / 推理的纯 CPU 自测(不加载 DiT,几秒钟)。

    PYTHONPATH=/opt/dlami/nvme/zhiyangdeng/ActionRoPE .venv/bin/python -m pytest tests/test_baseline_integration.py -q

1. dataset 样本带齐 action_inputs 的四个键,offset_tok = offset_px/32、delta_tok 差分、action_idx 经 MIRROR_ACTION 翻回真实方向,
   且方向与 offset 增量一致(同一条 clip 上余弦 > 0)。
2. baseline.ARMS / build_arm 对四个臂都能在一个假 dit 上建参数并零初始化通过。
3. detect_arm_from_keys 从 ckpt 键集识别臂名与 xattn 结构配置;train.py 的 parse_args 对新臂给出正确的 text_mode / mask_channel。
4. infer.py 的标签约定:velocity_to_label(mirror=False) 与 dataset 的 MIRROR_ACTION 翻转给出同一套 action_idx。
"""
from __future__ import annotations

import os
import sys

import numpy as np
import pytest
import torch
import torch.nn as nn

os.environ.setdefault("DIFFSYNTH_SKIP_DOWNLOAD", "True")
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, ROOT)

from actionrope import geometry as G  # noqa: E402
from actionrope.prompts import MIRROR_ACTION  # noqa: E402
from baseline import ARM_NAMES, ARMS, TEXT_ARMS, arm_text_mode, build_arm, detect_arm_from_keys, make_action_inputs  # noqa: E402

TRAIN_DIR = os.path.join(ROOT, "data/latent/train_eybx")


class _FakeDiT(nn.Module):
    """只带各臂 install 用得着的几个属性:dim / blocks / patch_embedding / time_projection(小尺寸,CPU)。"""

    def __init__(self, dim=64, n_blocks=4):
        super().__init__()
        self.dim = dim
        self.blocks = nn.ModuleList([nn.Identity() for _ in range(n_blocks)])
        self.patch_embedding = nn.Conv3d(48, dim, (1, 2, 2), stride=(1, 2, 2))
        self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))


@pytest.fixture(scope="module")
def dataset():
    from actionrope.dataset import AropeLatentDataset
    if not os.path.isdir(TRAIN_DIR):
        pytest.skip("没有训练 latent 目录")
    return AropeLatentDataset(TRAIN_DIR, text_mode="scene", limit=6, min_rt_ratio=0.8, scene_dropout=0.0, verbose=False)


def test_dataset_action_inputs(dataset):
    from actionrope.dataset import arope_collate
    batch = arope_collate([dataset[0], dataset[1]])
    assert batch["offset_px"].shape == batch["offset_tok"].shape == batch["delta_tok"].shape == (2, 21, 2)
    assert batch["action_idx"].shape == (2, 21) and batch["action_idx"].dtype == torch.int64
    assert torch.equal(batch["offset_tok"] * 32, batch["offset_px"])
    assert torch.all(batch["offset_tok"][:, 0] == 0) and torch.all(batch["delta_tok"][:, 0] == 0)
    assert torch.allclose(batch["delta_tok"][:, 1:], batch["offset_tok"][:, 1:] - batch["offset_tok"][:, :-1])
    # action_idx = MIRROR_ACTION[.pt 的 actions]
    d = torch.load(dataset.items[0][1], map_location="cpu", weights_only=False)
    assert batch["action_idx"][0].tolist() == [MIRROR_ACTION[int(a)] for a in d["actions"]]
    # 翻回真实方向后与 offset 增量同向(镜像标签会给出 cos < 0)
    cos = []
    for i in range(len(dataset)):
        s = dataset[i]
        dt, ai = s["delta_tok"].numpy(), s["action_idx"].numpy()
        for k in range(1, 21):
            if ai[k] != 0 and np.linalg.norm(dt[k]) > 0.05:
                cos.append(G.ACTION_DIRS[ai[k]] @ (dt[k] / np.linalg.norm(dt[k])))
    assert cos and float(np.mean(cos)) > 0.8, np.mean(cos)


def test_build_arm_all_arms():
    assert set(ARM_NAMES) == {"arope", "plain", "linear", "xattn", "prompt", "adaln"}
    assert set(ARMS) == {"linear", "xattn", "prompt", "adaln"}
    assert all(arm_text_mode(a) == ("scene_action" if a in TEXT_ARMS else "scene") for a in ARM_NAMES)
    assert build_arm("arope", _FakeDiT()) is None and build_arm("plain", _FakeDiT()) is None
    expect = {"linear": 4 * 2 * 64, "prompt": 0, "adaln": (128 * 64 + 64) + (64 * 64 + 64) + (64 * 384 + 384)}
    for name in ARMS:
        dit = _FakeDiT()
        kw = {"blocks": [0, 1], "window_frames": 1, "enable_mouse": False, "heads_num": 4} if name == "xattn" else None
        arm = build_arm(name, dit, kw)
        assert arm.name == name and arm.zero_init_check(), name
        assert all(p.dtype == torch.float32 for p in arm.parameters())
        if name in expect:
            assert arm.n_new_params() == expect[name], (name, arm.n_new_params())
        # state_dict 键加 arm. 前缀后与 DiT 键不冲突,且能按键集认回来
        keys = list(arm.state_dict())
        assert all(not k.startswith("dit.") for k in keys)
        detected, dkw = detect_arm_from_keys(keys, {k: tuple(v.shape) for k, v in arm.state_dict().items()})
        assert detected == (None if name == "prompt" else name)
        if name == "xattn":
            assert dkw == {"blocks": [0, 1], "enable_mouse": False, "enable_keyboard": True, "window_frames": 1, "hidden_size": 128}
            arm2 = build_arm("xattn", _FakeDiT(), {**dkw, "heads_num": 4})
            arm2.load_state_dict(arm.state_dict(), strict=True)


def test_parse_args_and_detect():
    from actionrope.train import parse_args
    a = parse_args(["--arm", "xattn", "--output", "x", "--arm_kwargs", '{"enable_mouse": false, "window_frames": 1}'])
    assert a.text_mode == "scene" and a.mask_channel is False and a.arm_kwargs == {"enable_mouse": False, "window_frames": 1}
    assert a.text_table.endswith("text_table_actionrope.pt")
    a = parse_args(["--arm", "prompt", "--output", "x"])
    assert a.text_mode == "scene_action" and a.mask_channel is False and a.text_table.endswith("text_table_eybx_mirror.pt")
    a = parse_args(["--arm", "arope", "--output", "x"])
    assert a.mask_channel is True and a.text_mode == "scene"
    assert detect_arm_from_keys([]) == (None, {})
    assert detect_arm_from_keys(["action_embedders.3.weight"]) == ("linear", {})
    assert detect_arm_from_keys(["embedder.0.weight", "proj.1.bias"], {"embedder.0.weight": (3072, 64)}) == ("adaln", {"use_delta": False})
    with pytest.raises(ValueError):
        detect_arm_from_keys(["something.weight"])


def test_make_action_inputs_matches_infer_labels():
    from actionrope.infer import parse_actions, velocity_to_label
    ws = np.array([69.0, 46.0])
    frame_off, vel = parse_actions("right:1.0:10,up:1.0:11", ws)
    off = torch.tensor(G.frames_to_cells(frame_off), dtype=torch.float32).unsqueeze(0)
    labels = [velocity_to_label(v, mirror=False) for v in vel]
    ai = make_action_inputs(off, torch.tensor([labels]))
    assert set(ai) == {"offset_px", "offset_tok", "delta_tok", "action_idx"}
    assert ai["action_idx"].tolist() == [[7] * 10 + [1] * 11]
    assert torch.equal(ai["offset_tok"] * 32, ai["offset_px"]) and torch.all(ai["delta_tok"][:, 0] == 0)
    # 真实向右 ⇒ dx > 0,标签 7(right)与 ACTION_DIRS 同向;镜像标签是 3
    assert ai["delta_tok"][0, 5, 0] > 0 and velocity_to_label(vel[0], mirror=True) == 3