Stochastic Mirror Descent on Overparameterized Nonlinear Models: Convergence, Implicit Regularization, and Generalization
Navid Azizan, Sahin Lale, Babak Hassibi
Introduction
Deep learning has demonstrably enjoyed a great deal of success in a wide variety of tasks . Despite its tremendous success, the reasons behind the good performance of these methods on unseen data is not fully understood (and, arguably, remains somewhat of a mystery). While the special deep architecture of these models seems to be important to the success of deep learning, the architecture is only part of the story, and it has been now widely recognized that the optimization algorithms used to train these models, typically stochastic gradient descent (SGD) and its variants, also play a key role in learning parameters that generalize well.
Since these deep models are highly overparameterized, they have a lot of capacity, and can fit to virtually any (even random) set of data points . In other words, these highly overparameterized models can “interpolate” the data, so much so that this regime has been called the “interpolating regime” . In fact, on a given dataset, the loss function typically has (infinitely) many global minima, which however can have drastically different generalization properties (many of them perform very poorly on the test set). Which minimum among all the possible minima we choose in practice is determined by the initialization and the optimization algorithm that we use for training the model.
Since the loss functions of deep neural networks are non-convex and sometimes even non-smooth, in theory, one may expect the optimization algorithms to get stuck in local minima or saddle points. In practice, however, such simple stochastic descent algorithms almost always reach zero training error, i.e., a global minimum of the training loss . More remarkably, even in the absence of any explicit regularization, dropout, or early stopping , the global minima obtained by these algorithms seem to generalize quite well to unseen data (contrary to many other global minima). It has been also observed that even among different optimization algorithms, i.e., SGD and its variants, there is a discrepancy in the solutions achieved by different algorithms and their generalization capabilities . Therefore, it is important to ask the question
Which global minima do these algorithms converge to?
In this paper, we study the family of stochastic mirror descent (SMD) algorithms, which includes the popular SGD algorithm. For any choice of potential function, there is a corresponding mirror descent algorithm. We show that, for overparameterized nonlinear models, if one initializes close enough to the manifold of parameter vectors that interpolates the data, then the SMD algorithm for any particular potential converges to a global minimum that is approximately the closest one to the initialization, in Bregman divergence corresponding to the potential. Furthermore, in highly overparameterized models, this closeness of the initialization comes for free, something that is occasionally referred to as “the blessing of dimensionality.” For the special case of SGD, this means that it converges to a global minimum which is approximately the closest one to the initialization in the usual Euclidean sense.
We perform extensive systematic experiments on various initializations, various mirror algorithms for the MNIST and CIFAR-10 datasets using the existing off-the-shelf deep neural network architectures, and we measure all the pairwise distances in different Bregman divergences. We found that every single result is exactly consistent with the hypothesis. Indeed, in all our experiments, the global minimum achieved by any particular mirror descent algorithm is the closest, among all other global minima obtained by other mirrors and other initializations, to its initialization in the corresponding Bregman divergence. In particular, the global minimum obtained by SGD from any particular initialization is closest to the initialization in Euclidean sense, both among the global minima obtained by different mirrors and among the global minima obtained by different initializations.
How well do different mirrors perform in practice?
In Section 2, we review the family of mirror descent algorithms and briefly revisit the linear overparameterized case. Section 3 provides our main theoretical results, which are (1) convergence of SMD, under reasonable conditions, to a global minimum, and (2) proximity of the obtained global minimum to the closest point from initialization in Bregman divergence. Our proofs are remarkably simple and are based on a powerful fundamental identity that holds for all SMD algorithms in a general setting. We comment on the related work in Section 4. In Section 5, we provide our experimental results, which consists of two parts, (1) testing the theoretical claims about the distances for different mirrors and different initializations, and (2) assessing the generalization properties of different mirrors. The proofs of the theoretical results and more details on the experiments are relegated to the appendix.
Background and Warm-Up
is the set of global minima, and every parameter vector in renders the loss on each data point zero, i.e., . The loss function is often attempted to be minimized by stochastic gradient descent (SGD), which is defined as
assuming the data is indexed randomly (for , one can cycle through the data or select them at random).
2 Stochastic Mirror Descent
Stochastic mirror descent (SMD), first introduced by Nemirovski and Yudin , is one of the most widely used families of algorithms for stochastic optimization , which includes the popular stochastic gradient descent (SGD) as a special case. Consider a strictly convex differentiable function , called the potential function. Then SMD updates are defined as
Note that, due to the strict convexity of , the gradient defines an invertible map, so the recursion in (3) yields a unique at each iteration, and thus is a well-defined update. Compared to classical SGD, rather than update the weight vector along the direction of the negative gradient, the update is done in the “mirrored” domain determined by the invertible transformation . Mirror descent was originally conceived to exploit the geometrical structure of the problem by choosing an appropriate potential. Note that SMD reduces to SGD when , since the gradient is simply the identity map.
Alternatively, the update rule (3) can be expressed as
is the Bregman divergence with respect to the potential function . Note that is non-negative, convex in its first argument, and that, due to strict convexity, iff .
Different choices of the potential function yield different optimization algorithms, which will potentially have different implicit biases. A few examples follow.
Gradient Descent. For the potential function , the Bregman divergence is , and the update rule reduces to that of SGD.
Exponentiated Gradient Descent. For , the Bregman divergence becomes the unnormalized relative entropy (Kullback-Leibler divergence) , which corresponds to the exponentiated gradient descent (aka the exponential weights) algorithm .
-norms Algorithm. For any -norm squared potential function , with , the algorithm will reduce to the so-called -norms algorithm .
Sparse Mirror Descent. For , the algorithm reduces to sparse mirror descent, which is used in compressed sensing .
3 Overparameterized Linear Models
Overparameterized (or underdetermined) linear models have been recently studied in many papers due to their simplicity, and there are interesting insights than one can obtain from them. In this case, the model is , the set of global minima is , and the loss is . The following result characterizes the solution that SMD converges to, in the linear overparameterized setting .
Consider a linear overparameterized model. For sufficiently small step size, i.e., for any for which is convex, and for any initialization , the SMD iterates converge to
Note that the step size condition, i.e., the convexity of , depends on both the loss and the potential function. For the case of SGD, , and the condition reduces to . In that case, is simply .
Theoretical Results
In this section, we provide our main theoretical results. In particular, we show that for highly overparameterized nonlinear models, if initialized close enough to the set , (1) SMD converges to a global minimum, (2) the global minimum obtained by SMD is approximately the closest one to the initialization in Bregman divergence corresponding to the potential.
It has been argued in several recent papers that in highly overparameterized neural networks, any random initialization is close to , with high probability (see also the discussion in Section A.4 of the supplementary material). Therefore, it is reasonable to make the following assumption about the initialization.
This assumption states that, while certainly need not be convex, since is a minimizer of , the initial point is close to so that (see Fig. 2 for an illustration).
Our second assumption states that in this local region, the first and second derivatives of the model are bounded.
This is again a mild assumption, which is assumed in other related works such as as well. The following theorem states that under Assumption 1, SMD converges to a global minimum.
All the iterates remain in .
Note that, while convergence (to some point) with decaying step size is almost trivial, this result establishes converges to the solution set with fixed step size. Furthermore, the convergence is deterministic, and is not in expectation or with high probability. For example, this result also applies to the case where we cycle through the data deterministically.
We should also remark that the choice of distance in the definition of the “ball” was important to be the Bregman divergence with respect to and in that particular order. In fact, one cannot guarantee that SMD gets closer to (i.e. does not get farther from) at every step in the usual Euclidean sense. At some steps, it may get farther from in other senses, while getting closer in .
Denote the global minimum that is closest to the initialization in Bregman divergence by , i.e.,
Recall that in the linear case, this was what SMD converges to. We show that in the nonlinear case, under Assumptions 1 and 2, SMD converges to a point which is “very close” to .
Define . Under the assumptions of Theorem 3, and Assumption 2, the following holds.
In other words, if we start with an initialization that is away from (in Bregman divergence), we converge to a point that is away from the , the closest point on .
2 Proof Technique: Fundamental Identity of SMD
The main tool used for the proofs is a fundamental identity that holds for SMD in a very general setting.
This identity allows one to prove the results in a remarkably simple and direct way. Due to space limitations, the proofs are relegated to the supplementary material.
The ideas behind this identity are related to estimation theory , which was originally developed in the 1990’s in the context of robust control theory. In fact, it has connections to the minimax optimality of SGD, which was shown by for linear models, and recently extended to nonlinear models and general mirrors by .
Related Work
There have been many efforts in the past few years to study deep learning from an optimization perspective, e.g., . While it is not possible to review all the contributions here, we comment on the ones that are most closely related to our results. We highlight the distinctions between our results and those.
Many recent papers have studied the so-called “overparameterized” setting, or the “interpolating” regime, which is common in deep learning . All these results, similar to our work, have assumptions for being close to the solution space (global minima), which is perhaps reasonable in highly overparameterized models, as we also argued in Section A.4 of the supplementary material. However, most of these results are limited to (S)GD and do not generalize to more general mirrors.
Furthermore, even for the case of SGD, our results are stronger than those in the literature, in the sense that not only do we show convergence to a global minimum, but we also show that and are close. In fact, showed that for SGD, is close to (i.e. bounded by a constant factor of) . Our Theorem 4 states that not only are these two distances close, but and are also actually close (), something that could not be inferred from the previous work.
As mentioned before, there have been a number of results on characterizing the implicit regularization properties of different algorithms in different contexts . The closest ones to our results, which concern mirror descent, are the works of . The authors in consider linear overparameterized models, and show that if SMD happens to converge to a global minimum, then that global minimum should be the one that is closest in Bregman divergence to the initialization, which can be shown by writing the KKT conditions. In fact, they do not provide any conditions for convergence and whether it converges with a fixed step size or not. In the authors’ earlier work , the condition on the step size for which SMD converges to the aforementioned global minimum was derived, for linear models. Our results in this paper are for nonlinear models, and we show that, under the specified conditions on the step size, these algorithms with a fixed step size converge to the mentioned global minimum, which had not been shown in any of the previous work. Furthermore, assuming every data point is revisited again after some steps, the convergence we establish is deterministic, and not in expectation or with high probability.
Experimental Results
In this section, we provide our experimental results, which consist of two main parts. In the first part, we evaluate the theoretical claims by running systematic experiments for different initializations and different mirrors, and evaluating the distances between the global minima achieved and the initializations, in different Bregman divergences. In the second part, we assess the generalization error of different mirrors, which correspond to different regularizers, in order to understand which regularizer performs better.
We measure the distances between the initializations and the global minima obtained from different mirrors and different initializations, in different Bregman divergences. Table 1, and Table 2 show some examples among different mirrors and different initializations, respectively. Fig. 5 shows the distances between a particular initial point and all the final points obtained from different initializations and different mirrors (the distances are often orders of magnitude different, so we show them in logarithmic scale). The global minimum achieved by any mirror from any initialization is the closest in the correct Bregman divergence, among all mirrors, among all initializations, and among both. This trend is very consistent among all our experiments, which can be found in Appendix B.
2 Distribution of the Weights of the Network
3 Generalization Errors of Different Mirrors
References
Appendix A Proofs of the Theoretical Results
In this section, we prove the main theoretical results. The proofs are based on a fundamental identity about the iterates of SMD, which holds for all mirrors and all overparametereized (even nonlinear) models (Lemma 6). We first prove this identity, and then use it to prove the convergence and implicit regularization results.
Let us start by expanding the Bregman divergence based on its definition
By plugging the SMD update rule into this, we can write it as
Using the definition of Bregman divergence for and , i.e., and , we can express this as
Expanding the last term using , and following the definition of from (7) for and , we have
Note that for all , we have . Therefore, for all
Combining the second and the last terms in the right-hand side leads to
for all , which concludes the proof. ∎
A.2 Convergence of SMD to the Interpolating Set
Now that we have proved Lemma 6, we can use it to prove our main results, in a remarkably simple fashion. Let us first prove the convergence of SMD to the set of solutions.
All the iterates remain in .
First we show that all the iterates wil remain in . Recall the identity of SMD from Lemma 6:
which holds for all . If is in the region , we know that the last term is non-negative. Furthermore, if the step size is small enough that is strictly convex, the second term is a Bregman divergence and is non-negative. Since the loss is non-negative, is always non-negative. As a result, we have
This implies that , which means is in too. Since is in , will be in , and therefore, will be in , and similarly all the iterates will remain in .
Next, we prove that the iterates converge and . If we sum up identity (9) for all , the first terms on the right- and left-hand side cancel each other telescopically, and we have
Since , we have If we take , the sum still has to remain bounded, i.e.,
Since the step size is small enough that is strictly convex for all , the first term is non-negative. The second term is non-negative because of the non-negativity of the loss. Finally, the last term is non-negative because for all . Hence, all the three terms in the summand are non-negative, and because the sum is bounded, they should go to zero as . In particular,
implies , i.e., convergence (), and further
This implies that all the individual losses are going to zero, and since every data point is being revisited after some steps, all the data points are being fit. Therefore, . ∎
A.3 Closeness of the Final Point to the Regularized Solution
In this section, we show that with the additional Assumption 2 (which is equivalent to having bounded Hessian in ), not only do the iterates remain in and converge to the set , but also they converge to a point which is very close to (the closest solution to the initial point, in Bregman divergence). The proof is again based on our fundamental identity for SMD.
Define . Under the assumptions of Theorem 3, and Assumption 2, the following holds.
which holds for all . Summing the identity for all , we have
for all . Note that the only terms in the right-hand side which depend on are the first one and the last one . In what follows, We will argue that, within , the dependence on in the last term is weak and therefore is close to .
To further spell out the dependence on in the last term, let us expand
By Taylor expansion of around and using Taylor’s theorem (Lagrange’s mean-value form), we have
for some in the convex hull of and . Since for all , it follows that
for all . Plugging this into (26), we have
for all . Finally, by plugging this back into the identity (24), we have
for all . Note that this can be expressed as
for all , where does not depend on :
From Theorem 3, we know that . Therefore, by plugging it into equation (31), and using the fact that , we have
Further, again since all the iterates are in , it follows that and . As a result the difference of the two terms, i.e., \big{[}(w^{*}-w_{i-1})^{T}H_{f_{i}}(w^{\prime\prime}_{i})(w^{*}-w_{i-1})-(w_{\infty}-w_{i-1})^{T}H_{f_{i}}(w^{\prime}_{i})(w_{\infty}-w_{i-1})\big{]}, is also , and we have
The term in parentheses is non-negative by definition of . The second term is non-negative by convexity of . Since both terms are non-negative and their sum is , each one of them is at most , i.e.
The proof is a straightforward application of Theorem 4. Note that we have
In particular, by plugging in and , we have and . Subtracting the two equations from each other yields
which along with the application of Theorem 4 concludes the proof. ∎
A.4 Closeness to the Interpolating Set in Highly Overparameterized Models
As we mentioned earlier, it has been argued in a number of recent papers that for highly overparameterized models, any random initial point is, whp, close to the solution set . In the highly overparameterized regime, , and so the dimension of the manifold , which is , is very large. For simplicity, we outline an argument for the case of Euclidean distance, bearing in mind that a similar argument can be used for general Bregman divergence. Note that the distance of an arbitrarily chosen to is given by
where and . This can be approximated by
where is the Jacobian matrix. The latter optimization can be solved to yield
Note that is an matrix consisting of the sum of outer products. When the are sufficiently random, and , it is not unreasonable to assume that whp
since is -dimensional. The above implies that is close to and hence .
Appendix B More Details on the Experimental Results
In order to evaluate the claim, we run systematic experiments on some standard deep learning problems.
Datasets. We use the standard MNIST and CIFAR-10 datasets.
Architectures. For MNIST, we use a 4-layer convolutional neural network (CNN) with 2 convolution layers and 2 fully connected layers. The convolutional layers and the fully connected layers are picked wide enough to obtain trainable parameters. Since MNIST dataset has 60,000 training samples, the number of parameters is significantly larger than the number of training data points, and the problem is highly overparameterized. For the CIFAR-10 dataset, we use the standard ResNet-18 architecture without any modifications. CIFAR-10 has 50,000 training samples and with the total number of parameters in ResNet-18, the problem is again highly overparameterized.
Loss Function. We use the cross-entropy loss as the loss function in our training. We train the models from different initializations, and with different mirror descents from each particular initialization, until we reach training accuracy, i.e., until we hit .
Initialization. We randomly initialize the parameters of the networks around zero (). We choose 6 independent initializations for the CNN, and 8 for ResNet-18, and for each initialization, we run the following 4 different SMD algorithms.
where denotes the -th element of the vector.
We use a fixed step size . The step size is chosen to obtain convergence to global minima.
We provide the distances from final points (global minima) obtained by different algorithms from the same initialization, measured in different Bregman divergences for MNIST classification task using a standard CNN. Note that in all tables the smallest element in each row is on the diagonal, which means the point achieved by each mirror has the smallest Bregman divergence to the initialization corresponding to that mirror, among all mirrors. Tables 3, 4, 5, 6, 7, 8 depict these results for 6 different initializations. The rows are the distance metrics used as the Bregman Divergences with specified potentials. The columns are the global minima obtained using specified SMD algorithms.
B.1.2 Closest Minimum for Different Initilizations with Fixed Mirror
We provide the pairwise distances between different initial points and the final points (global minima) obtained by using fixed SMD algorithms in MNIST dataset using a standard CNN. Note that the smallest element in each row is on the diagonal, which means the closest final point to each initialization, among all the final points, is the one corresponding to that point. Tables 9, 10, 11 and 12 depict these results for 4 different SMD algorithms. The rows are the initial points and the columns are the final points corresponding to each initialization.
B.1.3 Closest Minimum for Different Initilizations and Different Mirrors
Now we assess the pairwise distances between different initial points and final points (global minima) obtained by all different initilizations and all different mirrors (Table 8). The smallest element in each row is exactly the final point obtained by that mirror from that initialization, among all the mirrors and all the initial points.
B.2 CIFAR-10 Experiments
We provide the distances from final points (global minima) obtained by different algorithms from the same initialization, measured in different Bregman divergences for CIFAR-10 classification task using ResNet-18. Note that in all tables the smallest element in each row is on the diagonal, which means the point achieved by each mirror has the smallest Bregman divergence to the initialization corresponding to that mirror, among all mirrors. Tables 13, 14, 15, 16, 17, 18, 19, 20 depict these results for 8 different initializations. The rows are the distance metrics used as the Bregman Divergences with specified potentials. The columns are the global minima obtained using specified SMD algorithms.
B.2.2 Closest Minimum for Different Initilizations with Fixed Mirror
We provide the pairwise distances between different initial points and the final points (global minima) obtained by using fixed SMD algorithms in CIFAR-10 dataset using ResNet-18. Note that the smallest element in each row is on the diagonal, which means the closest final point to each initialization, among all the final points, is the one corresponding to that point. Tables 21, 22, 23, 24 depict these results for 4 different SMD algorithms. The rows are the initial points and the columns are the final points corresponding to each initialization.
B.2.3 Closest Minimum for Different Initilizations and Different Mirrors
Now we assess the pairwise distances between different initial points and final points (global minima) obtained by all different initilizations and all different mirrors (Table 8). The smallest element in each row is exactly the final point obtained by that mirror from that initialization, among all the mirrors and all the initial points.
B.3 Distribution of the Final Weights of the Network
B.4 Generalization Errors of Different Mirrors/Regularizers
In this section, we compare the performance of the SMD algorithms discussed before on the test set. This is important for understanding the effect of different regularizers on the generalization of deep networks.