RefDiffNet / extract_prebackbone_weights.py
ShashwatGupta23's picture
Cursor
RefDiffNet: standalone A11_CA prebackbone Gradio Space
5673379
Raw
History Blame Contribute Delete
4.34 kB
#!/usr/bin/env python3
"""Extract prebackbone-only weights from a full YOLO best.pt checkpoint."""
from __future__ import annotations
import argparse
import os
import sys
from pathlib import Path
import torch
from a11_ca import build_prebackbone
_HERE = Path(__file__).resolve().parent
def _clean_state_dict(state: dict) -> dict:
"""Drop thop profiling keys (total_ops, total_params) not in nn.Module."""
skip = ("total_ops", "total_params")
return {k: v for k, v in state.items() if not any(k == s or k.endswith(f".{s}") for s in skip)}
def _resolve_ultralytics_root() -> Path | None:
env = os.environ.get("ULTRALYTICS_ROOT", "").strip()
candidates = []
if env:
candidates.append(Path(env).expanduser())
candidates.extend((_HERE.parent / "ultralytics", _HERE / "vendor"))
for base in candidates:
if (base / "ultralytics" / "__init__.py").exists():
return base
if base.name == "ultralytics" and (base / "__init__.py").exists():
return base.parent
return None
def _load_full_checkpoint(full_ckpt: Path) -> tuple[dict, object]:
"""Load best.pt; uses VYOLO ultralytics fork when present (pickled EMA model)."""
root = _resolve_ultralytics_root()
if root is not None:
root_str = str(root.resolve())
if root_str not in sys.path:
sys.path.insert(0, root_str)
from ultralytics.nn.tasks import torch_safe_load
ckpt, _ = torch_safe_load(str(full_ckpt))
return ckpt, ckpt.get("ema") or ckpt.get("model")
raise SystemExit(
f"Cannot unpickle {full_ckpt} without the VYOLO ultralytics fork.\n"
"One-time extraction from the training repo:\n"
" cd /path/to/VYOLO/hf_prebackbone_demo\n"
" ULTRALYTICS_ROOT=../ultralytics python extract_prebackbone_weights.py --ckpt ../ultralytics/.../best.pt\n"
"Inference only needs weights/prebackbone_a11_ca.pt (no ultralytics)."
)
def extract(
full_ckpt: Path,
out_path: Path,
prebackbone_name: str | None = None,
) -> Path:
ckpt, model = _load_full_checkpoint(full_ckpt)
if model is None:
raise RuntimeError(f"No model/ema in {full_ckpt}")
train_args = ckpt.get("train_args") or {}
name = (prebackbone_name or train_args.get("prebackbone") or "A11_CA").upper()
channels = int(train_args.get("channels", 3) or 3)
pb = getattr(model, "prebackbone", None)
if pb is None:
raise RuntimeError(f"No prebackbone submodule in {full_ckpt}")
state = _clean_state_dict(pb.state_dict())
# Verify keys match standalone architecture before saving
standalone = build_prebackbone(name, channels=channels)
if standalone is None:
raise RuntimeError(f"Unsupported prebackbone type: {name}")
expected = set(standalone.state_dict().keys())
got = set(state.keys())
if expected != got:
missing = expected - got
extra = got - expected
raise RuntimeError(
f"State dict mismatch for {name}: missing={sorted(missing)[:5]}, extra={sorted(extra)[:5]}"
)
payload = {
"prebackbone": name,
"channels": channels,
"state_dict": state,
"source_checkpoint": str(full_ckpt.resolve()),
}
out_path.parent.mkdir(parents=True, exist_ok=True)
torch.save(payload, out_path)
n_params = sum(t.numel() for t in state.values())
size_mb = out_path.stat().st_size / (1024 * 1024)
print(f"Saved {out_path} ({size_mb:.2f} MB, {n_params:,} parameters, type={name})")
return out_path
def main() -> None:
p = argparse.ArgumentParser(description="Extract prebackbone-only weights from best.pt")
p.add_argument(
"--ckpt",
type=Path,
default=_HERE / "weights" / "best.pt",
help="Full YOLO checkpoint (best.pt)",
)
p.add_argument(
"--out",
type=Path,
default=_HERE / "weights" / "prebackbone_a11_ca.pt",
help="Output path for prebackbone-only weights",
)
p.add_argument("--name", type=str, default=None, help="Override prebackbone type (default: from train_args)")
args = p.parse_args()
if not args.ckpt.exists():
raise SystemExit(f"Checkpoint not found: {args.ckpt}")
extract(args.ckpt, args.out, args.name)
if __name__ == "__main__":
main()