ActionRoPE / code /tests /check_ckpt.py
teawhite's picture
add docs+code
880dff9 verified
Raw History Blame Contribute Delete
5.3 kB
#!/usr/bin/env python3
"""核对 train.py 导出的 safetensors 能否原样装回 WanModel(推理脚本走的就是这条路)。
.venv/bin/python tests/check_ckpt.py outputs/<run>/step-N.safetensors --arm arope
.venv/bin/python tests/check_ckpt.py outputs/<run>/step-N.safetensors # 臂按 ckpt metadata / 键集自动识别
做三件事:
1. 键集:原模型 + install_arope(mask_channel) 后 load_state_dict(strict=True) 必须成功;
plain 臂的 ckpt 不许出现 arope_mask_embedding.*。baseline 臂(linear / xattn / adaln)的 `arm.*` 键
按 baseline.build_arm 建臂后 strict 装回,并报告臂参数里非零的张量数(零初始化的臂训过后应当有变化)。
2. 权重确实变了:与原始权重比较,报告改动的张量数与最大 |Δ|(优化器没生效时这里全是 0)。
3. dtype / 文件大小 / metadata。
全程 CPU,不占卡。
"""
from __future__ import annotations
import argparse
import json
import os
import sys
os.environ.setdefault("DIFFSYNTH_SKIP_DOWNLOAD", "True")
import torch
from safetensors.torch import load_file
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, ROOT)
from actionrope.arope import MASK_EMBEDDING_NAME, install_arope # noqa: E402
from actionrope.train import MODEL_DIR, load_dit # noqa: E402
from baseline import ARM_NAMES, build_arm, detect_arm_from_ckpt, split_arm_keys # noqa: E402
def main():
ap = argparse.ArgumentParser()
ap.add_argument("ckpt")
ap.add_argument("--arm", choices=list(ARM_NAMES), default=None, help="不给 ⇒ 按 ckpt metadata / arm.* 键集识别")
ap.add_argument("--no_mask_channel", action="store_true")
ap.add_argument("--model_dir", default=MODEL_DIR)
ap.add_argument("--json_out", default=None)
args = ap.parse_args()
full_sd = load_file(args.ckpt)
try:
from safetensors import safe_open
with safe_open(args.ckpt, "pt") as f:
metadata = f.metadata() or {}
except Exception:
metadata = {}
arm_name, arm_kwargs = detect_arm_from_ckpt(full_sd, metadata)
if args.arm is None:
args.arm = arm_name or "arope"
elif arm_name and arm_name != args.arm and {arm_name, args.arm} != {"plain", "prompt"}:
raise SystemExit(f"--arm {args.arm} 与 ckpt 识别出的臂 {arm_name!r} 不符")
mask_channel = args.arm == "arope" and not args.no_mask_channel
sd, arm_sd = split_arm_keys(full_sd)
dit = load_dit(args.model_dir, device="cpu")
base = {k: v.clone() for k, v in dit.state_dict().items()}
install_arope(dit, mask_channel=mask_channel)
expected = set(dit.state_dict().keys())
keys = set(sd)
mask_keys = {k for k in keys if k.startswith(MASK_EMBEDDING_NAME)}
res = {
"ckpt": args.ckpt, "arm": args.arm, "mask_channel": mask_channel,
"size_gb": os.path.getsize(args.ckpt) / 1e9, "n_tensors": len(sd),
"missing": sorted(expected - keys), "unexpected": sorted(keys - expected),
"mask_keys": sorted(mask_keys),
"dtypes": sorted({str(v.dtype) for v in sd.values()}),
}
if args.arm != "arope":
assert not mask_keys, f"{args.arm} 臂 ckpt 不应含 {mask_keys}"
else:
assert mask_keys == {f"{MASK_EMBEDDING_NAME}.weight", f"{MASK_EMBEDDING_NAME}.bias"} if mask_channel else not mask_keys
dit.load_state_dict(sd, strict=True) # 键集或形状不符会在这里炸
res["strict_load_ok"] = True
# baseline 臂:按 ckpt 反推的配置建臂,strict 装回;零初始化的臂训过后 proj 之类应当已非零
arm = build_arm(args.arm, dit, arm_kwargs)
res.update({"arm_kwargs": arm_kwargs, "n_arm_tensors": len(arm_sd),
"n_new_params": arm.n_new_params() if arm is not None else 0})
if arm is not None:
arm.load_state_dict(arm_sd, strict=True)
res["arm_strict_load_ok"] = True
res["arm_nonzero_tensors"] = sum(int(bool((v != 0).any())) for v in arm_sd.values())
res["arm_absmax"] = max(v.float().abs().max().item() for v in arm_sd.values()) if arm_sd else 0.0
elif arm_sd:
raise SystemExit(f"ckpt 含 {len(arm_sd)} 个 arm.* 键,但 --arm {args.arm} 没有参数")
n_changed, max_abs, sum_abs, n_el = 0, 0.0, 0.0, 0
for k, v0 in base.items():
d = (sd[k].float() - v0.float()).abs()
m = d.max().item()
if m > 0:
n_changed += 1
max_abs = max(max_abs, m)
sum_abs += d.sum().item()
n_el += d.numel()
res.update({"n_base_tensors": len(base), "n_changed_tensors": n_changed,
"max_abs_delta": max_abs, "mean_abs_delta": sum_abs / max(n_el, 1)})
if mask_channel:
res["mask_weight_absmax"] = sd[f"{MASK_EMBEDDING_NAME}.weight"].float().abs().max().item()
try:
from safetensors import safe_open
with safe_open(args.ckpt, "pt") as f:
res["metadata"] = f.metadata()
except Exception as e: # metadata 不是硬要求
res["metadata_error"] = str(e)
print(json.dumps(res, indent=1, ensure_ascii=False))
if args.json_out:
with open(args.json_out, "w", encoding="utf-8") as fh:
json.dump(res, fh, indent=1, ensure_ascii=False)
if __name__ == "__main__":
main()