Download tables/scripts/checkpoint_selection_audit.py from sae-anon/sae-null-step: direct link, hf CLI and curl.
- Browser
- Download file 7.29 kB
-
https://huggingface.co/sae-anon/sae-null-step/resolve/main/tables/scripts/checkpoint_selection_audit.py
- Command line
-
hf download hf://sae-anon/sae-null-step/tables/scripts/checkpoint_selection_audit.py
-
curl -L -o checkpoint_selection_audit.py https://huggingface.co/sae-anon/sae-null-step/resolve/main/tables/scripts/checkpoint_selection_audit.py
7.29 kB
| """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() | |