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",
]