sae-null-step / tables /scripts /checkpoint_selection_audit.py
anonymous
Add tables
ba7ef34 verified
Raw History Blame Contribute Delete
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()