Towards understanding how momentum improves generalization in deep learning
Samy Jelassi, Yuanzhi Li
Introduction
It is commonly accepted that adding momentum to an optimization algorithm is required to optimally train a large-scale deep network. Most of the modern architectures maintain during the training process a heavy momentum close to 1 (Krizhevsky et al., 2012; Simonyan & Zisserman, 2014; He et al., 2016; Zagoruyko & Komodakis, 2016). Indeed, it has been empirically observed that architectures trained with momentum outperform those which are trained without (Sutskever et al., 2013). Several papers have attempted to explain this phenomenon. From the optimization perspective, (Defazio, 2020) assert that momentum yields faster convergence of the training loss since, at the early stages, it cancels out the noise from the stochastic gradients. On the other hand, (Leclerc & Madry, 2020) empirically observes that momentum yields faster training convergence only when the learning rate is small. While these works shed light on how momentum acts on neural network training, they fail to capture the generalization improvement induced by momentum (Sutskever et al., 2013). Besides, the noise reduction property of momentum advocated by (Defazio, 2020) contradicts the observation that, in deep learning, having a large noise in the training improves generalization (Li et al., 2019; HaoChen et al., 2020). To the best of our knowledge, there is no existing work which theoretically explains how momentum improves generalization in deep learning. Therefore, this paper aims to close this gap and addresses the following question:
Why does momentum improve generalization? What is the underlying mechanism of momentum improving generalization in deep learning?
In computer vision, practitioners usually train their architectures with stochastic gradient descent with momentum (SGD+M). It is therefore natural to investigate whether the generalization improvement induced by momentum is tied to the stochasticity of the gradient. We train a VGG-19 (Simonyan & Zisserman, 2014) using SGD, SGD+M, gradient descent (GD) and GD with momentum (GD+M) on the CIFAR-10 image classification task. To further isolate the regularization effect of momentum, we turn off data augmentation and batch normalization. Figure 1 displays the training loss and test accuracy of the four models. Not only momentum improves generalization in the full batch setting but the generalization improvement increases as the batch size is larger. Motivated by this empirical observation, we focus on the contribution of momentum in gradient descent. We emphasize that this setting allows to isolate the contribution of momentum on generalization since the stochastic gradient noise influences generalization (Li et al., 2019; HaoChen et al., 2020).
Given the success of momentum in different deep learning tasks such as image classification (Simonyan & Zisserman, 2014; He et al., 2016) or language modelling (Vaswani et al., 2017; Devlin et al., 2018), we start our investigation by raising the following question:
Does momentum unconditionally improve generalization in deep learning?
We respond in the negative to this question through the following synthetic binary classification example. We consider a Gaussian dataset where each data-point is sampled from a standard normal distribution. We generate the labels using multiple teacher networks. Starting from the same initialization, we train several student networks on this dataset using GD and GD+M and compare their test accuracies in Table 1. Whether the target function is simple (linear) or complex (neural network), momentum does not improve generalization for any of the student networks. The same observation holds for SGD/SGD+M as shown in Appendix A. Therefore, momentum does not always lead to a higher generalization in deep learning. Instead, such benefit seems to heavily depend on both the structure of the data and the learning problem.
Motivated by the aforementioned observations, this paper aims to determine the underlying mechanism produced by momentum to improve generalization. Our work is a first step to formally understand the role of momentum in deep learning. Our contributions are divided as follows:
In Section 2, we empirically confirm that momentum consistently improves generalization when using different architectures on a wide range of batch sizes and datasets. We also observe that as the batch size increases, momentum contributes more significantly to generalization.
In Section 3, we introduce our synthetic data structure and learning problem to theoretically study the contribution of momentum to generalization.
In Section 4, we present our main theorems along with the intermediate lemmas. We theoretically show that a 1-hidden layer neural network trained with GD+M on our synthetic dataset is able to generalize better than the same model trained with GD. Above all, we rigorously characterize the mechanism by which momentum improves generalization. A sketch of the proof is presented in Section 5 and Section 6.
The previous experiments suggest that momentum improves generalization in CIFAR-10 while it does not for Gaussian datasets. This means that this generalization improvement must be specific to the data structure and the learning problem. In Section 3, we devise a binary classification problem where the data are linearly separated by a hyperplane directed by the vector as depicted in Figure 2. We refer to this vector as the feature and the goal is to learn it. Each data-point is a vector constituted of a single signal patch equal to and of multiple noise patches. For , we assume that with probability , the sampled data-point has large margin i.e. while it has small margin i.e. with probability . The noise patches are Gaussian random vectors with small variance. We underline that all the examples share the same feature but differ in their margins. Our dataset can be viewed as an extreme simplification of real-world object-recognition datasets with data of different level of difficulty. Indeed, images are divided into signal patches that are helpful for the classification such as the nose of a dog and noise patches e.g. the background of an image that are uninformative. Besides, the signal patch may be strong i.e. the feature is clearly visible or weak when the feature is indistinguishable e.g. in a car image, the wheel feature is more or less visible depending on the orientation of the car.
This paper proposes a theory to explain why momentum improves generalization. The following informal theorems characterize the generalization of a 1-hidden layer convolutional neural network trained with GD and GD+M on the aforedescribed dataset. They dramatically simplify Theorem 4.1 and Theorem 4.2 but highlight the intuitions.
There exists a dataset of size such that a 1-hidden layer (over-parameterized) convolutional network trained with GD:
initially only learns the large margin data.
has small gradient after learning these data.
memorizes the remaining small margin data from the examples.
The model thus reaches zero training loss and well-classifies the large margin data at test. However, it fails to classify the small margin data because of the memorization step during training.
There exists a dataset of size such that a one-hidden layer (over-parameterized) convolutional network trained with GD+M:
initially only learns the large margin data.
has large historical gradients that contain the feature present in small margin data.
keeps learning the feature in the small margin data using its momentum historical gradients.
The model thus reaches zero training error and perfectly classify large and small margin data at test.
Theorem 1.1 and Theorem 1.2 indicate that since the large margin data are dominant, the two models learn in priority these examples to decrease their training losses. Since the training loss is the logistic one, this implies that the gradient terms stemming from the large margin data thus become negligible. Consequently, the current gradient becomes a sum of the small margin data gradients. Thus, it is in the direction of (signal patch) and Gaussian vectors (noise patches). Since , the current gradient is noisy. Therefore, the GD model keeps decreasing its training loss and memorizes the small margin data. On the other hand, contrary to GD, GD+M updates its weights using a weighted average of the historical gradients. In particular, it has large past gradients (stemming from large margin data) that are in the direction . Therefore, even though the current gradient is noisy, the GD+M uses its historical gradients to learn the small margin data since all the examples share the same feature. We name this process historical feature amplification and believe that it is key to understand why momentum improves generalization.
Our theory relies on the ability of momentum to well-classify small margin data. We first perform experiments in our theoretical setting described in Section 3. We set the dimension to , the number of training examples to , the test examples to . Regarding the architecture, we set the number of neurons to and the number of patches to . The parameters are set as in Section 3. We refer to stochastic gradient descent optimizer with full batch size as GD/GD+M. Note that for each optimizer, we grid-search over stepsizes to find the best one in terms of test accuracy. We trained the models for 50 epochs. We set the momentum parameter to 0.9. We apply a linear decay learning rate scheduling during training. Figure 3 shows that the models trained with GD and GD+M get zero training loss and well-classify large-margin data at test time. Contrary to GD, GD+M well-classifies small margin data.
To further validate our theory, we artificially generate small-margin data in CIFAR-10. We first randomly sample 10% of the training and test images. As displayed in 4(c), for each image, we randomly shuffle the RGB channels. We train a VGG-19 without data augmentation nor batch normalization. While the GD and GD+M models reach 100% training accuracy, Figure 4 shows that GD+M gets higher test accuracy than GD. Above all, GD+M generalizes better than GD on small-margin data as the accuracy drop factor for GD+M is while for GD, this drop factor is .
Related Work
A long line of work consists in understanding the convergence speed of momentum methods when optimizing non-convex functions. (Mai & Johansson, 2020; Liu et al., 2020; Cutkosky & Mehta, 2020; Defazio, 2020) show that SGD+M reaches a stationary point as fast as SGD under diverse assumptions. Besides, (Leclerc & Madry, 2020) empirically shows that momentum accelerates neural network training for small learning rates and slows it down otherwise. Our paper differs from these works as we work in the batch setting and theoretically investigate the generalization benefits brought by momentum (and not the training ones).
Momentum-based methods such as SGD+M, RMSProp (Tieleman & Hinton, 2012) and Adam (Kingma & Ba, 2014) are standard in deep learning training since the seminal work of (Sutskever et al., 2013). Although it is known that momentum improve generalization in deep learning, only a few works formally investigate the role of momentum in generalization. (Leclerc & Madry, 2020) empirically report that momentum yields higher generalization when using a large learning rate. However, they assert that this benefit can be obtained by applying an even larger learning rate on vanilla SGD. We suspect that this is due to data augmentation and batch normalization (Ioffe & Szegedy, 2015) which are known to bias the algorithm’s generalization (Bjorck et al., 2018). To our knowledge, our work is the first that theoretically investigates the generalization of momentum in deep learning.
Numerical performance of momentum
To evaluate the contribution of momentum to generalization, we conducted extensive experiments on CIFAR-10 and CIFAR-100 (Krizhevsky et al., 2009). We used VGG-19 (Simonyan & Zisserman, 2014) and Resnet-18 (He et al., 2016) as architectures. In this section, we only present the plots obtained with VGG-19 and invite the reader to look at Appendix A for the Resnet-18 experiments.
In all of our experiments, we refer to the stochastic gradient descent optimizer with batch size 128 as SGD/SGD+M and the optimizer with full batch size as GD/GD+M. We turn off data augmentation and batch normalization to isolate the contribution of momentum to the optimization. Note that for each algorithm, we grid-search over stepsizes and momentum parameter to find the best one in terms of test accuracy. We train the models for 300 epochs. The stepsize is decayed by a factor 10 at epochs 190 and 265 during training. All the results are averaged over 5 seeds.
Momentum improves generalization. Figure 5 shows the performance of GD, GD+M, SGD and SGD+M when training a VGG-19 on CIFAR-100. We observe that GD+M/SGD+M consistently outperform GD/SGD. Besides, we highlight that the generalization improvement induced is more significant for GD than for SGD. Similar observations hold for Resnet-18 (see Appendix A).
Influence of batch size. 6(a) shows the test accuracy of a VGG-19 trained on CIFAR-10 with the stochastic gradient descent optimizer on a wide range of batch sizes. We compare the generalization obtained with momentum and without. We remark that momentum does not improve generalization when the batch size is tiny. However, as the batch size increases, the gap between the momentum curve and the no momentum one widens.
Batch normalization and data augmentation. Practitioners usually add batch normalization and data augmentation when training their architectures. 6(b) displays the test accuracy obtained when training a VGG-19 with these two regularizers. We remark that they inhibit the generalization improvement of momentum for small and middle range batch sizes. For large batch sizes, momentum slightly improves generalization. Additional experiments on the influence of batch normalization and data augmentation are in Appendix A.
Setting and algorithms
In this section, we introduce our theoretical setting to analyze the implicit bias of momentum. We first formally define the data distribution sketched in the introduction and the neural network model we use to learn it. We finally present the GD and GD+M algorithms.
We define a data distribution where each sample consists in an input and a label such that:
Uniformly sample the label from
Signal patch: one patch satisfies
is distributed as with probability and otherwise.
Noisy patches: for .
When , if the loss is convex, then there is a unique global optimal solution, so the choice of optimization algorithm does not matter. In our case, due to the non-convexity of the training objective, GD+M converges to a different (approximate) global optimal compared to GD, with better generalization properties.
We assess the quality of a predictor using the classical 0-1 loss used in binary classification. Given a sample the individual test (classification) error is defined as While measures the error of on an individual data-point, we are interested in the test error that measures the average loss over data points generated from and defined as
We solve the training problem (P) using GD and GD+M. GD is defined for by
where is the learning rate. On the other hand, GD+M is defined by the update rule
where and is the momentum factor. We now detail how to set parameters in GD and GD+M.
Our 3.1 matches with the parameters used in practice as the weights are generally initialized from Gaussian with small variance and momentum is set close to 1 (Sutskever et al., 2013).
Main results
We now formally state our main theorems regarding the generalization of models trained using (GD) and (GD+M) in the setting described in Section 3. We first introduce some notations.
Let , , and Our analysis tracks the -th weight of the network, the gradient of with respect to , the momentum gradient defined by . We introduce the projection of these objects on the feature and noise patches :
– Projection on : .
– Projection on .
– Total noise:
– Maximum signal: .
Lastly, we define the negative sigmoid
We now provide our first result which states that the learner model trained with GD does not generalize well on .
Assume that we run GD on P for iterations with parameters set as in 3.1. With high probability, the weights learned by GD
Intuitively, the training process of the GD model is described as follows. Given and our choice of parameters for , the gradient points mainly in the direction of . Therefore, GD eventually learns the feature in (Lemma 5.1) and the gradients from quickly become small. Afterwards, the gradient is dominated by the gradients from (Lemma 5.2). Because has small margin, the full gradient is now directed by the noisy patches. It implies that GD memorizes noise in (Lemma 5.4). Since these gradients also control the amount of remaining feature to be learned (Lemma 5.3), we conclude that the GD model partially learns the feature and introduces a huge noise component in the learned weights. We provide a proof sketch of Theorem 4.1 in Section 5. On the other hand, the model trained with GD+M generalizes well on .
Assume that we run GD+M on (P) for iterations with parameters set as in 3.1. With high probability, the weights learned by GD+M
Intuitively, the GD+M model follows this training process. Similarly to GD, it first learns the feature in (Lemma 6.1). Contrary to GD, the momentum gradient is still highly correlated with after this step (Lemma 6.2). Indeed, the key difference is that momentum accumulates historical gradients. Since these gradients were accumulated when learning , the direction of momentum gradient is highly biased towards . Therefore, the GD+M model amplifies the feature from these historical gradients to learn the feature in small margin data (Lemma 6.3). Subsequently, the gradient becomes small (Lemma 6.4) and the GD+M model manages to ignore the noisy patches (Lemma 6.5) and learns the feature from both and We provide a proof sketch of Theorem 4.2 in Section 6.
Our analysis is built upon a decomposition of the updates (GD) and (GD+M) on and . The projection of the vanilla and momentum gradients along these directions are
and
and
We now define the projected updates as follows:
Analysis of GD
In this section, we provide a proof sketch for Theorem 4.1 that reflects the behavior of GD with . A more detailed proof extending to can be found in the Appendix.
At the beginning of the learning process, the gradient is mostly dominated by the gradients coming from the samples. Since these data have large margin, the gradient is thus highly correlated with and increases as shown in the following Lemma.
For all and , (1) is simplified as:
Lemma 5.3 implies that quantifying the decrease rate of provides an estimate on the quantity of feature learnt by the model. We remark that for some . We thus need to determine whether the feature or the noise terms dominates in the sigmoid.
We now show that the total correlation between the weights and the noise in data increases until being large.
Let , and . For , (2) is simplified as:
By Lemma 5.4, the noise dominates in . Consequently, the algorithm memorizes the data which implies a fast decay of .
Combining Lemma 5.5 and Lemma 5.3, we prove that GD partially learns the feature.
Lemma 5.4 and Lemma 5.6 respectively yield the first two items in Theorem 4.1. Bounds on the training loss and test errors are obtained by plugging these results in (P) and (TE).
Analysis of GD+M
In this section, we provide a proof sketch for Theorem 4.2 that reflects the behavior of GD+M with . A proof extending to can be found in the Appendix.
Similarly to GD, by our initialization choice, the early gradients and so, momentum gradients are large. They are in the span of and therefore, the GD+M model also increases its correlation with .
Contrary to GD, GD+M has a large momentum that contains after Step 1.
Lemma 6.2 hints an important distinction between GD and GD+M: while the current gradient along is small at time the momentum gradient stores historical gradients that are spanned by . It amplifies the feature present in previous gradients to learn the feature in .
With this fast convergence, Lemma 6.4 implies that the correlation of the weights with the noisy patches does not have enough time to increase and thus, remains small.
Lemma 6.3 and Lemma 6.5 respectively yield the two first items in Theorem 4.2.
Discussion
Our work is a first step towards understanding the algorithmic regularization of momentum and leaves room for improvements. We constructed a data distribution where historical feature amplification may explain the generalization improvement of momentum. However, it would be interesting to understand whether this phenomenon is the only reason or whether there are other mechanisms explaining momentum’s benefits. An interesting setting for this question is NLP where momentum is used to train large models as BERT (Devlin et al., 2018). Lastly, our analysis is in the batch setting to isolate the generalization induced by momentum. It would be interesting to understand how the stochastic noise and the momentum together contribute to the generalization of a neural network.
References
Appendix A Additional Experiments
In this section, we present additional experiments to further strengthen our empirical results. We verify on Resnet-18 that momentum induces generalization improvement when trained without batch normalization and data augmentation. We then check that when these two regualizers are used, momentum does not improve generalization. We then confirm on Resnet-18 that the generalization improvement gets larger as the batch size increases. Then, we provide the performance of SGD and SGD+M on the Gaussian experiment introduced in the introduction. Lastly, we give additional plots showing that momentum allows to well-classify small margin data as mentioned at the end of the introduction.
Figure 7 displays the training loss and test accuracy obtained by training a Resnet-18 on CIFAR-10 and CIFAR-100. Similarly to the case where we trained a VGG-19, momentum significantly improves generalization whether in the stochastic case (SGD) or in the full batch setting (GD).
A.2 Influence of the batch size
Figure 8 shows the test accuracy obtained with a Resnet-18 using the stochastic gradient descent optimizer on CIFAR-10. Similarly to the VGG-19 experiment in Section 2, the generalization improvement induced by momentum gets larger as the batch size increases.
A.3 Influence of batch normalization and data augmentation
As mentioned in Section 2, batch normalization and data augmentation significantly reduce the generalization improvement induced by momentum. We further confirm this observation in Figure 9 and Figure 10.
A.4 Synthetic Gaussian data experiments
We provide a complete table with mean and standard deviations obtained by using different student networks to learn the Gaussian synthetic experiment mentioned in the introduction.
A.5 Additional justification for the theory
In this section, we present further experiments to consolidate the experiment on the artificially decimated CIFAR-10 dataset described in the introduction.
In 11(a), we observe that using a Resnet-18, momentum still improves generalization on the small margin images.In 11(d) and 12(b), we see that using stochastic updates lead SGD to classify small margin images as well as SGD+M. Lastly, Figure 13 and Figure 14 show that batch normalization and data augmentation also reduce the generalization improvement of momentum: GD/SGD perform similarly as well as GD+M/SGD+M on the small margin data.
Appendix B Additional related work
GD+M (a.k.a. heavy ball or Polyak momentum) consists in using an exponentially weighted average of the past gradients to update the weights. For convex functions near a strict twice-differentiable minimum, GD+M is optimal regarding local convergence rate (Polyak, 1963, 1964; Nemirovskij & Yudin, 1983; Nesterov, 2003). However, it may fail to converge globally for general strongly convex twice-differentiable functions (Lessard et al., 2015) and is no longer optimal for the class of smooth convex functions. In the stochastic setting, GD+M is more sensitive to noise in the gradients; that is, to preserve their improved convergence rates, significantly less noise is required (d’Aspremont, 2008; Schmidt et al., 2011; Devolder et al., 2014; Kidambi et al., 2018). Finally, other momentum methods are extensively used for convex functions such as Nesterov’s accelerated gradient (Nesterov, 1983). Our paper focuses on the use of GD+M and contrary to the aforementioned papers, our setting is non-convex. Besides, we mainly focus on the generalization of the model learned by GD and GD+M when both methods converge to global optimal. Contrary to the non-convex case, generalization is disentangled from optimization for (strictly) convex functions.
The question we address concerns algorithmic regularization which characterizes the generalization of an optimization algorithm when multiple global solutions exist in over-parametrized models (Soudry et al., 2018; Lyu & Li, 2019; Ji & Telgarsky, 2019; Chizat & Bach, 2020; Gunasekar et al., 2018; Arora et al., 2019). This regularization arises in deep learning mainly due to the non-convexity of the objective function. Indeed, this latter potentially creates multiple global minima scattered in the space that vastly differ in terms of generalization. Algorithmic regularization is induced by and depends on many factors such as learning rate and batch size (Goyal et al., 2017; Hoffer et al., 2017; Keskar et al., 2016; Smith et al., 2018), initialization (Allen-Zhu & Li, 2020), adaptive step-size (Kingma & Ba, 2014; Neyshabur et al., 2015; Wilson et al., 2017), batch normalization (Arora et al., 2018; Hoffer et al., 2019; Ioffe & Szegedy, 2015) and dropout (Srivastava et al., 2014; Wei et al., 2020). However, none of these works theoretically analyzes the regularization induced by momentum.
Appendix C Notations
In this section, we introduce the different notations used in the proofs. We start by defining the notations that appear for GD and GD+M. We first consider the case when , we will extend the proof to in section H
Our paper rely on the notions of signal and noise components of the iterates.
– Signal intensity: if and otherwise.
– Signal: for
– Noise: for and
– Max noise:
– Total noise:
We also use the following notations when dealing with the loss function and its gradient.
– Noise loss: .
– Full derivative:
– Gradient on signal: for
– Gradient on noise: for , and
C.2 Notations specific to GD+M
We now introduce the notations that only appear in the proofs involving GD+M.
– Momentum gradient oracle: for
– Signal momentum: for
– Noise momentum: for , and
Appendix D Induction hypotheses
We prove our main result using an induction. More specifically, we make the following assumptions for every time
Throughout the training process using GD for , we maintain that:
(Large signal data have small noise component). For every , for every and we maintain:
(Small signal data have large noise component). For every , for every and we have:
Throughout the training process using GD for , the signal component is bounded for every as
Throughout the training process using GD for , we maintain:
Throughout the training process using GD+M for , for every , for every , we have that:
Throughout the training process using GD+M for , for , we have that:
In what follows, we assume these induction hypotheses for to prove our generalization results. We then prove these hypotheses for
Appendix E Gradients and updates
In this section, we first derive the gradient of the loss . We then provide its projection on (signal gradient) and on (noise gradient). We first derive the gradient of the loss
For and , the gradient of the loss with respect to is:
. We derive with respect to and obtain:
By rewriting (9), we obtain the desired result. ∎
To track the signal learnt by our models, we compute the signal gradient which is the projection of the gradient on
For all and , the signal gradient is:
We obtain the desired result by projecting the gradient from Lemma E.1 on and using ∎
E.2 Noise gradient
To prove the memorization of GD and the non-memorization of GD+M, we also need to compute the noise gradient which is the projection of the gradient on
For all , and and , the noise gradient is:
Similarly to Lemma E.2, we obtain the desired result by projecting the gradient from Lemma E.1 on and using ∎
Intuitively, (10) means that the sum of the sigmoid terms for all time steps is bounded (up to a logarithmic dependence).
Appendix F Learning with GD
In this section, we detail the proofs of the lemmas in Section 5 and Theorem 4.1. We first characterize the dynamics of the signal in subsection F.1. We then analyze the dynamics of the noise in subsection F.2 and show the memorization of the GD model. We finally prove Theorem 4.1 in subsection F.3 and the induction hypotheses in subsection F.4.
To track the amount of signal learnt by GD, we make use of the following update.
For all and , the signal update (1) is equal:
The signal update is obtained by using (1) and the signal gradient (Lemma E.2). This yields
To obtain the desired upper bound, we apply the same reasoning as above to bound the term. ∎
For , we know that for all , we have . Therefore, we have
Plugging (15) in the left-hand side of (11) yields the desired lower bound.
We now prove Lemma 5.1 that quantifies the amount of signal learnt by GD when the derivative is large.
Let . From Lemma F.2, the signal update for is
where and are respectively defined as:
We now prove Lemma 5.2. It states that since the signal has significantly increased, the derivative is now small. Before proving this result, we introduce an auxiliary Lemma.
From Lemma F.3, we deduce an upper bound on :
On the other hand, using Lemma E.2, the signal difference is bounded as:
We now bound (20) by a loss term by applying Lemma K.20. Using Lemma 5.1 and D.2, we have:
From Lemma I.7, we have the convergence rate of . We use it to bound
The bound on is obtained by using its definition . ∎
We earlier proved that after iterations, the signal learnt by the GD model significantly increases until making small. We therefore need to rewrite the signal update in this case.
For , the maximal signal updates as:
From the signal update given by Lemma F.1, we know that:
To obtain the desired result, we need to prove for :
By using D.1 and D.2, (27) is bounded as:
Using Remark 1, the sigmoid term in (F.1.2) becomes small when . To summarize, we have:
Besides, we use D.2 to bound in the right-hand side of (25). ∎
We now show that once is small, the amount of learnt signal is controlled by .
Let . From Lemma F.4, we know that:
Let . We now sum up (30) for and obtain:
We now plug the bound on from Lemma 5.2 in (31). This implies:
F.2 Memorization process of GD
Lemma 5.2 shows that after iterations, the gradient is controlled by . In this section, we show that this yields the GD model to memorize.
Using Lemma F.1, we simplify the noise update.
Let all , , and . Then, with probability at least , the noise update (2) is bounded as
Let , and . From Lemma E.3, we know that the noise update satisfies:
We now apply Lemma K.5 and Lemma K.7 to respectively bound and in (34) and obtain the desired result. ∎
In the next lemma, we further simplify the noise update from Lemma F.5.
Let , and . Our starting point is Lemma I.4 which states that:
Lemma F.7 indicates that is not non-decreasing but overall, this quantity gets large over time. We now want to determine the time where one of the becomes large.
Lemma F.7 indicates that the noise iterate satisfies for :
We proved in the previous section that after iterations, the amount of noise memorized by the GD model significantly increases. We want to show that after this phase, is well-controlled.
On the other hand, from Lemma I.4, we know that for all :
Moreover, by applying Lemma K.20 to (49), we have:
We now apply Lemma I.8 to bound the loss in (50).
Using Lemma F.9, we can obtain a bound on the sum over time of derivatives .
Combining the bound on from Lemma F.9 and (52) yields:
We have thus a control on the sum over time of . We can make use of Lemma 5.3 to get the final control on the signal iterate
Let From Lemma 5.3, we know that the signal is bounded as
We plug the bound from Lemma 5.5 to bound the last term in the right-hand side of (54).
F.3 Proof of Theorem 4.1
We proved that the weights learnt by GD satisfy for
We now bound the training and test error achieved by GD at time
Train error. Lemma I.8 provides a convergence bound on the training loss.
Test error. Let be a datapoint. We remind that where and for We bound the test error as follows:
We now want to compute the probability terms in (58) and (59). We remind that is given by
We now apply Lemma 5.6 in (60) and obtain:
Let . Similarly, by applying Lemma 5.6, is bounded as:
Therefore, using (122), we upper bound the test error (120) as:
Since is taken uniformly from we further simplify (63) as:
F.4 Proof of the GD induction hypotheses
To prove Theorem 4.1, we used the induction hypotheses stated in Appendix D. The goal of this section is to prove them for
We prove here the main hypotheses we made on the noise when using GD.
Let’s start with the upper bound for . Using Lemma I.6, Lemma 5.5 and D.1, we deduce from (66) that:
which proves the induction hypothesis for Regarding the lower bound, using D.1 and Lemma 5.5, we deduce from (66) that:
which proves the induction hypothesis for
Using D.1, we bound and in (69). We obtain:
We now apply Lemma I.5 and Lemma 5.5 to bound the derivative terms in (71).
Using the same type of reasoning as for the upper bound, one can show that (72) yields:
(73) shows the induction hypothesis for
We prove the induction hypotheses for the signal
Appendix G Learning with GD+M
In this section, we prove the Lemmas in Section 6 and Theorem 4.2.
To track the amount of signal learnt by GD, we make use of the following update.
For all and , the signal momentum in (3) is equal to:
By definition of the momentum update, we have: We project this update onto and use Lemma E.2 to get:
From Lemma G.1, we can simplify the momentum update as:
For , we know that for all , we have . Thus, we have:
Plugging (78) and (80) in (76) yields the desired result. ∎
We now prove Lemma 6.1 that quantifies the signal learnt by GD when is non-zero.
By Lemma G.2, the signal update for satisfies:
We now show that contrary to GD, GD+M still has a large momentum in the direction. In other words, we want to show that is still large after iterations. Given that the small margin and large margin data share the same feature , this large momentum helps to learn .
Before proving such result, we need some intermediate lemmas.
Using the momentum update rule, we know that:
Let’s define . We start by summing the GD+M update (3) for to get
Applying Lemma G.3 to bound the momentum gradient, we further bound (83) to get:
We now use the fact that in (84) to get:
Since with , we linearize the right-hand side in (85) to obtain:
Using Lemma G.4, we can therefore show that once we learn still stays large.
Using Lemma G.4, we bound in (87) and get:
Since the signal momentum is large (Lemma 6.2), we want to argue that GD+M keeps learning the feature to eventually have a large signal.
Let such that From the signal momentum update, we deduce:
We now apply Lemma 6.2 to bound in (89) and get:
We would like to find the time such that is a constant factor i.e. such that
Let . Using (3) update rule, we have
where we used the fact that in (94). Plugging (93) in (94) yields the desired bound.
G.2 GD+M does not memorize
Lemma 6.3 implies that after iterations, the learnt signal is very large. We would like to show that this implies that the full derivative quickly decreases (Lemma 6.4) which implies that the GD+M cannot memorize (Lemma 6.5). Before proving Lemma 6.4, we need an auxiliary lemma that connects the signal momentum and the full derivative .
For , the signal momentum is bounded as
From Lemma G.1 we know that the signal momentum is equal to
We finally apply Lemma 6.3 to bound in (96) to obtain the desired result. ∎
Lemma G.5 provides an upper bound on since:
We now would like to give a convergence rate on the iterates Since Lemma J.9 gives a rate on the loss function, we connect the momentum increment with a loss term. Applying Lemma G.1, we have:
We now show that for , we have:
Indeed, by using Lemma 6.3 and D.5, we have:
Given our choice of , and , we finally bound (102) as:
We now apply Lemma K.20 to link (104) with a loss term. By Lemma 6.3 and D.5, we have:
Therefore, applying Lemma K.20 in (104) gives:
We finally apply Lemma J.9 to bound the loss term in (107) and get the desired result. ∎
After iterations, the gradient is now very small and the noise component learnt by GD+M stays very small.
We sum up (108) for and obtain:
We apply the triangle inequality in (109) and obtain:
We now use D.4 to bound in (110):
We now plug the bound on given by Lemma J.6 and obtain:
Given the values of , and , we can deduce that
Plugging (113) in (112) proves the induction hypothesis for
G.3 Proof of Theorem 4.2
We proved that the weights learnt by GD+M satisfy for
We now bound the training and test error achieved by GD+M at time
Train error. Lemma J.9 provides a convergence bound on the fake loss. Indeed, we know that:
Using Lemma K.24 along with D.4, we lower bound the loss term in (115) by the true loss.
Combining (115) and (116), we obtain a bound on the training loss.
Test error. Let be a datapoint. We remind that where and for We bound the test error as follows:
We now want to compute the probability terms in (119) and (120). We remind that is given by
We now apply Lemma 6.3, (121) is finally bounded as:
Therefore, using (122), we upper bound the test error (120) as:
Since is uniformly sampled from we further simplify (123) as:
We know that . Therefore, is the cube of a centered Gaussian.This random variable is symmetric. By Lemma K.1, we know that is also symmetric. Therefore, we simplify (124) as:
From Lemma K.14, we know that is -subGaussian. Therefore, by applying Lemma K.3, (125) is further bounded by:
Using the fact that in (126) finally yields:
G.4 Proof of the GD+M induction hypotheses
We prove the induction hypotheses for the signal
We sum up (128) for and obtain:
We apply the triangle inequality in (129) and obtain:
We now use D.5 to bound in (130):
We now plug the bound on given by Lemma J.3. We have:
Appendix H Extension to λ>0𝜆0\lambda>0
After iteration , by Lemma I.8 and Lemma J.9, we know that for GD:
On the other hand, for we know that:
Appendix I Technical lemmas for GD
This section presents the technical lemmas needed in Appendix F. These lemmas mainly consists in different rewritings of GD.
I.2 Signal lemmas
In this section, we present a lemma that bounds the sum over time of the GD increment.
Let such that Then, the derivative is bounded as:
Let such that We now sum up (135) for and get:
Subcase 1: From Lemma 5.3, we know that:
Subcase 2: From Lemma 5.3, we know that:
We therefore managed to prove that in all the cases, (140) holds.
I.3 Noise lemmas
In this section, we present the technical lemmas needed in subsection F.2. The following lemma bounds the projection of the GD increment on the noise.
Let , and . Let such that Then, the noise update (2) satisfies
Let , and . We set up the following induction hypothesis:
Let’s first show this hypothesis for From Lemma F.5, we have:
Now, we apply D.3 to bound in (143) and obtain:
Therefore, the induction hypothesis is verified for Now, assume (LABEL:eq:noiseindhypoth) for Let’s prove the result for We start by summing up the noise update from Lemma F.5 for which yields:
We apply D.3 to bound in (LABEL:eq:noiseupd20) and obtain:
To bound the first term in the right-hand side of (LABEL:eq:noiseupd202), we use the induction hypothesis (LABEL:eq:noiseindhypoth). Plugging this inequality in (LABEL:eq:noiseupd202) yields:
By rearranging the terms, we finally have:
which proves the induction hypothesis for
Now, let’s simplify the sum terms in (LABEL:eq:noiseindhypoth). Since , by definition of a geometric sequence, we have:
Plugging (151) in (LABEL:eq:noiseindhypoth) yields
Now, let’s simplify the second sum term in (152). Indeed, we have:
where we used (151) in the last inequality. Plugging (153) in (152) gives the final result. ∎
Let , and . Let such that Then, the noise update (2) satisfies
On the other hand we know from Lemma I.5 that:
By applying D.1, (157) is eventually bounded as:
By combining (51) and (158) we deduce that for all and :
I.4 Convergence rate of the training loss using GD
In this section, we prove that when using GD, the training loss converges sublinearly in our setting.
Let . Run GD with learning rate for iterations. Then, the loss sublinearly converges to zero as:
Let From Lemma F.1, we know that the signal update is lower bounded as:
Let’s now assume by contradiction that for , we have:
From the (3) update, we know that is a non-decreasing sequence which implies that is also non-decreasing. Since is non-increasing, this implies that for , we have:
Plugging (164) in the update (162) yields for :
Let . We now sum (165) for and obtain:
Given the values of , we finally have:
Let . Run GD with learning rate for iterations. Then, the loss sublinearly converges to zero as:
We first apply the classical descent lemma for smooth functions (Lemma K.18). Since is smooth, we have:
Lemma I.9 provides a lower bound on the gradient. We plug it in (170) and get:
Applying Lemma K.19 to (171) yields the aimed result. ∎
I.4.3 Auxiliary lemmas for the proof of Lemma I.8
To obtain the convergence rate in Lemma I.8, we used the following auxiliary lemma.
Let . Run GD for iterations. Then, the norm of gradient is lower bounded as follows:
Let . To obtain the lower bound, we project the gradient on the the signal and on the noise.
Since , we lower bound as
By successively applying Lemma E.2 and Lemma I.1, is lower bounded as
For a fixed and , we know that is lower bounded as
Combining (172), (175), (173) and (176) and using we thus bound as:
We now sum up (177) for and apply Cauchy-Schwarz inequality to get:
We apply Lemma I.1 to further lower bound (178) and get:
Using Lemma I.10, Lemma I.11 and Lemma I.12 we have:
Plugging (180), (181) and (182) in (179) yields:
Finally, we use Lemma I.13 and lower bound (183) by . This gives the aimed result.
We now present auxiliary lemmas that link the gradient terms with their corresponding loss.
Let Run GD for iterations. Then, we have:
Therefore, we can apply Lemma K.20 and get the lower bound:
Let Run GD for iterations. Then, we have:
We again verify that the conditions of Lemma K.20 are met. By using D.1, D.2 and Lemma 5.1, we have:
Lastly, we want to link the loss term in (187) with . By applying D.1 and Lemma K.24 in (187), we finally get:
Combining (187) and (188) yields the aimed result. ∎
Let Run GD for iterations. Then, we have:
We again verify that the conditions of Lemma K.20 are met. Using D.1, D.2 and Lemma 5.4, we have:
Lastly, we want to link the loss term in (LABEL:eq:neknfce) with . By applying D.1 and Lemma K.24 in (LABEL:eq:neknfce), we finally get:
Combining (LABEL:eq:neknfce) and (191) yields the aimed result. ∎
Let Run GD for for iterations. Then, we have:
we need to lower bound . By successively applying Lemma K.24 and D.1, we obtain:
By successively applying Lemma K.24 and D.1, we obtain:
Combining (193) and (194) yields the aimed result.
Lastly, to obtain Lemma I.8, we need to bound which is given by the next lemma.
For , the gradient of the loss projected on the normalized noise satisfies with probability for :
Projecting the gradient (given by Lemma E.1) on yields:
Since is a unit Gaussian vector, using Lemma K.8, we bound the right-hand side of (LABEL:eq:Grbd2) with probability , as:
Now, using Lemma Lemma K.10 , we can further lower bound the left-hand side of (LABEL:eq:Grbd3) as:
Remark that . By applying Lemma K.9, we have:
Appendix J Auxiliary lemmas for GD+M
This section presents the auxiliary lemmas needed in Appendix G.
J.2 Signal lemmas
In this section, we present the auxiliary lemmas needed to prove D.5. We first rewrite the (3) update to take into account the case where the signal becomes large.
For , the maximal signal momentum is bounded as:
Let . Using the signal momentum given by Lemma G.1, we know that:
To obtain the desired result, we need to prove for :
By using D.4 and D.5, (203) is bounded as:
Using Remark 1, the sigmoid term in (J.2) becomes small when . To summarize, we have:
A similar reasoning implies for :
Plugging (202) and (206) in (201) yields the aimed result. ∎
For , the sum of maximal signal momentum is bounded as:
Let . From Lemma J.2, the signal momentum is bounded as:
For , we have and for , . From Lemma J.8 and Lemma 6.4, we can bound and . Therefore, (209) is further bounded as:
We now use Lemma K.25 to bound the sum terms in (210). We have:
We now sum (211) for . Using the geometric sum inequality and obtain:
We plug and in (212). This yields the desired result. ∎
J.3 Noise lemmas
In this section, we present the technical lemmas to prove Lemma 6.5.
Run GD+M on the loss function Let , . At a time , the noise momentum is bounded with probability as:
Let and . Combining the (4) update rule and Lemma E.3 to get the noise gradient , we obtain
Using Lemma K.5 and Lemma K.7, (LABEL:eq:diffmoms1) becomes with probability
We upper bound the second term in (LABEL:eq:diffmoms2) by again using D.4:
Let . The noise momentum is bounded as
Let From Lemma J.4, we know that:
We unravel the recursion (217) rule for and obtain:
For , the sum of noise momentum is bounded as:
Let . We first apply Lemma J.5 and obtain:
Using the bound from Lemma 6.4, (218) becomes
For , we have . Plugging these two bounds in (219) implies:
We now use Lemma K.25 to bound the sum terms in (220). We have:
We now sum (221) for . Using the geometric sum inequality we obtain:
We finally use the harmonic series inequality in (222) to obtain the desired result. ∎
J.4 Convergence rate of the training loss using GD+M
In this section, we prove that when using GD+M, the training loss converges sublinearly in our setting.
For Using GD+M with learning rate , the loss sublinearly converges to zero as
Let Using Lemma J.11, we bound the signal momentum as:
We now plug (225) in the signal update (3).
We now apply Lemma K.22 to lower bound (226) by loss terms. We have:
Let’s now assume by contradiction that for , we have:
From the (3) update, we know that is a non-decreasing sequence which implies that is also non-decreasing for . Since is non-increasing, this implies that for , we have:
Plugging (229) in the update (228) yields for :
We now sum (230) for and obtain:
Given the values of , we finally have:
We now link the bound on the loss to the derivative
The proof is similar to the one of Lemma 6.4.
For Using GD+M with learning rate , the loss sublinearly converges to zero as
Let From Lemma J.10, we know that the signal gradient is bounded as for
By combining (236) and (238), we finally obtain:
We now plug (239) in the signal update (3).
We now apply Lemma K.22 to lower bound (240) by loss terms. We have:
Let’s now assume by contradiction that for , we have:
From the (3) update, we know that is a non-decreasing sequence which implies that is also non-decreasing for . Since is non-increasing, this implies that for , we have:
Plugging (243) in the update (241) yields for :
We now sum (244) for and obtain:
Given the values of , we finally have:
J.4.3 Auxiliary lemmas
We now provide an auxiliary lemma needed to obtain (J.9).
Let Then, the signal gradient decreases i.e. for
The proof is similar to the one of Lemma J.10.
Appendix K Useful lemmas
In this section, we provide the probabilistic and optimization lemmas and the main inequalities used above.
In this section, we introduce the probabilistic lemmas used in the proof.
The sum of of symmetric random variables is symmetric.
Let Let and respetively be - and -subGaussian random variables. Then, is -subGaussian random variable.
Let Let be a -subGaussian random variable. Then, we have:
We know that the is -Lipschitz and by applying Theorem K.1, we therefore have::
By rewriting (253) and using Lemma K.4, we have with probability
By squaring (254) and using , we obtain the aimed result. ∎
We know that the is -Lipschitz and by applying Theorem K.1, we therefore have:
We use Lemma K.4 and set in (255) to finally get:
Let’s define We first remark that is a sub-exponential random variable. Indeed, the generating moment function is:
where we used for in the last inequality. Therefore, by definition of a sub-exponential variable, we have:
Since and (256) is bounded as:
Let and We know that the pdf of in polar coordinates is Therefore, the generating moment function of is:
(258) indicates that is a sub-Gaussian random variable of parameter . By definition, it satisfies
Setting in (259) yields that we have with probability
Let i.i.d. vectors from Then, with probability , we have:
We know that for , we have:
Therefore, using the law of total probability and (261), we have:
Since , we also have
Therefore by setting , we obtain:
Doing the similar reasoning for the lower bound yields:
Let i.i.d. vectors from Then, with probability we have:
To show the result, it’s enough to upper bound the following probability:
By using the law of total probability we have:
where we used Lemma K.6 in (269). Using Lemma K.6 again, we can simplify (269) as:
K.1.2 Anti-concentration of Gaussian polynomials
Let be a degree polynomial and be i.i.d. Gaussian univariate random variables. Then, the following holds for all .
Since the decomposition of a polynomial in the monomial basis is unique, we can equate the coefficients of and and obtain:
Setting in (273) yields the desired result.
K.1.3 Properties of the cube of a Gaussian
Let . Then, is -subGaussian.
By definition of the moment generating function, we have:
We know that . Therefore, is the cube of a centered Gaussian. From Lemma K.13, is -subGaussian. Using Lemma K.2, we deduce that is -subGaussian. Applying again Lemma K.2, we finally obtain that is -subGaussian. ∎
K.2 Tensor Power Method Bound
In this subsection we establish a lemma for comparing the growth speed of two sequences of updates of the form . This technique is reminiscent of the classical analysis of the growth of eigenvalues on the (incremental) tensor power method of degree and is stated in full generality in (Allen-Zhu & Li, 2020).
Let be a positive sequence defined by the following recursions
where is the initialization and .Let such that Then, the time such that for all is:
We use the fact that in (274) and obtain:
Now, we want to bound . Using again the recursion and , we have:
Combining (275) and (276), we get a bound on
Now, let’s find a bound for . Starting from the recursion and using the fact that for we have:
On the other hand, by using we upper bound as follows.
Besides, we know that . Therefore, we upper bound as
We now sum (281) for , use (277) and obtain:
Lastly, we know that satisfies this implies that we can set in (282). ∎
Let be a positive sequence defined by the following recursion
where and is the initialization. Assume that Let such that Then, the time such that is upper bounded as:
By assumption, we know that This implies that for all for all Plugging this in (284) yields:
Now, we want to upper bound . Using (283), we deduce that:
Combining the two equations in (287) yields
Since is the first time where , we have . Plugging this in (288) leads to:
Finally, using (289) in (286) and gives an upper bound on
Now, let’s find a bound for . Starting from the recursion, we have:
We substract the two equations in (291), use for and obtain:
On the other hand, from the recursion, we have the following inequalities:
We substract the two equations in (293), use and upper bound as follows.
Besides, we know that . Therefore, we upper bound as
We now sum (296) for , use and then (290) to obtain:
Lastly, we know that satisfies this implies that we can set in (297). ∎
K.2.2 Bounds for GD+M
Let Let and be positive sequences defined by the following recursions
Let We want to prove the following induction hypotheses:
After iterations, we have:
After , we have:
Let’s first prove (TPM-1) and (TPM-2) for First, by using the momentum update, we have:
Setting and using , we have Plugging this in (298) yields (TPM-1) for
Regarding (TPM-2), we use the iterate update to have:
where we used and (298) to obtain (299). Since is the first time where we further simplify (299) to obtain:
We therefore obtained (TPM-2) for Let’s now assume (TPM-1) and (TPM-2) for . We now want to prove these induction hypotheses for First, by using the momentum update, we have:
From (TPM-2) for , we know that for . Therefore, (301) becomes:
From (TPM-1), we know that for Therefore, we simplify (302) as:
When we set as in (TPM-1), we have Moreover, since , we have . Using these two observations, (303) is thus equal to:
We therefore proved (TPM-1) for Now, let’s prove (TPM-2). We use the iterates update and obtain:
where we used and (304) in the last inequality. Since is the first time where we further simplify (305) to obtain:
Let’s now obtain an upper bound on We have:
Finally, we choose such that or equivalently, . Plugging this choice in yields the desired bound. ∎
K.3 Optimization lemmas
By applying the definition of smooth functions and the GD update, we have:
Setting in (308) leads to the expected result.
Let . Let be a non-negative sequence that satisfies the recursion: for Then, it is bounded at a time as
Let . By multiplying each side of the recursion by , we get:
Besides, the update rule indicates that is non-increasing i.e. Using this fact in (310) yields:
Now, we sum up (311) for and obtain:
Inverting (312) yields the expected result. ∎
K.4 Other useful lemmas
We apply Lemma K.21 to the sequence and obtain:
We apply Lemma K.24 to further simplify (313).
We remark that the term inside the exponential in (314) can be bounded as:
Lastly, we need to bound the term in the middle in (316). On one hand, we have:
Besides, since is non-decreasing, we have the following lower bound:
Finally, we obtain the desired result by combining (316), (319) and (322).
We upper bound (323) by successively applying and for all :
where we used for all in (323). By applying the rearrangement inequality to (324), we obtain:
We obtain the final bound by applying Lemma K.22 to (325).
We lower bound (323) by using and :
We obtain the final bound by applying Lemma K.22 to (326). ∎
Let be a non-negative sequence. Let Assume that Then, there exists a time such that
Assume by contradiction that for all , . By summing up , we obtain This contradicts the assumption that
Let Then, the following inequalities holds:
Successively using the inequalities and for in (329) yields:
This proves item 1 of the Lemma. Let’s now prove item 2. Using for and , we know that:
Since is non-decreasing, applying to (330) proves item 2.
In Appendix J, we need to bound the sum for We derive such bound here.
given our choice of . Let We split the sum in two parts as as follows.
where we used the harmonic series inequality , and in (333).