File size: 4,343 Bytes
add24bd
5673379
add24bd
 
 
 
5673379
add24bd
 
 
 
 
5673379
 
add24bd
 
 
 
 
 
 
 
 
5673379
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
add24bd
 
 
 
 
5673379
 
 
add24bd
 
 
 
 
 
 
5673379
add24bd
 
 
5673379
 
 
 
 
 
 
 
 
 
 
 
 
add24bd
 
 
 
 
 
 
 
 
5673379
add24bd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/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()