Global optimality conditions for deep neural networks
Chulhee Yun, Suvrit Sra, Ali Jadbabaie
Introduction
Since the advent of AlexNet (Krizhevsky et al., 2012), deep neural networks have surged in popularity, and have redefined the state-of-the-art across many application areas of machine learning and artificial intelligence, such as computer vision, speech recognition, and natural language processing. However, a concrete theoretical understanding of why deep neural networks work well in practice remains elusive. From the perspective of optimization, a significant barrier is imposed by the nonconvexity of training neural networks. Moreover, it was proved by Blum & Rivest (1988) that training even a 3-node neural network to global optimality is NP-Hard in the worst case, so there is little hope that neural networks have properties that make global optimization tractable.
Despite the difficulties of optimizing weights in neural networks, the empirical successes suggest that the local minima of their loss surfaces could be close to global minima; and several papers have recently appeared in the literature attempting to provide a theoretical justification for the success of these models. For example, by relating neural networks to spherical spin-glass models from statistical physics, Choromanska et al. (2015) provided some empirical evidence that the increase of size of neural networks makes local minima close to global minima.
Another line of results (Yu & Chen, 1995; Soudry & Carmon, 2016; Xie et al., 2016; Nguyen & Hein, 2017) provides conditions under which a critical point of the empirical risk is a global minimum. Such results roughly involve proving that if full rank conditions of certain matrices (as well as some additional technical conditions) are satisfied, derivative of the risk being zero implies loss being zero. However, these results are obtained under restrictive assumptions; for example, Nguyen & Hein (2017) require the width of one of the hidden layers to be as large as the number of training examples. Soudry & Carmon (2016) and Xie et al. (2016) require the product of widths of two adjacent layers to be at least as large as the number of training examples, meaning that the number of parameters in the model must grow rapidly as we have more training data available. Another recent paper (Haeffele & Vidal, 2017) provides a sufficient condition for global optimality when the neural network is composed of subnetworks with identical architectures connected in parallel and a regularizer is designed to control the number of parallel architectures.
Towards obtaining a more precise characterization of the loss-surfaces, a valuable conceptual simplification of deep nonlinear networks is deep linear neural networks, in which all activation functions are linear and the output of the entire network is a chained product of weight matrices with the input vector. Although at first sight a deep linear model may appear overly simplistic, even its optimization is nonconvex, and only recently theoretical results on this problem have started emerging. Interestingly, already in 1989, Baldi & Hornik (1989) showed that some shallow linear neural networks have no local minima. More recently, Kawaguchi (2016) extended this result to deep linear networks and proved that any local minimum is also global while any other critical point is a saddle point. Subsequently, Lu & Kawaguchi (2017) provided a simpler proof that any local minimum is also global, with fewer assumptions than (Kawaguchi, 2016). Motivated by the success of deep residual networks (He et al., 2016a; b), Hardt & Ma (2017) investigated loss surfaces of deep linear residual networks and showed every critical point is a global minimum in a near-identity region; subsequently, Bartlett et al. (2017) extended this result to a nonlinear function space setting.
Inspired by this recent line of work, we study deep linear and nonlinear networks, in settings either similar to or more general than existing work. We summarize our main contributions below.
We provide both necessary and sufficient conditions for a critical point of the empirical risk to be a global minimum. Specifically, Theorem 2.1 shows that if the hidden layers are wide enough, then a critical point of the risk function is a global minimum if and only if the product of all parameter matrices is full-rank. In Theorem 2.2, we consider the case where some hidden layers have smaller width than both the input and output layers, and again provide necessary and sufficient conditions for global optimality. In comparison, Kawaguchi (2016) only proves that every critical point of the risk is either a global minimum or a saddle; it is an “existence” result without any computational implication. In contrast, we present efficiently checkable conditions for distinguishing the two different types of critical points; we can even use these conditions while running optimization to test whether the critical points we encounter are saddle points or not, if desired. It is also worth noting that such tests are intractable for general nonconvex optimization (Murty & Kabadi, 1987).
Under the same assumption as (Hardt & Ma, 2017) on the data distribution, namely, a linear model with Gaussian noise, we can modify Theorem 2.1 to handle the population risk. As a corollary, we not only recover Theorem 2.2 in (Hardt & Ma, 2017), but also extend it to a strictly larger set, while removing their assumption that the true underlying linear model has a positive determinant.
Motivated by (Bartlett et al., 2017), we extend our results on deep linear networks to obtain sufficient conditions for global optimality in deep nonlinear networks, although only via a function space view; these are presented in Theorems 4.1 and 4.2.
Global optimality conditions for deep linear neural networks
In this section, we describe the problem formulation and notations for deep linear neural networks, state main results (Theorems 2.1 and 2.2), and explain their implication.
where is a shorthand notation for the tuple .
We assume that and , and that and have full ranks. These assumptions are common when we consider supervised learning problems with deep neural networks (e.g. Kawaguchi (2016)). We also assume that the singular values of are all distinct, which is made for notational simplicity and can be relaxed without too much difficulty.
2 Necessary and sufficient conditions for global optimality
We now present two main theorems for deep linear neural networks. The theorems describe two sets, one for the case and the other for , inside which every critical point of is a global minimum. Moreover, the sets have another remarkable property that every critical point outside of these sets is a saddle point. Previous works (Kawaguchi, 2016; Lu & Kawaguchi, 2017) showed that any critical point is either a global minimum or a saddle point, without providing any condition to distinguish between the two; here, we take a step further and partition the domain of into two sets clearly delineating one set which only contains global minima and the other set with only saddle points.
If , define the following set
Then, every critical point of in is a global minimum. Moreover, every critical point of in is a saddle point.
If , define the following set
Then, every critical point of in is a global minimum. Moreover, every critical point of in is a saddle point.
Theorems 2.1 and 2.2 provide necessary and sufficient conditions for a critical point of to be globally optimal. From an algorithmic perspective, they provide easily checkable conditions, which we can use to determine if the critical point the algorithm encountered is a global optimum or not. Given that is nonconvex, it is interesting to have such efficient tests for global optimality, which is not possible in general (Murty & Kabadi, 1987).
In Hardt & Ma (2017), the authors consider minimizing population risk of linear residual networks:
where . They assume that is drawn from a zero-mean distribution with a fixed covariance matrix, and where is iid standard Gaussian noise and is the true underlying matrix with . With these assumptions they prove that whenever for all , any critical point is a global minimum (Hardt & Ma, 2017, Theorem 2.2).
Under the same assumptions on data distribution, we can slightly modify Theorem 2.1 to derive a population risk counterpart, and in fact notice that the result proved in Hardt & Ma (2017) is a corollary of this modification because having for all is a sufficient condition for having full rank. Moreover, notice that we can remove the assumption which was required by Hardt & Ma (2017). We state this special case as a corollary:
We also note in passing that the classical problem of matrix factorization is a special case of deep linear neural networks, so our theorems can also be directly applied.
The previous result (Kawaguchi, 2016) assumed and showed that: 1) every local minimum is a global minimum, and 2) any other critical point is a saddle point. A subsequent paper by Lu & Kawaguchi (2017) proved 1) without the assumption , but as far as we know there is no result showing 2) in the case of . We provide the proof for this case in Lemma B.1. In fact, we propose an alternative proof technique for handling degenerate critical points, which is much simpler than the technique presented by Kawaguchi (2016).
Analysis of deep linear networks
In this section, we provide proofs for Theorems 2.1 and 2.2.
We first analyze the globally optimal solution of a “relaxation” of , which turns out to be very useful while proving Theorems 2.1 and 2.2. Consider the relaxed risk function
This means that if there exists such that , then is a global minimum of the function . This observation is very important in proofs; we will show that inside certain sets, any critical point of must satisfy , where is a global optimum of . This proves that , thus showing that is a global minimum of .
By restating this observation as an optimization problem, the solution of problem in (1) is bounded below by the minimum value of the following:
In case where , (2) is actually an unconstrained optimization problem. Note that is a convex function of , so any critical point is a global minimum. By differentiating and setting the derivative to zero, we can easily get the unique globally optimal solution
In case of , the problem becomes non-convex because of the rank constraint, but its exact solution can still be computed easily. We present the solution of this case as a proposition and defer the proof to Appendix C due to its technicalities.
Suppose . Then the optimal solution to (2) is
which is the orthogonal projection of onto the column space of .
2 Partial derivatives of L(W)𝐿𝑊L(W)
By simple matrix calculus, we can calculate the derivatives of with respect to , for . We present the result as the following lemma, and defer the details to Appendix C.
The partial derivative of with respect to is given as
We also state an elementary lemma which proves useful in our proofs, whose proof we defer to Appendix C.
3 Proof of Theorem 2.1
We prove Theorem 2.1, which addresses the case . First, recall that the set defined in Theorem 2.1 is
As seen in (3), the unique minimum point of has rank . So, no point can be a global minimum of . Therefore, by Kawaguchi (2016, Theorem 2.3.(iii)) and Lemma B.1, any critical point in must be a saddle point.
For the rest of our proof, we need to consider two cases: and . If , both cases work. The outline of the proof is as follows: we define a new set , show that any critical point in the set is a global minimum, and then show that every is also in for some . This proves that any critical point of in is also a critical point in for some , hence a global minimum.
The following proposition proves the first step:
Assume that . For any , define the following set:
Then any critical point of in is a global minimum point.
By the above inequality, any critical point in satisfies
which means that . The product is the unique globally optimal solution (3) of the relaxed problem in (2), so is a global minimum point of .
and the rest of the proof flows in a similar way as the previous case. ∎
For any point , there exists an such that .
Define a new set , a “limit” version (as ) of , as
We show that by showing that . Consider
Then any must have , so . Thus, any is also in , so either or , depending on the cases. Then, we can set
We always have because the matrices are full rank, and we can see that . ∎
4 Proof of Theorem 2.2
In this section we prove Theorem 2.2, which tackles the case . Note that this assumption also implies that .
The globally optimal point of the relaxed problem (2) has rank , as seen in (4). Thus, any point outside of cannot be a global minimum. Then, by Kawaguchi (2016, Theorem 2.3.(iii)) and Lemma B.1, it follows that any critical point in must be a saddle point. The remaining proof considers points in .
For this section, let us introduce some additional notations to ease presentation. Define
so that . Notice that and are identity matrices.
However, we have and , so the ranks are all identically . Also,
but it was just shown that the these spaces have the same dimensions, which equals , meaning
Using this observation, we can now state a proposition showing necessary and sufficient conditions for a tuple to be a critical point of .
A tuple is a critical point of if and only if and .
(If part) implies that , so , for . Similarly, implies , so for .
(Only if part) We have for all . This means that
Now recall that and are identity matrices, so and , which proves and . ∎
A critical point of is a global minimum point if and only if .
Since is a critical point, by Proposition 3.6 we have . Also note from the definitions of ’s and ’s that , so
Comparing this with (4), is a global minimum solution if and only if
This equation holds if and only if , meaning that they are projecting onto the same subspace. The projection matrix is onto , while is onto . From this, we conclude that is a global minimum point if and only if . ∎
From Proposition 3.7, we can define the set that appeared in Theorem 2.2, and conclude that every critical point of in is a global minimum, and any other critical points are saddle points.
Extension to deep nonlinear neural networks
In this section, we present some sufficient conditions for global optimality for deep nonlinear neural networks via a function space view. Given a smooth nonlinear function that maps input to output, Bartlett et al. (2017) described a method to decompose it into a number of smooth nonlinear functions where ’s are close to identity. Using Fréchet derivatives of the population risk with respect to each function , they showed that when all ’s are close to identity, any critical point of the population risk is a global minimum. One can see that these results are direct generalization of Theorems 2.1 and 2.2 of Hardt & Ma (2017) to nonlinear networks and utilize the classical “small gain” arguments often used in nonlinear analysis and control (Khalil, 1996; Zames, 1966). Motivated by this result, we extended Theorem 2.1 to deep nonlinear neural networks and obtained sufficient conditions for global optimality in function space.
where the constant denotes the variance that is independent of . Note that if almost surely, the first term in vanishes and the optimal value of is .
Define the function spaces as the following:
where are defined for all . Assume that , and that we are optimizing with . In other words, the functions in are differentiable and show sublinear growth starting from 0. Notice that , because a composition of differentiable functions is also differentiable, and a composition of sublinear functions is also sublinear. We also assume that for all , which is identical to the assumption in Theorem 2.1.
2 Sufficient conditions for global optimality
Here, we present two theorems which give sufficient conditions for a critical point ( for all ) in the function space to be a global optimum. The proofs are deferred to Appendix A.
Consider the case . If there exists such that
then any critical point of , in terms of , is a global minimum.
Consider the case . Assume that there exists some such that and . If there exist such that
is twice-differentiable,
then any critical point of , in terms of , is a global minimum.
Note that these theorems give sufficient conditions, whereas Theorems 2.1 and 2.2 provide necessary and sufficient conditions. So, if the sets we are describing in Theorems 4.1 and 4.2 do not contain any critical point, the claims would be vacuous. We ensure that there are critical points in the sets, by presenting the following proposition, whose proof is also deferred to Appendix A.
For each of Theorems 4.1 and 4.2, there exists at least one global minimum solution of satisfying the conditions of the theorem.
Theorems 4.1 and 4.2 state that in certain sets of , any critical point in function space a global minimum. However, this does not imply that any critical point for a fixed sigmoid or arctan network is a global minimum. As noted in (Bartlett et al., 2017), there is a downhill direction in function space at any suboptimal point, but this direction might be orthogonal to the function space represented by a fixed network, and may hence result in local minima in the parameter space of the fixed architecture.
Understanding the connection between the function space and parameter space of commonly used architectures is an open direction for future research, and we believe that these results can be good initial steps from the theoretical point of view. For example, we can see that one of the sufficient conditions for global optimality is the Jacobian matrix being full rank. Given that a nonlinear function can locally be linearly approximated using Jacobians, this connection is already interesting. An extension of the function space viewpoint to cover different architectures or design new architectures (that have “better” properties when viewed via the function space view) should also be possible and worth studying.
Acknowledgments
This research project was supported in parts by DARPA DSO’s Fundamental Limits of Learning program.
References
Appendix A Analysis of deep nonlinear networks
In this section, we introduce additional notation that is used in the proofs. To emphasize that the Fréchet derivative is a linear functional that outputs a real number, we will write in an inner-product form . This notation also helps avoiding confusion coming from multiple parentheses and square brackets.
A.2 Fréchet Derivatives
By definition of Fréchet derivatives, we have
where is the direction of perturbation and is the directional derivative along that direction . From the definition of ,
This equation (6) will be used in the proof of Theorems 4.1 and 4.2.
A.3 Proof of Theorem 4.1
From (6), consider . For any ,
Let . Since has full row rank by assumption, is invertible. Then define a particular direction
Moreover, if we decompose with SVD, , is of the form and
From this we can see that if we have a critical point of , then implies , which means that the critical point is a global minimum of .
A.4 Proof of Theorem 4.2
Recall that by assumption we have such that and . Consider , then for any ,
By the assumption that is invertible and ,
A.5 Proof of Proposition 4.3
(Theorem 4.2) It is given that we have such that and . Set , where the first components are and the rest are zero. All the rest of are set as in (7). Then, it can be easily checked that for all and all the conditions of the theorem are satisfied.
Appendix B Deferred Lemma
For this lemma, we separate the proof into two cases: and . The crux of the proof is to show that any critical point cannot be a local maximum. Then, any critical point is either a local minimum or a saddle point, so the conclusion of this lemma follows.
In case of , we use some of the results in Kawaguchi (2016) and examine the Hessian of with respect to , where denotes vectorization of matrix . Let be the partial derivative of with respect to in numerator layout. It was shown by Kawaguchi (2016, Lemma 4.3) that the Hessian matrix
where denotes the Kronecker product of two matrices. Notice that is positive semidefinite. Since is full rank, whenever there exists a strictly positive eigenvalue in , which means that there exists an increasing direction. So cannot be a local maximum.
The case where requires a bit more careful treatment. Note that this case corresponds to where we have degenerate critical points, which are in many cases much harder to handle.
For any arbitrary , we describe a procedure that perturbs the matrices by perturbations sampled from Frobenius norm balls of radius centered at , which we will denote as , . Let be the uniform distribution over the ball . The algorithm goes as the following:
Sample , and define .
If , stop and return .
First, recall that the set of rank-deficient matrices have Lebesgue measure zero, so for any sample , has full rank with probability 1. If we proceed the for loop until , we have a full-rank with probability 1, which means that the algorithm must return with probability 1. Notice that before and after the -th iteration, we have
This means that if we define , then . Also, notice that
and notice that they are all in the neighborhood of , that is, the Cartesian product of -radius balls centered at . Moreover, we have
from which we can see that at least one of or must hold. This shows that for any , there is a point in -neighborhood of with a strictly greater function value . This proves that cannot be a local maximum. ∎
Appendix C Deferred Proofs
In case of , we can decompose the loss function in the following way:
Let us take a close look into the last term in the RHS. Note that is the orthogonal projection of onto , so each row of must be in . Also,
It is right-multiplied with some matrix, so its columns must lie in . By the fact that ,
Now, (2) becomes a problem of minimizing subject to the rank constraint . The optimal solution for this is obtained when is the -rank approximation of . Then, -rank approximation of can be expressed as , where is unique due to our assumption that all singular values are distinct. Therefore,
is the unique global minimum solution of (2) when .
C.2 Proof of Lemma 3.2
C.3 Proof of Lemma 3.3
1. Since , . Then
2. Since , . Then