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