#!/usr/bin/env python3 """核对 train.py 导出的 safetensors 能否原样装回 WanModel(推理脚本走的就是这条路)。 .venv/bin/python tests/check_ckpt.py outputs//step-N.safetensors --arm arope .venv/bin/python tests/check_ckpt.py outputs//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()