Instructions to use teawhite/ActionRoPE with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Wan2.2
How to use teawhite/ActionRoPE with Wan2.2:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
Download code/tests/test_infer.py from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 8.19 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/test_infer.py
- Command line
-
hf download hf://teawhite/ActionRoPE/code/tests/test_infer.py
-
curl -L -o test_infer.py https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/test_infer.py
8.19 kB
| """推理脚本的自测(一张 H200 足够:DiT bf16 ~10 GB + VAE,峰值 ~16 GB)。 | |
| CUDA_VISIBLE_DEVICES=4 /opt/dlami/nvme/zhiyangdeng/ActionRoPE/.venv/bin/python -m pytest \ | |
| /opt/dlami/nvme/zhiyangdeng/ActionRoPE/tests/test_infer.py -s -v | |
| 三件事各自独立,可用 -k 单跑: | |
| a. VAE 尺度:用 pipe.vae 编码一条 clip 的 mp4,与 data/latent/pool 同名 latent 比相对误差 | |
| —— 训练 latent 与推理首帧 latent 必须同一尺度(mean/std 归一化)。 | |
| b. 位移实测函数在 GT clip 上自检:帧 0→80 的实测背景位移 ≈ −frame_offset_px[80]。 | |
| c. 端到端 smoke:基座权重 --replay 2 步采样,视频可解码、81 帧、832×480。 | |
| 实测数字追加写到 $AROPE_TEST_NUMBERS(默认 $AROPE_TEST_DIR=outputs/test_artifacts 下的 json),供报告汇总。 | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import os | |
| import numpy as np | |
| import pytest | |
| import torch | |
| os.environ.setdefault("DIFFSYNTH_SKIP_DOWNLOAD", "True") | |
| ROOT = "/opt/dlami/nvme/zhiyangdeng/ActionRoPE" | |
| # 测试产物只落在项目内(outputs/test_artifacts),不写任何机器相关的临时目录 | |
| 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, "infer_numbers.json")) | |
| # val_eybx 里 21 个 cell 动作全为 7、rt_ratio=1、sidecar valid 的 clip(tidal_flats,户外纹理多,SIFT 好测) | |
| CLIP = os.environ.get("AROPE_TEST_CLIP", "clip_Eybx_200000958_000273") | |
| import sys # noqa: E402 | |
| sys.path.insert(0, ROOT) | |
| from actionrope import geometry as G # noqa: E402 | |
| from actionrope import infer as I # noqa: E402 | |
| def record(**kv): | |
| data = {} | |
| try: | |
| with open(NUMBERS_PATH) as fp: | |
| data = json.load(fp) | |
| except (FileNotFoundError, json.JSONDecodeError): | |
| pass | |
| data.update(kv) | |
| 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, default=float) | |
| print("\n[numbers]", json.dumps(kv, ensure_ascii=False, default=float)) | |
| def sidecar(): | |
| return torch.load(I.SIDECAR, map_location="cpu", weights_only=False) | |
| def gt_frames(): | |
| return I.read_frames(I.find_clip_mp4(CLIP)) | |
| # ---------------------------------------------------------------- 纯 CPU ---- | |
| def test_parse_actions(sidecar): | |
| ws = np.array(sidecar["meta"]["walk_speed_px_s"]) | |
| fo, vel = I.parse_actions("right:1.0:21", ws) | |
| assert fo.shape == (81, 2) and np.allclose(fo[0], 0) | |
| assert np.allclose(fo[80], [ws[0] * 80 / 16, 0]) # i/16 × speed × 方向 | |
| assert np.allclose(fo[16], [ws[0], 0]) | |
| off = G.frames_to_cells(fo) | |
| assert np.allclose(off[0], 0) and off[1, 0] > 0 | |
| # cell 0–9 = 帧 0–36 向右(36 步),cell 10–20 = 帧 37–80 向左(44 步) | |
| fo2, _ = I.parse_actions("right:1.0:10,left:1.0:11", ws) | |
| assert fo2[36, 0] == pytest.approx(ws[0] * 36 / 16) and fo2[80, 0] == pytest.approx(-ws[0] * 8 / 16) | |
| fo3, _ = I.parse_actions("71,0:10,-71,0:11", ws) # px/s 直给,逗号形式 | |
| assert fo3[36, 0] == pytest.approx(71 * 36 / 16) | |
| with pytest.raises(ValueError): | |
| I.parse_actions("right:1.0:20", ws) | |
| # plain 臂标签:真实向右 ⇒ 数据集的镜像标签是 3 "moving left";literal 则是 7 | |
| assert I.velocity_to_label(np.array([71.0, 0]), mirror=True) == 3 | |
| assert I.velocity_to_label(np.array([71.0, 0]), mirror=False) == 7 | |
| assert I.velocity_to_label(np.zeros(2), mirror=True) == 0 | |
| def test_compose_prompts(): | |
| p = I.compose_cell_prompts("tidal_flats", "burn:wills_farm:8:2", [7] * 21) | |
| assert p[7].endswith("strewn with wreckage, the player is moving right.") | |
| assert "burning away into a warm farmland hamlet" in p[8] and "burning away into" in p[9] | |
| assert p[10] == "In a 2.5D top down view, in a warm farmland hamlet, the player is moving right." | |
| # 三种串都在文本表里(转场表 342 条 + 纯场景 18 条),不必回退 T5 | |
| table = torch.load(I.TEXT_TABLE["arope"], map_location="cpu", weights_only=False) | |
| from actionrope.prompts import strip_action | |
| assert all(strip_action(x) in table for x in p) | |
| table_eybx = torch.load(I.TEXT_TABLE["plain"], map_location="cpu", weights_only=False) | |
| assert all(x in table_eybx for x in p) | |
| def test_gt_displacement(sidecar, gt_frames): | |
| """b. GT clip 上实测 ≈ −frame_offset_px[80](sidecar 已乘 gain,符号为玩家真实位移)。""" | |
| fo = sidecar["clips"][CLIP]["frame_offset_px"].numpy() | |
| expect = -fo[80] | |
| meas = I.measure_displacement(I.to_gray(gt_frames)) | |
| assert meas["sift"] is not None | |
| m = np.array(meas["sift"]) | |
| record(gt_clip=CLIP, gt_expected_bg_shift_80=expect.tolist(), gt_measured_sift=m.tolist(), | |
| gt_measured_phase=meas.get("phase_direct"), gt_ratio_x=float(m[0] / expect[0])) | |
| assert np.sign(m[0]) == np.sign(expect[0]) | |
| assert abs(m[0] - expect[0]) <= 0.10 * abs(expect[0]) + 5 # 单 clip 增益散布 ±10% | |
| assert abs(m[1] - expect[1]) <= 15 | |
| # ---------------------------------------------------------------- GPU ---- | |
| def models(): | |
| return I.load_models(None, False, "cuda") | |
| def pipe(models): | |
| return models[0] | |
| def test_vae_scale(pipe, gt_frames): | |
| """a. 推理侧 VAE 编码与数据集 latent 同尺度。bf16 编码 ~0.8%,fp32 ~0.3%(数据集是 fp32 编码后存 bf16); | |
| 推理首帧默认 fp32 编码,与训练时看到的帧 0 latent 只差 bf16 存储舍入。""" | |
| ref = torch.load(os.path.join(I.DATA, "latent", "pool", CLIP + ".pt"), map_location="cpu", weights_only=False)["latent"].float() | |
| z = I.encode_video(pipe, gt_frames, tiled=False)[0].float().cpu() | |
| assert z.shape == ref.shape == (48, 21, 30, 52) | |
| rel = ((z - ref).norm() / ref.norm()).item() | |
| z32 = I.encode_video(pipe, gt_frames, tiled=False, fp32=True)[0].float().cpu() | |
| rel32 = ((z32 - ref).norm() / ref.norm()).item() | |
| z0 = I.encode_first_frame(pipe, gt_frames[0])[0, :, 0].float().cpu() # 默认 fp32 | |
| rel0 = ((z0 - ref[:, 0]).norm() / ref[:, 0].norm()).item() | |
| z0_bf16 = I.encode_first_frame(pipe, gt_frames[0], fp32=False)[0, :, 0].float().cpu() | |
| rel0_bf16 = ((z0_bf16 - ref[:, 0]).norm() / ref[:, 0].norm()).item() | |
| # 首帧单独编码 == 整段同精度编码的 cell 0(因果 VAE);z0 已 cast 到 bf16,所以整段也 cast 后比 | |
| same_as_video = ((z0 - z32[:, 0].to(torch.bfloat16).float()).norm() / z32[:, 0].norm()).item() | |
| assert pipe.vae.model.encoder.conv1.weight.dtype == torch.bfloat16 # fp32 编码后 VAE 已切回 bf16 | |
| dec = I.decode_video(pipe, ref.unsqueeze(0).to(torch.bfloat16).cuda(), tiled=False) | |
| ps = [I.psnr(dec[i], gt_frames[i]) for i in range(81)] | |
| record(vae_encode_rel_l2_bf16=rel, vae_encode_rel_l2_fp32=rel32, | |
| vae_first_frame_rel_l2_cell0_fp32=rel0, vae_first_frame_rel_l2_cell0_bf16=rel0_bf16, | |
| vae_first_frame_vs_video_cell0=same_as_video, vae_ref_std=ref.std().item(), | |
| vae_roundtrip_psnr_frame0=ps[0], vae_roundtrip_psnr_mean=float(np.mean(ps))) | |
| assert rel < 0.02 and rel0_bf16 < 0.02 # 同一 mean/std 尺度;bf16 舍入量级 | |
| assert rel32 < 0.006 and rel0 < 0.006 # fp32 编码只剩 bf16 存储舍入(~0.3%) | |
| assert same_as_video == 0.0 # 首帧单独编码 == 整段编码的 cell 0(因果 VAE) | |
| assert ps[0] > 28 | |
| def test_smoke_replay(models): | |
| """c. 端到端:基座 + --replay,2 步采样,产出可解码的 81 帧 832×480 mp4。""" | |
| out = os.path.join(TEST_DIR, "test_smoke_replay.mp4") | |
| info = I.main(["--replay", CLIP, "--steps", "2", "--out", out], models=models) | |
| frames = I.read_frames(out) | |
| assert frames.shape == (81, 480, 832, 3) | |
| assert info["psnr_frame0"] > 28 | |
| record(smoke_replay_steps=2, smoke_sample_sec_2steps=info["sample_sec"], | |
| smoke_peak_mem_gb=info["peak_mem_gb"], smoke_measured=info["measured"]["sift"]) | |