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
File size: 8,939 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 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 | """baseline 动作臂的注册表与装配入口(baseline/SPEC.md「训练 / 推理接入」)。
from baseline import ARMS, build_arm
arm = build_arm("linear", dit) # 建参数 → install(dit) → 搬到 dit 的 device / dtype
五臂里 `arope`(本项目)与 `plain`(= prompt 臂旧名)不走 ActionArm:arope 的动作在 RoPE 里,
plain 与 prompt 行为完全相同(train.py 里 plain 仍按旧路径 arm=None 跑,prompt 走 PromptArm 的零钩子,
两者前向逐位一致;保留 plain 是为了旧 ckpt / 旧脚本不改)。
ckpt 约定:arm 参数以 `arm.` 前缀与 DiT 权重写在同一个 safetensors 里(train.py::export_dit_weights),
safetensors metadata 里记 `arm=<name>`;`detect_arm_from_keys` 只看键集就能判断臂名与 xattn 的结构配置,
metadata 缺失(手工拼的 ckpt)时推理仍能自动装对。
"""
from __future__ import annotations
import json
from typing import Iterable
import torch
from baseline.base import ActionArm
# 各臂的类延迟导入:prompt.py 会把 data/code 加进 sys.path、xattn.py 依赖 einops,
# 不用某个臂时不必付这些 import 的代价,也避免一个臂坏了拖垮别的臂。
_ARM_MODULES = {
"linear": ("baseline.linear", "LinearArm"),
"xattn": ("baseline.xattn", "XAttnArm"),
"prompt": ("baseline.prompt", "PromptArm"),
"adaln": ("baseline.adaln", "AdaLNArm"),
}
# 不走 ActionArm 的两个臂(train.py / infer.py 的 --arm 也接受它们)
NATIVE_ARMS = ("arope", "plain")
# 全部合法臂名(CLI choices 用)
ARM_NAMES = tuple(NATIVE_ARMS) + tuple(_ARM_MODULES)
# 动作走文本的臂:text_mode=scene_action(动作词翻回真实方向);其余剥动作从句
TEXT_ARMS = ("plain", "prompt")
# ckpt 里 arm 参数的前缀
ARM_PREFIX = "arm."
class _LazyArms(dict):
"""ARMS[name] → 臂类;第一次取时才 import 对应模块。"""
def __missing__(self, name):
if name not in _ARM_MODULES:
raise KeyError(f"未知动作臂 {name!r},可选 {list(_ARM_MODULES)}(arope / plain 不走 ActionArm)")
mod_name, cls_name = _ARM_MODULES[name]
import importlib
cls = getattr(importlib.import_module(mod_name), cls_name)
self[name] = cls
return cls
def __contains__(self, name):
return name in _ARM_MODULES
def __iter__(self):
return iter(_ARM_MODULES)
def __len__(self):
return len(_ARM_MODULES)
ARMS = _LazyArms()
def arm_text_mode(name: str) -> str:
"""臂 → dataset 的 text_mode:动作走文本的臂保留动作从句,其余剥掉。"""
return "scene_action" if name in TEXT_ARMS else "scene"
def parse_arm_kwargs(spec) -> dict:
"""--arm_kwargs 的 JSON 串(或已是 dict / None)→ dict。"""
if spec is None or spec == "":
return {}
if isinstance(spec, dict):
return dict(spec)
kw = json.loads(spec)
if not isinstance(kw, dict):
raise ValueError(f"--arm_kwargs 应为 JSON 对象,得到 {spec!r}")
return kw
def build_arm(name: str, dit, arm_kwargs: dict | str | None = None) -> ActionArm:
"""按名字建臂、install(dit)、搬到 dit 的 device / dtype。arope / plain 没有 ActionArm ⇒ 返回 None。
install 由各臂自己按 dit.dim / 层数建参数并 .to(dit);这里再统一 .to 一次是兜底:
某个臂的 install 若只建了参数没搬卡(例如 prompt 臂无参数),DeepSpeed 下混着 CPU 参数会炸。
"""
if name in NATIVE_ARMS:
return None
cls = ARMS[name]
arm = cls(**parse_arm_kwargs(arm_kwargs))
arm.install(dit)
ref = dit.patch_embedding.weight
arm.to(device=ref.device, dtype=ref.dtype)
return arm
# --------------------------------------------------------------------------
# ckpt 键集 → 臂
# --------------------------------------------------------------------------
def split_arm_keys(sd: dict) -> tuple[dict, dict]:
"""safetensors 的 state_dict → (DiT 部分, arm 部分(已去掉 'arm.' 前缀))。"""
dit_sd, arm_sd = {}, {}
for k, v in sd.items():
if k.startswith(ARM_PREFIX):
arm_sd[k[len(ARM_PREFIX):]] = v
else:
dit_sd[k] = v
return dit_sd, arm_sd
def detect_arm_from_keys(arm_keys: Iterable[str], shapes: dict | None = None) -> tuple[str | None, dict]:
"""去掉前缀后的 arm 键集 → (臂名, 构造 kwargs)。没有 arm 键 ⇒ (None, {}):可能是 arope / plain / prompt,需看 metadata。
各臂的键名互不重叠(linear: action_embedders.*;xattn: action_modules.*;adaln: embedder.* / proj.*),
xattn 的结构配置能从形状反推:keyboard_attn_kv.weight 的 in_features = hidden_size(128)·window_frames,
有无 mouse_mlp.* = enable_mouse,有无 keyboard_embed.* = enable_keyboard,blocks = action_modules.<i> 的 i 集合。
"""
keys = list(arm_keys)
if not keys:
return None, {}
heads = {k.split(".")[0] for k in keys}
if heads == {"action_embedders"}:
return "linear", {}
if heads == {"action_modules"}:
blocks = sorted({int(k.split(".")[1]) for k in keys})
kw: dict = {"blocks": blocks}
kw["enable_mouse"] = any(".mouse_mlp." in k for k in keys)
kw["enable_keyboard"] = any(".keyboard_embed." in k for k in keys)
if shapes is not None and kw["enable_keyboard"]:
kv_key = f"action_modules.{blocks[0]}.keyboard_attn_kv.weight"
emb_key = f"action_modules.{blocks[0]}.keyboard_embed.2.weight"
if kv_key in shapes and emb_key in shapes:
hidden = int(shapes[emb_key][0])
kw["window_frames"] = int(shapes[kv_key][1]) // hidden
kw["hidden_size"] = hidden
elif shapes is not None and kw["enable_mouse"]:
mlp_key = f"action_modules.{blocks[0]}.mouse_mlp.0.weight"
if mlp_key in shapes:
# in_features = 2·window_frames + dim;dim 从 proj_mouse 的 out_features 拿
proj_key = f"action_modules.{blocks[0]}.proj_mouse.weight"
dim = int(shapes[proj_key][0]) if proj_key in shapes else 3072
kw["window_frames"] = (int(shapes[mlp_key][1]) - dim) // 2
return "xattn", kw
if heads <= {"embedder", "proj"}:
kw = {}
if shapes is not None and "embedder.0.weight" in shapes:
n_in = int(shapes["embedder.0.weight"][1])
# n_in = n_axes · freq_dim_per_axis;默认 freq 32 ⇒ 4 轴 128 / 2 轴 64
kw["use_delta"] = n_in >= 4 * 32
return "adaln", kw
raise ValueError(f"无法从 arm 键集识别臂:顶层键 {sorted(heads)}")
def detect_arm_from_ckpt(sd: dict, metadata: dict | None = None) -> tuple[str | None, dict]:
"""(state_dict, safetensors metadata) → (臂名, 构造 kwargs)。
metadata 里的 arm 名优先(它能区分 arope / plain / prompt 这三个无 arm 键的臂);有 arm 键时
以键集反推的结构配置为准,并核对与 metadata 一致。metadata 缺失且无 arm 键 ⇒ (None, {}),由调用方决定默认。
"""
_, arm_sd = split_arm_keys(sd)
shapes = {k: tuple(v.shape) for k, v in arm_sd.items()}
name_keys, kw = detect_arm_from_keys(arm_sd.keys(), shapes)
name_meta = (metadata or {}).get("arm")
if name_meta and name_keys and name_meta != name_keys:
raise ValueError(f"ckpt metadata 说 arm={name_meta!r},键集却是 {name_keys!r}")
if metadata and metadata.get("arm_kwargs"):
# 训练时的构造参数(JSON)比形状反推更完整(例如 heads_num)
try:
kw = {**parse_arm_kwargs(metadata["arm_kwargs"]), **kw}
except (ValueError, json.JSONDecodeError):
pass
return name_meta or name_keys, kw
def make_action_inputs(offset_px: torch.Tensor, action_idx: torch.Tensor | None = None) -> dict:
"""offset_px [b, 21, 2] (+ action_idx [b, 21]) → SPEC 的 action_inputs dict(offset_tok / delta_tok 由此派生)。
dataset.py 与 infer.py 共用,保证两边 delta_tok 的定义一致(cell 0 = 0,其余为相邻 cell 之差)。
"""
from actionrope.arope import PX_PER_TOKEN
off_px = torch.as_tensor(offset_px, dtype=torch.float32)
off_tok = off_px / PX_PER_TOKEN
delta = torch.zeros_like(off_tok)
delta[..., 1:, :] = off_tok[..., 1:, :] - off_tok[..., :-1, :]
out = {"offset_px": off_px, "offset_tok": off_tok, "delta_tok": delta}
if action_idx is not None:
out["action_idx"] = torch.as_tensor(action_idx, dtype=torch.int64)
return out
__all__ = [
"ARMS", "ARM_NAMES", "NATIVE_ARMS", "TEXT_ARMS", "ARM_PREFIX", "ActionArm",
"build_arm", "arm_text_mode", "parse_arm_kwargs", "split_arm_keys",
"detect_arm_from_keys", "detect_arm_from_ckpt", "make_action_inputs",
]
|