amirali1985 commited on
Commit
33aa037
·
1 Parent(s): 157a25b

add auto-interp table to appendix

Browse files
Files changed (1) hide show
  1. app.py +28 -204
app.py CHANGED
@@ -21,20 +21,7 @@ MODEL_REPO = "thoughtworks/arithmetic-sorl"
21
  # LaTeX scratchpad content — edit here, copy from dashboard into Overleaf
22
  # ═══════════════════════════════════════════════════════════════════
23
 
24
- LATEX_ARITHMETIC_SETUP = r"""% ── Arithmetic case study ──────────────────────────────────────────────────
25
-
26
- \begin{figure}[t]
27
- \centering
28
- \includegraphics[width=\linewidth]{figures/fig_arithmetic_example.pdf}
29
- \caption{Addition $959{,}271 + 040{,}756 = 1{,}000{,}027$ — a four-deep carry cascade.
30
- Each answer-digit position is annotated with its Quirke subtask label (coloured box)
31
- and the \sorl{} abstraction token assigned by the model (dashed purple box).
32
- The carry chain ($d_1$--$d_5$) is highlighted; \sorl{} uses a consistent token
33
- (\texttt{t2}) at cascade positions and distinct tokens elsewhere.}
34
- \label{fig:arithmetic-example}
35
- \end{figure}
36
-
37
- \subsection{Case study: six-digit addition and subtraction}
38
  \label{sec:arithmetic}
39
 
40
  Six-digit addition and subtraction provides a setting where the internal
@@ -80,191 +67,7 @@ The margin grows with cascade depth, consistent with explicit carry/borrow routi
80
  See Appendix \ref{app:arithmetic} further details on SORL interpretability, including a demonstration of auto-interp ~\citep{bills2023language_models_explain_neurons}, token specializations and polysemantic tokens.
81
  """
82
 
83
- LATEX_FIGURE_EXAMPLE = r"""% fig_arithmetic_example.tex
84
- % Usage in paper: \input{figures/fig_arithmetic_example/fig_arithmetic_example.tex}
85
- % Required packages: tikz, xcolor
86
- %
87
- % Shows 959271 + 040756 = 1000027 (4-deep carry cascade).
88
- % Token assignments from model add_sub_sorl_v1_abs30_K1_100K (K=1, abs30).
89
-
90
- \begin{figure}[t]
91
- \centering
92
- \begin{tikzpicture}[
93
- % ── node styles ─────────────────────────────────────────────────────
94
- digit/.style={
95
- draw=#1!60!gray, fill=#1, rounded corners=2pt,
96
- minimum width=1.05cm, minimum height=0.62cm,
97
- font=\small\bfseries, inner sep=2pt, text=#1!20!black,
98
- },
99
- subtask/.style={
100
- draw=#1!60!gray, fill=#1, rounded corners=2pt,
101
- minimum width=1.05cm, minimum height=0.55cm,
102
- font=\footnotesize, inner sep=2pt, text=#1!20!black,
103
- },
104
- token/.style={
105
- draw=violet!55, fill=violet!8, rounded corners=2pt, dashed,
106
- minimum width=1.05cm, minimum height=0.55cm,
107
- font=\footnotesize\bfseries, inner sep=2pt, text=violet!70!black,
108
- },
109
- rowlabel/.style={font=\footnotesize\itshape, text=gray!70!black, anchor=east},
110
- poslabel/.style={font=\scriptsize, text=gray!60!black},
111
- carry/.style={->, >=stealth, thick, color=orange!70!red,
112
- shorten <=3pt, shorten >=3pt},
113
- ]
114
-
115
- % ── colours (Quirke subtask families) ───────────────────────────────────
116
- \colorlet{cSA}{green!22!white}
117
- \colorlet{cSC}{yellow!48!white}
118
- \colorlet{cUC}{blue!22!white}
119
- \colorlet{cUS}{blue!36!white}
120
-
121
- % ── column spacing ───────────────────────────────────────────────────────
122
- \def\cs{1.45} % inter-column distance (cm)
123
-
124
- % ── data (d0 = MSB/overflow on left, d6 = LSB on right) ─────────────────
125
- % d0 d1 d2 d3 d4 d5 d6
126
- % --- 9 5 9 2 7 1 (Addend A)
127
- % --- 0 4 0 7 5 6 (Addend B)
128
- % 1 0 0 0 0 2 7 (Answer)
129
- % UC US US US US SC SA (Subtask)
130
- % t2 t2 t6 t2 t1 t16 t3 (SoRL token)
131
-
132
- % ── position labels ──────────────────────────────────────────────────────
133
- \foreach \i/\lbl in {0/$d_0$,1/$d_1$,2/$d_2$,3/$d_3$,4/$d_4$,5/$d_5$,6/$d_6$}{
134
- \node[poslabel] at (\i*\cs, 3.15) {\lbl};
135
- }
136
-
137
- % ── row labels ───────────────────────────────────────────────────────────
138
- \node[rowlabel] at (-0.65, 2.5) {Addend $A$};
139
- \node[rowlabel] at (-0.65, 1.8) {Addend $B$};
140
- \node[rowlabel] at (-0.65, 0.85) {Answer};
141
- \node[rowlabel] at (-0.65, 0.1) {Subtask};
142
- \node[rowlabel] at (-0.65,-0.65) {Token};
143
-
144
- % ── operand digits (d1..d6; d0 is the overflow, has no operand digits) ───
145
- \foreach \i/\d in {1/9,2/5,3/9,4/2,5/7,6/1}{
146
- \node[font=\small, text=gray!30!black] at (\i*\cs, 2.5) {\d};
147
- }
148
- \node[font=\small\bfseries, text=gray!50!black] at (-0.4*\cs, 1.8) {$+$};
149
- \foreach \i/\d in {1/0,2/4,3/0,4/7,5/5,6/6}{
150
- \node[font=\small, text=gray!30!black] at (\i*\cs, 1.8) {\d};
151
- }
152
-
153
- % ── horizontal rule ──────────────────────────────────────────────────────
154
- \draw[gray!50, thin] (-0.55*\cs, 1.38) -- (6.55*\cs, 1.38);
155
-
156
- % ── answer digit boxes ───────────────────────────────────────────────────
157
- \node[digit=cUC] (a0) at (0*\cs, 0.85) {1};
158
- \node[digit=cUS] (a1) at (1*\cs, 0.85) {0};
159
- \node[digit=cUS] (a2) at (2*\cs, 0.85) {0};
160
- \node[digit=cUS] (a3) at (3*\cs, 0.85) {0};
161
- \node[digit=cUS] (a4) at (4*\cs, 0.85) {0};
162
- \node[digit=cSC] (a5) at (5*\cs, 0.85) {2};
163
- \node[digit=cSA] (a6) at (6*\cs, 0.85) {7};
164
-
165
- % ── subtask boxes ────────────────────────────────────────────────────────
166
- \node[subtask=cUC] at (0*\cs, 0.1) {UC};
167
- \node[subtask=cUS] at (1*\cs, 0.1) {US};
168
- \node[subtask=cUS] at (2*\cs, 0.1) {US};
169
- \node[subtask=cUS] at (3*\cs, 0.1) {US};
170
- \node[subtask=cUS] at (4*\cs, 0.1) {US};
171
- \node[subtask=cSC] at (5*\cs, 0.1) {SC};
172
- \node[subtask=cSA] at (6*\cs, 0.1) {SA};
173
-
174
- % ── SoRL token boxes ─────────────────────────────────────────────────────
175
- \foreach \i/\t in {0/t2,1/t2,2/t6,3/t2,4/t1,5/t16,6/t3}{
176
- \node[token] at (\i*\cs, -0.65) {\texttt{\t}};
177
- }
178
-
179
- % ── carry arrows (cascade flows right→left: d5→d4→d3→d2→d1→d0) ──────────
180
- \foreach \fr/\to in {5/4, 4/3, 3/2, 2/1, 1/0}{
181
- \draw[carry] (a\fr.west) -- (a\to.east);
182
- }
183
-
184
- % ── cascade bracket + label ──────────────────────────────────────────────
185
- \draw[orange!60!red, thin]
186
- (a5.north west) -- ++(0, 0.22)
187
- -- (a0.north east) -- ++(0,-0.22);
188
- \node[font=\scriptsize\itshape, text=orange!60!red]
189
- at (2.5*\cs, 1.35) {carry cascade};
190
-
191
- % ── legend ───────────────────────────────────────────────────────────────
192
- \matrix[
193
- matrix of nodes,
194
- nodes={font=\scriptsize, inner sep=2pt, anchor=west},
195
- row sep=1pt, column sep=4pt,
196
- anchor=south east,
197
- ] at (6*\cs + 0.6, -1.05) {
198
- \node[digit=cSA, minimum width=0.45cm, minimum height=0.3cm,
199
- font=\scriptsize] {}; &
200
- \node {SA --- simple add}; &
201
- \node[digit=cUC, minimum width=0.45cm, minimum height=0.3cm,
202
- font=\scriptsize] {}; &
203
- \node {UC --- uses carry}; \\
204
- \node[digit=cSC, minimum width=0.45cm, minimum height=0.3cm,
205
- font=\scriptsize] {}; &
206
- \node {SC --- generates carry}; &
207
- \node[digit=cUS, minimum width=0.45cm, minimum height=0.3cm,
208
- font=\scriptsize] {}; &
209
- \node {US --- cascade}; \\
210
- \node[token, minimum width=0.45cm, minimum height=0.3cm,
211
- font=\scriptsize] {}; &
212
- \node[text=violet!70!black] {\sorl{} token}; & & \\
213
- };
214
-
215
- \end{tikzpicture}
216
- \caption{%
217
- Six-digit addition $959{,}271 + 040{,}756 = 1{,}000{,}027$, a four-deep
218
- carry cascade. At each answer-digit position \sorl{} assigns one
219
- abstraction token (bottom row). Tokens \texttt{t2} and \texttt{t6}
220
- cluster on cascade positions (UC/US); \texttt{t16} marks the carry
221
- source (SC); \texttt{t3} marks the trivial position (SA).
222
- Token assignments from model \texttt{add\_sub\_sorl\_v1\_abs30\_K1\_100K}
223
- (K=1, 30-token codebook).%
224
- }
225
- \label{fig:arithmetic-example}
226
- \end{figure}
227
- """
228
-
229
- LATEX_TABLE_UNDERSIZED = r"""% tab:undersized-wins — SoRL vs SFT on undersized architectures
230
- % Generated by arithmetic/paper/results/result_low_data_wins/run.py
231
- % Requires: \usepackage{booktabs}, \usepackage{xcolor}
232
 
233
- \begin{table}[t]
234
- \centering
235
- \small
236
- \begin{tabular}{llrrrr}
237
- \toprule
238
- Architecture & Data & Baseline & SoRL & Gap & C6 gap \\
239
- \midrule
240
- \texttt{1L/2H/256d} & 10K & 10\% & \textbf{19\%} & \textcolor{green!50!black}{\textbf{+9\%}} & \textcolor{green!50!black}{\textbf{+18\%}} \\
241
- & 25K & 32\% & 26\% & $-7\%$ & \textcolor{green!50!black}{\textbf{+10\%}} \\
242
- & 50K & 44\% & \textbf{65\%} & \textcolor{green!50!black}{\textbf{+21\%}} & \textcolor{green!50!black}{\textbf{+34\%}} \\
243
- & 100K & 49\% & \textbf{65\%} & \textcolor{green!50!black}{\textbf{+16\%}} & \textcolor{green!50!black}{\textbf{+31\%}} \\
244
- \midrule
245
- \texttt{1L/3H/510d} & 10K & 36\% & \textbf{52\%} & \textcolor{green!50!black}{\textbf{+16\%}} & \textcolor{green!50!black}{\textbf{+30\%}} \\
246
- & 25K & 46\% & \textbf{60\%} & \textcolor{green!50!black}{\textbf{+14\%}} & \textcolor{green!50!black}{\textbf{+22\%}} \\
247
- & 50K & 53\% & \textbf{72\%} & \textcolor{green!50!black}{\textbf{+19\%}} & \textcolor{green!50!black}{\textbf{+38\%}} \\
248
- & 100K & 67\% & \textbf{83\%} & \textcolor{green!50!black}{\textbf{+16\%}} & \textcolor{green!50!black}{\textbf{+26\%}} \\
249
- \midrule
250
- \texttt{2L/1H/128d} & 10K & 16\% & \textbf{36\%} & \textcolor{green!50!black}{\textbf{+21\%}} & \textcolor{green!50!black}{\textbf{+39\%}} \\
251
- & 25K & 40\% & \textbf{55\%} & \textcolor{green!50!black}{\textbf{+15\%}} & \textcolor{green!50!black}{\textbf{+23\%}} \\
252
- & 50K & 59\% & \textbf{87\%} & \textcolor{green!50!black}{\textbf{+28\%}} & \textcolor{green!50!black}{\textbf{+50\%}} \\
253
- & 75K & 75\% & \textbf{87\%} & \textcolor{green!50!black}{\textbf{+12\%}} & \textcolor{green!50!black}{\textbf{+5\%}} \\
254
- & 100K & 73\% & \textbf{95\%} & \textcolor{green!50!black}{\textbf{+22\%}} & \textcolor{green!50!black}{\textbf{+33\%}} \\
255
- \bottomrule
256
- \end{tabular}
257
- \caption{\sorl{} ($K{=}1$, $|\mathcal{A}|{=}30$) vs.\ \sft{} baseline on
258
- undersized architectures across data sizes.
259
- \textbf{Gap} = overall accuracy gain; \textbf{C6 gap} = gain on
260
- 6-deep carry cascades (the hardest split).
261
- \sorl{} wins in \textbf{12 of 13} (architecture, data-size) pairs;
262
- the single exception is \texttt{1L/2H/256d} at 25K, where the model
263
- is undertrained (accuracy still rising at epoch 20).
264
- \sorl{} wins on C6 in \textbf{all 13} configurations.}
265
- \label{tab:undersized-wins}
266
- \end{table}
267
- """
268
 
269
  LATEX_APPENDIX = r"""\section{Arithmetic case study: interpretability analysis}
270
  \label{app:arithmetic}
@@ -592,13 +395,33 @@ variable. The specialist tokens concentrate at mid-sequence positions ($d_2$-$d_
592
  We implement a light version of the automated interpretation procedure of \citet{bills2023language}.
593
  For each active token, we collect the $N{=}10$ examples from the evaluation set where the model assigned it with highest softmax confidence,
594
  then ask \texttt{claude-haiku} to produce a one-sentence role description.
 
595
 
596
- % [PLACEHOLDER — run experiments/11_auto_interp/run.py to generate table.tex]
597
- % Then paste the output of table.tex here:
598
- %
599
- % \input{experiments/11_auto_interp/table.tex}
600
- %
601
- % Expected columns: Token | Top subtask (purity) | Mean conf. | Auto-interpretation
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
602
 
603
  \begin{tcolorbox}[colback=gray!6, colframe=gray!40,
604
  fonttitle=\bfseries\small, title={Finding \#7: Automated interpretation matches ground-truth subtask labels},
@@ -647,6 +470,7 @@ then ask \texttt{claude-haiku} to produce a one-sentence role description.
647
 
648
  Evaluation uses fixed-length autoregressive decoding (no teacher forcing):
649
  the model generates answer digits $d_0 \to d_6$ using its own predictions, with abstraction tokens inserted via the \sorl{} search-then-recurse procedure (matching training). Accuracy is measured on 100 held-out examples per split (seed 42; \texttt{thoughtworks/arithmetic-sorl-data}).
 
650
  """
651
 
652
 
 
21
  # LaTeX scratchpad content — edit here, copy from dashboard into Overleaf
22
  # ═══════════════════════════════════════════════════════════════════
23
 
24
+ LATEX_ARITHMETIC_SETUP = r"""\subsection{Case study: six-digit addition and subtraction}
 
 
 
 
 
 
 
 
 
 
 
 
 
25
  \label{sec:arithmetic}
26
 
27
  Six-digit addition and subtraction provides a setting where the internal
 
67
  See Appendix \ref{app:arithmetic} further details on SORL interpretability, including a demonstration of auto-interp ~\citep{bills2023language_models_explain_neurons}, token specializations and polysemantic tokens.
68
  """
69
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
70
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
71
 
72
  LATEX_APPENDIX = r"""\section{Arithmetic case study: interpretability analysis}
73
  \label{app:arithmetic}
 
395
  We implement a light version of the automated interpretation procedure of \citet{bills2023language}.
396
  For each active token, we collect the $N{=}10$ examples from the evaluation set where the model assigned it with highest softmax confidence,
397
  then ask \texttt{claude-haiku} to produce a one-sentence role description.
398
+ Table~\ref{tab:auto-interp} shows results for the 8 highest-confidence tokens.
399
 
400
+ \begin{table}[ht]
401
+ \centering\small
402
+ \begin{tabular}{clrp{5.5cm}}
403
+ \toprule
404
+ Token & Top subtask & Conf. & Auto-interpretation \\
405
+ \midrule
406
+ \texttt{t0} & UC (47\%) & 1.00 & Token t0 marks the tens digit position in addition problems, regardless of carry state or sum value. \\
407
+ \texttt{t2} & UC (70\%) & 0.99 & Token t2 outputs the ones digit (0) when adding two numbers whose ones digits sum to 10 or more. \\
408
+ \texttt{t1} & UC (30\%) & 0.99 & This token routes to the fourth digit position during addition when a carry from the previous position must be incorporated. \\
409
+ \texttt{t3} & UC (44\%) & 0.94 & Token t3 routes to the hundreds position (d3) when processing carries from the tens column in addition. \\
410
+ \texttt{t5} & MD (65\%) & 0.93 & Token t5 routes cases where the ones digit result is 0, spanning multiple subtasks and operations. \\
411
+ \texttt{t8} & MD (26\%) & 0.91 & Token t8 activates when processing the tens digit (d2) across addition/subtraction with various carry states. \\
412
+ \texttt{t10} & UB (41\%) & 0.88 & Token t10 routes subtraction problems requiring borrow propagation at mid-to-late digit positions. \\
413
+ \texttt{t6} & UC (27\%) & 0.88 & Token t6 routes cases where the ones digit result is 0, regardless of operation or carry state. \\
414
+ \bottomrule
415
+ \end{tabular}
416
+ \caption{Automated interpretation of the 8 highest-confidence \sorl{} abstraction tokens
417
+ (\`{a} la \citealt{bills2023language}).
418
+ For each token, the 10 examples with highest softmax confidence are shown to an LLM,
419
+ which produces a one-sentence role description.
420
+ \textbf{Conf.} = mean softmax probability of the assigned token.
421
+ High-confidence specialists receive crisp, position- and operation-specific descriptions;
422
+ polysemantic tokens (not shown) produce broader descriptions.}
423
+ \label{tab:auto-interp}
424
+ \end{table}
425
 
426
  \begin{tcolorbox}[colback=gray!6, colframe=gray!40,
427
  fonttitle=\bfseries\small, title={Finding \#7: Automated interpretation matches ground-truth subtask labels},
 
470
 
471
  Evaluation uses fixed-length autoregressive decoding (no teacher forcing):
472
  the model generates answer digits $d_0 \to d_6$ using its own predictions, with abstraction tokens inserted via the \sorl{} search-then-recurse procedure (matching training). Accuracy is measured on 100 held-out examples per split (seed 42; \texttt{thoughtworks/arithmetic-sorl-data}).
473
+
474
  """
475
 
476