ActionRoPE / code /baseline /__init__.py
teawhite's picture
add docs+code
880dff9 verified
Raw History Blame Contribute Delete
8.94 kB
"""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",
]