Download code/baseline/__init__.py from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 8.94 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/baseline/__init__.py
- Command line
-
hf download hf://teawhite/ActionRoPE/code/baseline/__init__.py
-
curl -L -o __init__.py https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/baseline/__init__.py
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", | |
| ] | |