Omnigrok: Grokking Beyond Algorithmic Data

Ziming Liu, Eric J. Michaud, Max Tegmark

Introduction

Generalization lies at the heart of machine learning. A good machine learning model should arguably be able to generalize fast, and behave in a smooth/predictable way under changes of (hyper)parameters. Grokking, the phenomenon where the model generalizes long after overfitting the training set, has raised interesting questions after it was observed on algorithmic datasets by [Power et al., 2022]:

The origin of grokking: Why is generalization much delayed after overfitting?

The prevalence of grokking: Can grokking occur on datasets other than algorithmic datasets?

This paper aims to answer these questions by analyzing neural loss landscapes:

Grokking is caused by the mismatch between training and test loss landscapes. Specifically, (reduced) training and test losses plotted against model weight norm resemble "L" and "U", respectively, as shown in Figure 1(b). We refer to this phenomenon as the "LU mechanism", which we elaborate on in Section 2 and 3.

Yes. Indeed, we demonstrate grokking for a wide range of machine learning tasks in Section 4, including image classification, sentiment analysis and molecule property prediction. Grokking signals observed for these tasks are usually less dramatic than for algorithmic datasets, which we attribute to representation learning in Section 5.

Partial answers to Q1 are provided in recent studies: Liu et al. attribute grokking to the slow formation of good representations, Thilak et al. attempts to link grokking to the slingshot mechanism of adaptive optimizers, and Barak et al. uses Fourier gap to describe hidden progress. This paper aims to understand grokking through the lens of neural loss landscapes. Our landscape analysis is able to explain many aspects of grokking: data size dependence, weight decay dependence, emergence of representations, etc.

The paper is organized as follows: In Section 2, we review background on generalization, and introduce the LU mechanism. In Section 3, we show how the LU mechanism leads to grokking for a toy teacher-student setup. In Section 4, we show that the intuition gained from the toy problem can transfer to realistic datasets (MNIST, IMDb reviews and QM9), for which we also observe grokking, although in a slightly non-standard setup where it is relatively weak. In Section 5, we discuss why grokking is more dramatic for algorithmic datasets than on others (e.g., MNIST), by comparing their loss landscapes. As a byproduct, we find that training with constrained weight norm can almost eliminate grokking. We review related work in Section 6 and summarize our conclusions in Section 7. Code is available at https://github.com/KindXiaoming/Omnigrok.

The LU mechanism for grokking

In practice, we perform the constrained minimization by rescaling the model weights back to their original norm after each unconstrained optimization step. We will see that this reduced 1D loss landscape, which is easy to visualize, captures important features related to grokking. Throughout the paper, our model is initialized by multiplying a factor α≡w/w0\alpha\equiv w/w_{0} to the standard initialization The standard initialization means the default one in PyTorch., where w0w_{0} and ww are the weight norm of the network before and after multiplying α\alpha.

It is well known in statistics that generalization error has a "U" shape against model capacity, which is usually attributed to the bias-variance trade-off. Although this common wisdom was challenged by the observation of double descent [Nakkiran et al., 2021], the "U" curve can be recovered from a double descent simply by changing the x-axis from the number of model parameters NN to the 2-norm of model parameters w≡∣∣w∣∣2w\equiv||\mathbf{w}||_{2} [Ng and Ma, 2022]. Although the LU mechanism may remind readers of related phenomena [Schoenholz et al., 2016, Yang and Schoenholz, 2017, Nakkiran et al., 2021], their setups are not exactly the same as ours. More importantly, our focus and contribution is to understand grokking, a brand new generalization puzzle.

Grokking dynamics We identify the "LU mechanism" as the cause of grokking. If the weight norm is initialized to be large (e.g., the black square in the w>wcw>w_{c} region), the model first quickly moves to a nearby overfitting solution by minimizing the training loss. Without any regularization, the model will stay where it is, because the gradient of the training loss is almost zero along the valley of overfitting solutions, so generalization does not happen. Fortunately, there are usually explicit and/or implicit regularizations that can drive the weight vector towards the Goldilocks zone w≈wcw\approx w_{c}. When the regularization magnitude is non-zero but small, the radial motion can be (arbitrarily) slow. If weight decay is the only source of regularization, and training loss is negligible after overfitting, then weight decay γ\gamma causes w(t)≈exp(−γt)w0w(t)\approx{\rm exp}(-\gamma t)w_{0}, when w0>wcw_{0}>w_{c}, so it takes time t≈ln(w0/wc)/γ∝γ−1t\approx{\rm ln}(w_{0}/w_{c})/\gamma\propto\gamma^{-1} to generalize. A small γ\gamma results in a huge generalization delay (i.e., grokking). The dependence on regularization magnitudes is illustrated in Figure 1(b): no generalization at all happens for γ=0\gamma=0, small γ\gamma leads to slow generalization (grokking), and large γ\gamma leads to faster generalization γ\gamma should not be too large, otherwise it will bring the weights to a trivial solution w=0\mathbf{w}=\mathbf{0}. . The above analysis only applies to large initializations w>wcw>w_{c}. Small initializations w<wcw<w_{c} can always generalize fast ww should not be too small to harm optimization., regardless of regularization.

Why isn’t grokking commonly observed? The standard initialization schemes typically initialize ww no larger than wcw_{c}. However, if we increase initialization scales (explicitly or implicitly), grokking can appear. In Section 3 and 4, we find that explicitly increasing initialization weight norm can induce grokking. In Section 5, we argue for algorithmic datasets because (shown in Figure 6(d))

i.e., a proper initialization for a bad representation is effectively too large for a good representation, leading to grokking. Take the addition (base pp) for example: with the good (linear) representation or a bad (random) representation, the decoder needs to learn to classify O(p)O(p) or O(p2)O(p^{2}) examples, respectively.

Grokking for a teacher-student setup

To illustrate how the LU mechanism results in grokking, we employ a toy teacher-student setup. The teacher and the student share the same architecture (a 5-100-100-5 MLP with tanh activation), but are initialized with different seeds. The student network is initialized with the standard initialization (the default one in PyTorch) but each weight is rescaled by the same factor α≡w/w0\alpha\equiv w/w_{0}, where w0w_{0} and ww are the weight norm of the student network before and after rescaling. The teacher network is initialized standardly, i.e., αteacher=1\alpha_{\rm teacher}=1. Inputs and outputs have dimensions din=5d_{\rm in}=5 and dout=5d_{\rm out}=5, respectively. We generate Ntrain=100N_{\rm train}=100 training and Ntest=100N_{\rm test}=100 test samples by first drawing inputs from the standard Gaussian distribution N(0,Idin×din)N(0,\mathbf{I}_{d_{\rm in}\times d_{\rm in}}), and then feed the input data to the teacher to generate output labels. The student network is trained with the Adam optimizer (learning rate 3×10−43\times 10^{-4}) for 10510^{5} steps.

Training dynamics Our problem is a regression task, but we can imitate the behavior of a classification task by manually setting a threshold θ=0.01\theta=0.01 and defining a sample to be correctly “classified" if the prediction error is less than θ\theta. We study the dynamics of training and test accuracy. We run experiments with two initializations α=0.5\alpha=0.5 (small) and α=2.0\alpha=2.0 (large), and three weight decays γ=0\gamma=0 (no reg), γ=0.03\gamma=0.03 (small reg) and γ=1\gamma=1 (large reg). As shown in Figure 2(b) (bottom), small initialization runs always generalize fast regardless of regularization. Large initialization runs (top) depended on weight decay: no regularization fails to generalize, small regularization generalizes slowly (grokking), while large regularization generalizes faster.

For the large initialization α=2.0\alpha=2.0, we do a finer sweep of γ\gamma in [0.03,1][0.03,1]. We compute the number of steps and weight norm ww when training or test accuracy reaches 95%. As shown in Figure 2(c), the time (number of steps) to reach 95% training accuracy is independent of weight decay γ\gamma, while the time to reach 95% test accuracy is inversely proportional to the weight decay, as we derived above for the LU mechanism.

Omnigrok: Grokking for more interesting tasks

We now analyze loss landscapes and search for grokking for several more interesting datasets, and see that the insights obtained from our toy model can transfer to these datasets. We report the main results here, with experiment details included in Appendix A.

Image classification We visualize loss landscapes of MNIST [Deng, 2012] to verify the LU mechanism, and study the dependence on training data size. Similar to the teacher-student case, we reduce losses and errors (one minus accuracy) to two variables (weight norm ww and data size NN) by minimizing over angular directions of weights, i.e.,

shown in Figure 3 (a)(b). The reduced loss landscape reveals three things: (1) Larger initializations lead to grokking. Point A in Figure 3 corresponds to the standard initialization (α=1)(\alpha=1), which has low training and test errors, hence no grokking. When increasing the weight norm from A to B, training error is seen to remain low while test error rises. To generalize, implicit or explicit regularization such as weight decay then brings the weight norm down, leading to delayed generalization (grokking) if regularization is small. (2) Larger datasets lead to de-grokking. Comparing B and C in Figure 3, C is seen to have larger training size than B and lower test error. Larger data size NN makes the Goldilocks zone broader, reducing or eliminating grokking even for large weight initializations. (3) Critical data size can be defined. As reported in [Power et al., 2022, Liu et al., 2022], we see that there exists a critical training set size below which generalization is impossible. The effective theory analysis in [Liu et al., 2022] only applies to algorithmic datasets, but not to other datasets with unknown optimal representations. The loss landscape analysis presented is this work can apply to all supervised-learning tasks. As shown in Figure 3 (b), the contours of constant test error are thumb-like, and the tip of the thumb determines the minimum amount of data required for generalization.

Guided by the landscape analysis, we make two nonstandard decisions to induce grokking on MNIST: (1) we reduce the size of the training set from 60k to 1k samples (by taking a random subset) and (2) we increase the scale of the weight initialization distribution (by multiplying the initial weights, sampled with Kaiming uniform initialization, by a constant α>1\alpha>1). With these modifications to the training set size and initialization scale, we train a depth-3 width-200 MLP with ReLU activations with the AdamW optimizer using MSE loss with one-hot targets. We find that the network quickly fits the training set, and test accuracy improves much later, as shown in Figure 3(d), just as in the stereotypical grokking learning first observed in algorithmic datasets. Figure 3(e) shows the effect of training set size on time to generalization for MNIST. We find a result similar to what [Power et al., 2022] observed, namely that generalization time increases rapidly once one approaches a certain critical data set size. We also include the learning phase diagram in Appendix LABEL:app:mnist_pd.

Sentiment analysis of text We look for grokking using LSTMs [Hochreiter and Schmidhuber, 1997] for IMDb dataset [Maas et al., 2011]. Similar to Eq. (3), we reduce training and test losses to depend on only the weight norm ww and data size NN. We show the reduced training and test error in Figure 4 (a)(b). For large data size say the full dataset, training and test errors have similar "U" shapes In principle, reduced training losses should be non-increasing (”L”), but optimization issues may occur for too large initializations [Schoenholz et al., 2016]., so one cannot create grokking via the "LU" mechanism. For small data size, say 1k, however, the mismatch between training and test errors makes it possible to create grokking via large initializations. In Figure 4 (c), we initialize weights larger (α=6\alpha=6) with weight decay 1, overfitting is complete within 10210^{2} steps, but generalization does not start until around 10310^{3} steps. Note that the generalization "jump" is not as sharp as on algorithmic datasets [Power et al., 2022] or MNIST, but at least generalization is delayed here. By contrast, if we use the standard initialization (α=1\alpha=1) with no weight decay, generalization happens early on during training, and does not improve much after overfitting.

Molecules We search for grokking using the graph convolutional neural network (GCNN) for QM9 dataset [Ramakrishnan et al., 2014]. Similar to Eq. (3), we define the reduced training/test losses, which are only dependent on weight norm ww and data size NN. As shown in Figure 5(a)(b), when data size is large, training and test losses have similar "U" shapes, hence grokking is impossible via the "LU mechanism". When data size is small, training and test losses mismatch somewhere in the region α=w/w0>1\alpha=w/w_{0}>1, making grokking possible. Indeed, shown in Figure 5(d), there is a sharp drop in test loss around 10410^{4} steps if initialization is 3 times larger than standard, while standard initialization does not lead to grokking. Note that zero weight decay is applied in both cases, implying the existence of implicit regularizations.

Representation is key to grokking

In Section 4, we showed that increasing initialization scales can make grokking happen for standard ML tasks. However, this seems a bit artificial and does not explain why standard initialization leads to grokking on algorithmic datasets, but not on standard ML datasets, say MNIST. The key difference is how much the task relies on representation learning. For the MNIST dataset, the quality of representation determines whether the test accuracy is 95% or 100%; by constrast in algorithmic datasets, the quality of representation determines whether test accuracy is random guess (bad representation) or 100% (good representation). So overfitting (under a bad representation) has a more dramatic effect on algorithmic datasets, i.e., the model weights increase quickly during overfitting but test accuracy remains low. During overfitting, model weight norm is much larger than at initialization, but then drops below the initialization norm when the model generalizes, shown in Figure 7(a), and also observed by [Nanda et al., 2023]. As a byproduct, we are able to eliminate grokking by constraining the model on a small weight norm sphere, shown in Figure 7(b).

In the following, we will compare algorithmic datasets (Section 5.1) to MNIST (Section 5.2). We show how their loss landscapes depend on representations differently, and how the difference leads to different outcomes (grokking or not).

Setup We take the toy addition setup in [Liu et al., 2022], where each input digit 0≤i≤p−10\leq i\leq p-1 (output label 0≤k≤2(q−1)0\leq k\leq 2(q-1)) is embedded as a vector Ei\mathbf{E}_{i} (Yk\mathbf{Y}_{k}). A decoder MLP is employed to predict Yk=Dec(Ei+Ej) (k=i+j)\mathbf{Y}_{k}={\rm Dec}(\mathbf{E}_{i}+\mathbf{E}_{j})\ (k=i+j). In the setup of grokking, both the decoder and the input representations R≡{Ei\mathbf{R}\equiv\{\mathbf{E}_{i}} are trainable, with learning rates ηD\eta_{D} and ηR\eta_{R}, respectively; in the setup of landscape analysis, only decoder is trainable, as we explain below. Training and test losses depend on three factors: (i) representation R\mathbf{R}, (ii) weight norm ww and (iii) weight direction w^\hat{\mathbf{w}}. As in previous sections, we can optimize w^\hat{\mathbf{w}} by minimizing the training loss on constant weight norm spheres. We further reduce the high-dimensional representations to 1D by interpolating in a particular direction:

where Rlinear\mathbf{R}_{\rm linear} refers to the linear representation in which number kk is embedded to Ek=[k,0,⋯ ,0]\mathbf{E}_{k}=[k,0,\cdots,0], Rrandom\mathbf{R}_{\rm random} is drawn from Gaussian distributions, i.e, Ek∼N(0,I)\mathbf{E}_{k}\sim N(\mathbf{0},\mathbf{I}), and m∈m\in is a scalar interpolating between Rlinear\mathbf{R}_{\rm linear} and Rrandom\mathbf{R}_{\rm random}, that we term representation messiness because R=Rlinear\mathbf{R}=\mathbf{R}_{\rm linear} when m=0m=0, and R=Rrandom\mathbf{R}=\mathbf{R}_{\rm random} when m=1m=1. After these reductions, both training and test losses become functions of two variables, representation messiness mm and weight norm ww:

Grokking dynamics In region II, the dynamics is slow (for small γ\gamma) due to nearly vanishing gradients. By contrast, the dynamics in region I is relatively fast. As we will explain, dynamics is also slow on the boundary of I and II, and grokking is the consequence of traversing region II and/or the boundary.

The slow dynamics from C to E is the origin of grokking. During this period, the model first moves in the −w-w direction with a velocity vv over the distance L1=L−hcotθL_{1}=L-h{\rm cot}\theta, and then moves along the boundary with a velocity v′v^{\prime} over the distance L2=h/sinθL_{2}=h/{\rm sin}\theta. So the total time is t=L1/v+L2/v′=(L+htanθ)/(ηDγ)t=L_{1}/v+L_{2}/v^{\prime}=(L+h{\rm tan}\theta)/(\eta_{D}\gamma). This formula agrees with the observation that large weight decays γ\gamma and/or larger decoder learning rates ηD\eta_{D} can make generalization happen faster [Power et al., 2022, Liu et al., 2022]. Besides, the path manifests intriguing multiple descent of test loss, shown in Figure 6(c).

The above picture is supported by a transformer experiment: Figure 7(a), shows how model norm changes over time and we see that there is an initial increase in weight norm, which peaks during overfitting, but then drops during the period of generalization to be lower than the initialization norm. For this experiment, we used the setup of [Nanda et al., 2023], training a 1-layer transformer on modular addition (p=113p=113). The model width dmodel=128d_{\text{model}}=128, with 4 attention heads, and dmlp=512d_{\text{mlp}}=512 with ReLU activations. We train with AdamW with a learning rate of 0.001 and weight decay γ=1\gamma=1.

Dependence of grokking on training data size Another important observation in [Power et al., 2022] is that grokking happens faster for larger training size. Our landscape analysis can also explain the data size dependence. In Figure 6(e), we show the contours (training loss = 0.02) for different training sizes (25, 35, 45, 55). The contours of training size 45 and 55 both connect to the green star, meaning that generalization will eventually happen. However, the slopes of the contours are different, i.e., θ55<θ45\theta_{55}<\theta_{45}. Since t=(L+htanθ)/(ηDγ)t=(L+h{\rm tan}\theta)/(\eta_{D}\gamma) increases as θ\theta increases, we have t55<t45t_{55}<t_{45}, i.e, more training data leads to faster grokking. For training size 35 and 25, the contours do not connect to the green star, so generalization will not happen, no matter how long the training will be run.

De-grokking by constraining weight norm Guided by our understanding of grokking from the LU mechanism, we find that we can control grokking in the setting where it was first observed by Power et al. – transformers trained on algorithmic tasks. As shown in Figure 7(b), reducing the initialization scale and constraining optimization to hold model weight norm constant over training brings train accuracy and test accuracy learning curves together, almost eliminating grokking.

2 MNIST

We now study how training and test losses depend on representation messiness in the MNIST dataset. We denote the 28×2828\times 28 images as the raw representation Rraw\mathbf{R}_{\rm raw}. We construct a linearly separable representation Rlinear\mathbf{R}_{\rm linear} by assigning input representations proportional to their label yiy_{i}, for example, an image of a 2 is represented by a matrix with all elements being 2. Similar to Eq. (4), we use m∈m\in to interpolate between Rraw\mathbf{R}_{\rm raw} and Rlinear\mathbf{R}_{\rm linear}:

Comparing Figure 6 and 8, we see that the (strong) dependence of test performance on the representation is the key to grokking: the dependence on representation is strong for algorithmic datasets, so grokking happens. By contrast, the dependence is weak for MNIST, so grokking does not happen.

Discussion: grokking on language models? We conjecture that grokking is more easily observed in tasks where generalization relies heavily on learning good representations (from scratch). This seems to imply the possibility of grokking on language tasks where word embeddings are key to generalization. However, we have not yet observed clear grokking signals for large language models, perhaps because: (i) the structure of languages is complicated, so the "optimal representation" for language might be much "messier" than algorithmic representations. (ii) Pre-training avoids learning representations from scratch, hence helps reduce possible grokking.

Relation to Related works

Grokking was first observed for algorithmic datasets by [Power et al., 2022]. Several formal or informal attempts have been made to understand grokking: (a) [Liu et al., 2022] attributes grokking to the slow formation of good representations. (b) [Shah, 2021] suggests that generalizable solutions achieve lower loss than overfitting solutions, providing a training signal encouraging generalization. (c) [Nanda et al., 2023] suggests grokking is a phase change due to limited data and regularization. (d) [Barak et al., 2022] suggests that generalization is due not to random search, but to hidden progress of SGD to gradually amplify a Fourier gap. (e) [Thilak et al., 2022] links grokking to the "Slingshot mechanism" specific to adaptive optimizers. (f) [Millidge, 2022] describes training as a random walk over parameters. Our conclusion supports (a)(b)(c)(d), but does not necessarily negate (e)(f).

Double descent is the phenomenon that performance first gets worse and then gets better as we increase the model size, data size, training epochs or regularization [Nakkiran et al., 2021, Yilmaz and Heckel, 2022]. The typical "U" shape of test loss in this paper does not conflict with double descent, because we are plotting the weight norm instead of the number of model parameters [Ng and Ma, 2022]. However, the "U"-shape should better be considered as empirically common rather than provably universal. In fact, the interaction between properties of data and inductive biases of learning algorithms can be more complicated than double descent [Chen et al., 2021, d’Ascoli et al., 2020].

Initialization From the optimization perspective, initializations are usually based on the "edge of chaos" idea such that variance of features and gradients should be preserved in the forward and backward pass [Glorot and Bengio, 2010, He et al., 2015, Bahri et al., 2020, Yang and Schoenholz, 2017, Jing et al., 2017], or based on analyzing Jacobians and/or Hessians [Skorski et al., 2020]. From the generalization perspective, it was shown that large initializations overfit data easily but result in poor generalization [Xu et al., 2019, Zhang et al., 2020], which agrees with our LU mechanism.

Weight decay regularization is a standard trick in machine learning and has various effects on optimization and generalization [Zhang et al., 2018, Van Laarhoven, 2017]. In particular, [Lewkowycz and Gur-Ari, 2020] observes that it takes t∝1/λt\propto 1/\lambda training steps to reach maximum test performance. This is strikingly similar to the grokking time t∝1/λt\propto 1/\lambda we derived from the LU mechanism.

Conclusions

This study elucidates the grokking phenomenon from the perspective of loss landscapes. Our conclusions are: (i) grokking originates from the mismatch between training and test losses ("LU" mechanism). (ii) grokking can happen in various models for a wide range of datasets, although the grokking signature is usually most dramatic for algorithmic datasets. (iii) The dramaticness of grokking depends on how much the task relies on learning representations. This work not only reveals the mechanism of grokking, but also shows that reduced landscape analysis is a useful tool for characterizing data-model interaction and representation learning.

Acknowledgement

We thank Wenxian Shi, Niklas Nolte, Ouail Kitouni and Mike Williams for helpful discussions. This work was supported by The Casey and Family Foundation, the Foundational Questions Institute, the Rothberg Family Fund for Cognitive Science, the NSF Graduate Research Fellowship (Grant No. 2141064), and the NSF AI Institute for Artificial Intelligence and Fundamental Interactions (IAIFI) through NSF Grant No. PHY-2019786.

References

Appendix A Experiment details

Sentiment analysis of text IMDb [Maas et al., 2011] includes 50k movie reviews to be classified as being positive or negative. To pre-process the data, we extract the 1000 most frequent words and tokenize each review into an array of token indices. Less frequent words are ignored, and each review array is padded to length 500. We adopt the LSTM model [Hochreiter and Schmidhuber, 1997] to perform the classification, with two layers, embedding dimension 64, and hidden dimension 128. We use the Adam optimizer [Kingma and Ba, 2014] with learning rate 0.001 to minimize the binary cross entropy loss. We hold back 25% of the dataset for testing.

Molecules QM9 is a database for small molecules and their properties. We use a graph convolutional neural network (GCNN) to predict the isotropic polarizability. The GCNN contains 2 convolutional layers with ReLU activation, followed by a linear layer. We use the Adam optimizer with learning rate 0.001 to minimize the MSE loss. We split the dataset into 50/50 train/test.

MNIST We train width-200 depth-3 ReLU MLPs on the MNIST dataset with MSE loss. We use the AdamW optimizer with a learning rate of 0.001 and a batch size of 200.

Appendix B Reduced loss for modular addition with transformers

In Figure 9 we show reduced loss landscape plots for transformers trained on modular addition. We use the setup of Nanda et al. and train a 1-layer transformer on modular addition (p=113p=113) with dmodel=128d_{\text{model}}=128, 4 attention heads, and dmlp=512d_{\text{mlp}}=512 with ReLU activations. We train with a learning rate of 0.001 while constraining model weight norm, for a variety of α\alpha and a variety of train set fractions. The LU shape holds for α∈[0.1,4]\alpha\in[0.1,4] (some optimization issue may be responsible for the rise in train loss for α>4\alpha>4). We see the critical train set size is approximately 0.25, in line with earlier studies on grokking.

Appendix C time to generalize versus weight decay

In our discussion of the “LU mechanism” as an explanation for grokking in Section 2, we predicted that the training time required for a model to generalize should be t∝γ−1t\propto\gamma^{-1} where γ\gamma is the weight decay. To test this, we perform a grid search over weight decays γ\gamma and plot the number of training steps required for models to reach a specified level of test accuracy in Figure 10(a)-10(b). We also show full training curves for these runs in Figure 10(c)-10(d). We perform experiments in two setups:

Transformer on modular addition: We use the replication of grokking from Nanda et al. and train a 1-layer transformer on modular addition (p=113p=113 and a train set fraction of 0.30.3) where dmodel=128d_{\text{model}}=128, with 4 attention heads, dmlp=512d_{\text{mlp}}=512, ReLU activations, and an AdamW learning rate of 0.001. From Figure 10(a), we find that t∝γ−1t\propto\gamma^{-1} holds across roughly two orders of magnitude of tt and γ\gamma. There is some seed dependence on the generalization time (some seeds consistently require longer to generalize), but for each seed (corresponding to a particular model initialization) the relation t∝γ−1t\propto\gamma^{-1} appears to fit the data well.

ReLU MLP on MNIST: We train ReLU MLPs on MNIST as described in Appendix A. We use an α=9.0\alpha=9.0 and train on a reduced training set of 1000 samples to delay generalization / induce grokking. From Figure 10(b), we find that for γ\gamma roughly between 0.1 and 1.0 the relation t∝γ−1t\propto\gamma^{-1} holds. Very high values of weight decay seem to mess with optimization. On the other hand, with very low weight decay the model generalizes faster than naively expected, perhaps due to implicit regularization.

Appendix D Section 5.1 setup

Architecture Similar to Liu et al. , the decoder architecture is an MLP with hard coded addition. Each input symbol ii is encoded to a scalar EiE_{i}. Each output symbol kk is represented by a 30D random vector Y^k\hat{\mathbf{Y}}_{k}. We consider addition with base pp, so input 0≤i,j≤p−10\leq i,j\leq p-1 and output 0≤k=i+j≤2(p−1)0\leq k=i+j\leq 2(p-1). We denote representation as R={E0,E1⋯ ,Ep−1}\mathbf{R}=\{E_{0},E_{1}\cdots,E_{p-1}\}. The MLP has two hidden layers, with neurons 1-200-200-30 in each layer and ReLU activations. Given a training sample (Ei,Ej)→Yk(E_{i},E_{j})\to\mathbf{Y}_{k} where i+j=ki+j=k, the prediction of the MLP decoder is

and the loss function being the mean squared error (MSE) between Yk\mathbf{Y}_{k} and Y^k\hat{\mathbf{Y}}_{k}, and w\mathbf{w} being the decoder weight. Although the common setup of grokking is to make both the representation R\mathbf{R} and the decoder w\mathbf{w} trainable, we will freeze part of them for easier analysis. This is where it could be a bit confusing, so we explicitly distinguish three setups: landscape analysis, reduced trajectory analysis and full trajectory analysis. Each setup have different subset of trainable parameters, as shown in Table 1.

Landscape analysis Both the representation R\mathbf{R} and weight norm ww are fixed. Only the weight direction w^\hat{\mathbf{w}} is trainable. The representation R\mathbf{R} is fixed according to Eq. (4), which is dependent on mm, the representation messiness. The decoder has fixed weight norm ww, but the weight direction w^\hat{\mathbf{w}} is trainable. For each fixed (w,m)(w,m), we minimize training loss over w^\hat{\mathbf{w}} to get

and define reduced training and test loss, as in Eq. (5). The minimization is implemented by the Adam optimizer with learning rate 10−310^{-3} for 10410^{4} steps. Although (w,m)(w,m) are not trainable, we repeat the above minimization independently for different (w,m)(w,m). In Figure 6 (a)(b)(d), the background heatmaps belong to landscape analysis.

Reduced trajectory analysis is a “thought experiment" based on landscape analysis. Since full trajectory analysis can be intractable due to too high dimensions, we try to reduce the trajectory anaysis to 2D, by making two assumptions about the real dynamics: (1) Scale separation: the dynamics of w^\hat{\mathbf{w}} is much faster than the dynamics along ww and along mm, such that w^(t)=w^∗(w(t),m(t))\hat{\mathbf{w}}(t)=\hat{\mathbf{w}}^{*}(w(t),m(t)) is valid at every moment during training. (2) Representation evolution is linear, i.e., interpolating between initial random Gaussian and final linear representation. With these two assumptions, the training dynamics is effectively reduced to 2D, depending only on (w,m)(w,m), obeying Eq. (6). In Figure 6 (a)(b)(c), the path A→E{\rm A}\to{\rm E} belongs to reduced trajectory analysis.

Admittedly the reduced trajectory may deviate from the full trajectory since the assumptions may not be met, but it can shed light on the full trajectory: the weight norm first increases and then increases, and the decrease of weight norm is highly correlated with generalization (please see Appendix LABEL:app:weight_evolution and Figure 7.

Appendix E MNIST experiments with cross entropy loss

To respond to a reviewer’s concern that our use of the MSE loss is the “secret" to get grokking on MNIST (Figure 3), we reran our experiments with the cross entropy (CE) loss. The results are qualitatively similar, with some quantitative differences.

Comparing Figure 3 (MSE) and Figure 11 (CE), we notice the they are qualitatively similar: (1) for small datasets, the reduced training error and test error resemble an “L" and “U" against the weight norm, respectively; (2) for large datasets, the “U" becomes more like “L", i.e., the mismatch between the reduced training and test error is small. However, a quantitative difference exist: CE produces a broader “Goldilocks zone" (the weight range where generalization happens) than MSE. This implies that to induce grokking with CE, we need to increase the weight norm to a larger value (say α=100\alpha=100).

We are able to observe delayed generalization during trianing on MNIST with cross entropy loss, but doing so requires a higher α\alpha than was necessary when using MSE loss, as predicted by the reduced loss landscapes in Figure 11. Figure 12 shows training trajectories from a 3-layer ReLU MLP on MNIST trained with cross entropy loss with α=100\alpha=100 and D=200D=200. We see that test accuracy rises to 30-40% early in training, then plateaus for an extended period, before increasing to ≈\approx75% while train accuracy remains at 100%. While the dynamics are not as clean as with MSE loss, since test accuracy first plateaus at better-than-random accuracy, we think it is still fair to classify these dynamics as “grokking” due to the improvement in generalization late in training after a plateau.