sae-null-step / tables /scripts /make_numbers.py
anonymous
Upload tables/scripts/make_numbers.py with huggingface_hub
4ea7bda verified
Raw History Blame Contribute Delete
11.3 kB
"""Emit the numbers the body quotes, as LaTeX macros, from the run outputs.
One file, \input by paper.tex:
sections/numbers.tex scalar macros (\PCTONE, \DIRECTIONPARA, ...) and
\TABBODY, the rows of Table 2
Anything here is a number the reviewer can trace: it is read from the CSVs the
runs wrote, never typed. Re-run after any rerun and diff the two files.
python3 scripts/make_numbers.py
"""
import itertools
import math
import pathlib
import numpy as np
import pandas as pd
import torch
HERE = pathlib.Path(__file__).resolve().parent
PAPER = HERE.parent
DRIFT = PAPER.parent / "sae_rl" / "drift_run"
LAYERS = [6, 12, 18]
OLD = {6: "jr_L6_topk", 12: "jr_L12_topk", 18: "jumprelu_sft_topk", 23: "jr_L23_topk"}
METRICS = PAPER.parent / "sae_rl" / "sae_rl" / "results" / "sae_checkpoint_metrics.csv"
def fidelity_bounds():
"""The flat-fidelity envelope at layers 6-18, rounded OUTWARD.
Rounding outward matters. An earlier draft wrote "loss recovery never falls
below 0.937" from a minimum of 0.93660, which rounds to 0.937 and is
therefore false as a bound. floor/ceil at three or four decimals keeps the
printed number a bound rather than a rounding of one.
"""
import math
m = pd.read_csv(METRICS)
s = m[m.layer.isin([6, 12, 18])]
return dict(
n=len(s),
nmse=math.ceil(s.nmse.max() * 1e4) / 1e4,
frac_rec=math.floor(s.frac_rec.min() * 1e3) / 1e3,
dead=math.ceil(s.dead_latents_pct.max() * 10) / 10,
)
def chain_table(layer):
"""Real and null chains at one layer, replication rows preferred."""
frames = []
new = DRIFT / f"real_arm_L{layer}" / "real_arm_replicate.csv"
have = set()
if new.exists():
n = pd.read_csv(new, converters={"chain": str})
have = set(map(tuple, n[["chain", "seed"]].drop_duplicates().values))
frames.append(n[["chain", "seed", "step", "dec_cos", "kept_epoch"]])
old = DRIFT / OLD[layer] / "jumprelu_sft_chain.csv"
if old.exists():
o = pd.read_csv(old, converters={"chain": str})
o = o[(o.arm == "TopK_k64") & (o.chain != "base")]
o = o[[tuple(r) not in have for r in o[["chain", "seed"]].values]]
frames.append(o[["chain", "seed", "step", "dec_cos", "kept_epoch"]])
return pd.concat(frames, ignore_index=True)
def arm(df, chain, step):
return 1.0 - df[(df.chain == chain) & (df.step == step)].dec_cos.to_numpy()
def table_and_shares():
"""Table rows, plus everything the body says about spread and separation.
The separation statistic is per cell, not global: end-to-end drift compounds
over seven steps, so its seed ranges are an order of magnitude wider than the
first transition's, and one global "largest range" compared against one global
"smallest gap" would be a false claim in both directions. What holds in every
cell is that the two arms' seed ranges are disjoint, and the ratio of the gap
between their means to the wider of the two ranges is the honest summary.
"""
rows, s1, se, counts, resid, ratios, disjoint = [], [], [], [], {}, [], []
for L in LAYERS:
df = chain_table(L)
cells = {}
for tag, step in (("1", 1), ("7", 7)):
r, n = arm(df, "real", step), arm(df, "null", step)
cells[tag] = (r, n)
resid[(L, tag)] = np.array([rr - nn for rr in r for nn in n])
wider = max(r.max() - r.min(), n.max() - n.min())
ratios.append((r.mean() - n.mean()) / max(wider, 1e-12))
disjoint.append(bool(r.min() > n.max() or r.max() < n.min()))
r1, n1 = cells["1"]
r7, n7 = cells["7"]
s1.append(100 * n1.mean() / r1.mean())
se.append(100 * n7.mean() / r7.mean())
counts.append((len(r1), len(n1), len(r7), len(n7)))
rows.append(f"L{L} & {r1.mean():.3f} & {n1.mean():.3f} & {s1[-1]:.0f}\\% & & "
f"{r7.mean():.3f} & {n7.mean():.3f} & {se[-1]:.0f}\\% \\\\")
return rows, s1, se, counts, (ratios, all(disjoint)), resid
def direction(layer=18):
"""Pairwise dictionary agreement within and across the two arms."""
wdir = DRIFT / f"real_arm_L{layer}" / "weights"
if not wdir.exists():
return None
out = {}
for step in (1, 7):
W = {}
for f in sorted(wdir.glob(f"*_step{step}.pt")):
key = (f.stem.split("_")[0], int(f.stem.split("_")[1][1:]))
W[key] = torch.nn.functional.normalize(torch.load(f).float(), dim=-1)
if not W:
continue
real = sorted(k for k in W if k[0] == "real")
null = sorted(k for k in W if k[0] == "null")
if len(real) < 2 or len(null) < 2:
continue
def cos(a, b):
return float((W[a] * W[b]).sum(-1).mean())
rr = [cos(a, b) for a, b in itertools.combinations(real, 2)]
nn = [cos(a, b) for a, b in itertools.combinations(null, 2)]
rn = [cos(a, b) for a in real for b in null]
out[step] = dict(n_real=len(real), n_null=len(null),
rr=np.array(rr), nn=np.array(nn), rn=np.array(rn),
separated=min(min(rr), min(nn)) > max(rn))
return out or None
def rng(a):
return f"${a.mean():.3f}$ $[{a.min():.3f},{a.max():.3f}]$"
def main():
rows, s1, se, counts, (ratios, all_disjoint), resid = table_and_shares()
# The rows travel as a macro rather than an \input file: \input inside a
# tabular leaves a token in front of \bottomrule, whose \noalign is then
# misplaced. The body must end with \\ and no trailing space for the same
# reason, which a macro gives for free.
tabbody = "\\newcommand{\\TABBODY}{%\n" + "\n".join(rows) + "}"
L = []
L.append("% Generated by scripts/make_numbers.py -- do not edit by hand.")
L.append(tabbody)
# Round a reported range OUTWARD, never with round(). "86--93%" when the
# observed floor is 85.8% claims a stronger lower bound than we measured,
# and "84--94%" when the observed ceiling is 94.5% claims a tighter upper
# bound. Same rule as fidelity_bounds().
def rng_pct(v):
return (f"{math.floor(min(v))}\\text{{--}}{math.ceil(max(v))}\\%")
L.append(f"\\newcommand{{\\PCTONE}}{{{rng_pct(s1)}}}")
L.append(f"\\newcommand{{\\PCTEND}}{{{rng_pct(se)}}}")
if not all_disjoint:
raise SystemExit("a cell's real and null seed ranges overlap -- "
"the separation claim in Table 1's caption is not true")
# Is the headline share an artifact of reading the first transition, where
# the model has moved least? Report the share at every step so the reader
# can see it is flat.
per_step = []
for lay in LAYERS:
df = chain_table(lay)
sh = []
for st in range(1, 8):
r, n = arm(df, "real", st), arm(df, "null", st)
if len(r) and len(n):
sh.append(100 * n.mean() / r.mean())
if sh:
per_step.append((lay, math.floor(min(sh)), math.ceil(max(sh))))
L.append("\\newcommand{\\SHAREBYSTEP}{" + ", ".join(
f"${lo}\\text{{--}}{hi}\\%$ at layer {lay}" for lay, lo, hi in per_step) + "}")
# The body quotes the envelope over every step and layer; the per-layer
# breakdown above is kept for the appendix.
if per_step:
blo = min(lo for _, lo, _ in per_step)
bhi = max(hi for _, _, hi in per_step)
L.append(f"\\newcommand{{\\SHAREBAND}}{{${blo}\\text{{--}}{bhi}\\%$}}")
print("null share across all seven steps: " + ", ".join(
f"L{lay} {lo}-{hi}%" for lay, lo, hi in per_step))
f = fidelity_bounds()
L.append(f"\\newcommand{{\\NEVALS}}{{{f['n']}}}")
L.append(f"\\newcommand{{\\NMSEMAX}}{{{f['nmse']:.4f}}}")
L.append(f"\\newcommand{{\\FRACRECMIN}}{{{f['frac_rec']:.3f}}}")
L.append(f"\\newcommand{{\\DEADMAX}}{{{f['dead']:.1f}}}")
L.append(f"\\newcommand{{\\SEPRATIO}}{{{math.floor(min(ratios))}$ to ${math.ceil(max(ratios))}}}")
# NSEEDS describes every row of Table 1, so it has to be the range across
# layers, not layer 18's counts. L6 and L12 carry 3 null chains, L18 has 4;
# printing "4/4" would overstate replication for two of the three rows.
reals = {c[0] for c in counts} | {c[2] for c in counts}
nulls = {c[1] for c in counts} | {c[3] for c in counts}
fmt = lambda x: f"${min(x)}$" if len(x) == 1 else f"${min(x)}$ to ${max(x)}$"
L.append(f"\\newcommand{{\\NSEEDS}}{{{fmt(reals)} real and {fmt(nulls)} null "
"chains per layer}")
rerun_dead = {}
for lay in LAYERS:
csvp = DRIFT / f"real_arm_L{lay}" / "real_arm_replicate.csv"
if csvp.exists():
v = pd.read_csv(csvp)["dead_frac"] * 100
rerun_dead[lay] = (math.floor(v.min()), math.ceil(v.max()))
if 18 in rerun_dead and 12 in rerun_dead:
L.append("\\newcommand{\\RERUNDEAD}{$%d$ to $%d\\%%$ at layer 18 and $%d$ to $%d\\%%$ "
"at layer 12}" % (*rerun_dead[18], *rerun_dead[12]))
print(f"rerun dead fractions: L18 {rerun_dead[18]}, L12 {rerun_dead[12]}")
q = resid[(18, "1")]
L.append(f"\\newcommand{{\\RESIDL}}{{${q.mean():.4f}$ "
f"$[{q.min():.4f},{q.max():.4f}]$}}")
L.append(f"\\newcommand{{\\RESIDRATIO}}{{{math.floor(ratios[LAYERS.index(18)*2])}}}")
d = direction(18)
if d and 7 in d:
end, one = d[7], d.get(1, d[7])
L.append("\\newcommand{\\DIRECTIONPARA}{%\n"
f"At layer 18, over {end['n_real']} real and {end['n_null']} null chains, two "
f"real endpoints agree at a mean matched-latent cosine of "
f"${end['rr'].mean():.3f}$, two null endpoints at ${end['nn'].mean():.3f}$, and "
f"a real endpoint with a null endpoint at only ${end['rn'].mean():.3f}$. Every "
"within-arm pair therefore exceeds every real--null pair, both after the first "
"step and after the seventh (Figure~\\ref{fig:main}b). Changes in the model "
"activations affect the direction of SAE movement, even though retraining "
"accounts for most of its magnitude.}")
print(f"direction: step7 rr={end['rr'].mean():.4f} nn={end['nn'].mean():.4f} "
f"rn={end['rn'].mean():.4f} separated={end['separated']}")
else:
L.append("\\newcommand{\\DIRECTIONPARA}{}")
print("WARNING: no direction data yet; macros left empty")
(PAPER / "sections" / "numbers.tex").write_text("\n".join(L) + "\n")
for L_, a, b, c in zip(LAYERS, s1, se, counts):
print(f"L{L_}: first {a:.1f}% (real n={c[0]}, null n={c[1]}) "
f"end {b:.1f}% (real n={c[2]}, null n={c[3]})")
print(f"fidelity envelope at L6-18 over {f['n']} evals: NMSE <= {f['nmse']}, "
f"frac_rec >= {f['frac_rec']}, dead <= {f['dead']}%")
print(f"gap / wider seed range, per cell: "
f"{', '.join(f'{r:.1f}' for r in ratios)} (all disjoint: {all_disjoint})")
q = resid[(18, "1")]
print(f"L18 first-transition residual: {q.mean():.5f} [{q.min():.5f}, {q.max():.5f}]")
print(f"-> PCTONE {min(s1):.0f}-{max(s1):.0f}% PCTEND {min(se):.0f}-{max(se):.0f}%")
if __name__ == "__main__":
main()