On the Power of Over-parametrization in Neural Networks with Quadratic Activation

Simon S. Du, Jason D. Lee

Introduction

Neural networks have achieved a remarkable impact on many applications such computer vision, reinforcement learning and natural language processing. Though neural networks are successful in practice, their theoretical properties are not yet well understood. Specifically, there are two intriguing empirical observations that existing theories cannot explain.

Optimization: Despite the highly non-convex nature of the objective function, simple first-order algorithms like stochastic gradient descent are able to minimize the training loss of neural networks. Researchers have conjectured that the use of over-parametrization (Livni et al., 2014; Safran and Shamir, 2017) is the primary reason why local search algorithms can achieve low training error. The intuition is over-parametrization alters the loss function to have a large manifold of globally optimal solutions, which in turn allows local search algorithms to more easily find a global optimal solution.

Generalization: From the statistical point of view, over-parametrization may hinder effective generalization, since it greatly increases the number of parameters to the point of having number of parameters exceed the sample size. To address this, practitioners often use explicit forms of regularization such as weight decay, dropout, or early stopping to improve generalization. However in the non-convex setting, theoretically, we do not have a good quantitative understanding on how these regularizations help generalization for neural network models.

In this paper, we provide new theoretical insights into the optimization landscape and generalization ability of over-parametrized neural networks. Specifically we consider the neural network of the following form:

In our setting, we fix the second layer to be a=(1,…,1)\mathbf{a}=\left(1,\ldots,1\right). Although it is simpler than the case where the second layer is not fixed, the effect of over-parameterization can be studied in this setting as well because we do not have any restriction on the number of hidden nodes.

We focus on quadratic activation function σ(z)=z2\sigma\left(z\right)=z^{2}. Though quadratic activations are rarely used in practice, stacking multiple such two-layer blocks can be used to simulate higher-order polynomial neural networks and sigmodial activated neural networks (Livni et al., 2014; Soltani and Hegde, 2017).

In practice, we have nn training samples {xi,yi}i=1n\left\{\mathbf{x}_{i},y_{i}\right\}_{i=1}^{n} and solve the following optimization problem to learn a neural network

To improve the generalization ability, we often add explicit regularization. In this paper, we focus on a particular regularization technique, weight decay for which we slightly change the gradient descent algorithm to

where λ\lambda is the decay rate. Note this algorithm is equivalent to applying the gradient descent algorithm on the regularized loss

In this setup, we make the following theoretical contributions to explain why over-parametrization helps optimization and still allows for generalization.

We analyze two kinds of over-parameterization. First we show that for

Second, we consider another form of over-parametrization,

This condition on the amount of over-parameterization is much milder than k≥nk\geq n, a condition used in many previous papers (Nguyen and Hein, 2017a, b). Further in practice, k(k+1)/2>nk(k+1)/2>n is a much milder requirement than k≥dk\geq d, since if k≈2nk\approx\sqrt{2n} and n<<d2n<<d^{2} then k<<dk<<d. In this setting, we consider the perturbed version of the Problem (2):

where C\mathbf{C} is a random positive semidefinite matrix with arbitrarily small Frobenius norm. We show that if k(k+1)2>n\frac{k(k+1)}{2}>n, Problem (3) also has the desired properties that all local minima are global and all saddle points are strict with probability 11. Since C\mathbf{C} has small Frobenius norm, the optimal value of Problem (3) is very close to that of Problem (2). See Section 3 for the precise statement.

To prove this surprising fact, we bring forward ideas from smoothed analysis in constructing the perturbed loss function (3), which we believe is useful for analyzing the landscape of non-convex losses.

Weight-decay Helps Generalization.

We show because of weight-decay, the optimal solution of Problem (2) also generalizes well. The major observation is weight-decay ensures the solution of Problem (2) has low Frobenius norm, which is equivalent to matrix W⊤W\mathbf{W}^{\top}\mathbf{W} having low nuclear norm (Srebro et al., 2005). This observation allows us to use theory of Rademacher complexity to directly obtain quantitative generalization bounds. Our theory applies to a wide range of data distribution and in particular, does not need to assume the model is realizable. Further, the generalization bound does not depend on the number of epochs SGD runs or the number of hidden nodes.

To sum up, in this paper we justify the following folklore.

Over-parametrization allows us to find global optima and with weight decay, the solution also generalizes well.

2 Organization

This paper is organized as follows. In Section 2 we introduce necessary background and definitions. In Section 3 we present our main theorems on why over-parametrization helps optimization when k≥dk\geq d or k(k+1)2>n\frac{k(k+1)}{2}>n. In Section 4, we give quantitative generalization bounds to explain why weight decay helps generalization in the presence of over-parametrization. In Section 5, we prove our main theorems. We conclude and list future works in Section 6.

3 Related Works

Neural networks have enjoyed great success in many practical applications (Krizhevsky et al., 2012; Dauphin et al., 2016; Silver et al., 2016). To explain this success, many works have studied the expressiveness of neural networks. The expressive ability of shallow neural network dates back to 90s (Barron, 1994). Recent results give more refined analysis on deeper models (Bölcskei et al., 2017; Telgarsky, 2016; Wiatowski et al., 2017).

However, from the point of view of learning theory, it is well known that training a neural network is hard in the worst case (Blum and Rivest, 1989). Despite the worst-case pessimism, local search algorithms such as gradient descent are very successful in practice. With some additional assumptions, many works tried to design algorithms that provably learn a neural network (Goel et al., 2016; Sedghi and Anandkumar, 2014; Janzamin et al., 2015). However these algorithms are not gradient-based and do not provide insight on why local search algorithm works well.

Focusing on gradient-based algorithms, a line of research (Tian, 2017; Brutzkus and Globerson, 2017; Zhong et al., 2017a, b; Li and Yuan, 2017; Du et al., 2017b, c) analyzed the behavior of (stochastic) gradient descent with a structural assumption on the input distribution. The major drawback of these papers is that they all focus on the regression setting with least-squares loss and further assume the model is realizable meaning the label is the output of a neural network plus a zero mean noise, which is unrealistic. In the case of more than one hidden unit, the papers of (Li and Yuan, 2017; Zhong et al., 2017b) further require a stringent initialization condition to recover the true parameters.

Finding the optimal weights of a neural network is non-convex problem. Recently, researchers found that if the objective functions satisfy the following two key properties: (1) all local minima are global and (2) all saddle points and local maxima are strict, then first order method like gradient descent (Ge et al., 2015; Jin et al., 2017; Levy, 2016; Du et al., 2017a; Lee et al., 2016) can find a global minimum.

We now turn our attention to generalization ability of learned neural networks. It is well known that the classical learning theory cannot explain the generalization ability because VC-dimension of neural networks is large (Harvey et al., 2017; Zhang et al., 2016). A line of research tries to explain this phenomenon by studying the implicit regularization from stochastic gradient descent algorithm (Hardt et al., 2015; Pensia et al., 2018; Mou et al., 2017; Brutzkus et al., 2017; Li et al., 2017). However, the generalization bounds of these papers often depend on the number of epochs SGD runs, which is large in practice. Another direction is to study the generalization ability based on the norms of weight matrices in neural networks (Neyshabur et al., 2015, 2017a, 2017b; Bartlett et al., 2017; Liang et al., 2017; Golowich et al., 2017; Dziugaite and Roy, 2017; Wu et al., 2017). Our theorem on generalization ability also uses this idea but is more specialized to the network architecture (1).

After the initial submission of this manuscript, we became aware of concurrent work of (Bhojanapalli et al., 2018), which also considered the smoothed analysis technique to solve semi-definite programs in penalty form. The mathematical techniques in our work and (Bhojanapalli et al., 2018) are similar, but the focus is on two distinct problems of solving semi-definite programs and quadratic activation neural networks.

Preliminaries

We use bold-faced letters for vectors and matrices. For a vector v\mathbf{v}, we use ∥v∥2\left\|\mathbf{v}\right\|_{2} to denote the Euclidean norm. For a matrix M\mathbf{M}, we denote ∥M∥2\left\|\mathbf{M}\right\|_{2} the spectral norm and ∥M∥F\left\|\mathbf{M}\right\|_{F} the Frobenius norm. We let N(M)\mathcal{N}\left(\mathbf{M}\right) to denote the left null-space of M\mathbf{M}, i.e.

We use Σ(M)\Sigma\left(M\right) to denote the set of matrices with Frobenius norm bounded by MM and Σ1(1)\Sigma_{1}\left(1\right) to denote the set of rank-11 matrices with spectral norm bounded by 11. We also denote Sd\mathcal{S}^{d} the set of d×dd\times d symmetric positive semidefinite matrices.

In this paper, we characterize the landscape of over-parameterized neural networks. More specifically we study the properties of critical points of empirical loss. Here for a loss function L(W)L\left(\mathbf{W}\right), a critical point W∗\mathbf{W}^{*} satisfies ∇L(W∗)=0\nabla L\left(\mathbf{W}^{*}\right)=0. A critical point can be a local minimum or a saddle point.We do not differentiate between saddle points and local maxima in this paper. If W∗\mathbf{W}^{*} is a local minimum, then there is a neighborhood OO around W∗\mathbf{W}^{*} such that L(W∗)≤L(W)L\left(\mathbf{W}^{*}\right)\leq L\left(\mathbf{W}\right) for all W∈O\mathbf{W}\in O. If W∗\mathbf{W}^{*} is a saddle point, then for all neighborhood OO around W∗\mathbf{W}^{*}, there is a W∈O\mathbf{W}\in O such that L(W)<L(W∗)L\left(\mathbf{W}\right)<L\left(\mathbf{W}^{*}\right).

Ideally, we would like a loss function that satisfies the following two geometric properties.

If a loss function L(⋅)L\left(\cdot\right) satisfies Property 2.1 and Property 2.2, recent algorithmic advances in non-convex optimization show randomly initialized gradient descent algorithm or perturbed gradient descent can find a global minimum (Lee et al., 2016; Ge et al., 2015; Jin et al., 2017; Du et al., 2017a).

Lastly, standard applications of Rademacher complexity theory will be used to derive generalization bounds.

Given a sample S=(x1,…,xn)S=\left(\mathbf{x}_{1},\ldots,\mathbf{x}_{n}\right), the empirical Rademacher complexity of a function class F\mathcal{F} is defined as

Overparametrization Helps Optimization

In this section we present our main results on explaining why over-parametrization helps local search algorithms find a global optimal solution. We consider two kinds of over-parameterization, k≥dk\geq d and k(k+1)2>n\frac{k(k+1)}{2}>n. We begin with the simpler case when k≥dk\geq d.

The above result states that given an arbitrary data set, the optimization landscape has benign properties that facilitate finding globally optimal neural networks. In particular, by setting the last layer to be the average pooling layer, all local minima are global minima and all saddles have a direction of negative curvature. This in turn implies that gradient descent on the first layer weights, when initialized at random, converges to a global optimum. These desired properties hold as long as the hidden layer is wide (k≥dk\geq d).

An interesting and perhaps surprising aspect of Theorem 3.1 is its generality. It applies to arbitrary data set of any size with any convex differentiable loss function.

Now we consider the second case when k(k+1)2>n\frac{k(k+1)}{2}>n. As mentioned earlier, in practice this is often a milder requirement than k≥dk\geq d, and one of the main novelties of this paper.

obeys Property 2.1 and Property 2.2. Further, any global optimal solution W^\widehat{\mathbf{W}}of Problem (4) satisfies

Similar to Theorem 3.1, Theorem 3.2 states that if k(k+1)2>n\frac{k(k+1)}{2}>n, then for an arbitrary data set, the perturbed objective function (3) has the desired properties that enable local search heuristics to find globally optimal solution for a general class of loss functions. Further, we can choose this perturbation to be arbitrarily small so the minimum of (3) is close to (2).

The proof of theorem is inspired by a line of literature started by Pataki (1998, 2000); Burer and Monteiro (2003); Boumal et al. (2016). In summary, Boumal et al. (2016) showed that for “almost all” semidefinite programs, every local minima of the rank rr non-convex formulation of an SDP is a global minimum of the original SDP. However, this theorem applies with the important caveat of only applying to semidefinite programs that do not fall into a measure zero set. Our primary contribution is to develop a procedure that exploits this by a) constructing a perturbed objective to avoid the measure zero set, b) proving that the perturbed objective has Property 2.1 and 2.2, and c) showing the optimal value of the perturbed objective is close to the original objective. Further note that the analysis of (Boumal et al., 2016) does not apply since our loss functions, such as the logistic loss, are not semi-definite representable. We refer readers to Section 5.2 for more technical insights.

Weight-decay Helps Generalization

To derive the generalization bound, we first recall the classical generalization bound based on Rademacher complexity bound (c.f. Theorem 2 of (Koltchinskii and Panchenko, 2002)).

Assume each data point is sampled i.i.d from some distribution P\mathcal{P}, i.e.,

where C>0C>0 is an absolute constant and RS(Σ(M))R_{S}\left(\Sigma\left(M\right)\right) is the Rademacher complexity of Σ(M)\Sigma\left(M\right).

With Theorem 4.1 at hand, we only need to bound the Rademacher complexity of Σ(M)\Sigma\left(M\right). Note that Rademacher complexity is a distribution dependent quantity. If the data is arbitrary, we cannot have any guarantee. We begin with a theorem for bounded input domain.

Combining Theorem 4.1 and Theorem 4.2 we can obtain a generalization bound.

Under the same assumptions of Theorem 4.1 and Theorem 4.2, we have

While Theorem 4.2 is a valid bound, it is rather pessimistic because we only assume x\mathbf{x} is bounded. Consider the following scenario in which each input is sampled from a standard Gaussian distribution xi∼N(0,I)\mathbf{x}_{i}\sim N\left(0,\mathbf{I}\right). Then ignoring the logarithmic factors, using standard Gaussian concentration bound we can show with high probability ∥xi∥2=O~(d)\left\|\mathbf{x}_{i}\right\|_{2}=\widetilde{O}\left(\sqrt{d}\right). O~(⋅)\widetilde{O}(\cdot) hides logarithmic factors. Plugging in this bound we have

Note in this bound, it has a quadratic dependency on the dimension, so we need to have n=Ω(d2)n=\Omega\left(d^{2}\right) to have a meaningful bound.

In fact, for specific distributions like Gaussian using Theorem 5.2, we can often derive a stronger generalization bound.

Suppose xi∼N(0,I)\mathbf{x}_{i}\sim N(0,\mathbf{I}) for i=1,…,ni=1,\ldots,n. If the number of samples satisfies n≥dlog⁡dn\geq d\log d, we have with high probability that Rademacher complexity satisfies

Again, combining Theorem 4.1 and Corollary 4.1 we obtain the following generalization bound for Gaussian input distribution

Under the same assumptions of Theorem 4.1 and Corollary 4.1, we have

Comparing Theorem 4.4 with generalization bound (5), Theorem 4.4 has an O(d)O(\sqrt{d}) advantage. Theorem 4.4 has the d/n\sqrt{d/n} dependency, which is the usual parametric rate. Further in practice, number of training samples and input dimension are often of the same order for common datasets and architectures (Zhang et al., 2016).

Corollary 4.1 is a special case of the more general Theorem 5.2 which only requires a bound on the fourth moment, ∥∑i=1n(xixi⊤)2∥2≤s\left\|\sum_{i=1}^{n}(x_{i}x_{i}^{\top})^{2}\right\|_{2}\leq s. In general, our theorems suggest if the Frobenius norm of weight matrix W\mathbf{W} is small and the input x\mathbf{x} is sampled from a benign distribution with controlled 4th moments, then we have good generalization.

As a concrete scenario, consider a favorable setting where the true data can be correctly classified by a small network using only k0≪kk_{0}\ll k hidden units. The weights W∗W^{*} are non-zero only in the first k0k_{0} rows and max⁡j∈[k0]∥ej⊤W∗∥2≤R\max_{j\in[k_{0}]}\left\|e_{j}^{\top}W^{*}\right\|_{2}\leq R. From Theorem 4.4 to reach generalization gap of ϵ\epsilon, we have sample complexity of n≥1ϵ2C2L2R4dk02n\geq\frac{1}{\epsilon^{2}}C^{2}L^{2}R^{4}dk_{0}^{2}, which only depends on the effective number of hidden units k0≪kk_{0}\ll k. The same result can be reached for more general input distributions by using Theorem 5.2 in place of Theorem 4.4.

Proofs

Our proofs of over-parametrization helps optimization build upon existing geometric characterization on matrix factorization. We first cite a useful Theorem by Haeffele et al. (2014).Theorem 2 of (Haeffele et al., 2014) assumes W\mathbf{W} is a local minimum, but scrutinizing its proof, we can see that the assumption can be relaxed to ∇L(W)=0 and ∇2L(W)≽0\nabla L\left(\mathbf{W}\right)=\mathbf{0}\text{ and }\nabla^{2}L\left(\mathbf{W}\right)\succcurlyeq\mathbf{0}.

We prove Property 2.1 and Property 2.2 simultaneously by showing if a W\mathbf{W} satisfy

To prove Theorem 3.1, the key idea is to consider the follow reference optimization problem.

Using Equation (7), we know M=W⊤W\mathbf{M}=\mathbf{W}^{\top}\mathbf{W} achieves the global minimum in Problem (8). The proof is thus complete. ∎

2 Proof of Theorem 3.2

We first prove LC(W)L_{\mathbf{C}}\left(\mathbf{W}\right) satisfies Property 2.1 and Property 2.2. Similar to the proof of Theorem 3.1, we prove these two properties simultaneously by showing if a W\mathbf{W} satisfy

Now define M=S(v∗)+λI+C\mathbf{M}=\mathbf{S}\left(\mathbf{v}^{*}\right)+\lambda\mathbf{I}+\mathbf{C} with

The key idea is to use these two conditions to upper bound the dimension of C\mathbf{C}. To this end, we first define the set

Note because C\mathbf{C} and W^⊤W^\widehat{\mathbf{W}}^{\top}\widehat{\mathbf{W}} are both positive semidefinite, we have ⟨C,W^⊤W^⟩≥0\langle\mathbf{C},\widehat{\mathbf{W}}^{\top}\widehat{\mathbf{W}}\rangle\geq 0. Thus L(W^)≤L(W∗)+δ∥W∗∥F2L\left(\widehat{\mathbf{W}}\right)\leq L\left(\mathbf{W}^{*}\right)+\delta\left\|\mathbf{W}^{*}\right\|_{F}^{2}. ∎

3 Proof of Theorem 4.2 and Theorem 4.4

Our proof is inspired by (Srebro and Shraibman, 2005) which exploits the structure of nuclear norm bounded space. We first prove a general Theorem that only depends on the property of the fourth-moment of input random variables.

Suppose the input random variable satisfies ∥∑i=1n(xixi⊤)2∥2≤s\left\|\sum_{i=1}^{n}\left(\mathbf{x}_{i}\mathbf{x}_{i}^{\top}\right)^{2}\right\|_{2}\leq s. Then the Rademacher complexity of Σ(M)\Sigma\left(M\right) is bounded by

For a given set of inputs S={xi}i=1nS=\left\{\mathbf{x}_{i}\right\}_{i=1}^{n} in our context, we can write Rademacher complexity as

Since Rademacher complexity does not change when taking convex combinations, we can first bound Rademacher complexity of the class of rank-11 matrices with spectral norm bounded by 11 and then take convex hull and scale by MM. Note for W∈Σ1(1)\mathbf{W}\in\Sigma_{1}\left(1\right), we can write W=vw⊤\mathbf{W}=\mathbf{v}\mathbf{w}^{\top} with ∥w∥2≤1\left\|\mathbf{w}\right\|_{2}\leq 1 and ∥v∥2=1\left\|\mathbf{v}\right\|_{2}=1. Using this expression, we can obtain an explicit formula of Rademacher complexity.

Applying Rademacher matrix series expectation bound (Theorem 4.6.1 of (Tropp et al., 2015)), we have

Now taking the convex hull and, scaling by MM we obtain the desired result. ∎

With Theorem 5.2 at hand, for different distributions, we only need to bound ∥∑i=1n(xixi⊤)2∥2\left\|\sum_{i=1}^{n}\left(\mathbf{x}_{i}\mathbf{x}_{i}^{\top}\right)^{2}\right\|_{2}.

Since we assume ∥xi∥2≤b\left\|\mathbf{x}_{i}\right\|_{2}\leq b, we directly have

Plugging this bound in Theorem 5.2 we obtain the desired inequality. ∎

To prove Corollary 4.1, we use Theorem 5.2 and Lemma 4.7 in (Soltanolkotabi et al., 2017) to upper bound s=∥∑i=1n∥xi∥2xixi⊤∥s=\left\|\sum_{i=1}^{n}\left\|x_{i}\right\|^{2}x_{i}x_{i}^{\top}\right\|. By letting A=IA=I in Lemma 4.7, we find that

with probability at least 1−Cd1-\frac{C}{d}. This completes the proof of Corollary 4.1.

Using this bound in Theorem 5.2 comletes the proof of Theorem 4.4. ∎

Conclusion and Future Works

In this paper we provided new theoretical results on over-parameterized neural networks. Using smoothed analysis, we showed as long as the number of hidden nodes is bigger than the input dimension or square root of the number of training data, the loss surface has benign properties that enable local search algorithms to find global minima. We further use the theory of Rademacher complexity to show the learned neural can generalize well.

Our next step is consider neural networks with other activation functions and how over-parametrization allows for efficient local-search algorithms to find near global minimzers. Another interesting direction to extend our results to deeper model.

Acknowledgment

S.S.D. was supported by NSF grant IIS1563887, AFRL grant FA8750-17-2-0212 and DARPA D17AP00001. J.D.L. acknowledges support of the ARO under MURI Award W911NF-11-1-0303. This is part of the collaboration between US DOD, UK MOD and UK Engineering and Physical Research Council (EPSRC) under the Multidisciplinary University Research Initiative.

References