"""Audit: how much of the null-step story is best-validation-checkpoint selection? `05_train_sae.py` saves the best-val-MSE epoch and restores it at the end, and it makes the cosine LR schedule depend on the *requested* epoch budget. Both facts are recorded per SAE in `config["best_epoch"]`. This script reads that field back and crosses it against measured drift. Three things it establishes, all of which bear on claims in the paper: 1. Budget sweep. Every layer-18 sweep run at 15+ epochs restored epoch 10, so the sweep varies the LR-schedule horizon over a near-constant number of *retained* epochs. It is not a dose-response in optimizer steps. 2. "Discrete basins". Null-step drift is almost perfectly predicted by the selected epoch, at every layer. The multimodality is checkpoint selection, not basins. 3. Epoch-matched shares. Comparing each real transition only against null seeds that selected the same epoch removes the layer-12 anomaly and tightens the headline. python3 checkpoint_selection_audit.py --root /path/to/drift_run For a size-limited supplement without the checkpoint tensors, pass the exported ``drift_out/checkpoint_selection.csv`` with ``--metadata``. A full checkout can create that file with ``--export-metadata``. """ import argparse import glob import os from pathlib import Path import re import pandas as pd REAL = {"sft": "saes_local/sae_sft/sae_sft_step29_layer{L}.pt", "flexible": "saes_local/sae_flexible/sae_ppo_step10_layer{L}.pt", "strict": "saes_local/sae_strict/k64/sae_ppo_step10_layer{L}.pt"} def best_epoch(path): import torch return torch.load(path, map_location="cpu", weights_only=False).get("config", {}).get("best_epoch") def relative_key(path, root): return path.relative_to(root).as_posix() def null_seeds(root, L, metadata_paths=()): paths = {root / f"null_ctrl/sae_instruct_base_layer{L}_nullstep.pt"} paths.update(Path(p) for p in glob.glob( str(root / f"null_seeds/sae_instruct_base_layer{L}_null_s*.pt"))) pat = re.compile(rf"^null_seeds/sae_instruct_base_layer{L}_null_s\d+\.pt$") paths.update(root / p for p in metadata_paths if pat.match(p)) out = [] for p in paths: m = re.search(r"_null_s(\d+)", str(p)) out.append(("s%d" % int(m.group(1)) if m else "null_s0", p)) return sorted(out, key=lambda x: int(re.search(r"\d+", x[0]).group(0))) def main(): ap = argparse.ArgumentParser() ap.add_argument("--root", type=Path, default=None, help="drift_run directory (auto-detected in the full repository)") ap.add_argument("--metadata", type=Path, default=None, help="checkpoint_selection.csv exported from a full artifact") ap.add_argument("--export-metadata", type=Path, default=None, help="write relative checkpoint paths and retained epochs for a lite artifact") args = ap.parse_args() if args.root is not None: root = args.root.expanduser().resolve() else: here = Path(__file__).resolve() candidates = [here.parent.parent, here.parent / "drift_run", here.parent.parent / "drift_run", Path.cwd()] root = next((p for p in candidates if (p / "drift_out/null_replicates.csv").exists()), None) if root is None: ap.error("could not find drift_run; pass --root /path/to/drift_run") rep = pd.read_csv(root / "drift_out/null_replicates.csv") metadata_path = args.metadata if metadata_path is None and (root / "drift_out/checkpoint_selection.csv").exists(): metadata_path = root / "drift_out/checkpoint_selection.csv" epoch_map = {} if metadata_path is not None: meta = pd.read_csv(metadata_path) epoch_map = dict(zip(meta.relative_path.astype(str), meta.best_epoch.astype(int))) seen = {} def epoch(path): key = relative_key(path, root) value = epoch_map.get(key) if value is None: if not path.exists(): raise FileNotFoundError(f"missing checkpoint and metadata row: {key}") value = best_epoch(path) seen[key] = int(value) return int(value) metadata_paths = set(epoch_map) print("=" * 72) print("1. Budget sweep at layer 18: requested epochs vs the epoch actually kept") print("=" * 72) budget_paths = {Path(p) for p in glob.glob( str(root / "null_budget/sae_instruct_base_layer18_e*_s*.pt"))} budget_pat = re.compile(r"^null_budget/sae_instruct_base_layer18_e\d+_s\d+\.pt$") budget_paths.update(root / p for p in metadata_paths if budget_pat.match(p)) for p in sorted(budget_paths, key=lambda q: (int(re.search(r"_e(\d+)_", str(q)).group(1)), str(q))): req = int(re.search(r"_e(\d+)_", str(p)).group(1)) print(f" {os.path.basename(p):44s} requested={req:>3} kept={epoch(p)}") print(" -> the sweep varies the LR-schedule horizon, not the retained step count.\n") print("=" * 72) print("2. Null-step drift vs selected epoch") print("=" * 72) for L in (6, 12, 18, 23): rows = [] for lab, p in null_seeds(root, L, metadata_paths): d = rep[(rep.layer == L) & (rep.kind == "nullstep") & (rep.label == lab)] rows.append({"label": lab, "epoch": epoch(p), "drift": float(d.drift.iloc[0])}) df = pd.DataFrame(rows).sort_values("drift") g = df.groupby("epoch").drift.agg(["size", "mean", "min", "max"]).round(4) print(f" L{L}:") print(" " + g.to_string().replace("\n", "\n ")) if df.epoch.nunique() > 1: print(f" Spearman(epoch, drift) = {df.epoch.corr(df.drift, method='spearman'):.3f}") print() print("=" * 72) print("3. Share of the real first transition, pooled vs epoch-matched") print("=" * 72) for L in (6, 12, 18, 23): nd = pd.DataFrame([ {"epoch": epoch(p), "drift": float(rep[(rep.layer == L) & (rep.kind == "nullstep") & (rep.label == lab)].drift.iloc[0])} for lab, p in null_seeds(root, L, metadata_paths)]) print(f" L{L}:") for ch, pat in REAL.items(): p = root / pat.format(L=L) d = rep[(rep.layer == L) & (rep.kind == "real") & (rep.label == ch)] key = relative_key(p, root) if (not os.path.exists(p) and key not in epoch_map) or not len(d): continue e, rd = epoch(p), float(d.drift.iloc[0]) m = nd[nd.epoch == e] pooled = 100 * nd.drift.mean() / rd matched = 100 * m.drift.mean() / rd if len(m) else float("nan") print(f" {ch:9s} kept_epoch={e:>3} real_drift={rd:.4f} " f"pooled={pooled:6.1f}% epoch-matched={matched:6.1f}% (n_matched={len(m)})") print() if args.export_metadata is not None: out = args.export_metadata.expanduser() out.parent.mkdir(parents=True, exist_ok=True) pd.DataFrame([{"relative_path": p, "best_epoch": e} for p, e in sorted(seen.items())]).to_csv(out, index=False) print(f"wrote {out}") if __name__ == "__main__": main()