Instructions to use teawhite/ActionRoPE with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Wan2.2
How to use teawhite/ActionRoPE with Wan2.2:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
Download code/tests/check_ckpt.py from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 5.3 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/check_ckpt.py
- Command line
-
hf download hf://teawhite/ActionRoPE/code/tests/check_ckpt.py
-
curl -L -o check_ckpt.py https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/tests/check_ckpt.py
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() | |