File size: 4,331 Bytes
ffdcfe7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Probe every final export under a run root, in parallel across GPUs.

Reports AS-20K probe mAP against the patching's own random-init floor. That
ratio -- not reconstruction loss -- is the gate. This project has twice been
burned by treating MAE val_loss as a quality signal: cluster-weight-channels.md
measures it at rho=0.72 with the train/eval distribution-match channel, and
regmix-regression-256.md found it uncorrelated with downstream score at
rho=-0.034 over 256 mixtures.
"""
import argparse
import json
import queue
import subprocess
import threading
from pathlib import Path

REPO = "/workspace/code/eat-map-regmix"
FLOORS = {"p16": 0.03861, "p4x64": 0.04867}


def latest_export(run: Path) -> Path | None:
    exports = sorted((run / "exports").glob("step_*")) if (run / "exports").exists() else []
    return exports[-1] if exports else None


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("root")
    ap.add_argument("--gpus", default="0,1,2,3")
    ap.add_argument("--all-steps", action="store_true", help="probe every export, not just the last")
    # The probe is CPU-bound, not GPU-bound: feature extraction is ~15 s but the
    # ridge solve and the 527-class mAP run on BLAS, which by default grabs one
    # thread per core. Four concurrent probes then ask for 512 threads on 128
    # cores and spend 90% of wall time descheduling. Capping threads and raising
    # concurrency is worth ~8x here.
    ap.add_argument("--per-gpu", type=int, default=2, help="concurrent probes per GPU")
    ap.add_argument("--threads", type=int, default=12, help="BLAS threads per probe")
    args = ap.parse_args()

    jobs: queue.Queue = queue.Queue()
    for run in sorted(Path(args.root).iterdir()):
        if not run.is_dir():
            continue
        exports = (sorted((run / "exports").glob("step_*")) if args.all_steps
                   else [latest_export(run)])
        for exp in [e for e in exports if e]:
            if not (exp / "probe.json").exists():
                jobs.put((run.name, exp))

    total = jobs.qsize()
    print(f"{total} exports to probe")
    lock = threading.Lock()

    def worker(gpu):
        while True:
            try:
                name, exp = jobs.get_nowait()
            except queue.Empty:
                return
            threads = str(args.threads)
            proc = subprocess.run(
                ["python", "-m", "eatmap.cli.probe", "--export", str(exp)],
                cwd=REPO, capture_output=True,
                env={**__import__("os").environ, "CUDA_VISIBLE_DEVICES": str(gpu),
                     "PYTHONPATH": REPO,
                     "OMP_NUM_THREADS": threads, "OPENBLAS_NUM_THREADS": threads,
                     "MKL_NUM_THREADS": threads, "NUMEXPR_NUM_THREADS": threads},
            )
            with lock:
                if proc.returncode:
                    print(f"  FAIL {name} {exp.name}: {proc.stderr.decode()[-200:]}")
            jobs.task_done()

    threads = [threading.Thread(target=worker, args=(int(g),))
               for g in args.gpus.split(",") for _ in range(args.per_gpu)]
    for t in threads:
        t.start()
    for t in threads:
        t.join()

    rows = []
    for run in sorted(Path(args.root).iterdir()):
        if not run.is_dir():
            continue
        for exp in sorted((run / "exports").glob("step_*")):
            pj = exp / "probe.json"
            if not pj.exists():
                continue
            d = json.loads(pj.read_text())
            tag = "p4x64" if "p4x64" in run.name else "p16"
            floor = FLOORS[tag]
            status = json.loads((run / "status.json").read_text()) if (run / "status.json").exists() else {}
            rows.append((run.name, int(exp.name.split("_")[1]), d["probe/map"],
                         d["probe/map"] / floor, status.get("val/loss", float("nan"))))

    print(f"\n{'run':<24}{'step':>8}{'probe mAP':>12}{'x floor':>10}{'val loss':>10}")
    for name, step, m, ratio, vl in sorted(rows, key=lambda r: -r[2]):
        print(f"{name:<24}{step:>8}{m:>12.5f}{ratio:>10.2f}{vl:>10.4f}")
    Path(args.root, "probe_summary.json").write_text(json.dumps(
        [{"run": n, "step": s, "probe_map": m, "x_floor": r, "val_loss": v}
         for n, s, m, r, v in rows], indent=1))


if __name__ == "__main__":
    main()