On the training dynamics of deep networks with $L_2$ regularization

Aitor Lewkowycz, Guy Gur-Ari

Introduction

Machine learning models are commonly trained with L2L_{2} regularization. This involves adding the term 12λ∥θ∥22\frac{1}{2}\lambda\|\theta\|_{2}^{2} to the loss function, where θ\theta is the vector of model parameters and λ\lambda is a hyperparameter. In some cases, the theoretical motivation for using this type of regularization is clear. For example, in the context of linear regression, L2L_{2} regularization increases the bias of the learned parameters while reducing their variance across instantiations of the training data; in other words, it is a manifestation of the bias-variance tradeoff. In statistical learning theory, a “hard” variant of L2L_{2} regularization, in which one imposes the constraint ∥θ∥2≤ϵ\|\theta\|_{2}\leq\epsilon, is often employed when deriving generalization bounds.

In deep learning, the use of L2L_{2} regularization is prevalent and often leads to improved performance in practical settings [Hinton, 1986], although the theoretical motivation for its use is less clear. Indeed, it well known that overparameterized models overfit far less than one may expect [Zhang et al., 2016], and so the classical bias-variance tradeoff picture does not apply [Neyshabur et al., 2017, Belkin et al., 2018, Geiger et al., 2020]. There is growing understanding that this is caused, at least in part, by the (implicit) regularization properties of stochastic gradient descent (SGD) [Soudry et al., 2017]. The goal of this paper is to improve our understanding of the role of L2L_{2} regularization in deep learning.

We study the role of L2L_{2} regularization when training over-parameterized deep networks, taken here to mean networks that can achieve training accuracy 1 when trained with SGD. Specifically, we consider the early stopping performance of a model, namely the maximum test accuracy a model achieves during training, as a function of the L2L_{2} parameter λ\lambda. We make the following observations based on the experimental results presented in the paper.

The number of SGD steps until a model achieves maximum performance is t∗≈cλt_{*}\approx\frac{c}{\lambda}, where cc is a coefficient that depends on the data, the architecture, and all other hyperparameters. We find that this relationship holds across a wide range of λ\lambda values.

If we train with a fixed number of steps, model performance peaks at a certain value of the L2L_{2} parameter. However, if we train for a number of steps proportional to λ−1\lambda^{-1} then performance improves with decreasing λ\lambda. In such a setup, performance becomes independent of λ\lambda for sufficiently small λ\lambda. Furthermore, performance with a small, non-zero λ\lambda is often better than performance without any L2L_{2} regularization.

Figure 1a shows the performance of an overparameterized network as a function of the L2L_{2} parameter λ\lambda. When the model is trained with a fixed steps budget, performance is maximized at one value of λ\lambda. However, when the training time is proportional to λ−1\lambda^{-1}, performance improves and approaches a constant value as we decrease λ\lambda.

As we demonstrate in the experimental section, these observations hold for a variety of training setups which include different architectures, data sets, and optimization algorithms. In particular, when training with vanilla SGD (without momentum), we observe that the number of steps until maximum performance depends on the learning rate η\eta and on λ\lambda as t∗≈c′η⋅λt_{*}\approx\frac{c^{\prime}}{\eta\cdot\lambda}. The performance achieved after this many steps depends only weakly on the choice of learning rate.

We present two practical applications of these observations. First, we propose a simple way to predict the optimal value of the L2L_{2} parameter, based on a cheap measurement of the coefficient cc. Figure 1b compares the performance of models trained with our predicted L2L_{2} parameter with that of models trained with a tuned parameter. In this realistic setting, we find that our predicted parameter leads to performance that is within 0.4% of the tuned performance on CIFAR-10, at a cost that is marginally higher than a single training run. As shown below, we also find that the predicted parameter is consistently within an order of magnitude of the optimal, tuned value.

As a second application we propose AutoL2L_{2}, a dynamical schedule for the L2L_{2} parameter. The idea is that large L2L_{2} values achieve worse performance but also lead to faster training. Therefore, in order to speed up training one can start with a large L2L_{2} value and decay it during training (this is similar to the intuition behind learning rate schedules). In Figure 1c we compare the performance of a model trained with AutoL2L_{2} against that of a tuned but constant L2L_{2} parameter, and find that AutoL2L_{2} outperforms the tuned model both in speed and in performance.

Learning rate schedules.

Our empirical observations apply in the presence of learning rate schedules. In particular, Figure 1a shows that the test accuracy remains approximately the same if we scale the training time as 1/λ1/\lambda. As to our applications, in section 3 we propose an algorithm for predicting the optimal L2L_{2} value in the presence of learning rate schedules, and the predicted value gives comparable performance to the tuned result. As to the AutoL2L_{2} algorithm, we find that in the presence of learning rate schedules it does not perform as well as a tuned but constant L2L_{2} parameter. We leave combining AutoL2L_{2} with learning rate schedules to future work.

Theoretical contribution.

Finally, we turn to a theoretical investigation of the empirical observations made above. As a first attempt at explaining these effects, consider the following argument based on the loss landscape. For overparameterized networks, the Hessian spectrum evolves rapidly during training [Sagun et al., 2017, Gur-Ari et al., 2018, Ghorbani et al., 2019]. After a small number of training steps with no L2L_{2} regularization, the minimum eigenvalue is found to be close to zero. In the presence of a small L2L_{2} term, we therefore expect that the minimal eigenvalue will be approximately λ\lambda. In quadratic optimization, the convergence time is inversely proportional to the smallest eigenvalue of the Hessian In linear regression with L2L_{2} regularization, optimization is controlled by a linear kernel K=XTX+λIK=X^{T}X+\lambda I, where XX is the sample matrix and II is the identity matrix in parameter space. Optimization in each kernel eigendirection evolves as e−γte^{-\gamma t} where γ\gamma is the corresponding eigenvalue. When λ>0\lambda>0 and the model is overparameterized, the lowest eigenvalue of the kernel will be typically close to λ\lambda, and therefore the time to convergence will be proportional to λ−1\lambda^{-1}. , see Ali et al. for a recent discussion . Based on this intuition, we may then expect that convergence time will be proportional to λ−1\lambda^{-1}. The fact that performance is roughly constant for sufficiently small λ\lambda can then be explained if overfitting can be mostly attributed to optimization in the very low curvature directions [Rahaman et al., 2018, Wadia et al., 2020]. Now, our empirical finding is that the time it takes the network to reach maximum accuracy is proportional to λ−1\lambda^{-1}. In some cases this is the same as the convergence time, but in other cases (see for example Figure 4a) we find that performance decays after peaking and so convergence happens later. Therefore, the loss landscape-based explanation above is not sufficient to fully explain the effect.

To gain a better theoretical understanding, we consider the setup of an infinitely wide neural network trained using gradient flow. We focus on networks with positive-homogeneous activations, which include deep networks with ReLU activations, fully-connected or convolutional layers, and other common components. By analyzing the gradient flow update equations of such networks, we are able to show that the performance peaks at a time of order λ−1\lambda^{-1} and deteriorates thereafter. This is in contrast to the performance of linear models with L2L_{2} regularization, where no such peak is evident. These results are consistent with our empirical observations, and may help shed light on the underlying causes of these effects.

According to known infinite width theory, in the absence of explicit regularization, the kernel that controls network training is constant [Jacot et al., 2018]. Our analysis extends the known results on infinitely wide network optimization, and indicates that the kernel decays in a predictable way in the presence of L2L_{2} regularization. We hope that this analysis will shed further light on the observed performance gap between infinitely wide networks which are under good theoretical control, and the networks trained in practical settings [Arora et al., 2019, Novak et al., 2019, Wei et al., 2018, Lewkowycz et al., 2020].

Related works.

L2L_{2} regularization in the presence of batch-normalization [Ioffe and Szegedy, 2015] has been studied in [van Laarhoven, 2017, Hoffer et al., 2018, Zhang et al., 2018]. These papers discussed how the effect of L2L_{2} on scale invariant models is merely of having an effective learning rate (and no L2L_{2}). This was made precise in Li and Arora where they showed that this effective learning rate is ηeff=ηe2ηλt\eta_{\rm eff}=\eta e^{2\eta\lambda t} (at small learning rates). Our theoretical analysis of large width networks will have has the same behaviour when the network is scale invariant. Finally, in parallel to this work, Li et al. carried out a complementary analysis of the role of L2L_{2} regularization in deep learning using a stochastic differential equation analysis. Their conclusions regarding the effective learning rate in the presence of L2L_{2} regularization are consistent with our observations.

Experiments

We now turn to an empirical study of networks trained with L2L_{2} regularization. In this section we present results for a fully-connected network trained on MNIST, a Wide ResNet [Zagoruyko and Komodakis, 2016] trained on CIFAR-10, and CNNs trained on CIFAR-10. The experimental details are in SM A. The empirical findings discussed in section 1.1 hold across this variety of overparameterized setups.

Figure 2 presents experimental results on fully-connected and Wide ResNet networks. Figure 3 presents experiments conducted on CNNs. We find that the number of steps until optimal performance is achieved (defined here as the minimum time required to be within .5%.5\% of the maximum test accuracy) scales as λ−1\lambda^{-1}, as discussed in Section 1.1. Our experiments span 66 decades of η⋅λ\eta\cdot\lambda (larger η,λ\eta,\lambda won’t train at all and smaller would take too long to train). Moreover, when we evolved the networks until they have reached optimal performance, the maximum test accuracy for smaller L2L_{2} parameters did not get worse. We compare this against the performance of a model trained with a fixed number of epochs, reporting the maximum performance achieved during training. In this case, we find that reducing λ\lambda beyond a certain value does hurt performance.

While here we consider the simplified set up of vanilla SGD and no data augmentation, our observations also hold in the presence of momentum and data augmentation, see SM C.2 for more experiments. We would like to emphasize again that while the smaller L2L_{2} models can reach the same test accuracy as its larger counterparts, models like WRN28-10 on CIFAR-10 need to be trained for a considerably larger number of epochs to achieve this.The longer experiments ran for 5000 epochs while one usually trains these models for ∼\sim300 epochs.

Learning rate schedules.

So far we considered training setups that do not include learning rate schedules. Figure 1a shows the results of training a Wide ResNet on CIFAR-10 with a learning rate schedule, momentum, and data augmentation. The schedule was determined as follows. Given a total number of epochs TT, the learning rate is decayed by a factor of 0.20.2 at epochs {0.3⋅T,0.6⋅T,0.9⋅T}\{0.3\cdot T,0.6\cdot T,0.9\cdot T\}. We compare training with a fixed TT against training with T∝λ−1T\propto\lambda^{-1}. We find that training with a fixed budget leads to an optimal value of λ\lambda, below which performance degrades. On the other hand, training with T∝λ−1T\propto\lambda^{-1} leads to improved performance at smaller λ\lambda, consistent with our previous observations.

Applications

We now discuss two practical applications of the empirical observations made in the previous section.

We observed that the time t∗t_{*} to reach maximum test accuracy is proportional to λ−1\lambda^{-1}, which we can express as t∗≈cλt_{*}\approx\frac{c}{\lambda}. This relationship continues to hold empirically even for large values of λ\lambda. When λ\lambda is large, the network attains its (significantly degraded) maximum performance after a relatively short amount of training time. We can therefore measure the value of cc by training the network with a large L2L_{2} parameter until its performance peaks, at a fraction of the cost of a normal training run.

Based on our empirical observations, given a training budget TT we predict that the optimal L2L_{2} parameter can be approximated by λpred=c/T\lambda_{\rm pred}=c/T. This is the smallest L2L_{2} parameter such that model performance will peak within training time TT. Figure 1b shows the result of testing this prediction in a realistic setting: a Wide ResNet trained on CIFAR-10 with momentum=0.9=0.9 , learning rate η=0.2\eta=0.2 and data augmentation. The model is first trained with a large L2L_{2} parameter for 2 epochs in order to measure cc, and we find c≈0.0066c\approx 0.0066, see figure 4a. We then compare the tuned value of λ\lambda against our prediction for training budgets spanning close to two orders of magnitude, and find excellent agreement: the predicted λ\lambda’s have a performance which is rather close to the optimal one. Furthermore, the tuned values are always within an order of magnitude of our predictions see figure 4b.

So far we assumed a constant learning rate. In the presence of learning rate schedules, one needs to adjust the prediction algorithm. Here we address this for the case of a piecewise-constant schedule. For compute efficiency reasons, we expect that it is beneficial to train with a large learning rate as long as accuracy continues to improve, and to decay the learning rate when accuracy peaks. Therefore, given a fixed learning rate schedule, we expect the optimal L2L_{2} parameter to be the one at which accuracy peaks at the time of the first learning rate decay. Our prediction for the optimal parameter is then λpred=c/T1\lambda_{\rm pred}=c/T_{1}, where T1T_{1} is the time of first learning rate decay, and the coefficient cc is measured as before with a fixed learning rate. In our experiments, this prediction is consistently within an order of magnitude of the optimal parameter, and gives comparable performance. For example, in the case of Figure 1a with T=200T=200 epochs and T1=0.3TT_{1}=0.3T, we find λpred≈0.0001\lambda_{\rm pred}\approx 0.0001 (leading to test accuracy 0.9600.960), compared with the optimal value 0.00050.0005 (with test accuracy 0.9670.967).

We now turn to another application, based on the observation that models trained with larger L2L_{2} parameters reach their peak performance faster. It is therefore plausible that one can speed up the training process by starting with a large L2L_{2} parameter, and decaying it according to some schedule. Here we propose to choose the schedule dynamically by decaying the L2L_{2} parameter when performance begins to deteriorate. See SM E for further details.

AutoL2L_{2} is a straightforward implementation of this idea: We begin training with a large parameter, λ=0.1\lambda=0.1, and we decay it by a factor of 10 if either the empirical loss (the training loss without the L2L_{2} term) or the training error increases. To improve stability, immediately after decaying we impose a refractory period during which the parameter cannot decay again. Figure 1c compares this algorithm against the model with the optimal L2L_{2} parameter. We find that AutoL2L_{2} trains significantly faster and achieves superior performance. See SM E for other architectures.

In other experiments we have found that this algorithm does not yield improved results when the training procedure includes a learning rate schedule. We leave the attempt to effectively combine learning rate schedules with L2L_{2} schedules to future work.

Theoretical results

We say that the network function is kk-homogeneous if fαθ(x)=αkfθ(x)f_{\alpha\theta}(x)=\alpha^{k}f_{\theta}(x) for any α>0\alpha>0. As an example, a fully-connected network with LL layers and ReLU or linear activations is LL-homogeneous. Networks made out of convolutional, max-pooling or batch-normalization layers are also kk-homogeneous.Batch normalization is often implemented with an ϵ\epsilon parameter meant to prevent numerical instabilities. Such networks are only approximately homogeneous. See Li and Arora for a discussion of networks with homogeneous activations.

Dyer and Gur-Ari presented a conjecture that allows one to derive the large width asymptotic behavior of the network function, the Neural Tangent Kernel, as well as of combinations involving higher-order derivatives of the network function. The conjecture was shown to hold for networks with polynomial activations [Aitken and Gur-Ari, 2020], and has been verified empirically for commonly used activation functions. In what follows, we will assume the validity of this conjecture. The following is our main theoretical result.

Consider a kk-homogeneous network, and assume that the network obeys the correlation function conjecture of Dyer and Gur-Ari . In the infinite width limit, the network function ft(x)f_{t}(x) and the kernel Θt(x,x′)\Theta_{t}(x,x^{\prime}) evolve according to the following equations at training time tt.

The proof hinges on the following equation, which holds for kk-homogeneous functions: ∑μθμ∂μ∂ν1⋯∂νmf(x)=(k−m)∂ν1⋯∂νmf(x)\sum_{\mu}\theta_{\mu}\partial_{\mu}\partial_{\nu_{1}}\cdots\partial_{\nu_{m}}f(x)=(k-m)\partial_{\nu_{1}}\cdots\partial_{\nu_{m}}f(x). This equation allows us to show that the only effect of L2L_{2} regularization at infinite width is to introduce simple terms proportional to λ\lambda in the gradient flow update equations for both the function and the kernel.

We refer the reader to the SM for the proof. We mention in passing that the case k=0k=0 corresponds to a scaling-invariant network function which was studied in Li and Arora . In this case, training with L2L_{2} term is equivalent to training with an exponentially increasing learning rate.

For commonly used loss functions, and for k>1k>1, we expect that the solution obeys lim⁡t→∞ft(x)=0\lim_{t\to\infty}f_{t}(x)=0. We will prove that this holds for MSE loss, but let us first discuss the intuition behind this statement. At late times the exponent in front of the first term in (1) decays to zero, leaving the approximate equation df(x)dt≈−λkf(x)\frac{df(x)}{dt}\approx-\lambda kf(x) and leading to an exponential decay of the function to zero. Both the explicit exponent in the equation, and the approximate late time exponential decay, suggest that this decay occurs at a time tdecay∝λ−1t_{\rm decay}\propto\lambda^{-1}. Therefore, we expect that the minimum of the empirical loss to occur at a time proportional to λ−1\lambda^{-1}, after which the bare loss will increase because the function is decaying to zero. We observe this behaviour empirically for wide fully-connected networks and for Wide ResNet in the SM.

Furthermore, notice that if we include the kk dependence, the decay time scale is approximately tdecay∝(kλ)−1t_{\rm decay}\propto(k\lambda)^{-1}. Models with a higher degree of homogeneity (for example deeper fully-connected networks) will converge faster.

We now focus on MSE loss and solve the gradient flow equation (1) for this case.

Here, ya:=(e^a)Tyy_{a}:=(\hat{e}_{a})^{T}y. At late times, lim⁡t→∞ft(x)=0\lim_{t\to\infty}f_{t}(x)=0 on the training set.

The properties of the solution (3) depend on whether the ratio γa/λ\gamma_{a}/\lambda is greater than or smaller than 1, as illustrated in Figure 5. When γa/λ>1\gamma_{a}/\lambda>1, the function approaches the label mode ymode=yay_{\rm mode}=y_{a} at a time that is of order 1/γa1/\gamma_{a}. This behavior is the same as that of a linear model, and represents ordinary learning. Later, at a time of order λ−1\lambda^{-1} the mode decays to zero as described above; this late time decay is not present in the linear model. Next, when γa/λ<1\gamma_{a}/\lambda<1 the mode decays to zero at a time of order λ−1\lambda^{-1}, which is the same behavior as that of a linear model.

Let us now return to infinitely wide networks. These behave like linear models with a fixed kernel when λ=0\lambda=0, but as we have seen when λ>0\lambda>0 the kernel decays exponentially. Nevertheless, we argue that this decay is slow enough such that the training dynamics follow that of the linear model (obtained by setting k=1k=1 in eq. (1)) up until a time of order λ−1\lambda^{-1}, when the function begins decaying to zero. This can be seen in Figure 5c, which compares the training curves of a linear and a 2-layer network using the same kernel. We see that the agreement extends until the linear model is almost fully trained, at which point the 2-layer model begins deteriorating due to the late time decay. Therefore, if we stop training the 2-layer network at the loss minimum, we end up with a trained and regularized model. It would be interesting to understand how the generalization properties of this model with decaying kernel differ from those of the linear model.

Finite-width network.

Theorem 1 holds in the strict large width, fixed λ\lambda limit for NTK parameterization. At large but finite width we expect (1) to be a good description of the training trajectory at early times, until the kernel and function because small enough such that the finite-width corrections become non-negligible. Our experimental results imply that this approximation remains good until after the minimum in the loss, but that at late times the function will not decay to zero; see for example Figure 5c. See the SM for further discussion for the case of deep linear models. We reserve a more careful study of these finite width effects to future work.

Discussion

In this work we consider the effect of L2L_{2} regularization on overparameterized networks. We make two empirical observations: (1) The time it takes the network to reach peak performance is proportional to λ\lambda, the L2L_{2} regularization parameter, and (2) the performance reached in this way is independent of λ\lambda when λ\lambda is not too large. We find that these observations hold for a variety of overparameterized training setups; see the SM for some examples where they do not hold. We expect the peak performance to depend on λ\lambda and η\eta, but not on other quantities such as the initialization scale. We verify this empirically in SM F.

Motivated by these observations, we suggest two practical applications. The first is a simple method for predicting the optimal L2L_{2} parameter at a given training budget. The performance obtained using this prediction is close to that of a tuned L2L_{2} parameter, at a fraction of the training cost. The second is AutoL2L_{2}, an automatic L2L_{2} parameter schedule. In our experiments, this method leads to better performance and faster training when compared against training with a tuned L2L_{2} parameter. We find that these proposals work well when training with a constant learning rate; we leave an extension of these methods to networks trained with learning rate schedules to future work.

We attempt to understand the empirical observations by analyzing the training trajectory of infinitely wide networks trained with L2L_{2} regularization. We derive the differential equations governing this trajectory, and solve them explicitly for MSE loss. The solution reproduces the observation that the time to peak performance is of order λ−1\lambda^{-1}. This is due to an effect that is specific to deep networks, and is not present in linear models: during training, the kernel (which is constant for linear models) decays exponentially due to the L2L_{2} term.

Acknowledgments

The authors would like to thank Yasaman Bahri, Ethan Dyer, Jaehoon Lee, Behnam Neyshabur, and Sam Schoenholz for useful discussions. We especially thank Behnam for encouraging us to use our scaling law observations to come up with a schedule for the L2L_{2} parameter.

References

Supplementary material

Appendix A Experimental details

We are using JAX [Bradbury et al., 2018].

All the models except for section C.4 have been trained with Softmax loss normalized as L({x,y}B)=12k∣B∣∑(x,y)∈B,iyilog⁡pi(x),pi(x)=efi(x)∑jefj(x){\cal L}(\{x,y\}_{B})=\frac{1}{2k|B|}\sum_{(x,y)\in B,i}y_{i}\log p_{i}(x),p_{i}(x)=\frac{e^{f^{i}(x)}}{\sum_{j}e^{f^{j}(x)}}, where kk is the number of classes and yiy^{i} are one-hot targets.

All experiments that compare different learning rates and L2L_{2} parameters use the same seed for the weights at initialization and we consider only one such initialization (unless otherwise stated) although we have not seen much variance in the phenomena described. We will be using standard normalization with LeCun initialization W∼N(0,σw2Nin),b∼N(0,σb2)W\sim{\cal N}(0,\frac{\sigma_{w}^{2}}{N_{in}}),b\sim{\cal N}(0,\sigma_{b}^{2}).

Batch Norm: we are using JAX’s Stax implementation of Batch Norm which doesn’t keep track of training batch statistics for test mode evaluation.

Data augmentation: denotes flip, crop and mixup.

WRN: Wide Resnet 28-10 [Zagoruyko and Komodakis, 2016] with has batch-normalization and batch size 10241024 (per device batch size of 128128), σw=1,σb=0\sigma_{w}=1,\sigma_{b}=0. Trained on CIFAR-10.

FC: Fully connected, three hidden layers with width 20482048 and ReLU activation and batch size 512512,σw=2,σb=0\sigma_{w}=\sqrt{2},\sigma_{b}=0. Trained on 512 samples of MNIST.

CNN: We use the following architecture: Conv1(300)→Act→Conv2(300)→Act→MaxPool((6,6), ’VALID’)→Conv1(300)→Act→Conv2(300)→MaxPool((6,6), ’VALID’)→Flatten()→Dense(500)→Dense(10)\text{Conv}_{1}(300)\rightarrow\text{Act}\rightarrow\text{Conv}_{2}(300)\rightarrow\text{Act}\rightarrow\text{MaxPool((6,6), 'VALID')}\rightarrow\text{Conv}_{1}(300)\rightarrow\text{Act}\rightarrow\text{Conv}_{2}(300)\rightarrow\text{MaxPool((6,6), 'VALID')}\rightarrow\text{Flatten()}\rightarrow\text{Dense}(500)\rightarrow\text{Dense}(10). Dense(n)\text{Dense}(n) denotes a fully-connected layer with output dimension nn. Conv1(n),Conv2(n)\text{Conv}_{1}(n),\text{Conv}_{2}(n) denote convolutional layers with ’SAME’ or ’VALID’ padding and nn filters, respectively; all convolutional layers use (3,3)(3,3) filters. MaxPool((2,2), ’VALID’) performs max pooling with ’VALID’ padding and a (2,2) window size. Act denotes the activation: ‘(Batch-Norm →\rightarrow) ReLU ’ depending on whether we use Batch-Normalization or not. We use batch size 128128, σw=2,σb=0\sigma_{w}=\sqrt{2},\sigma_{b}=0. Trained on CIFAR-10 without data augmentation.

The WRN experiments are run on v3-8 TPUs and the rest on P100 GPUs.

Here we describe the particularities of each figure. Whenever we report performance for a given time budget, we report the maximum performance during training which does not have to happen at the end of training.

Figure 1a WRN trained using momentum=0.9=0.9, data augmentation and a learning rate schedule where η(t=0)=0.2\eta(t=0)=0.2 and then decays η→0.2η\eta\rightarrow 0.2\eta at {0.3⋅T,0.6⋅T,0.9⋅T}\{0.3\cdot T,0.6\cdot T,0.9\cdot T\}, where TT is the number of epochs. We compare training with a fixed T=200T=200 training budget, against training with T(λ)=0.1/λT(\lambda)=0.1/\lambda. This was chosen so that T(0.0005)=200T(0.0005)=200.

Figures 1b, 4, S4. WRN trained using momentum=0.9=0.9, data augmentation and η=0.2\eta=0.2 for λ∈(5⋅10−6,10−5,5⋅10−5,0.0001,0.0002,0.0004,0.001,0.002)\lambda\in(5\cdot 10^{-6},10^{-5},5\cdot 10^{-5},0.0001,0.0002,0.0004,0.001,0.002). The predicted λ\lambda performance of 1b was computed at λ=0.0066/T∈(0.000131,6.56⋅10−5,2.63⋅10−5,1.31⋅10−5,6.56⋅10−6,4.38⋅10−6,3.28⋅10−6)\lambda=0.0066/T\in(0.000131,6.56\cdot 10^{-5},2.63\cdot 10^{-5},1.31\cdot 10^{-5},6.56\cdot 10^{-6},4.38\cdot 10^{-6},3.28\cdot 10^{-6}) for T∈(50,100,250,500,1000,1500,2000)T\in(50,100,250,500,1000,1500,2000) respectively.

Figures 1c,S9. WRN trained using momentum=0.9=0.9, data augmentation and η=0.2\eta=0.2, evolved for 200200 epochs. The AutoL2L_{2} algorithm is written explicitly in SM E and make measurements every 10 steps.

Figure 2a,b,c. FC trained using SGD 2ηλ\frac{2}{\eta\lambda} epochs with learning rate and L2L_{2} regularizations η∈(0.0025,0.01,0.02,0.025,0.03,0.05,0.08,0.15,0.3,0.5,1,1.5,2,5,10,25,50)\eta\in(0.0025,0.01,0.02,0.025,0.03,0.05,0.08,0.15,0.3,0.5,1,1.5,2,5,10,25,50), λ∈(0,10−5,0.0001,0.0005,0.001,0.005,0.01,0.05,0.1,0.5,1,5,10,20,50,100)\lambda\in(0,10^{-5},0.0001,0.0005,0.001,0.005,0.01,0.05,0.1,0.5,1,5,10,20,50,100). The λ=0\lambda=0 model was evolved for 106/η10^{6}/\eta epochs which is more than the smallest λ\lambda.

Figure 2d,e,f. WRN trained using SGD without data augmentation for 0.1ηλ\frac{0.1}{\eta\lambda} epochs for the following hyperparameters η∈(0.0125,0.025,0.05,0.1,0.2,0.4,0.8,1.6),λ∈(0,1.5625⋅10−5,6.25⋅10−5,1.25⋅10−4,2.5⋅10−4,5⋅10−4,10−3,2⋅10−3,4⋅10−3,8⋅10−3,0.016)\eta\in(0.0125,0.025,0.05,0.1,0.2,0.4,0.8,1.6),\lambda\in(0,1.5625\cdot 10^{-5},6.25\cdot 10^{-5},1.25\cdot 10^{-4},2.5\cdot 10^{-4},5\cdot 10^{-4},10^{-3},2\cdot 10^{-3},4\cdot 10^{-3},8\cdot 10^{-3},0.016), as long as the total number of epochs was ≤4000\leq 4000 epochs (except for η=0.2,λ=6.25⋅10−5\eta=0.2,\lambda=6.25\cdot 10^{-5} which was evolved for 80008000 epochs). We evolved the λ=0\lambda=0 models for 1000010000 epochs.

Figure S7. Fully connected depth 33 and width 6464 trained on CIFAR-10 with batch size 512512, η=0.1\eta=0.1 and cross-entropy loss.

Figure S8. ResNet-50 trained on ImageNet with batch size 8192, using the implementation in https://github.com/tensorflow/tpu.

Figure 5 (a,b) plots ftf_{t} in equation 3 with k=2k=2 (for 2-layer) and k=1k=1 (for linear), for different values of γ\gamma and λ=0.01\lambda=0.01. (c) The empirical kernel of a 2−2-layer ReLU network of width 5,000 was evaluated on 200200-samples of MNIST with even/odd labels. The linear, 2−2-layer curves come from evolving equation 1 with the previous kernel and setting k=1,k=2k=1,k=2, respectively . The experimental curve comes from training the 2-layer ReLU network with width 10510^{5} and learning rate η=0.01\eta=0.01 (the time is step×η\text{step}\times\eta).

Figure 3a,b,c. CNN without BN trained using SGD for 0.01ηλ\frac{0.01}{\eta\lambda} epochs for the following hyperparameters η=0.01,λ∈(0,5⋅10−5,0.0001,0.0005,0.001,0.01,0.05,0.1,0.25,0.5,1,2)\eta=0.01,\lambda\in(0,5\cdot 10^{-5},0.0001,0.0005,0.001,0.01,0.05,0.1,0.25,0.5,1,2). with λ=0\lambda=0 was evolved for 2100021000 epochs.

Figure 3d,e,f. CNN with BN trained using SGD for a time 0.01ηλ\frac{0.01}{\eta\lambda} for the following hyperparameters η=0.01,λ=0,5⋅10−5,0.0001,0.0005,0.001,0.01,0.05,0.1,0.25,0.5,1,2\eta=0.01,\lambda=0,5\cdot 10^{-5},0.0001,0.0005,0.001,0.01,0.05,0.1,0.25,0.5,1,2. The model with λ=0\lambda=0 was evolved for 95009500 epochs, which goes beyond where all the other λ\lambda’s have peaked.

Figure S5. FC trained using SGD and MSE loss for 1ηλ\frac{1}{\eta\lambda} epochs and the following hyperparameters η∈(0.001,0.005,0.01,0.02,0.035,0.05,0.15,0.3),λ∈(10−5,0.0001,0.0005,0.001,0.005,0.01,0.05,0.1,0.5,1,5,50,100)\eta\in(0.001,0.005,0.01,0.02,0.035,0.05,0.15,0.3),\lambda\in(10^{-5},0.0001,0.0005,0.001,0.005,0.01,0.05,0.1,0.5,1,5,50,100). For λ=0\lambda=0, it was trained for 105/η10^{5}/\eta epochs.

Rest of SM figures. Small modifications of experiments in previous figures, specified explicitly in captions.

Appendix B Details of theoretical results

In this section we prove the main theoretical results. We begin with two technical lemmas that apply to kk-homogeneous network functions, namely network functions fθ(x)f_{\theta}(x) that obey the equation faθ(x)=akfθ(x)f_{a\theta}(x)=a^{k}f_{\theta}(x) for any input xx, parameter vector θ\theta, and a>0a>0.

Let fθ(x)f_{\theta}(x) be a kk-homogeneous network function. Then ∑μθμ∂μ∂ν1⋯∂νmf(x)=(k−m)∂ν1⋯∂νmf(x)\sum_{\mu}\theta_{\mu}\partial_{\mu}\partial_{\nu_{1}}\cdots\partial_{\nu_{m}}f(x)=(k-m)\partial_{\nu_{1}}\cdots\partial_{\nu_{m}}f(x).

We prove by induction on mm. For m=0m=0, we differentiate the homogeneity equation with respect to aa.

The cluster graph of CC has mm vertices; we denote by nen_{e} (non_{o}) the number of even (odd) components in the graph (we refer the reader to Dyer and Gur-Ari for a definition of the cluster graph and other terminology used in this proof). By assumption, ne+(no−m)/2≤−1n_{e}+(n_{o}-m)/2\leq-1.

Notice that each CbC_{b} is obtained from CC by replacing a derivative tensor ∂f\partial f with d(∂f)/dtd(\partial f)/dt inside the expectation value. Let us see how this affects the cluster graph. For any derivative tensor ∂μ1…μaf(x):=∂af(x)/∂θν1⋯∂θνa\partial_{\mu_{1}\dots\mu_{a}}f(x):=\partial^{a}f(x)/\partial\theta^{\nu^{1}}\cdots\partial\theta^{\nu^{a}}, we have

In the last step we used lemma 1. We now compute how replacing the derivative tensor ∂f\partial f by each of the terms in the last line of (S6) affects the cluster graph, and specifically the combination ne+(no−m)/2n_{e}+(n_{o}-m)/2.

We now turn to the proof of Theorems 1 and 2.

A straightforward calculation leads to the following gradient flow equations for the network function and kernel.

The solution of this set of equations (labelled by mm) is the same as for the any-time equation ddtΘt(x,x′)=−2(k−1)λΘt(x,x′)\frac{d}{dt}\Theta_{t}(x,x^{\prime})=-2(k-1)\lambda\Theta_{t}(x,x^{\prime}), and the solution is given by

The evolution of the kernel eigenvalues, and the fact that its eigenvectors do not evolve, follow immediately from (2). The solution (3) can be verified directly by plugging it into (1) after projecting the equation on the eigenvector e^a\hat{e}_{a}. Finally, the fact that the function decays to zero at late times can be seen from (3) as follows. From the assumption k≥2k\geq 2, notice that exp⁡ ⁣[−γa(t′)2(k−1)λ−(k−2)λt′]≤1\exp\!\left[-\frac{\gamma_{a}(t^{\prime})}{2(k-1)\lambda}-(k-2)\lambda t^{\prime}\right]\leq 1 when t′≥0t^{\prime}\geq 0. Therefore, we can bound each mode as follows.

Therefore, lim⁡t→∞∣fa(x;t)∣=0\lim_{t\to\infty}|f_{a}(x;t)|=0. ∎

For completeness we now write down the solution (3) in functional form, for x∈Sx\in S in the training set.

Here, exp⁡(⋅)\exp(\cdot) is a matrix exponential, and Θt\Theta_{t} is a matrix of size Nsamp×NsampN_{\rm samp}\times N_{\rm samp}.

Let’s consider a deep linear model f(x)=βWL....W0.xf(x)=\beta W_{L}....W_{0}.x, with β=n−L/2\beta=n^{-L/2} for NTK normalization and β=1\beta=1 for standard normalization. The gradient descent equation will be:

Evolution will stop when the fixed point ( ΔW=0\Delta W=0) is reached:

Now, we would like to show that, at the fixed point

We can be very explicit if we consider L=1L=1 and one sample with x,y=1,1x,y=1,1 for MSE loss. The fixed point has a logit:

which is only different from 0,10,1 for fixed λ2n\lambda^{2}n.

Appendix C More on experiments

We can see how the time it takes to rech training accuracy 11 depends very mildly on λ\lambda, and for small enough learning rates it scales like 1/η1/\eta.

C.2 More WRN experiments

We can also study the previous in the presence of momentum and data augmentation. These are the experiments that we used in figure 4, evolved until convergence. As discussed before, in the presence of momentum the t∗t_{*} depends on η\eta, so we will fixed the learning rate η=0.2\eta=0.2.

Here we give more details about the optimal L2L_{2} prediction of section 3. Figure S3 illustrates how performance changes as a function of λ\lambda for different time budgets with the predicted λ\lambda marked with a dashed line. If one wanted to be more precise, from figure 2 we see that while the scaling works across λ\lambda’s, generally lower λ\lambda’s have a scaling ∼2\sim 2 times higher than the larger λ\lambda’s. One could try to get a more precise prediction by multiplying cc by two, csmallλ∼2clargeλc_{\text{small}\lambda}\sim 2c_{\text{large}\lambda}, see figure S4. We reserve a more detailed analysis of this more fine-grained prescription for the future.

C.4 MSE and the catapult effect

In Lewkowycz et al. it was argued that, in the absence of L2L_{2} when training a network with SGD and MSE loss, high learning rates have a rather different final accuracy, due to the fact that at early times they undergo the "catapult effect". However, this seems to contradict with our story around 1.1 where we argue that performance doesn’t depend strongly on η\eta. In figure S5, we can see how, while when stopped at training accuracy 11, performance depends strongly on the learning rate, this is no longer the case in the presence of L2L_{2} if we evolve it for ttestt_{test}. We also show how the training MSE loss has a minimum after which it increases.

C.5 Dynamics of loss and accuracy

In figure S6 we illustrate the training curves of the experiments we have discussed in the main text and SM.

Appendix D Examples of setups where the scalings don’t work

We will consider a couple of setups which don’t exhibit the behaviour described in the main text: a 3 hidden layer, width 64 fully-connected network trained on CIFAR-10 and ResNet-50 trained on ImageNet. We attribute this difference to deviations from the overparametrized/large width regime. In this situation, the optimal test accuracy with respect to λ\lambda has a maximum at some λopt≠0\lambda_{\rm opt}\not=0.

For the FC experiment, the time it takes to reach this maximum accuracy scales like 1/λ1/\lambda for λ≳λopt\lambda\gtrsim\lambda_{opt}, but becomes constant (equal to the value for λ=0\lambda=0) for λ≲λopt\lambda\lesssim\lambda_{\rm opt}. This peak of the maximum test accuracy happens before the training accuracy reaches 11. Generically, we don’t observe that a network trained with cross-entropy and without regularization to have a peak in the test accuracy at a finite time.

We do not have as clear an understanding of the ImageNet experimental results because they involve a learning rate schedule. Performance for small λ\lambdas does not improve even if when evolving for a longer time. However, we do observe that performance is roughly constant when η⋅λ\eta\cdot\lambda is held fixed.

We require the loss/error to be bigger than its minimum value two measurements in a row (we make measurements every 5 steps), we do this to make sure that this increase is not due to a fluctuation. After decaying, we force λ\lambda to stay constant for a time 0.1/λ0.1/\lambda steps, we choose the refractory period to scale with 1/λ1/\lambda because this is the physical scale of the system.

To complement the AutoL2L_{2} discussion of section 3 we have done another experiment where the learning rate is decayed using the schedule described in 2. Here we see how while AutoL2L_{2} trains faster in the beginning, the optimal λ=0.0005\lambda=0.0005 outperforms it. We have not hyperparameter tuned the possible parameters of AutoL2L_{2}.

We have also applied AutoL2L_{2} to other setups, and in the absence of learning rate schedules beats the optimal L2L_{2} parameter. See figure S11.

Appendix F Different initialization scales

We have discussed the dependence on the time to convergence in η,λ\eta,\lambda. Another quantity which is relevant for learning is σw\sigma_{w} the scale of initialization, which for WRN we set to 11. We can repeat the WRN experiment of figure 1a for different σw\sigma_{w}. We see that the final performance depends very mildly on σw\sigma_{w}. This is what we expect when reaching equilibrium: the dependence of properties at initialization is eventually washed away.