Download code/tests/test_baseline_integration.py from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 6.99 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/test_baseline_integration.py
- Command line
-
hf download hf://teawhite/ActionRoPE/code/tests/test_baseline_integration.py
-
curl -L -o test_baseline_integration.py https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/test_baseline_integration.py
6.99 kB
| """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)) | |
| 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 | |