"""Read back the 15M LoRA calibration and report the knee, the lr boundary and sigma_ft.""" import glob import json import statistics as st from pathlib import Path CK = { "natural": "/workspace/runs/mae-validate/natural-lr0.0005/exports/step_00039833", "d7": "/workspace/runs/mae-validate/d7-lr0.0007/exports/step_00039833", "d185": "/workspace/runs/mae-validate/d185-lr0.0007/exports/step_00039833", } def load(ck, tag): p = Path(CK[ck]) / f"lora_{tag}.json" return json.loads(p.read_text()) if p.exists() else None print("[1] epoch knee (natural, lr 3e-3) probe mAP on this ckpt = 0.18675") print(f" {'epochs':>7}{'val mAP':>11}{'train mAP':>12}{'train s':>10}") knee = {} for ep in (1, 2, 3, 4, 6): d = load("natural", f"ep{ep}") if d: knee[ep] = d["lora/map"] print(f" {ep:>7}{d['lora/map']:>11.5f}{d['final_train_map']:>12.5f}" f"{d['train_seconds']:>10.1f}") best = max(knee, key=knee.get) print(f" -> knee at {best} epochs") print("\n[2] lr boundary (natural, 3 epochs)") for lr in (1e-3, 3e-3, 1e-2): d = load("natural", f"lr{lr:g}") if d: flag = " <-- diverged" if d["lora/map"] < 0.05 else "" print(f" lr {lr:<8g} val {d['lora/map']:.5f}{flag}") print("\n[3] sigma_ft (3 checkpoints x 3 seeds, 3 epochs, lr 3e-3)") grid = {} for ck in CK: vals = [load(ck, f"var-s{s}")["lora/map"] for s in (0, 1, 2) if load(ck, f"var-s{s}")] if vals: grid[ck] = vals print(f" {ck:<9}{' '.join(f'{v:.5f}' for v in vals)} " f"mean {st.mean(vals):.5f} sd {st.stdev(vals):.5f}") within = [st.variance(v) for v in grid.values() if len(v) > 2] sigma_ft = (sum(within) / len(within)) ** 0.5 means = [st.mean(v) for v in grid.values()] between = st.stdev(means) # Common-mode: how much of sigma_ft moves every checkpoint together, and so # cancels in a contrast that pins the fine-tune seed across arms. by_seed = [st.mean([grid[c][i] for c in grid]) for i in range(3)] common = st.stdev(by_seed) print(f"\n sigma_ft pooled within-checkpoint : {sigma_ft:.5f} (df={2*len(within)})") print(f" of which common-mode across ckpts: {common:.5f} (cancels under a pinned ft seed)") print(f" between-checkpoint sd : {between:.5f}") if between > sigma_ft: r1 = (between**2 - sigma_ft**2) / between**2 r3 = 3 * r1 / (1 + 2 * r1) print(f" single-run reliability : {r1:.3f} ({r3:.3f} at 3 seeds)") print(f"\n scale_base reference: sigma_ft 0.00233, reliability 0.31 (0.58 at 3 seeds)," f"\n knee 3 epochs, lr 3e-3, collapse between 3e-3 and 1e-2") Path("/workspace/analysis/lora_calibration_15m.json").write_text(json.dumps( {"epoch_knee": knee, "sigma_ft_grid": grid, "sigma_ft": sigma_ft, "common_mode": common, "between_ckpt_sd": between}, indent=1)) print("\nwrote /workspace/analysis/lora_calibration_15m.json")