Stability and Generalization of Learning Algorithms that Converge to Global Optima

Zachary Charles, Dimitris Papailiopoulos

Introduction

The recent success of training complex models at state-of-the-art accuracy in many common machine learning tasks has sparked significant interest and research in algorithmic machine learning. In practice, not only can these complex deep neural models yield zero training loss, they can also generalize surprisingly well . Although there has been significant recent work in analyzing the training loss performance of several learning algorithms, our theoretical understanding of their generalization properties falls far below what has been observed empirically.

A useful proxy for analyzing the generalization performance of learning algorithms is that of stability. A training algorithm is stable if small changes in the training set result in small differences in the output predictions of the trained model. In their foundational work, Bousquet and Elisseeff establish that stability begets generalization.

While there has been stability analysis for empirical risk minimizers , there are far fewer results for commonly used iterative learning algorithms. In a recent novel work, Hardt et al. establish stability bounds for SGD, and discuss algorithmic heuristics that provably increase the stability of SGD models. Unfortunately, generalizing their techniques to establish stability bounds for other first-order methods can be a strenuous task. Showing non-trivial stability for more involved algorithms like SVRG (even in the convex case), or SGD in more nuanced non-convex setups is far from straightforward. While provides a clean and elegant analysis that shows stability of SGD for non-convex loss functions, the result requires very small step-sizes. The step-size is small enough that one may require exponentially many steps for provable convergence to an approximate critical point, under standard smoothness assumptions (see subsection A.10). Generally, there seems to be a trade-off between convergence and stability of algorithms. In this work we show that under certain geometric assumptions on the loss function around global minima, we can actually leverage the convergence properties of an algorithm to prove that it is stable.

The goal of this work is to provide black-box and easy-to-use stability results for a variety of learning algorithms in non-convex settings. We show that this is in some cases possible by decoupling the stability of critical points and their proximity to models trained by iterative algorithms.

We establish that models trained by algorithms that converge to local minima are stable under the Polyak-Łojasiewicz (PL) and the quadratic growth (QG) conditions . Informally, these conditions assert that the suboptimality of a model is upper bounded by the norm of its gradient and lower bounded by its distance to the closest global minimizer. As we see in the following, these conditions are sufficient for stability and are general enough to yield useful bounds for a variety of settings.

Our results require weaker conditions compared to the state-of-the art, while recovering several prior stability bounds. For example in the authors require convexity, or strong convexity. Gonen and Shalev-Shwartz prove the stability of ERMs for nonconvex, but locally strongly convex loss functions obeying strict saddle inequalities . By contrast, we develop comparable stability results for a large class of functions, where no convexity, local convexity, or saddle point conditions are imposed. We note that although establishes the stability of SGD for smooth non-convex objectives, the stepsize selection can be prohibitively small for convergence. In our bounds, we make no assumptions on the hyper-parameters of the algorithms.

We use our black-box results to directly compare the generalization performance of popular first-order methods in general learning setups. While direct proofs of stability seem to require a substantial amount of algorithm-specific analysis, our results are derived from known convergence rates of popular algorithms. For strong convexity—a special case of the PL condition—we recover order-wise the stability bounds of Hardt et al. , but for a large family of optimization algorithms (e.g., SGD, GD, SVRG, etc). We show that many of these algorithms offer order-wise similar stability as saddle-point avoiding algorithms in nonconvex problems where all local minima are global . We finally show that while SGD and GD have analogous stability in the convex setting, this breaks down in the non-convex setting. We give an explicit example of a simple 1-layer neural network on which SGD is stable but GD is not. Such an example was theorized in (i.e., Figure 10 in the aforementioned paper); here we formalize the authors’ intuition. Our results offer yet another indication that SGD trained models can be more generalizable than full-batch GD ones.

Finally, we give examples of some machine learning scenarios where the PL condition mentioned above holds true. Adapting techniques from , we show that deep networks with linear activation functions are PL almost everywhere in the parameter space. Our theory allows us to derive results similar to those in about local/global minimizers in linear neural networks.

Prior work

The idea of stability analysis has been around for more than 30 years since the work of Devroye and Wagner . Bousquet and Elisseef defined several notions of algorithmic stability and used them to derive bounds on generalization error. Further work has focused on stability of randomized algorithms and the interplay between uniform convergence and generalization . Mukherjee et al. show that stability implies consistency of empirical risk minimization. Shalev-Shwartz et al. show that stability can also imply learnability is some problems.

Our work is heavily influenced by that of Hardt et al. that establish stability bounds for stochastic gradient descent (SGD) in the convex, strongly convex, and non-convex case. The work by Lin et al. shows that stability of SGD can be controlled by forms of regularization. In , the authors give stability bounds for SGD that are data dependent. Since they do not rely on worst-case arguments, they lead to smaller generalization error bounds than that in , but require assumptions on the underlying data. The work by Liu et al. gives a related notion of uniform hypothesis stability and show that it implies guarantees on the generalization error.

Stability is closely related to the notion of differential privacy introduced in . Roughly speaking, differential privacy ensures that the probability of observing any outcome from a statistical query changes if you modify any single dataset element. Dwork et al. later showed that differentially private algorithms generalize well . These generalization bounds were later improved by Nissim and Stemmer . Such generalization bounds are similar to those guaranteed by stability but often require different tools to handle directly.

Preliminaries

The generalization performance of the model can then be measured by the generalization gap:

For our purposes, ww will be the output of some (potentially randomized) learning algorithm A\mathcal{A}, trained on some data set SS. We will denote this output by A(S)\mathcal{A}(S).

Let us now define a related training set S′={z1,…,zi−1,zi′,zi+1,…,zn}S^{\prime}=\{z_{1},\ldots,z_{i-1},z_{i}^{\prime},z_{i+1},\ldots,z_{n}\}, where zi′∼Dz_{i}^{\prime}\sim\mathcal{D}. We then have the following notion of uniform stability that was first introduced in .

An algorithm A\mathcal{A} is uniformly ϵ\epsilon-stable, if for all data sets S,S′S,S^{\prime} differing in at most one example, we have

The expectation is taken with respect to the (potential) randomness of the algorithm A\mathcal{A}. Bousquet and Elisseeff establish that uniform stability implies small generalization gap .

Suppose A\mathcal{A} is uniformly ϵ\epsilon-stable. Then,

In practice, uniform stability may be too restrictive, since the bound above must hold for all zz, irrespective of its marginal distribution. The following notion of stability, while weaker, is still enough to control the generalization gap. Given a data set S={z1,…,zn}S=\{z_{1},\ldots,z_{n}\} and i∈{1,…,n}i\in\{1,\ldots,n\}, we define SiS^{i} as S\ziS\backslash z_{i}.

Note that this is a weaker notion than uniform stability, but one can still use it to establish non-trivial generalization bounds:

In the following, we derive stability bounds like the above, for models trained on empirical risk functions satisfying the PL and QG conditions. To do so, we will assume that the functions in question are LL-Lipschitz.

If ff is assumed to be differentiable, this is equivalent to saying that for all xx, ∥∇f(x)∥2≤L\|\nabla f(x)\|_{2}\leq L.

In a recent work, Karimi et al. in used the Polyak-Łojasiewicz condition to prove simplified nearly optimal convergence rates for several first-order methods. Notably, there are some non-convex functions that satisfy the PL condition. The condition is defined below.

Fix a set X\mathcal{X} and let f∗f^{*} denote the minimum value of ff on X\mathcal{X}. We will say that a function ff satisfies the Polyak-Łojasiewicz (PL) condition on X\mathcal{X}, if there exists μ>0\mu>0 such that for all x∈Xx\in\mathcal{X} we have

Note that for PL functions, every critical point is a global minimizer. While strong convexity implies PL, the reverse is not true. Moreover, PL functions are in general nonconvex (e.g., invex functions). We also consider a strictly larger family of functions that satisfy the quadratic growth condition.

We will say that a function ff satisfies the quadratic growth (QG) condition on a set X\mathcal{X}, if there exists μ>0\mu>0 such that for all x∈Xx\in\mathcal{X} we have

where xpx_{p} denotes the euclidean projection of xx onto the set of global minimizers of ff in X\mathcal{X} (i.e., xpx_{p} is the closest point to xx in X\mathcal{X} satisfying f(xp)=f∗f(x_{p})=f^{*}).

Both of these conditions have been considered in previous studies. The PL condition was first introduced by Polyak in , who showed that under this assumption, gradient descent converges linearly. The QG condition has been considered under various guises and can imply important properties about the geometry of critical points. For example, showed that local minima of nonlinear programs satisfying the QG condition are actually isolated stationary points. These kinds of geometric implications will allow us to derive stability results for large classes of algorithms.

Stability of Approximate Global Minima

In this section, we establish the stability of large classes of learning algorithms under the PL and QG conditions presented above. Our stability results are “black-box” in the sense that our bounds are decomposed as a sum of two terms: a term concerning the convergence of the algorithm to a global minimizer, and a term relevant to the geometry of the loss function around the global minima. Both terms are used to establish good generalization and provide some insights into the way that learning algorithms perform.

For a given data set SS, suppose we use an algorithm A\mathcal{A} to train some model ww. We let wSw_{S} denote the output of our algorithm on SS. The empirical training error on a data set SS is denoted fS(w)f_{S}(w) and is given by

We assume that each of these losses is LL-Lipschitz with respect to the parameters of the model. We are interested in conditions on fSf_{S} that allow us to make guarantees on the stability of A\mathcal{A}. As it turns out, the PL and QG condition will allow us to prove such results. Although it seems unclear if these conditions are reasonable, in our last section we show that they arise in a large number of machine learning settings, including in certain deep neural networks.

To analyze the performance machine learning algorithms, it often suffices to understand the algorithm’s behavior with respect to critical points. This requires knowledge of the convergence of the algorithm, and an understanding of the geometric properties of the loss function around critical points. As it turns out, the PL and QG conditions allow us to understand the geometry underlying the minima of our function. Let Xmin⁡\mathcal{X}_{\min} denote the set of global minima of fSf_{S}.

Assume that for all SS and w∈Xw\in\mathcal{X}, fSf_{S} is PL with parameter μ\mu. We assume that applying A\mathcal{A} to fSf_{S} produces output wSw_{S} that is converging to some global minimizer wS∗w_{S}^{*}. Then A\mathcal{A} has pointwise hypothesis stability with parameter ϵstab\epsilon_{stab} satisfying the following conditions.

Case 1: If for all SS, ∥wS−wS∗∥≤O(ϵA)\|w_{S}-w_{S}^{*}\|\leq O(\epsilon_{\mathcal{A}}) then

Case 2: If for all SS, ∣fS(wS)−fS(wS∗)∣≤O(ϵA′)|f_{S}(w_{S})-f_{S}(w_{S}^{*})|\leq O(\epsilon^{\prime}_{\mathcal{A}}) then

Case 3: If for all SS, ∥∇fS(wS)∥≤O(ϵA′′)\|\nabla f_{S}(w_{S})\|\leq O(\epsilon^{\prime\prime}_{\mathcal{A}}), then

Suppose our loss functions are PL and our algorithm A\mathcal{A} is an oracle that returns a global optimizer wS∗w_{S}^{*}. Then the terms ϵA,ϵA′,ϵA′′\epsilon_{\mathcal{A}},\epsilon^{\prime}_{\mathcal{A}},\epsilon^{\prime\prime}_{\mathcal{A}} above are all identical to , leading to the following corollary.

Let fSf_{S} satisfy the PL inequality with parameter μ\mu and let A(S)=argmin⁡w∈XfS(w)\mathcal{A}(S)=\text{arg}\min_{w\in\mathcal{X}}f_{S}(w). Then, A\mathcal{A} has pointwise hypothesis stability with

Bousquet and Ellisseef considered the stability of empirical risk minimizers where the loss function satisfied strong convexity. Their work implies that for λ\lambda-strongly convex functions, the empirical risk minimizer has we stability satisfying ϵstab≤L2λn\epsilon_{stab}\leq\frac{L^{2}}{\lambda n}. Since λ\lambda-strongly convex implies λ\lambda-PL, Corollary 3.4 generalizes their result, with only a constant factor loss.

A similar result to Theorem 3.1 can be derived for empirical risk functions satisfy the QG condition and are realizable, e.g., where zero training loss is achievable.

Case 1: If for all SS, ∥wS−wS∗∥≤O(ϵA)\|w_{S}-w_{S}^{*}\|\leq O(\epsilon_{\mathcal{A}}) then

Case 2: If for all SS, ∣fS(wS)−fS(wS∗)∣≤O(ϵA′)|f_{S}(w_{S})-f_{S}(w_{S}^{*})|\leq O(\epsilon^{\prime}_{\mathcal{A}}) then

Observe that unlike the case of PL empirical losses, QG empirical losses only allow for a O(1n)O(\frac{1}{\sqrt{n}}) convergence rate of stability. Moreover, similarly to our result for PL loss functions, the result of Theorem 3.3 holds even if we only have information about the convergence of A\mathcal{A} in expectation.

Finally, we can obtain the following Corollary for empirical risk minimizers.

Let fSf_{S} satisfy the QG inequality with parameter μ\mu and let A(S)=argmin⁡w∈XfS(w)\mathcal{A}(S)=\text{arg}\min_{w\in\mathcal{X}}f_{S}(w). Then, A\mathcal{A} has pointwise hypothesis stability with

2 Uniform Stability for PL/QG Loss Functions

Under a more restrictive setup, we can obtain similar bounds for uniform hypothesis stability, which is a stronger stability notion compare to its pointwise hypothesis variant. The usefulness of uniform stability compared to pointwise stability, is that it can lead to generalization bounds that concentrate exponentially faster with respect to the sample size nn.

As before, given a data set SS, we let denote wSw_{S} be the model that A\mathcal{A} outputs. Let πS(w)\pi_{S}(w) denote the closest optimal point of fSf_{S} to ww. We will denote πS(wS)\pi_{S}(w_{S}) by wS∗w_{S}^{*}. Let S,S′S,S^{\prime} be data sets differing in at most one entry. We will make the following technical assumption:

The empirical risk minimizers for fSf_{S} and fS′f_{S^{\prime}}, i.e., wS∗,wS′∗w_{S}^{*},w_{S^{\prime}}^{*} satisfy πS(wS′∗)=wS∗\pi_{S}(w_{S^{\prime}}^{*})=w_{S}^{*}, where πS(w)\pi_{S}(w) is the projection of ww on the set of empirical risk minimizers of fSf_{S}. Note that this is satisfied if for every data set SS, there is a unique minimizer wS∗w^{*}_{S}.

We would like to note that the above assumption is extremely strict, and in general does not apply to empirical losses with infinitely many global minima. To tackle the existence of infinitely many global minima, one could imagine designing A(S)\mathcal{A}(S) to output a structured empirical risk minimizer, e.g., one such that if A\mathcal{A} is applied on S′S^{\prime}, its projection on the optima of fSf_{S} would always yield back A(S)A(S). This could be possible, if A(S)A(S) corresponded to minimizing instead a regularized, or structure constrained cost function whose set of optimizers only contained a small subset of the global minima of fSf_{S}. Unfortunately, coming up with such a structured empirical risk minimizer for general nonconvex losses seems far from straightforward, and serves as an interesting open problem.

Assume that for all SS, fSf_{S} satisfies the PL condition with constant μ\mu, and suppose that Assumption 1 holds. Then A\mathcal{A} has uniform stability with parameter ϵstab\epsilon_{stab} satisfying the following conditions.

Case 1: If for all SS, ∥wS−wS∗∥≤O(ϵA)\|w_{S}-w_{S}^{*}\|\leq O(\epsilon_{\mathcal{A}}) then

Case 2: If for all SS, ∣fS(wS)−fS(wS∗)∣≤O(ϵA′)|f_{S}(w_{S})-f_{S}(w_{S}^{*})|\leq O(\epsilon^{\prime}_{\mathcal{A}}) then

Case 3: If for all SS, ∥∇fS(wS)∥≤O(ϵA′′s)\|\nabla f_{S}(w_{S})\|\leq O(\epsilon^{\prime\prime}_{\mathcal{A}}s), then

Since strong convexity is a special case of PL, this theorem implies that if we run enough iterations of a convergent algorithm A\mathcal{A} on a λ\lambda-strongly convex loss function, then we would expect uniform stability on the order of

In particular, this theorem recovers the stability estimates for ERMs and SGD applied to strongly convex functions proved in and , respectively.

In order to make this result more generally applicable, we would like to extend the theorem to a larger class of functions than just globally PL functions. If we assume boundedness of the loss function, then we can derive a similar result for globally QG functions. This leads us to the following theorem:

Case 1: If for all SS, ∥wS−wS∗∥≤O(ϵA)\|w_{S}-w_{S}^{*}\|\leq O(\epsilon_{\mathcal{A}}) then for all zz we have:

Case 2: If for all SS, ∣fS(wS)−fS∗∣≤O(ϵA′)|f_{S}(w_{S})-f_{S}^{*}|\leq O(\epsilon^{\prime}_{\mathcal{A}}) then for all zz we have:

By analogous reasoning to that in Remark 3.1, both Theorem 3.5 and Theorem 3.6 hold if you only have information about the output of A\mathcal{A} in expectation.

PL loss functions in practice

As the bounds above show, the PL and QG conditions are sufficient for algorithmic stability and therefore imply good generalization. In this section, we show that the PL condition actually arises in some interesting machine learning setups, including least squares minimization, strongly convex functions composed with piecewise linear functions, and neural networks with linear activation functions. A first step towards a characterization of PL loss functions was proved by Karimi et al. , which established that the composition of a strongly-convex function and a linear function results in a loss that satisfies the PL condition.

Let gg be strongly-convex with parameter λ\lambda, σ\sigma a leaky ReLU activation function with slopes c1c1 and c2c_{2}, and XX a matrix with minimum singular value σmin⁡(X)\sigma_{\min}(X). Let c=min⁡{∣c1∣,∣c2∣}c=\min\{|c_{1}|,|c_{2}|\}. Then f(w)=g(σ(Xw))f(w)=g(\sigma(Xw)) is PL almost everywhere with parameter μ=λσmin⁡(X)2c2\mu=\lambda\sigma_{\min}(X)^{2}c^{2}.

In particular, 1-layer neural networks with a squared error loss and leaky ReLU activations satisfy the PL condition. More generally, this holds for any piecewise-linear activation function with slopes {ci}i=1k\{c_{i}\}_{i=1}^{k}. As long as each slope is non-zero and XX is full rank, the result above shows that the PL condition is satisfied.

2 Linear Neural Networks

The results above only concern one layer neural networks. Given the prevalence of deep networks, we would like to say something about the associated loss function. As it turns out, we can prove that a PL inequality holds in large regions of the parameter space for deep linear networks.

Suppose that the WiW_{i} satisfy σmin⁡(Wi)≥τ>0\sigma_{\min}(W_{i})\geq\tau>0 for all ii. Then,

Combining Lemmas 4.2 and 4.3, we derive the following interesting corollary about when critical points are global minimizers. This result is not directly related to the work above, but gives an easy way to understand the landscape of critical points of deep linear networks.

Thematically similar results have been derived previously for 1 layer networks in and for deep neural networks in . In , Kawaguchi derives a similar result to ours for deep linear neural networks. Kawaguchi shows that every critical point is either a global minima or a saddle point. Our result, by contrast, implies that all full-rank critical points are global minima.

Lemmas 4.2 and 4.3 can also be combined to show that linear networks satisfy the PL condition in large regions of parameter space, as the follwing theorem says.

Stability of Some First-order Methods

We wish to apply our bounds from the previous section to popular convergent gradient-based methods. We consider SGD, GD, RCD, and SVRG. When we have LL-Lipschitz, μ\mu-PL loss functions fSf_{S} and nn training examples, Theorem 3.5 states that any learning algorithm A\mathcal{A} has uniform stability ϵstab\epsilon_{stab} satisfying

Here, ϵA\epsilon_{\mathcal{A}} refers to how quickly A\mathcal{A} converges to the optimal value of the loss function. Specifically, this holds if the algorithm produces a model wSw_{S} satisfying ∣fS(ws)−fS∗∣≤O(ϵA)|f_{S}(w_{s})-f_{S}^{*}|\leq O(\epsilon_{\mathcal{A}}). For example, if we want to guarantee that our algorithm has the same stability as SGD in the strongly convex case, then we need to determine how many iterations TT we need to perform such that

The convergence rates of SGD, GD, RCD, and SVRG have been studied extensively in the literature . The results are given below. When necessary to state the result, we assume a constant step-size of γ\gamma. Figure 1 below summarizes the values of ϵA\epsilon_{\mathcal{A}} for TT iterations of SGD, GD, RCD, and SVRG applied to λ\lambda-strongly convex loss functions and μ\mu-PL loss functions. Note that if Eq. (1) holds, then Corollary 3.4 implies that our algorithm is uniformly stable with parameter ϵstab=O(L2/μn)\epsilon_{stab}=O(L^{2}/\mu n). Moreover, this is the same stability as that of SGD for strongly-convex functions , and saddle point avoiding algorithms on strict-saddle loss functions .

We use the above convergence rates of these algorithms in the λ\lambda-strongly convex and μ\mu-PL settings to determine how many iterations are required such that we get stability that is O(L2/μn)O(L^{2}/\mu n). The results are summarized in Figure 2 below.

Note that in the μ\mu-PL case, although it is a nonconvex setup the above algorithms all exhibit the same stability for these values of TT. This is not the case in general: several studies have observed that small-batch SGD offers superior generalization performance compared to large-batch SGD, or full-batch GD, when training deep neural networks .

Unfortunately our bounds above, are not nuanced enough to capture the difference in generalization performance between mini-batch and large-batch SGD. Below, we will make this observation formal. Although SGD and GD can be equally stable for nonconvex problems satisfying the PL condition, there exist nonconvex problems where full-batch GD is not stable and SGD is stable.

The Instability of Gradient Descent

In , Hardt et al. proved bounds on the uniform stability of SGD. They also noted that GD does not appear to be provably as stable for the nonconvex case and sketched a situation in which this difference would appear. Due to the similarity of SGD and GD, one may expect similar uniform stability. While this is true in every convex setting as we show in the Appendix, subsection A.6 (without requiring strong convexity), this breaks down in the non-convex setting. Below we construct an explicit example where GD is not uniformly stable, but SGD is. This example formalizes the intuition given in .

Intuitively, this is a generalized quadratic model where the predicted label y^\hat{y} for a given xx is given by

The above predictive model can be described by the following 1-layer neural network using a quadratic and a linear activation function, denoted z2z^{2} and zz.

where zi=(−1,1)z_{i}=(-1,1) for 1≤i≤n−121\leq i\leq\frac{n-1}{2}, zi=(−1/2,1)z_{i}=(-1/2,1) for n−12<i≤n−1\frac{n-1}{2}<i\leq n-1. By construction, we have

Then fS(w)f_{S}(w), fS′(w)f_{S^{\prime}}(w) will approximately have the shape of the right-most function in Figure 4 above. However, recall that ddwg(w)=0\frac{d}{dw}g(w)=0 at w=w^w=\hat{w}. Therefore, there is some δ\delta, with 0<δ<ϵ0<\delta<\epsilon, such that for all w∈(w^−δ,w^+δ)w\in(\hat{w}-\delta,\hat{w}+\delta),

After enough iterations of gradient descent on fS,fS′f_{S},f_{S^{\prime}}, we will obtain models wSw_{S} and wS′w_{S^{\prime}} that are close to the distinct local minima in the right-most graph in Figure 4. This will hold as long as γ\gamma is not extremely large, in which case the steps of gradient descent could simply jump from one local minima to another. To ensure this does not happen, we restrict to γ≤1\gamma\leq 1.

For all nn, for all step-sizes 0<γ≤10<\gamma\leq 1, there is a KK such that for all k≥Kk\geq K, there are data sets S,S′S,S^{\prime} of size nn differing in one entry and a non-zero measure set of initial starting points such that if we perform kk iterations of gradient descent with step-size γ\gamma on SS and S′S^{\prime} to get outputs A(S),A(S′)\mathcal{A}(S),\mathcal{A}(S^{\prime}) then there is a z∗z^{*} such that

Theorem 6.1 establishes that there exist simple non-convex settings, for which the uniform stability of gradient descent does not decrease with nn. In light of the work in , where the authors show that for very conservative step-sizes, SGD is stable on non-convex loss function, we might wonder whether SGD is stable in this setting with moderate step-sizes. We know that gradient descent is not stable, by Theorem 6.1, for γ=1\gamma=1. For simplicity of analysis, we focus on the case where γ=1\gamma=1.

Suppose we run SGD on the above fS,fS′f_{S},f_{S^{\prime}} with step-size 11 and initialize near w^\hat{w} (we will be more concrete later about where we initialize). With probability n−1n\frac{n-1}{n}, the first iteration of SGD will use the same example for both SS and S′S^{\prime}, either z=(−1,1)z=(-1,1) or z=(−1/2,1)z=(-1/2,1). Computing derivatives at w^\hat{w} shows

In both cases, the slope is at least 0.40.4. Therefore, there is some η\eta such that for all w∈[w^−η,w^+η]w\in[\hat{w}-\eta,\hat{w}+\eta],

Therefore, continuing to run SGD in this setting, even if we now decrease the step-size, will eventually lead us to the same basins of fS(w),fS′(w)f_{S}(w),f_{S^{\prime}}(w). Let wS,wS′w_{S},w_{S^{\prime}} denote the outputs of SGD in this setting after enough steps so that we get convergence to within 1n\frac{1}{n} of a local minima. If our first sample zz was (−1,1)(-1,1), we will end up in the right basin, while if our first sample zz was (−1/2,1)(-1/2,1), we will end up in the left basin. In particular, for ϵ\epsilon small, z±z_{\pm} are close enough that the minima of fS,fS′f_{S},f_{S^{\prime}} are within 1n\frac{1}{n} of each other. Note that the minima w1,w2w_{1},w_{2} that wS,wS′w_{S},w_{S^{\prime}} are converging to are different. However, because they are in the same basin we know that for ϵ\epsilon small, z±z_{\pm} are close enough that ∥wS−wS′∥≤O(1n)\|w_{S}-w_{S^{\prime}}\|\leq O(\frac{1}{n}). Therefore, we have

Suppose that we initialize SGD in [w^−η,w^+η][\hat{w}-\eta,\hat{w}+\eta] with a step-size of γ=1\gamma=1. Let A(S),A(S′)\mathcal{A}(S),\mathcal{A}(S^{\prime}) denote the output of SGD after kk iteration for sufficiently large kk. For ∥z∥≤2\|z\|\leq 2,

This is in stark contrast to gradient descent, which is unstable in this setting. While work in suggests that this stability of SGD even in non-convex settings is a more general phenomenon, proving that this holds remains an open problem.

Conclusion

The success of machine learning algorithms in practice is often dictated by their ability to generalize. While recent work has developed great insight in to the training error of machine learning algorithms, much less is understood about their generalization error. Most work up to this point has either focused on specific algorithms or has made assumptions on the loss function (such as strong convexity) that may not be true in practice. By analyzing stability as coming from the convergence of an algorithm to global minima and the geometry surrounding them, we are able to derive much broader results. We develop easy-to-use stability results that encompass a general class of non-convex settings, and some of those can appear in interesting learning setups. Our bounds establish the stability for SGD, GD, SVRG, and RCD that is quantitatively comparable to more specialized results made in the past. Although our bounds are not nuanced enough to explain the generalization of mini-batch SGD compared to large-batch SGD, or full-batch GD, we hope that the generality of our nonconvex bounds serves as a step towards developing stability analyses that may help in understanding the generalization performance of practical machine learning algorithms.

There are still many exciting open problems concerning the stability and generalization of machine learning and optimization algorithms. We give a few below.

Stability for non-convex loss functions: While our results establish the stability of learning algorithms in some non-convex scenarios, it is unclear how to extend them directly to more general non-convex loss functions. Due to the wide variety of similar but not identical algorithms in machine learning, it would be particularly interesting to derive black-box results on the stability of learning algorithms for general non-convex loss functions, even when it concerns convergence to approximate local minima. In general, we expect these to be a function of the geometry of the loss function and the convergence of the algorithm, but it is quite possible that there are other factors that more directly control stability.

Generalization error of local minima: In general, we cannot expect a loss function to have only global minima. The question remains, among all local minima of a loss function, which one has the smallest generalization gap? The local minima with the smallest generalization error may not be a global minimizer. Even if we restrict to global minimizers, these may have different generalization errors. Is there a simple geometric characterization of their generalization error? This question may be connected to how sharp the loss function looks nearby the point. Steeper loss functions imply that small perturbations can greatly change the error, which suggests smaller generalization error. Theoretical results concerning this fact could be useful in stability analysis and the design of algorithms.

Generalization of SGD vs GD: We showed above that there are settings in which gradient descent is not uniformly stable, but SGD is. Empirically, SGD leads to small generalization error in neural networks . Theoretically, it is unclear how widespread this phenomenon is. Does SGD actually lead to more generalizable models than gradient descent? If so, why? How does this compare to other variants of SGD? Any theoretical results concerning the difference in generalization between SGD and other related algorithms could be extremely useful in the design of training algorithms and for understanding the empirical success of neural networks trained by SGD.

The geometry of critical points in real neural networks: We show above that linear neural networks obey relatively nice conditions on their local minima. In particular, as long as we restrict to weights that are full rank, all critical points are actually global minima. However, linear neural networks are not very exciting (e.g., they just correspond to linear models). Developing general theorems concerning the geometry underlying the critical points of neural networks is a challenging but very interesting open problem.

References

Appendix A Omitted Proofs

We have the following equivalent definition of PL due to Karimi et al.

The PL condition is equivalent to the condition that there is some constant μ>0\mu>0 such that for all xx, ∥∇f(x)∥≥μ∥xp−x∥\|\nabla f(x)\|\geq\mu\|x_{p}-x\|.

In this same paper, Karimi et al. show that PL functions also satisfy QG.

The PL condition implies the QG condition.

A.2 Proof of Theorem 3.1

Fix a training set SS and i∈{1,…,n}i\in\{1,\ldots,n\}. We will show pointwise hypothesis stability for all S,iS,i instead of for them in expectation. Let w1w_{1} denote the output of A\mathcal{A} on SS, and let w2w_{2} denote the output of A\mathcal{A} on SiS^{i}. Let w1∗w_{1}^{*} denote the critical point of fSf_{S} to which w1w_{1} is approaching, and w2∗w_{2}^{*} denote the critical point of fSif_{S^{i}} that w2w_{2} is approaching. We then have,

We first wish to bound the first and third terms of (3). The bound depends on the case in Theorem 3.1.

Case 2: As stated in Lemma A.2, the PL condition implies the QG condition. Therefore,

By assumption on case 2, ∣fS(w1)−fS(w1∗)∣≤O(ϵA′)|f_{S}(w_{1})-f_{S}(w_{1}^{*})|\leq O(\epsilon^{\prime}_{\mathcal{A}}). This implies

Case 3: By Lemma A.1, the PL condition on w1,w1∗w_{1},w_{1}^{*} implies that

Using the fact that fSf_{S} is LL-Lipschitz and the fact that ∥∇fS(wj)∥≤O(ϵA′′)\|\nabla f_{S}(w_{j})\|\leq O(\epsilon^{\prime\prime}_{\mathcal{A}}) by assumption on Case 3, we find

By the PL condition, we can find a local minima uu of fSf_{S} such that

Similarly, we can find a local minima vv of fSif_{S^{i}} such that

Note that since ∇fSi(w2∗)=0\nabla f_{S^{i}}(w_{2}^{*})=0, we get:

Similarly, since ∇fS(w1∗)=0\nabla f_{S}(w_{1}^{*})=0, we get:

Since all local minima of a PL function are global minima, we obtain

A.3 Proof of Theorem 3.3

Fix a training set SS and i∈{1,…,n}i\in\{1,\ldots,n\}. We will show pointwise hypothesis stability for all S,iS,i instead of for them in expectation. Let w1w_{1} denote the output of A\mathcal{A} on SS, and let w2w_{2} denote the output of A\mathcal{A} on SiS^{i}. Let w1∗w_{1}^{*} denote the critical point of fSf_{S} to which w1w_{1} is approaching, and w2∗w_{2}^{*} denote the critical point of fSif_{S^{i}} that w2w_{2} is approaching. We then have,

We first wish to bound the first and third terms of (8). The bound depends on the case in Theorem 3.3.

By assumption on case 2, ∣fS(w1)−fS(w1∗)∣≤O(ϵA′)|f_{S}(w_{1})-f_{S}(w_{1}^{*})|\leq O(\epsilon^{\prime}_{\mathcal{A}}). This implies

Note that ∣fS(w1∗)−fS(v)∣=0|f_{S}(w_{1}^{*})-f_{S}(v)|=0 by our realizability assumption. By assumption on w1∗,w2∗w_{1}^{*},w_{2}^{*}, we know that fS(w1∗)≤fS(w2∗)f_{S}(w_{1}^{*})\leq f_{S}(w_{2}^{*}) and fSi(w2∗)≤fSi(w1∗)f_{S^{i}}(w_{2}^{*})\leq f_{S^{i}}(w_{1}^{*}). Some simple analysis shows

A.4 Proof of Theorem 3.5

Let S1={z1,…,zn},S2={z1,…,zn−1,zn′}S_{1}=\{z_{1},\ldots,z_{n}\},S_{2}=\{z_{1},\ldots,z_{n-1},z_{n}^{\prime}\} be data sets of size nn differing only in one entry. Let wiw_{i} denote the output of A\mathcal{A} on data set SiS_{i} and let wi∗w_{i}^{*} denote wSi∗w_{S_{i}}^{*}. Let fi(w)=fSi(w)f_{i}(w)=f_{S_{i}}(w).

Note that by A1, we know that w1∗w_{1}^{*} is the closest optimal point of f1f_{1} to w2∗w_{2}^{*}. By the PL condition,

This bounds the second term of 12. The first and third terms must be bounded differently depending on the case.

Case 1: By assumption, for i=1,2i=1,2, ∥wi−wi∗∥=O(ϵA)\|w_{i}-w_{i}^{*}\|=O(\epsilon_{\mathcal{A}}), proving the result.

Case 2: As stated in Lemma A.2, PL implies QG. Therefore for i=1,2i=1,2,

Case 3: As mentioned above in Lemma A.1, the PL condition implies that for i=1,2i=1,2,

Since ∥∇fi(wi)∥≤O(ϵA′′)\|\nabla f_{i}(w_{i})\|\leq O(\epsilon^{\prime\prime}_{\mathcal{A}}), we get the desired result.∎

A.5 Proof of Theorem 3.6

Note that by A1, we know that w1∗w_{1}^{*} is the closest optimal point of f1f_{1} to w2∗w_{2}^{*}. By Lemma A.3,

This bounds the second term of 12. The first and third terms must be bounded differently depending on the case.

Case 1: By assumption, for i=1,2i=1,2, ∥wi−wi∗∥=O(ϵA)\|w_{i}-w_{i}^{*}\|=O(\epsilon_{\mathcal{A}}), proving the result.

Case 2: By the QG property, we get that for i=1,2i=1,2,

A.6 Stability of Gradient Descent for Convex Loss Functions

To prove the stability of gradient descent, we will assume that the underlying loss function is smooth.

In , Hardt et al. show the following theorem.

Performing similar analysis for gradient descent, we obtain the following theorem.

To prove this theorem, we use similar techniques to those in . We first consider the convex case.

So, if γtn≤2β\frac{\gamma_{t}}{n}\leq\frac{2}{\beta} for all tt, we get:

We now move to the λ\lambda-strongly convex case. For simplicity of analysis, we assume that we use a constant step size γ\gamma such that γ≤1/β\gamma\leq 1/\beta.

The proof remains the same, except when using co-coercivity. Under this assumption, some plug and play in an analogous fashion will show:

Note that if γ≤1β\gamma\leq\frac{1}{\beta} then the second term is nonnegative and one can show that this implies:

This implies uniform stability with parameter 2L2λn\frac{2L^{2}}{\lambda n}.∎

A.7 Proof of Theorem 4.1

For almost all ww, we can write σ(Xw)\sigma(Xw) as diag(b)Xw\text{diag}(b)Xw for a vector bb where bi(Xw)i=σ((Xw)i)b_{i}(Xw)_{i}=\sigma((Xw)_{i}) (this only exludes points ww such that (Xw)i(Xw)_{i} is on a cusp of the piecewise-linear function). Then in an open neighborhood of such an ww, we find:

For a given ww, let wpw_{p} be the closest global minima of ff (i.e., the closest point such that f∗=f(wp)f^{*}=f(w_{p})). By strong convexity of gg, we find:

Note that the minimum singular value of diag(b)\text{diag}(b) is the square root of the minimum eigenvalue of diag(b)2\text{diag}(b)^{2}. Since diag(b)2\text{diag}(b)^{2} has entries ci2c_{i}^{2} on the diagonal, we know that the minimum singular value is at least c=min⁡i{∣ci∣}c=\min_{i}\{|c_{i}|\}. Therefore we get:

A.8 Proof of Lemma 4.2

Using basic properties of the Frobenius norm and the definition of the pseudo-inverse, we have

This last step follows by basic properties of the pseudo-inverse. By the triangle inequality,

Note that YX+X−YYX^{+}X-Y is the component of YY that is orthogonal to the row-space of XX, while YX+XYX^{+}X is the projection of YY on to this row space. Therefore, YX+X−YYX^{+}X-Y is orthogonal to WX−YX+XWX-YX^{+}X with respect to the trace inner product. Therefore, the inequality above is actually an equality, that is

A.9 Proof of Lemma 4.3

Our proof uses similar techniques to that of Hardt and Ma . We wish to compute the gradient of ff with respect to a matrix WjW_{j}. One can show the following:

By assumption, σmin⁡(Wj)≥τ\sigma_{\min}(W_{j})\geq\tau. Therefore:

Taking the gradient with respect to all WiW_{i} we get:

In this subsection we show that a provable rate for SGD on smooth functions with learning rate proportional to O(1/t)O(1/t) might require a large number of iterations. While show stability of SGD in non-convex settings with such a step-size, their stability bounds grow close to linearly with the number of iterations. When exponentially many steps are taken, this no longer implies useful generalization bounds on SGD.

When a function f(x)=∑i=1nfi(x)f(x)=\sum_{i=1}^{n}f_{i}(x) is β\beta smooth on its domain, then the following holds:

Fix some tt and let x=xt+1=xt−γt∇fst(xt)x=x_{t+1}=x_{t}-\gamma_{t}\nabla f_{s_{t}}(x_{t}) and let y=xty=x_{t}. Here, x0x_{0} is set to some initial vector value, sts_{t} is a uniform iid sample from {1,…,n]\{1,\ldots,n], and γt=c/t\gamma_{t}=c/t for some constant c>0c>0. Then, due to the β\beta-smoothness of ff we have the following:

Taking expectation with respect to all random samples sts_{t}, yields

Summing the above inequality for all tt terms from 0 to TT we get

where C1C_{1} is a universal constant that depends only on cc. Observe that even if C2=0C_{2}=0, using the above simple bounding technique (a simplified version of the nonconvex convergence bounds of ), requires O(e−ϵ)O(e^{-\epsilon}) steps to reach error ϵ\epsilon, i.e., an exponentially large number of steps.

We note that the above bound does not imply that there does not exist a smooth function for which 1/t1/t stepsizes suffice for polynomial-time convergence (in fact there are several convex problems for wich 1/t1/t suffices for fast convergence). However, the above implies that when we are only assuming smoothness on a nonconvex function, it may be the case that there exist nonconvex problems where 1/t1/t implies exponentially slow convergence.

This implies that if the optimal model is at distance Ω(Md)\Omega(Md) from x0x_{0}, we would require at least O(eM⋅d)O(e^{M\cdot d}) iterations to reach it in expectation.