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 w∗\bm{w^{*}} 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 θw∗\theta\bm{w^{*}} and of multiple noise patches. For μ≪1\mu\ll 1, we assume that with probability 1−μ1-\mu, the sampled data-point has large margin i.e. θ=α≫1\theta=\alpha\gg 1 while it has small margin i.e. θ=β≪1\theta=\beta\ll 1 with probability μ\mu. 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 NN such that a 1-hidden layer (over-parameterized) convolutional network trained with GD:

initially only learns the (1−μ)N(1-\mu)N large margin data.

has small gradient after learning these data.

memorizes the remaining small margin data from the μN\mu N 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 NN such that a one-hidden layer (over-parameterized) convolutional network trained with GD+M:

initially only learns the (1−μ)N(1-\mu)N large margin data.

has large historical gradients that contain the feature w∗\bm{w^{*}} 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 βw∗\beta\bm{w^{*}} (signal patch) and Gaussian vectors g\mathbf{g} (noise patches). Since ∥βw∗∥2≪∥g∥2\|\beta\bm{w^{*}}\|_{2}\ll\|\mathbf{g}\|_{2}, 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 αw∗\alpha\bm{w^{*}}. 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 d=30d=30, the number of training examples to N=20000N=20000, the test examples to 20002000. Regarding the architecture, we set the number of neurons to m=5m=5 and the number of patches to P=5P=5. The parameters α,β,μ\alpha,\beta,\mu 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 79.47/53.30=1.4979.47/53.30=1.49 while for GD, this drop factor is 68.33/34.80=1.9668.33/34.80=1.96.

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 D\mathcal{D} where each sample consists in an input X\bm{X} and a label yy such that:

Uniformly sample the label yy from {−1,1}.\{-1,1\}.

Signal patch: one patch P(X)∈[P]P(X)\in[P] satisfies

cc is distributed as c=αyc=\alpha y with probability 1−μ1-\mu and c=βyc=\beta y otherwise.

Noisy patches: X[j]∼N(0,(Id−w∗w∗⊤)σ2),\bm{X}[j]\sim\mathcal{N}(0,(\mathbf{I}_{d}-\bm{w^{*}}\bm{w^{*\top}})\sigma^{2}), for j∈[P]\{P(X)}j\in[P]\backslash\{P(\bm{X})\}.

When λ>0\lambda>0, if the loss 1N∑i=1Nlog⁡(1+exp⁡(−yifW(Xi)))\frac{1}{N}\sum_{i=1}^{N}\log\left(1+\exp\left(-y_{i}f_{\bm{W}}(\bm{X}_{i})\right)\right) 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 W^\bm{\widehat{W}} using the classical 0-1 loss used in binary classification. Given a sample (X,y),(\bm{X},y), the individual test (classification) error is defined as L(X,y)=1{fW^(X)y<0}.\mathscr{L}(\bm{X},y)=\mathbf{1}\{f_{\bm{\widehat{W}}}(\bm{X})y<0\}. While L\mathscr{L} measures the error of fW^f_{\bm{\widehat{W}}} on an individual data-point, we are interested in the test error that measures the average loss over data points generated from D\mathcal{D} and defined as

We solve the training problem (P) using GD and GD+M. GD is defined for t≥0t\geq 0 by

where η>0\eta>0 is the learning rate. On the other hand, GD+M is defined by the update rule

where g(0)=0m×d\bm{g}^{(0)}=\bm{0}_{m\times d} and γ∈(0,1)\gamma\in(0,1) 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 r∈[m]r\in[m], i∈[N]i\in[N], j∈P\{P(Xi)}j\in P\backslash\{P(\bm{X}_{i})\} and t≥0.t\geq 0. Our analysis tracks wr(t)\bm{w}_{r}^{(t)} the rr-th weight of the network, ∇wrL^(W(t))\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)}) the gradient of L^\widehat{L} with respect to wr\bm{w}_{r}, gr(t)\bm{g}_{r}^{(t)} the momentum gradient defined by gr(t+1)=γgr(t)+(1−γ)∇wrL^(W(t))\bm{g}_{r}^{(t+1)}=\gamma\bm{g}_{r}^{(t)}+(1-\gamma)\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)}). We introduce the projection of these objects on the feature w∗\bm{w^{*}} and noise patches Xi[j]\bm{X}_{i}[j]:

– Projection on w∗\bm{w^{*}}: cr(t)=⟨wr(t),w∗⟩c_{r}^{(t)}=\langle\bm{w}_{r}^{(t)},\bm{w^{*}}\rangle.

– Projection on Xi[j]:\bm{X}_{i}[j]: Ξi,j,r(t)=⟨wr(t),Xi[j]⟩\Xi_{i,j,r}^{(t)}=\langle\bm{w}_{r}^{(t)},\bm{X}_{i}[j]\rangle.

– Total noise: Ξi(t)=∑r=1m∑j∈[P]\{P(Xi)}yi(Ξi,j,r(t))3.\Xi_{i}^{(t)}=\sum_{r=1}^{m}\sum_{j\in[P]\backslash\{P(\bm{X}_{i})\}}y_{i}(\Xi_{i,j,r}^{(t)})^{3}.

– Maximum signal: c(t)=max⁡r∈[m]crmax⁡(t)c^{(t)}=\max_{r\in[m]}c_{r_{\max}}^{(t)}.

Lastly, we define the negative sigmoid S(x)=1/(1+ex).\mathfrak{S}(x)=1/(1+e^{x}).

We now provide our first result which states that the learner model trained with GD does not generalize well on D\mathcal{D}.

Assume that we run GD on P for TT 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 ∣Z1∣≫∣Z2∣|\mathcal{Z}_{1}|\gg|\mathcal{Z}_{2}| and our choice of parameters for α,β,σ\alpha,\beta,\sigma, the gradient points mainly in the direction of w∗\bm{w^{*}}. Therefore, GD eventually learns the feature in Z1\mathcal{Z}_{1} (Lemma 5.1) and the gradients from Z1\mathcal{Z}_{1} quickly become small. Afterwards, the gradient is dominated by the gradients from Z2\mathcal{Z}_{2} (Lemma 5.2). Because Z2\mathcal{Z}_{2} has small margin, the full gradient is now directed by the noisy patches. It implies that GD memorizes noise in Z2\mathcal{Z}_{2} (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 D\mathcal{D}.

Assume that we run GD+M on (P) for TT 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 Z1\mathcal{Z}_{1} (Lemma 6.1). Contrary to GD, the momentum gradient is still highly correlated with w∗\bm{w^{*}} after this step (Lemma 6.2). Indeed, the key difference is that momentum accumulates historical gradients. Since these gradients were accumulated when learning Z1\mathcal{Z}_{1}, the direction of momentum gradient is highly biased towards w∗\bm{w^{*}}. 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 Z1\mathcal{Z}_{1} and Z2.\mathcal{Z}_{2}. 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 w∗\bm{w^{*}} and Xi[j]\bm{X}_{i}[j]. The projection of the vanilla and momentum gradients along these directions are

Gr(t)=⟨∇wrL^(W(t)),w∗⟩\mathscr{G}_{r}^{(t)}=\langle\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)}),\bm{w^{*}}\rangle and Gr(t)=⟨gr(t),w∗⟩.\mathcal{G}_{r}^{(t)}=\langle\bm{g}_{r}^{(t)},\bm{w^{*}}\rangle.

Gi,j,r(t)=⟨∇wrL^(W(t)),Xi[j]⟩\texttt{G}_{i,j,r}^{(t)}=\langle\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)}),\bm{X}_{i}[j]\rangle and Gi,j,r(t)=⟨gr(t),Xi[j]⟩.G_{i,j,r}^{(t)}=\langle\bm{g}_{r}^{(t)},\bm{X}_{i}[j]\rangle.

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 λ=0\lambda=0. A more detailed proof extending to λ>0\lambda>0 can be found in the Appendix.

At the beginning of the learning process, the gradient is mostly dominated by the gradients coming from the Z1\mathcal{Z}_{1} samples. Since these data have large margin, the gradient is thus highly correlated with w∗\bm{w^{*}} and cr(t)c_{r}^{(t)} increases as shown in the following Lemma.

For all r∈[m]r\in[m] and t≥0t\geq 0, (1) is simplified as:

Lemma 5.3 implies that quantifying the decrease rate of ν2(t)\nu_{2}^{(t)} provides an estimate on the quantity of feature learnt by the model. We remark that ν2(t)=S(β3∑s=1m(cs(t))3+Ξi(t))\nu_{2}^{(t)}=\mathfrak{S}(\beta^{3}\sum_{s=1}^{m}(c_{s}^{(t)})^{3}+\Xi_{i}^{(t)}) for some i∈Z2i\in\mathcal{Z}_{2}. 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 Z2\mathcal{Z}_{2} data increases until being large.

Let i∈Z2i\in\mathcal{Z}_{2}, j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m]r\in[m]. For t≥0t\geq 0, (2) is simplified as:

By Lemma 5.4, the noise Ξi(t)\Xi_{i}^{(t)} dominates in ν2(t)\nu_{2}^{(t)}. Consequently, the algorithm memorizes the Z2\mathcal{Z}_{2} data which implies a fast decay of ν2(t)\nu_{2}^{(t)}.

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 λ=0\lambda=0. A proof extending to λ>0\lambda>0 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 w∗\bm{w^{*}} and therefore, the GD+M model also increases its correlation with w∗\bm{w^{*}}.

Contrary to GD, GD+M has a large momentum that contains w∗\bm{w^{*}} after Step 1.

Lemma 6.2 hints an important distinction between GD and GD+M: while the current gradient along w∗\bm{w^{*}} is small at time T0,\mathcal{T}_{0}, the momentum gradient stores historical gradients that are spanned by w∗\bm{w^{*}}. It amplifies the feature present in previous gradients to learn the feature in Z2\mathcal{Z}_{2}.

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 λ=0\lambda=0, we will extend the proof to λ>0\lambda>0 in section H

Our paper rely on the notions of signal and noise components of the iterates.

– Signal intensity: θ=α\theta=\alpha if i∈Z1i\in\mathcal{Z}_{1} and β\beta otherwise.

– Signal: cr(t)=⟨w∗,wr(t)⟩c_{r}^{(t)}=\langle\bm{w^{*}},\bm{w}_{r}^{(t)}\rangle for r∈[m].r\in[m].

– Noise: Ξi,j,r(t)=⟨wr(t),Xi[j]⟩\Xi_{i,j,r}^{(t)}=\langle\bm{w}_{r}^{(t)},\bm{X}_{i}[j]\rangle for i∈[N]i\in[N] and j∈[P]\{P(Xi)}.j\in[P]\backslash\{P(\bm{X}_{i})\}.

– Max noise: Ξmax⁡(t)=max⁡i∈[N],j≠P(Xi),r∈[m]∣Ξi,j,r(t)∣2.{\Xi}_{\max}^{(t)}=\max_{i\in[N],j\not=P(\bm{X}_{i}),r\in[m]}|\Xi_{i,j,r}^{(t)}|^{2}.

– Total noise: Ξi(t)=∑r∈[m],j∈[P],j≠P(Xi)yi(Ξi,j,r(t))3.\Xi_{i}^{(t)}=\sum_{r\in[m],j\in[P],j\not=P(\bm{X}_{i})}y_{i}\left(\Xi_{i,j,r}^{(t)}\right)^{3}.

We also use the following notations when dealing with the loss function and its gradient.

– Noise loss: L^(t)(Ξi(t))=log⁡(1+exp⁡(−Ξi(t)))\widehat{\mathcal{L}}^{(t)}(\Xi_{i}^{(t)})=\log\left(1+\exp\left(-\Xi_{i}^{(t)}\right)\right).

– Full derivative: ν(t)=ν1(t)+ν2(t).\nu^{(t)}=\nu_{1}^{(t)}+\nu_{2}^{(t)}.

– Gradient on signal: Gr(t)=⟨∇wrL^(W(t)),w∗⟩\mathscr{G}_{r}^{(t)}=\langle\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)}),\bm{w^{*}}\rangle for r∈[m].r\in[m].

– Gradient on noise: Gi,j,r(t)=⟨∇wrL^(W(t)),Xi[j]⟩\texttt{G}_{i,j,r}^{(t)}=\langle\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)}),\bm{X}_{i}[j]\rangle for i∈[N]i\in[N], j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m].r\in[m].

C.2 Notations specific to GD+M

We now introduce the notations that only appear in the proofs involving GD+M.

– Momentum gradient oracle: gr(t)=γgr(t−1)+(1−γ)∇wrL^(W(t))\bm{g}_{r}^{(t)}=\gamma\bm{g}_{r}^{(t-1)}+(1-\gamma)\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)}) for r∈[m].r\in[m].

– Signal momentum: Gr(t):=⟨gr(t),w∗⟩\mathcal{G}_{r}^{(t)}:=\langle\bm{g}_{r}^{(t)},\bm{w^{*}}\rangle for r∈[m].r\in[m].

– Noise momentum: Gi,j,r(t)=⟨gr(t),Xi[j]⟩G_{i,j,r}^{(t)}=\langle\bm{g}_{r}^{(t)},\bm{X}_{i}[j]\rangle for i∈[N]i\in[N], j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m].r\in[m].

Appendix D Induction hypotheses

We prove our main result using an induction. More specifically, we make the following assumptions for every time t≤T.t\leq T.

Throughout the training process using GD for t≤Tt\leq T, we maintain that:

(Large signal data have small noise component). For every i∈Z1i\in\mathcal{Z}_{1}, for every j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m],r\in[m], we maintain:

(Small signal data have large noise component). For every i∈Z2i\in\mathcal{Z}_{2}, for every j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m],r\in[m], we have:

Throughout the training process using GD for t≤Tt\leq T, the signal component is bounded for every r∈[m]r\in[m] as

Throughout the training process using GD for t≤Tt\leq T, we maintain:

Throughout the training process using GD+M for t≤Tt\leq T, for every i∈[N]i\in[N], for every j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\}, we have that:

Throughout the training process using GD+M for t≤Tt\leq T, for r∈[m]r\in[m], we have that:

In what follows, we assume these induction hypotheses for t<Tt<T to prove our generalization results. We then prove these hypotheses for t+1.t+1.

Appendix E Gradients and updates

In this section, we first derive the gradient of the loss L^\widehat{L}. We then provide its projection on w∗\bm{\bm{w^{*}}} (signal gradient) and on Xi[j]\bm{X}_{i}[j] (noise gradient). We first derive the gradient of the loss L^.\widehat{L}.

For t≥0t\geq 0 and r∈[m]r\in[m], the gradient of the loss L^\widehat{L} with respect to wr\bm{w}_{r} is:

. We derive L^\widehat{L} with respect to wr\bm{w}_{r} 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 w∗.\bm{w^{*}}.

For all t≥0t\geq 0 and r∈[m]r\in[m], the signal gradient is:

We obtain the desired result by projecting the gradient from Lemma E.1 on w∗\bm{w^{*}} and using Xi[j]⊥w∗.\bm{X}_{i}[j]\perp\bm{w^{*}}. ∎

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 ∇wrL^\nabla_{\bm{w}_{r}}\widehat{L} on Xi[j].\bm{X}_{i}[j].

For all t≥0t\geq 0, i∈[N]i\in[N] and j∈[P]\{P(Xi)}j\in[P]\backslash\{P(X_{i})\} and r∈[m]r\in[m], the noise gradient is:

Similarly to Lemma E.2, we obtain the desired result by projecting the gradient from Lemma E.1 on Xi[j]\bm{X}_{i}[j] and using Xi[j]⊥w∗.\bm{X}_{i}[j]\perp\bm{w^{*}}. ∎

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 cr(t)c_{r}^{(t)} in subsection F.1. We then analyze the dynamics of the noise Ξi,j,r(t)\Xi_{i,j,r}^{(t)} 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 t≥0t\geq 0 and r∈[m]r\in[m], 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 Z1\mathcal{Z}_{1} term. ∎

For t∈[0,T0]t\in[0,T_{0}], we know that for all s∈[m]s\in[m], we have cs(t)≤κm1/3αc_{s}^{(t)}\leq\frac{\kappa}{m^{1/3}\alpha}. 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 r∈[m]r\in[m]. From Lemma F.2, the signal update for t∈[0,T0]t\in[0,T_{0}] is

where AA and BB are respectively defined as:

We now prove Lemma 5.2. It states that since the signal c(t)c^{(t)} has significantly increased, the Z1\mathcal{Z}_{1} derivative ν1(t)\nu_{1}^{(t)} is now small. Before proving this result, we introduce an auxiliary Lemma.

From Lemma F.3, we deduce an upper bound on ν1(t)\nu_{1}^{(t)}:

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 L^(t)(α)\widehat{\mathcal{L}}^{(t)}(\alpha). We use it to bound ν1(t).\nu_{1}^{(t)}.

The bound on ν(t)\nu^{(t)} is obtained by using its definition ν(t)=ν1(t)+ν2(t)\nu^{(t)}=\nu_{1}^{(t)}+\nu_{2}^{(t)}. ∎

We earlier proved that after T0T_{0} iterations, the signal c(t)c^{(t)} learnt by the GD model significantly increases until making ν1(t)\nu_{1}^{(t)} small. We therefore need to rewrite the signal update in this case.

For t∈[T]t\in[T], the maximal signal c(t)c^{(t)} updates as:

From the signal update given by Lemma F.1, we know that:

To obtain the desired result, we need to prove for i∈Z1i\in\mathcal{Z}_{1}:

By using D.1 and D.2, (27) is bounded as:

Using Remark 1, the sigmoid term in (F.1.2) becomes small when αc(t)≥κ1/3\alpha c^{(t)}\geq\kappa^{1/3}. To summarize, we have:

Besides, we use D.2 to bound (c(t))2(c^{(t)})^{2} in the right-hand side of (25). ∎

We now show that once ν1(t)\nu_{1}^{(t)} is small, the amount of learnt signal is controlled by ν2(t)\nu_{2}^{(t)} .

Let τ∈[T0,T)\tau\in[T_{0},T). From Lemma F.4, we know that:

Let t∈[τ,T)t\in[\tau,T). We now sum up (30) for τ=T0,…,t\tau=T_{0},\dots,t and obtain:

We now plug the bound on ν1(t)\nu_{1}^{(t)} from Lemma 5.2 in (31). This implies:

F.2 Memorization process of GD

Lemma 5.2 shows that after T0T_{0} iterations, the gradient is controlled by ν2(t)\nu_{2}^{(t)}. In this section, we show that this yields the GD model to memorize.

Using Lemma F.1, we simplify the noise update.

Let all t≥0t\geq 0, i∈[N]i\in[N], j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m]r\in[m]. Then, with probability at least 1−o(1)1-o(1), the noise update (2) is bounded as

Let i∈[N]i\in[N], j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m]r\in[m]. From Lemma E.3, we know that the noise update satisfies:

We now apply Lemma K.5 and Lemma K.7 to respectively bound ∥Xi[j]∥22\|\bm{X}_{i}[j]\|_{2}^{2} and ⟨Xa[k],Xi[j]⟩\langle\bm{X}_{a}[k],\bm{X}_{i}[j]\rangle in (34) and obtain the desired result. ∎

In the next lemma, we further simplify the noise update from Lemma F.5.

Let i∈Z2i\in\mathcal{Z}_{2}, j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m]r\in[m]. Our starting point is Lemma I.4 which states that:

Lemma F.7 indicates that yiΞi,j,r(t)y_{i}\Xi_{i,j,r}^{(t)} is not non-decreasing but overall, this quantity gets large over time. We now want to determine the time T1T_{1} where one of the yiΞi,j,r(t)y_{i}\Xi_{i,j,r}^{(t)} becomes large.

Lemma F.7 indicates that the noise iterate satisfies for t∈[0,T1]t\in[0,T_{1}]:

We proved in the previous section that after T1T_{1} iterations, the amount of noise memorized by the GD model significantly increases. We want to show that after this phase, ν2(t)\nu_{2}^{(t)} is well-controlled.

On the other hand, from Lemma I.4, we know that for all r∈[m]r\in[m]:

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 Z2\mathcal{Z}_{2} derivatives .

Combining the bound on ∑j=T1Tν2(j)\sum_{j=T_{1}}^{T}\nu_{2}^{(j)} from Lemma F.9 and (52) yields:

We have thus a control on the sum over time of ν2(t)\nu_{2}^{(t)}. We can make use of Lemma 5.3 to get the final control on the signal iterate c(t).c^{(t)}.

Let t∈[T].t\in[T]. 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 r∈[m]r\in[m]

We now bound the training and test error achieved by GD at time T.T.

Train error. Lemma I.8 provides a convergence bound on the training loss.

Test error. Let (X,y)(X,y) be a datapoint. We remind that X=(X,…,X[P])\bm{X}=(\bm{X},\dots,\bm{X}[P]) where X[P(X)]=θyw∗\bm{X}[P(\bm{X})]=\theta y\bm{w^{*}} and X[j]∼N(0,σ2Id)\bm{X}[j]\sim\mathcal{N}(0,\sigma^{2}\mathbf{I}_{d}) for j∈[P]\{P(X)}.j\in[P]\backslash\{P(\bm{X})\}. We bound the test error as follows:

We now want to compute the probability terms in (58) and (59). We remind that (X,y)∼Z1,(\bm{X},y)\sim\mathcal{Z}_{1}, yfW(T)(X)yf_{\bm{W}^{(T)}}(\bm{X}) is given by

We now apply Lemma 5.6 in (60) and obtain:

Let (X,y)∼Z2(\bm{X},y)\sim\mathcal{Z}_{2}. Similarly, by applying Lemma 5.6, yfW(T)(X)yf_{\bm{W}^{(T)}}(\bm{X}) is bounded as:

Therefore, using (122), we upper bound the test error (120) as:

Since yy is taken uniformly from {−1,1},\{-1,1\}, 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 t+1.t+1.

We prove here the main hypotheses we made on the noise when using GD.

Let’s start with the upper bound yiΞi,j,r(t)y_{i}\Xi_{i,j,r}^{(t)} for i∈Z2i\in\mathcal{Z}_{2}. Using Lemma I.6, Lemma 5.5 and D.1, we deduce from (66) that:

which proves the induction hypothesis for t+1.t+1. Regarding the lower bound, using D.1 and Lemma 5.5, we deduce from (66) that:

which proves the induction hypothesis for t+1.t+1.

Using D.1, we bound yiΞi,j,r(0)y_{i}\Xi_{i,j,r}^{(0)} and (Ξa,k,r(τ))2(\Xi_{a,k,r}^{(\tau)})^{2} 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 t+1.t+1.

We prove the induction hypotheses for the signal cr(t).c_{r}^{(t)}.

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 t≥0t\geq 0 and r∈[m]r\in[m], the signal momentum in (3) is equal to:

By definition of the momentum update, we have: gr(t+1)=γgr(t)+(1−γ)∇wrL^(W(t)).\bm{g}_{r}^{(t+1)}=\gamma\bm{g}_{r}^{(t)}+(1-\gamma)\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)}). We project this update onto w∗\bm{w^{*}} and use Lemma E.2 to get:

From Lemma G.1, we can simplify the momentum update as:

For t∈[0,T0]t\in[0,\mathcal{T}_{0}], we know that for all s∈[m]s\in[m], we have cs(t)≤κm1/3αc_{s}^{(t)}\leq\frac{\kappa}{m^{1/3}\alpha} . 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 ν1(t)\nu_{1}^{(t)} is non-zero.

By Lemma G.2, the signal update for t∈[0,T0]t\in[0,\mathcal{T}_{0}] satisfies:

We now show that contrary to GD, GD+M still has a large momentum in the w∗\bm{w^{*}} direction. In other words, we want to show that −G(t)-\mathcal{G}^{(t)} is still large after T0\mathcal{T}_{0} iterations. Given that the small margin and large margin data share the same feature w∗\bm{w^{*}}, this large momentum helps to learn Z2\mathcal{Z}_{2}.

Before proving such result, we need some intermediate lemmas.

Using the momentum update rule, we know that:

Let’s define t′:=T0−11−γt^{\prime}:=\mathcal{T}_{0}-\frac{1}{\sqrt{1-\gamma}}. We start by summing the GD+M update (3) for τ=t′,…,T0\tau=t^{\prime},\dots,\mathcal{T}_{0} to get

Applying Lemma G.3 to bound the momentum gradient, we further bound (83) to get:

We now use the fact that T0−t′=11−γ\mathcal{T}_{0}-t^{\prime}=\frac{1}{\sqrt{1-\gamma}} in (84) to get:

Since γ=1−ε\gamma=1-\varepsilon with ε≪1\varepsilon\ll 1, we linearize the right-hand side in (85) to obtain:

Using Lemma G.4, we can therefore show that once we learn Z1,\mathcal{Z}_{1}, G(t)\mathcal{G}^{(t)} still stays large.

Using Lemma G.4, we bound (cr(τ))(c_{r}^{(\tau)}) 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 T1∈[T]\mathcal{T}_{1}\in[T] such that T0<T1.\mathcal{T}_{0}<\mathcal{T}_{1}. From the signal momentum update, we deduce:

We now apply Lemma 6.2 to bound −G(T1)-\mathcal{G}^{(\mathcal{T}_{1})} in (89) and get:

We would like to find the time T1\mathcal{T}_{1} such that γT1−T0\gamma^{\mathcal{T}_{1}-\mathcal{T}_{0}} is a constant factor a≤1a\leq 1 i.e. such that

Let t∈(T1,T]t\in(\mathcal{T}_{1},T]. Using (3) update rule, we have

where we used the fact that −G(τ)≥0-\mathcal{G}^{(\tau)}\geq 0 in (94). Plugging (93) in (94) yields the desired bound.

G.2 GD+M does not memorize

Lemma 6.3 implies that after T1\mathcal{T}_{1} 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 ν(t)\nu^{(t)}.

For t∈[T1,T]t\in[\mathcal{T}_{1},T], 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 c(t)c^{(t)} in (96) to obtain the desired result. ∎

Lemma G.5 provides an upper bound on ν(t)\nu^{(t)} since:

We now would like to give a convergence rate on the iterates G(t+1)−γG(t).\mathcal{G}^{(t+1)}-\gamma\mathcal{G}^{(t)}. 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 t∈[T1,T]t\in[\mathcal{T}_{1},T], we have:

Indeed, by using Lemma 6.3 and D.5, we have:

Given our choice of α\alpha, β\beta and μ^\hat{\mu}, 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 T1\mathcal{T}_{1} iterations, the gradient is now very small and the noise component learnt by GD+M stays very small.

We sum up (108) for s=T1,…,ts=\mathcal{T}_{1},\dots,t and obtain:

We apply the triangle inequality in (109) and obtain:

We now use D.4 to bound ∣Ξi,j,r(T1)∣|\Xi_{i,j,r}^{(\mathcal{T}_{1})}| in (110):

We now plug the bound on ∑s=T1t∣Gi,j,r(s+1)∣\sum_{s=\mathcal{T}_{1}}^{t}|G_{i,j,r}^{(s+1)}| given by Lemma J.6 and obtain:

Given the values of T1,\mathcal{T}_{1}, η\eta, γ\gamma and β\beta, we can deduce that

Plugging (113) in (112) proves the induction hypothesis for t+1.t+1.

G.3 Proof of Theorem 4.2

We proved that the weights learnt by GD+M satisfy for r∈[m]r\in[m]

We now bound the training and test error achieved by GD+M at time T.T.

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 (X,y)(\bm{X},y) be a datapoint. We remind that X=(X,…,X[P])\bm{X}=(\bm{X},\dots,\bm{X}[P]) where X[P(X)]=θyw∗\bm{X}[P(\bm{X})]=\theta y\bm{w^{*}} and X[j]∼N(0,σ2Id)\bm{X}[j]\sim\mathcal{N}(0,\sigma^{2}\mathbf{I}_{d}) for j∈[P]\{P(X)}.j\in[P]\backslash\{P(\bm{X})\}. We bound the test error as follows:

We now want to compute the probability terms in (119) and (120). We remind that yfWT(X)yf_{\bm{W}^{T}}(\bm{X}) 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 yy is uniformly sampled from {−1,1},\{-1,1\}, we further simplify (123) as:

We know that ⟨vs(T),X[j]⟩∼N(0,∥vs(T)∥22σ2)\langle\bm{v}_{s}^{(T)},\bm{X}[j]\rangle\sim\mathcal{N}(0,\|\bm{v}_{s}^{(T)}\|_{2}^{2}\sigma^{2}). Therefore, ⟨vs(T),X[j]⟩3\langle\bm{v}_{s}^{(T)},\bm{X}[j]\rangle^{3} is the cube of a centered Gaussian.This random variable is symmetric. By Lemma K.1, we know that ∑s=1m∑j≠P(X)⟨vs(T),X[j]⟩3\sum_{s=1}^{m}\sum_{j\neq P(\bm{X})}\langle\bm{v}_{s}^{(T)},\bm{X}[j]\rangle^{3} is also symmetric. Therefore, we simplify (124) as:

From Lemma K.14, we know that ∑s=1m∑j≠p⟨vs(T),X[j]⟩3\sum_{s=1}^{m}\sum_{j\neq p}\langle\bm{v}_{s}^{(T)},\bm{X}[j]\rangle^{3} is σ3P−1∑s=1m∥vs(T)∥26\sigma^{3}\sqrt{P-1}\sqrt{\sum_{s=1}^{m}\|\bm{v}_{s}^{(T)}\|_{2}^{6}}-subGaussian. Therefore, by applying Lemma K.3, (125) is further bounded by:

Using the fact that ∥vs(T)∥2≤1\|\bm{v}_{s}^{(T)}\|_{2}\leq 1 in (126) finally yields:

G.4 Proof of the GD+M induction hypotheses

We prove the induction hypotheses for the signal cr(t).c_{r}^{(t)}.

We sum up (128) for τ=T1,…,t\tau=\mathcal{T}_{1},\dots,t and obtain:

We apply the triangle inequality in (129) and obtain:

We now use D.5 to bound ∣cr(T1)∣|c_{r}^{(\mathcal{T}_{1})}| in (130):

We now plug the bound on ∑τ=T1t∣Gr(τ+1)∣\sum_{\tau=\mathcal{T}_{1}}^{t}|\mathcal{G}_{r}^{(\tau+1)}| given by Lemma J.3. We have:

Appendix H Extension to λ>0𝜆0\lambda>0

After iteration TT, by Lemma I.8 and Lemma J.9, we know that for GD:

On the other hand, for Ξi,j,r(t)\Xi_{i,j,r}^{(t)} 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 t,T∈[T]t,\mathscr{T}\in[T] such that T<t.\mathscr{T}<t. Then, the Z1\mathcal{Z}_{1} derivative is bounded as:

Let T,t∈[T]\mathscr{T},t\in[T] such that T<t.\mathscr{T}<t. We now sum up (135) for τ=T,…,t\tau=\mathscr{T},\dots,t and get:

Subcase 1: T<T0.\mathscr{T}<T_{0}. From Lemma 5.3, we know that:

Subcase 2: T>T0.\mathscr{T}>T_{0}. 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 i∈[N]i\in[N], j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m]r\in[m]. Let T,t∈[T]\mathscr{T},t\in[T] such that T<t.\mathscr{T}<t. Then, the noise update (2) satisfies

Let i∈[N]i\in[N], j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m]r\in[m]. We set up the following induction hypothesis:

Let’s first show this hypothesis for t=T.t=\mathscr{T}. From Lemma F.5, we have:

Now, we apply D.3 to bound (Ξa,k,r(T))2(\Xi_{a,k,r}^{(\mathscr{T})})^{2} in (143) and obtain:

Therefore, the induction hypothesis is verified for t=T.t=\mathscr{T}. Now, assume (LABEL:eq:noiseindhypoth) for t.t. Let’s prove the result for t+1.t+1. We start by summing up the noise update from Lemma F.5 for τ=T,…,t\tau=\mathscr{T},\dots,t which yields:

We apply D.3 to bound (Ξa,k,r(t))2(\Xi_{a,k,r}^{(t)})^{2} 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 t+1.t+1.

Now, let’s simplify the sum terms in (LABEL:eq:noiseindhypoth). Since P≪dP\ll\sqrt{d}, 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 i∈[N]i\in[N], j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m]r\in[m]. Let T,t∈[T]\mathscr{T},t\in[T] such that T<t.\mathscr{T}<t. 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 j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\} and r∈[m]r\in[m]:

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 t∈[T0,T]t\in[T_{0},T]. Run GD with learning rate η\eta for tt iterations. Then, the Z1\mathcal{Z}_{1} loss sublinearly converges to zero as:

Let t∈[T0,T].t\in[T_{0},T]. From Lemma F.1, we know that the signal update is lower bounded as:

Let’s now assume by contradiction that for t∈[T0,T]t\in[T_{0},T], we have:

From the (3) update, we know that cr(τ)c_{r}^{(\tau)} is a non-decreasing sequence which implies that ∑r=1m(αcr(τ))3\sum_{r=1}^{m}(\alpha c_{r}^{(\tau)})^{3} is also non-decreasing. Since x↦log⁡(1+exp⁡(−x))x\mapsto\log(1+\exp(-x)) is non-increasing, this implies that for s≤ts\leq t, we have:

Plugging (164) in the update (162) yields for s∈[T0,t]s\in[T_{0},t]:

Let t∈[T0,T]t\in[T_{0},T]. We now sum (165) for s=T0,…,ts=T_{0},\dots,t and obtain:

Given the values of T,η,α,μ^T,\eta,\alpha,\hat{\mu}, we finally have:

Let t∈[T1,T]t\in[T_{1},T]. Run GD with learning rate η∈(0,1/L)\eta\in(0,1/L) for tt iterations. Then, the loss sublinearly converges to zero as:

We first apply the classical descent lemma for smooth functions (Lemma K.18). Since L^(W)\widehat{L}(W) 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 t∈[T1,T]t\in[T_{1},T]. Run GD for tt iterations. Then, the norm of gradient is lower bounded as follows:

Let t∈[T1,T]t\in[T_{1},T]. To obtain the lower bound, we project the gradient on the the signal and on the noise.

Since ∥w∗∥2=1\|w^{*}\|_{2}=1, we lower bound ∥∇wrL^(W(t))∥22\|\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)})\|_{2}^{2} as

By successively applying Lemma E.2 and Lemma I.1, (Gr(t))2(\mathscr{G}_{r}^{(t)})^{2} is lower bounded as

For a fixed i∈Z2i\in\mathcal{Z}_{2} and j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\}, we know that ∥∇wrL^(W(t))∥22\|\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)})\|_{2}^{2} is lower bounded as

Combining (172), (175), (173) and (176) and using 2a2+2b2≥(a+b)2,2a^{2}+2b^{2}\geq(a+b)^{2}, we thus bound ∥∇wrL^(W(t))∥22\|\nabla_{\bm{w}_{r}}\widehat{L}(\bm{W}^{(t)})\|_{2}^{2} as:

We now sum up (177) for r=1,…,mr=1,\dots,m 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 L^(W(t))2\widehat{L}(\bm{W}^{(t)})^{2}. This gives the aimed result.

We now present auxiliary lemmas that link the gradient terms with their corresponding loss.

Let t∈[T1,T].t\in[T_{1},T]. Run GD for tt iterations. Then, we have:

Therefore, we can apply Lemma K.20 and get the lower bound:

Let t∈[T1,T].t\in[T_{1},T]. Run GD for tt 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 L^(t)(α)\widehat{\mathcal{L}}^{(t)}(\alpha). By applying D.1 and Lemma K.24 in (187), we finally get:

Combining (187) and (188) yields the aimed result. ∎

Let t∈[T1,T].t\in[T_{1},T]. Run GD for tt 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 L^(t)(Ξi(t))\widehat{\mathcal{L}}^{(t)}(\Xi_{i}^{(t)}). 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 t∈[0,T]t\in[0,T] Run GD for for tt iterations. Then, we have:

we need to lower bound L^(t)(α)\widehat{\mathcal{L}}^{(t)}(\alpha). 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 Gr(t)G_{r}^{(t)} which is given by the next lemma.

For r∈[m]r\in[m], the gradient of the loss L^(W(t))\widehat{L}(\bm{W}^{(t)}) projected on the normalized noise χ\textstyle\chi satisfies with probability 1−o(1)1-o(1) for r∈[m]r\in[m]:

Projecting the gradient (given by Lemma E.1) on χ\textstyle\chi yields:

Since 1N∑i∈Z2∑j≠P(Xi)Xi[j]∥1N∑b∈Z2∑l≠P(Xi)Xb[l]∥2\frac{\frac{1}{N}\sum_{i\in\mathcal{Z}_{2}}\sum_{j\neq P(\bm{X}_{i})}\bm{X}_{i}[j]}{\|\frac{1}{N}\sum_{b\in\mathcal{Z}_{2}}\sum_{l\neq P(\bm{X}_{i})}\bm{X}_{b}[l]\|_{2}} is a unit Gaussian vector, using Lemma K.8, we bound the right-hand side of (LABEL:eq:Grbd2) with probability 1−o(1)1-o(1), as:

Now, using Lemma Lemma K.10 , we can further lower bound the left-hand side of (LABEL:eq:Grbd3) as:

Remark that 1N∑b∈Z2∑l≠P(Xi)Xb[l]∼N(0,μ^PNσ2)\frac{1}{N}\sum_{b\in\mathcal{Z}_{2}}\sum_{l\neq P(\bm{X}_{i})}\bm{X}_{b}[l]\sim\mathcal{N}(0,\frac{\hat{\mu}P}{N}\sigma^{2}). 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 c(τ)c^{(\tau)} becomes large.

For t∈[T]t\in[T], the maximal signal momentum G(t)\mathcal{G}^{(t)} is bounded as:

Let t∈[T]t\in[T]. Using the signal momentum given by Lemma G.1, we know that:

To obtain the desired result, we need to prove for i∈Z1i\in\mathcal{Z}_{1}:

By using D.4 and D.5, (203) is bounded as:

Using Remark 1, the sigmoid term in (J.2) becomes small when αc(τ)≥κ1/3\alpha c^{(\tau)}\geq\kappa^{1/3}. To summarize, we have:

A similar reasoning implies for i∈Z2i\in\mathcal{Z}_{2}:

Plugging (202) and (206) in (201) yields the aimed result. ∎

For t∈[T1,T)t\in[\mathcal{T}_{1},T), the sum of maximal signal momentum is bounded as:

Let s∈[T1,T]s\in[\mathcal{T}_{1},T]. From Lemma J.2, the signal momentum is bounded as:

For τ∈[T0−1]\tau\in[\mathcal{T}_{0}-1], we have γs−τ≤γs−T0+1\gamma^{s-\tau}\leq\gamma^{s-\mathcal{T}_{0}+1} and for τ∈[T1−1]\tau\in[\mathcal{T}_{1}-1], γs−τ≤γs−T1+1\gamma^{s-\tau}\leq\gamma^{s-\mathcal{T}_{1}+1}. From Lemma J.8 and Lemma 6.4, we can bound ν1(τ)\nu_{1}^{(\tau)} and ν2(τ)\nu_{2}^{(\tau)}. 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 s=T1,…,ts=\mathcal{T}_{1},\dots,t. Using the geometric sum inequality ∑sγs≤1/(1−γ)\sum_{s}\gamma^{s}\leq 1/(1-\gamma) and obtain:

We plug ∑sγs≤1/(1−γ)\sum_{s}\sqrt{\gamma}^{s}\leq 1/(1-\sqrt{\gamma}) and ∑s=1t−T1+11/s≤log⁡(t)+1\sum_{s=1}^{t-\mathcal{T}_{1}+1}1/s\leq\log(t)+1 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 L^(W).\widehat{L}(\bm{W}). Let i∈[N]i\in[N], j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\}. At a time tt, the noise momentum is bounded with probability 1−o(1)1-o(1) as:

Let i∈[N]i\in[N] and j∈[P]\{P(Xi)}j\in[P]\backslash\{P(\bm{X}_{i})\}. Combining the (4) update rule and Lemma E.3 to get the noise gradient Gi,j,r(t)\texttt{G}_{i,j,r}^{(t)}, we obtain

Using Lemma K.5 and Lemma K.7, (LABEL:eq:diffmoms1) becomes with probability 1−o(1),1-o(1),

We upper bound the second term in (LABEL:eq:diffmoms2) by again using D.4:

Let t∈[T]t\in[T]. The noise momentum is bounded as

Let τ∈[T].\tau\in[T]. From Lemma J.4, we know that:

We unravel the recursion (217) rule for τ=0,…,t\tau=0,\dots,t and obtain:

For t∈[T1,T)t\in[\mathcal{T}_{1},T), the sum of noise momentum is bounded as:

Let s∈[T1,T)s\in[\mathcal{T}_{1},T). We first apply Lemma J.5 and obtain:

Using the bound from Lemma 6.4, (218) becomes

For τ∈[0,T1−1]\tau\in[0,\mathcal{T}_{1}-1], we have γs−1−τ≤γs−T1+1\gamma^{s-1-\tau}\leq\gamma^{s-\mathcal{T}_{1}+1}. 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 s=T1,…,ts=\mathcal{T}_{1},\dots,t. Using the geometric sum inequality ∑sγs≤1/(1−γ),\sum_{s}\gamma^{s}\leq 1/(1-\gamma), we obtain:

We finally use the harmonic series inequality ∑s=1t−T11/s≤1+log⁡(t)\sum_{s=1}^{t-\mathcal{T}_{1}}1/s\leq 1+\log(t) 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 t∈[T0,T]t\in[\mathcal{T}_{0},T] Using GD+M with learning rate η\eta, the loss sublinearly converges to zero as

Let t∈[T0,T].t\in[\mathcal{T}_{0},T]. 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 t∈[T0,T]t\in[\mathcal{T}_{0},T], we have:

From the (3) update, we know that cr(τ)c_{r}^{(\tau)} is a non-decreasing sequence which implies that ∑r=1m(αcr(τ))3\sum_{r=1}^{m}(\alpha c_{r}^{(\tau)})^{3} is also non-decreasing for τ∈[T]\tau\in[T]. Since x↦log⁡(1+exp⁡(−x))x\mapsto\log(1+\exp(-x)) is non-increasing, this implies that for s≤ts\leq t, we have:

Plugging (229) in the update (228) yields for s∈[T0,t]s\in[\mathcal{T}_{0},t]:

We now sum (230) for s=T0,…,ts=\mathcal{T}_{0},\dots,t and obtain:

Given the values of α,η,T\alpha,\eta,T, we finally have:

We now link the bound on the loss to the derivative ν1(t).\nu_{1}^{(t)}.

The proof is similar to the one of Lemma 6.4.

For t∈[T1,T]t\in[\mathcal{T}_{1},T] Using GD+M with learning rate η>0\eta>0, the loss sublinearly converges to zero as

Let t∈[T1,T].t\in[\mathcal{T}_{1},T]. From Lemma J.10, we know that the signal gradient is bounded as −G(t)≥−G(s)-\mathscr{G}^{(t)}\geq-\mathscr{G}^{(s)} for s∈[T1,t].s\in[\mathcal{T}_{1},t].

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 t∈[T1,T]t\in[\mathcal{T}_{1},T], we have:

From the (3) update, we know that cr(τ)c_{r}^{(\tau)} is a non-decreasing sequence which implies that ∑r=1m(θcr(τ))3\sum_{r=1}^{m}(\theta c_{r}^{(\tau)})^{3} is also non-decreasing for τ∈[T]\tau\in[T]. Since x↦log⁡(1+exp⁡(−x))x\mapsto\log(1+\exp(-x)) is non-increasing, this implies that for s≤ts\leq t, we have:

Plugging (243) in the update (241) yields for s∈[T1,t]s\in[\mathcal{T}_{1},t]:

We now sum (244) for s=T1,…,ts=\mathcal{T}_{1},\dots,t and obtain:

Given the values of α,β,η,T,μ^\alpha,\beta,\eta,T,\hat{\mu}, we finally have:

J.4.3 Auxiliary lemmas

We now provide an auxiliary lemma needed to obtain (J.9).

Let t∈[T1,T].t\in[\mathcal{T}_{1},T]. Then, the signal gradient decreases i.e. −G(s)≥−G(t)-\mathscr{G}^{(s)}\geq-\mathscr{G}^{(t)} for s∈[T1,t].s\in[\mathcal{T}_{1},t].

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 σ1,σ2>0.\sigma_{1},\sigma_{2}>0. Let XX and YY respetively be σ1\sigma_{1}- and σ2\sigma_{2}-subGaussian random variables. Then, X+YX+Y is σ1+σ2\sqrt{\sigma_{1}+\sigma_{2}}-subGaussian random variable.

Let t>0.t>0. Let XX be a σ\sigma-subGaussian random variable. Then, we have:

We know that the ∥⋅∥2\|\cdot\|_{2} is 11-Lipschitz and by applying Theorem K.1, we therefore have::

By rewriting (253) and using Lemma K.4, we have with probability 1−δ,1-\delta,

By squaring (254) and using (a+b)2≤a2+b2(a+b)^{2}\leq a^{2}+b^{2}, we obtain the aimed result. ∎

We know that the ∥⋅∥2\|\cdot\|_{2} is 11-Lipschitz and by applying Theorem K.1, we therefore have:

We use Lemma K.4 and set ϵ=σd2\epsilon=\frac{\sigma\sqrt{d}}{2} in (255) to finally get:

Let’s define Z:=⟨X,Y⟩.Z:=\langle\bm{X},\bm{Y}\rangle. We first remark that ZZ is a sub-exponential random variable. Indeed, the generating moment function is:

where we used log⁡(1−x)≥−x\log(1-x)\geq-x for x<1x<1 in the last inequality. Therefore, by definition of a sub-exponential variable, we have:

Since σ2d≤1\sigma^{2}d\leq 1 and ϵ∈,\epsilon\in, (256) is bounded as:

Let U:=X/∥X∥2\bm{U}:=\bm{X}/\|\bm{X}\|_{2} and Z:=⟨U,Y⟩.Z:=\langle\bm{U},\bm{Y}\rangle. We know that the pdf of U\bm{U} in polar coordinates is fU(θ)=Γ(d/2)2πd/2.f_{\bm{U}}(\theta)=\frac{\Gamma(d/2)}{2\pi^{d/2}}. Therefore, the generating moment function of ZZ is:

(258) indicates that ZZ is a sub-Gaussian random variable of parameter σ\sigma. By definition, it satisfies

Setting δ=2e−ϵ22σ2\delta=2e^{-\frac{\epsilon^{2}}{2\sigma^{2}}} in (259) yields that we have with probability 1−δ,1-\delta,

Let X1,…,Xn\bm{X}_{1},\dots,\bm{X}_{n} i.i.d. vectors from N(0,σ2I).\mathcal{N}(0,\sigma^{2}\mathbf{I}). Then, with probability 1−o(1)1-o(1), we have:

We know that for X1∼N(0,σ2d)\bm{X}_{1}\sim\mathcal{N}(0,\sigma^{2}d), we have:

Therefore, using the law of total probability and (261), we have:

Since ∑i=1nXi∼N(0,nσ2Id)\sum_{i=1}^{n}\bm{X}_{i}\sim\mathcal{N}(0,n\sigma^{2}\mathbf{I}_{d}), we also have

Therefore by setting t=3σ2dnt=\frac{3\sigma}{2}\sqrt{\frac{d}{n}}, we obtain:

Doing the similar reasoning for the lower bound yields:

Let X1,…,Xn\bm{X}_{1},\dots,\bm{X}_{n} i.i.d. vectors from N(0,σ2Id).\mathcal{N}(0,\sigma^{2}\mathbf{I}_{d}). Then, with probability 1−o(1),1-o(1), 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 P(x)=P(x1,…,xn)P(x)=P(x_{1},\dots,x_{n}) be a degree dd polynomial and x1,…,xnx_{1},\dots,x_{n} be i.i.d. Gaussian univariate random variables. Then, the following holds for all d,nd,n.

Since the decomposition of a polynomial in the monomial basis is unique, we can equate the coefficients of HH and PP and obtain:

Setting ϵ=1/2\epsilon=1/2 in (273) yields the desired result.

K.1.3 Properties of the cube of a Gaussian

Let X∼N(0,σ2)X\sim\mathcal{N}(0,\sigma^{2}). Then, X3X^{3} is σ3\sigma^{3}-subGaussian.

By definition of the moment generating function, we have:

We know that ⟨ws,X[j]⟩∼N(0,∥ws∥22σ2)\langle\bm{w}_{s},\bm{X}[j]\rangle\sim\mathcal{N}(0,\|\bm{w}_{s}\|_{2}^{2}\sigma^{2}). Therefore, ⟨ws,X[j]⟩3\langle\bm{w}_{s},\bm{X}[j]\rangle^{3} is the cube of a centered Gaussian. From Lemma K.13, ⟨ws,X[j]⟩3\langle\bm{w}_{s},\bm{X}[j]\rangle^{3} is σ3∥ws∥23\sigma^{3}\|\bm{w}_{s}\|_{2}^{3}-subGaussian. Using Lemma K.2, we deduce that ∑j=1P−1⟨ws,X[j]⟩3\sum_{j=1}^{P-1}\langle\bm{w}_{s},\bm{X}[j]\rangle^{3} is Pσ3∥ws∥23\sqrt{P}\sigma^{3}\|\bm{w}_{s}\|_{2}^{3}-subGaussian. Applying again Lemma K.2, we finally obtain that ∑s=1m∑j=1P−1⟨ws,X[j]⟩3\sum_{s=1}^{m}\sum_{j=1}^{P-1}\langle\bm{w}_{s},\bm{X}[j]\rangle^{3} is σ3P−1∑s=1m∥ws∥26\sigma^{3}\sqrt{P-1}\sqrt{\sum_{s=1}^{m}\|\bm{w}_{s}\|_{2}^{6}}-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 z(t+1)=z(t)+ηC(t)(z(t))2z^{(t+1)}=z^{(t)}+\eta C^{(t)}(z^{(t)})^{2}. This technique is reminiscent of the classical analysis of the growth of eigenvalues on the (incremental) tensor power method of degree 22 and is stated in full generality in (Allen-Zhu & Li, 2020).

Let {z(t)}t=0T\{z^{(t)}\}_{t=0}^{T} be a positive sequence defined by the following recursions

where z(0)>0z^{(0)}>0 is the initialization and m,M>0m,M>0.Let υ>0\upsilon>0 such that z(0)≤υ.z^{(0)}\leq\upsilon. Then, the time t0t_{0} such that zt≥υz_{t}\geq\upsilon for all t≥t0t\geq t_{0} is:

We use the fact that z(s)≥z(0))z^{(s)}\geq z^{(0))} in (274) and obtain:

Now, we want to bound z(T1)−z(0)z^{(T_{1})}-z^{(0)}. Using again the recursion and z(T1−1)≤2z(0)z^{(T_{1}-1)}\leq 2z^{(0)}, we have:

Combining (275) and (276), we get a bound on T1.T_{1}.

Now, let’s find a bound for TnT_{n}. Starting from the recursion and using the fact that z(s)≥2n−1z(0)z^{(s)}\geq 2^{n-1}z^{(0)} for s≥Tn−1s\geq T_{n-1} we have:

On the other hand, by using z(Tn−1)≤2nz(0)z^{(T_{n}-1)}\leq 2^{n}z^{(0)} we upper bound z(Tn)z^{(T_{n})} as follows.

Besides, we know that z(Tn−1)≥2n−1z(0)z^{(T_{n-1})}\geq 2^{n-1}z^{(0)}. Therefore, we upper bound z(Tn)−z(Tn−1)z^{(T_{n})}-z^{(T_{n-1})} as

We now sum (281) for n=2,…,nn=2,\dots,n, use (277) and obtain:

Lastly, we know that nn satisfies 2nz(0)≥υ2^{n}z^{(0)}\geq\upsilon this implies that we can set n=⌈log⁡(υ/z0)log⁡(2)⌉n=\left\lceil\frac{\log(\upsilon/z_{0})}{\log(2)}\right\rceil in (282). ∎

Let {z(t)}t=0T\{z^{(t)}\}_{t=0}^{T} be a positive sequence defined by the following recursion

where A,C>0A,C>0 and z(0)>0z^{(0)}>0 is the initialization. Assume that C≤z(0)/2.C\leq z^{(0)}/2. Let υ>0\upsilon>0 such that z(0)≤υ.z^{(0)}\leq\upsilon. Then, the time t0t_{0} such that z(t)≥υz^{(t)}\geq\upsilon is upper bounded as:

By assumption, we know that C≤z(0)/2.C\leq z^{(0)}/2. This implies that for all z(t)≥z(0)/2z^{(t)}\geq z^{(0)}/2 for all t≥0.t\geq 0. Plugging this in (284) yields:

Now, we want to upper bound z(T1)−z(0)z^{(T_{1})}-z^{(0)}. Using (283), we deduce that:

Combining the two equations in (287) yields

Since T1T_{1} is the first time where z(T1)≥z(0)z^{(T_{1})}\geq z^{(0)}, we have z(T1−1)≤z(0)z^{(T_{1}-1)}\leq z^{(0)}. Plugging this in (288) leads to:

Finally, using (289) in (286) and C=o(z(0))C=o(z^{(0)}) gives an upper bound on T1.T_{1}.

Now, let’s find a bound for TnT_{n}. Starting from the recursion, we have:

We substract the two equations in (291), use z(s)≥2n−2z^{(s)}\geq 2^{n-2} for s≥Tn−1s\geq T_{n-1} and obtain:

On the other hand, from the recursion, we have the following inequalities:

We substract the two equations in (293), use z(Tn−1)≤2n−1z(0)z^{(T_{n}-1)}\leq 2^{n-1}z^{(0)} and upper bound z(Tn)z^{(T_{n})} as follows.

Besides, we know that z(Tn−1)≥2n−2z(0)z^{(T_{n-1})}\geq 2^{n-2}z^{(0)}. Therefore, we upper bound z(Tn)−z(Tn−1)z^{(T_{n})}-z^{(T_{n-1})} as

We now sum (296) for n=2,…,nn=2,\dots,n, use C=o(z(0))C=o(z^{(0)}) and then (290) to obtain:

Lastly, we know that nn satisfies 2nz(0)≥υ2^{n}z^{(0)}\geq\upsilon this implies that we can set n=⌈log⁡(υ/z0)log⁡(2)⌉n=\left\lceil\frac{\log(\upsilon/z_{0})}{\log(2)}\right\rceil in (297). ∎

K.2.2 Bounds for GD+M

Let γ∈(0,1).\gamma\in(0,1). Let {c(t)}t≥0\{c^{(t)}\}_{t\geq 0} and {G(t)}\{\mathcal{G}^{(t)}\} be positive sequences defined by the following recursions

Let δ∈(0,1).\delta\in(0,1). We want to prove the following induction hypotheses:

After Tn=n1−γ+∑j=0n−2δ(δ+1)jη(1−e−1)α3c(0)∑τ=0je−(j−τ)(1+δ)2τT_{n}=\frac{n}{1-\gamma}+\sum_{j=0}^{n-2}\frac{\delta(\delta+1)^{j}}{\eta(1-e^{-1})\alpha^{3}c^{(0)}\sum_{\tau=0}^{j}e^{-(j-\tau)}(1+\delta)^{2\tau}} iterations, we have:

After Tn′=n1−γ+∑j=0n−1δ(δ+1)jη(1−e−1)α3c(0)∑τ=0je−(j−τ)(1+δ)2τT_{n}^{\prime}=\frac{n}{1-\gamma}+\sum_{j=0}^{n-1}\frac{\delta(\delta+1)^{j}}{\eta(1-e^{-1})\alpha^{3}c^{(0)}\sum_{\tau=0}^{j}e^{-(j-\tau)}(1+\delta)^{2\tau}}, we have:

Let’s first prove (TPM-1) and (TPM-2) for n=1.n=1. First, by using the momentum update, we have:

Setting T1=1/(1−γ)T_{1}=1/(1-\gamma) and using γ=1−ε\gamma=1-\varepsilon, we have 1−γ11−γ=1−exp⁡(log⁡(1−ε)/ε)=1−e−1.1-\gamma^{\frac{1}{1-\gamma}}=1-\exp(\log(1-\varepsilon)/\varepsilon)=1-e^{-1}. Plugging this in (298) yields (TPM-1) for n=1.n=1.

Regarding (TPM-2), we use the iterate update to have:

where we used c(T1)≥c(0)c^{(T_{1})}\geq c^{(0)} and (298) to obtain (299). Since T1′+1T_{1}^{\prime}+1 is the first time where c(t)≥(1+δ)c(0),c^{(t)}\geq(1+\delta)c^{(0)}, we further simplify (299) to obtain:

We therefore obtained (TPM-2) for n=1.n=1. Let’s now assume (TPM-1) and (TPM-2) for nn. We now want to prove these induction hypotheses for n+1.n+1. First, by using the momentum update, we have:

From (TPM-2) for nn, we know that c(t)≥(1+δ)nc(0)c^{(t)}\geq(1+\delta)^{n}c^{(0)} for t>Tn′t>T_{n}^{\prime}. Therefore, (301) becomes:

From (TPM-1), we know that −G(Tn′)≥(1−e−1)α3(c(0))2∑τ=0n−1e−(n−1−τ)(1+δ)2τ-\mathcal{G}^{(T_{n}^{\prime})}\geq(1-e^{-1})\alpha^{3}(c^{(0)})^{2}\sum_{\tau=0}^{n-1}e^{-(n-1-\tau)}(1+\delta)^{2\tau} for t≥Tn.t\geq T_{n}. Therefore, we simplify (302) as:

When we set Tn+1T_{n+1} as in (TPM-1), we have Tn+1−Tn′=11−γ.T_{n+1}-T_{n}^{\prime}=\frac{1}{1-\gamma}. Moreover, since γ=1−ε\gamma=1-\varepsilon, we have γ11−γ=e−1\gamma^{\frac{1}{1-\gamma}}=e^{-1}. Using these two observations, (303) is thus equal to:

We therefore proved (TPM-1) for n+1.n+1. Now, let’s prove (TPM-2). We use the iterates update and obtain:

where we used c(Tn+1)≥(δ+1)nc(0)c^{(T_{n+1})}\geq(\delta+1)^{n}c^{(0)} and (304) in the last inequality. Since Tn+1′+1T_{n+1}^{\prime}+1 is the first time where c(t)≥(1+δ)n+1c(0),c^{(t)}\geq(1+\delta)^{n+1}c^{(0)}, we further simplify (305) to obtain:

Let’s now obtain an upper bound on Tn′.T_{n}^{\prime}. We have:

Finally, we choose nn such that (1+δ)n≥υ(1+\delta)^{n}\geq\upsilon or equivalently, n=⌈log⁡(υ)log⁡(1+δ)⌉n=\left\lceil\frac{\log(\upsilon)}{\log(1+\delta)}\right\rceil. Plugging this choice in Tn\mathscr{T}_{n} yields the desired bound. ∎

K.3 Optimization lemmas

By applying the definition of smooth functions and the GD update, we have:

Setting η<1/L\eta<1/L in (308) leads to the expected result.

Let T≥0\mathscr{T}\geq 0. Let (xt)t>T(x_{t})_{t>\mathscr{T}} be a non-negative sequence that satisfies the recursion: x(t+1)≤x(t)−A(x(t))2,x^{(t+1)}\leq x^{(t)}-A(x^{(t)})^{2}, for A>0.A>0. Then, it is bounded at a time t>Tt>\mathscr{T} as

Let τ∈(T,t]\tau\in(\mathscr{T},t]. By multiplying each side of the recursion by (x(τ)x(τ+1))−1(x^{(\tau)}x^{(\tau+1)})^{-1}, we get:

Besides, the update rule indicates that x(τ)x^{(\tau)} is non-increasing i.e. x(τ+1)≤x(τ).x^{(\tau+1)}\leq x^{(\tau)}. Using this fact in (310) yields:

Now, we sum up (311) for τ=T,…,t−1\tau=\mathscr{T},\dots,t-1 and obtain:

Inverting (312) yields the expected result. ∎

K.4 Other useful lemmas

We apply Lemma K.21 to the sequence ai+δa_{i}+\delta 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 x↦x3x\mapsto x^{3} 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 ∑i=1nai>C−\sum_{i=1}^{n}a_{i}>C_{-} and ai>0a_{i}>0 for all ii:

where we used ai>0a_{i}>0 for all ii 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 ∑i=1nai≤C+\sum_{i=1}^{n}a_{i}\leq C_{+} and ∑i=1m∑j≠iai2aj\sum_{i=1}^{m}\sum_{j\neq i}a_{i}^{2}a_{j}:

We obtain the final bound by applying Lemma K.22 to (326). ∎

Let (x(t))t≥0(x^{(t)})_{t\geq 0} be a non-negative sequence. Let A>0.A>0. Assume that ∑τ=0Tx(τ)≤A.\sum_{\tau=0}^{T}x^{(\tau)}\leq A. Then, there exists a time T∈[T]\mathscr{T}\in[T] such that x(T)≤A/T.x^{(\mathscr{T})}\leq A/T.

Assume by contradiction that for all τ∈[T]\tau\in[T], x(τ)>A/Tx^{(\tau)}>A/T. By summing up xτx^{\tau}, we obtain ∑τ=0Tx(τ)>A.\sum_{\tau=0}^{T}x^{(\tau)}>A. This contradicts the assumption that ∑τ=0Tx(τ)≤A.\sum_{\tau=0}^{T}x^{(\tau)}\leq A.

Let x,y>0.x,y>0. Then, the following inequalities holds:

Successively using the inequalities log⁡(1+x)≤x\log(1+x)\leq x and x1+x≤log⁡(1+x)\frac{x}{1+x}\leq\log(1+x) for x>−1x>-1 in (329) yields:

This proves item 1 of the Lemma. Let’s now prove item 2. Using az≤1+(a−1)za^{z}\leq 1+(a-1)z for z∈(0,1)z\in(0,1) and a≥1a\geq 1, we know that:

Since log⁡\log is non-decreasing, applying log⁡\log to (330) proves item 2.

In Appendix J, we need to bound the sum ∑s=1tγt−ss\sum_{s=1}^{t}\frac{\gamma^{t-s}}{s} for γ<1.\gamma<1. We derive such bound here.

given our choice of γ\gamma. Let t≥2.t\geq 2. We split the sum in two parts as as follows.

where we used the harmonic series inequality ∑s=2T1/s≤log⁡(T)\sum_{s=2}^{\mathscr{T}}1/s\leq\log(\mathscr{T}), ∑u=0Tγu≤1/(1−γ)\sum_{u=0}^{\mathscr{T}}\gamma^{u}\leq 1/(1-\gamma) and ⌊t/2⌋≤t/2\lfloor t/2\rfloor\leq t/2 in (333).