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()