"""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