File size: 11,296 Bytes
ba7ef34 4ea7bda ba7ef34 4ea7bda ba7ef34 4ea7bda ba7ef34 4ea7bda ba7ef34 4ea7bda ba7ef34 4ea7bda ba7ef34 4ea7bda ba7ef34 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 | """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()
|