ActionRoPE / code /tests /test_baseline_integration.py
teawhite's picture
add docs+code
880dff9 verified
Raw History Blame Contribute Delete
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))
@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