"""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=`;`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 集合。 """ 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", ]