MIIII

MIIII

Copenhagen, Denmark

Talk on Mechanistic Interpretability of Deep Learning Models presented at the University of Copenhagen
MIIIINoah SyrkisJanuary 8, 20261 |Mechanistic interpretability (MI)?2 |Grokking is sudden Generalization3 |Modular arithmetic as a task4 |We see grokking on the task5 |Analyzing model embeddings6 |We see a temporary neural spike“This disgusting pile of matrices is just a poorly written elegant algorithm” — Neel Nanda11Not verbatim, but the gist of it1 |Mechanistic interpretability (MI)?▶Deep learning (DL) is sub-symbolic▶No clear map from params to math notation▶MI is about finding that map▶Step 1 train on task. Step 2 reverse engineer▶Turning black boxes … opaque?Figure 1: Activations of an MLP neuron trained on modular addition (𝑥0+𝑥1mod𝑝=𝑦)3 of 211.1 |MI Style Questions▶When does the model learn what?▶Are the learned mechanisms static?▶How are the mechanisms learned?▶How to write a learnt algo in math?↓𝑓(𝑥)=sin𝑤𝑒𝑥+cos𝑤𝑒𝑥Figure 2: Mapping model params to math4 of 212 |Grokking is sudden Generalization▶Grokking [1] is generalization after overfitting▶Mech. interpretability needs a mechanism▶Model params move from archive to algorithm▶Figure 3 shows example of train and eval curvesFigure 3: Example of the grokking5 of 213 |Modular arithmetic as a task▶In the following assume 𝑝 and 𝑞 to be prime▶Seminal [2] MI work uses Eq. 1.1 as task▶We created the strictly harder Eq. 1.2 task▶Eq. 1.2 is multitask and non-commutative𝑦=(𝑥0+𝑥1)mod𝑝(1.1)⃗𝑦=(𝑥0+𝑥1𝑝)mod𝑞∀𝑞<𝑝(1.2)6 of 213 |Modular arithmetic as a task▶Figure 4 shows a vis of a subset of the data▶On top we see all (𝑥0,𝑥1)-pairs for 𝑝=7▶Below (𝑥0+𝑥1𝑝)mod𝑞,𝑝=13,𝑞=11↓Figure 4: Visualizing 𝑋 for 𝑝=7 (top)and 𝑌 for 𝑞=11,𝑝=13 (bottom)7 of 214 |We see grokking on the task▶The model groks on 𝒯︀miiii (Figure 5)▶Final hyper-params are seen in Table 2▶GrokFast [3] posits gradient series is made of:1.A fast varying overfitting component2.A slow varying generalizing component▶Grokking is sped up1 by boosting the latterFigure 5: Training (top) and validation (bottom) accuracy during training on 𝒯︀miiii1Our model did not converge without GrokFast8 of 215 |Analyzing model embeddings▶Pos embs in Figure 6 shows commutativity▶Corr. is 0.95 for 𝒯︀nanda and −0.64 for 𝒯︀miiii▶Assumed to fully account for commutativityFigure 6: Positional embeddings for 𝒯︀nanda (top) and 𝒯︀miiii (bottom).9 of 215 |Analyzing model embeddings▶For 𝒯︀nanda token embs are linear comb of 5 freqs▶For 𝒯︀miiii more freqs indicate larger table▶Each task focuses on a unique prime (no over­lap)▶As per Figure 7 the embs of 𝒯︀miiii are saturatedFigure 7: 𝒯︀nanda (top) and 𝒯︀miiii (bottom) token embeddings in Fourier basis10 of 21Conclusion: Embs alone account for commutativity and multitask(edness?)6 |We see a temporary neural spike▶We plot neuron activation varying 𝑥0 and 𝑥1▶Activations are largely identical to 𝒯︀nandaFigure 8: Activations of first three neurons for 𝒯︀nanda (top) and 𝒯︀miiii (bottom)12 of 216 |We see a temporary neural spike▶Some freqs 𝜔 rise to significance (𝜔>𝜇+2𝜎)▶But how many? And at what points in time?Figure 9: FFT of activations of first three neurons for 𝒯︀nanda (top) and 𝒯︀miiii (bottom)13 of 21Figure 10: Number of neurons with active freq 𝜔 (rows) through time (cols)6 |We see a temporary neural spike▶Initial freqs coincide with solving 2, 3, 5 and 7▶Spike in active freqs during generalization▶Decrease in active freqs after generalizationepoch256102440961638465536freqs00101810Table 1: number active freqs 𝜔 through trainingFigure 11: Figure 10 (top) and validation accuracy from Figure 5 (bottom)15 of 216 |We see a temporary neural spike▶Previous work [2] shows final circuitry begins developing right away (no sudden phase shift)▶GrokFast [3] targets this circuitry, assuming associated gradient updates to be slow varying▶With the Ω-spike we observe temporarily useful structures (not part of final solution)▶We propose to modify GrokFast to allow dynamical targeting of temporarily useful circuitry16 of 21References[1]A. Power, Y. Burda, H. Edwards, I. Babuschkin, and V. Misra, “Grokking: Generalization Beyond Overfitting on Small Algorithmic Datasets.” Accessed: Feb. 13, 2024. [Online]. Available: http://arxiv.org/abs/2201.02177[2]N. Nanda, L. Chan, T. Lieberum, J. Smith, and J. Steinhardt, “Progress Measures for Grokking via Mechanistic Interpretability.” Accessed: Dec. 16, 2023. [Online]. Available: http://arxiv.org/abs/2301.05217[3]S. Lee and S. Kim, “Exploring Prime Number Classification: Achieving High Recall Rate and Rapid Convergence with Sparse Encoding.” Accessed: Feb. 13, 2024. [Online]. Available: http://arxiv.org/abs/2402.03363A |Hyperparametersrate𝜆wd𝑑lrheads110121325631044Table 2: Hyperparams for 𝒯︀miiii18 of 21B |Stochastic Signal ProcessingWe denote the weights of a model as 𝜃. The gradient of 𝜃 with respect to our loss function at time 𝑡 we denote 𝑔(𝑡). As we train the model, 𝑔(𝑡) varies, going up and down. This can be thought of as a stocastic signal. We can represent this signal with a Fourier basis. GrokFast posits that the slow varying frequencies contribute to grokking. Higer frequencies are then muted, and grokking is indeed accelerated.19 of 21C |Discrete Fourier TransformFunction can be expressed as a linear combination of cosine and sine waves. A similar thing can be done for data / vectors.20 of 21D |Singular Value DecompositionAn 𝑛×𝑚 matrix 𝑀 can be represented as a 𝑈Σ𝑉∗, where 𝑈 is an 𝑚×𝑚 complex unitary matrix, Σ a rectangular 𝑚×𝑛 diagonal matrix (padded with zeros), and 𝑉 an 𝑛×𝑛 complex unitary matrix. Multiplying by 𝑀 can thus be viewed as first rotating in the 𝑚-space with 𝑈, then scaling by Σ and then rotating by 𝑉 in the 𝑛-space.21 of 21