SGD learning on neural networks: leap complexity and saddle-to-saddle dynamics
Emmanuel Abbe, Enric Boix-Adsera, Theodor Misiakiewicz
Introduction
Deep learning has emerged as the standard approach to exploiting massive high-dimensional datasets. At the core of its success lies its capability to learn effective features with fairly blackbox architectures without suffering from the curse of dimensionality. To explain this success, two structural properties of data are commonly conjectured: (i) a low-dimensional structure that SGD-trained neural networks are able to adapt to; (ii) a hierarchical structure that neural networks can leverage with SGD training. In particular,
A line of work has investigated the sample complexity of learning with deep neural networks, decoupled from computational considerations. By directly considering global solutions of empirical risk minimization (ERM) problems over arbitrarily large neural networks and sparsity inducing norms, they showed that deep neural networks can overcome the curse of dimensionality on classes of functions with low-dimensional and hierarchical structures. However, this approach does not provide efficient algorithms: instead, a number of works have shown computational hardness of ERM problems and it is unclear how much this line of work can inform practical neural networks, which are trained using SGD and variants.
A line of work in computational learning theory has provided time- and sample-efficient algorithms for learning Boolean functions with low-dimensional structure, based on their sparse Fourier spectrum . However, these algorithms are a priori quite different from SGD-trained neural networks. While unconstrained architectures can emulate any efficient learning algorithms , it is unclear whether more ‘standard’ neural networks can succeed on these same classes of functions or whether they require additional structure that pertains to hierarchical properties.
Thus, an outstanding question emerges from the current state of affairs:
For neural networks satisfying “regularity assumptions” (e.g., fully-connected, isotropically initialized layers), are there structural properties of the data that govern the time complexity of SGD learning? How does SGD exploit these properties in its training dynamics?
Here the key points are: (i) the “regularity” assumption, which prohibits the use of unorthodox neural networks that can emulate generic PAC/SQ learning algorithms as in ; (ii) the requirement on the time complexity, which prohibits direct applications of infinite width, continuous time or infinite time analyses as in . We discuss in Section 1.3 the various works that made progress towards the above, in particular regarding single- and multi-index models. We now specify the setting of this paper.
We focus on the following class of data distributions. First of all, we consider IID inputs, i.e.,
and we focus on the case where is either or , although we expect that other distributions would admit a similar treatment. Incidentally, the latter distribution is of interest in reasoning tasks related to Boolean arithmetic or logic . We now make a key assumption on the target function, that of having a low latent dimension, i.e., where and
with the assumption that . In other words, the target function has a large ambient dimension but depends only on a finite number of latent coordinates. In the Gaussian case the coordinates are not known because of a possible rotation of the input, and in the Boolean case the coordinates are not known because of a possible permutation of the input.
Data with large ambient dimension but low latent dimension have long been a center of focus in machine learning and data science. It is known that kernel methods cannot exploit low latent dimension, i.e., it was proved in that any kernel method needs a number of features or samples satisfying
in order to learn a Boolean function as above with degree . In other words, for kernel methods controls the sample complexity irrespective of any potential additional structural properties of (e.g., hierarchical properties). On the other hand, it is known that this is not the limit for deep learning, which can break the curse, as discussed next.
Consider the following example: is drawn from the hypercube and the target function is -sparse, either
The first function is called a vanilla staircase of degree 4 . The second is a monomial of degree 4. Each of these functions induces a function class under the permutation of the variables (i.e., one can consider the class of all monomials on any 4 of the input variables, and similarly for staircases). One can verify that these function classes have similar approximation and statistical complexity because of the low-dimensional structure, but have different computational complexity because of the hierarchical structure. For example, under the Correlational Statistical Query (CSQ) model of computation , the first class has CSQ dimension versus for the second classSee Section 2 for more details on CSQ..
1 The leap complexity
We now define the leap complexity. Any function in can be expressed in the orthogonal basis of , i.e., the Hermite or Fourier-Walsh basis for and respectively,
In words, a function is leap- if its non-zero monomials can be ordered in a sequence such that each time a monomial is added, the support of grows by at most new coordinates, where each new coordinate is counted with multiplicity in the Gaussian case (and the 1-norm collapses to the cardinality of the difference set in the Boolean case). Note that the definition of leap- functions on the hypercube generalizes the definition of functions with the merged-staircase property (leap-1 functions) from .
2 Summary of our contributions
This paper puts forward a general conjecture characterizing the time complexity of SGD-learning on regular neural networks with isotropic data of low latent dimension. The key quantity that emerges to govern the complexity is the leap (Definition 1). This gives a formal measure of “hierarchy” in target functions, going beyond spectrum sparsity and emerging from the study of SGD-trained regular networks. The paper then proves a specialization of the conjecture to a representative class of functions on Gaussian inputs, but for 2-layer neural networks and with certain technical assumptions on how SGD is run. The two main innovations of the proof are (i) a full control of the time complexity of SGD learning on a fully-connected network (without infinite width or continuous time approximations); (ii) going beyond one-step gradient analyses and showing that the leap controls the entire learning trajectory due to a sequential learning mechanism (saddle-to-saddle). We also provide experimental evidence towards the more general conjecture with vanilla SGD and derive CSQ lower-bounds for noisy GD that match our achievability bounds.
We believe that the conjecture (in particular the time complexity scaling) holds for more general architectures than those with isotropically-initialized layers, as long as enough ‘regularity’ assumptions are present at initialization (prohibiting the type of ‘emulation networks’ used in ). Note that it is not enough to ask for only the first layer to be initialized with a rotationally-invariant distribution, as this may be handled by using emulation networks on the subsequent layers, but weaker invariances of subsequent layers (e.g., permutation subgroups) may suffice.
The characterization obtained in this paper implies a relatively simple picture for learning low-dimensional functions with SGD on neural networks:
SGD on regular neural networks implicitly implements a form of ‘adaptive curriculum’ learning. SGD first picks up low-level features that are computationally and statistically easier to learn, and by picking up these low level features, it makes the learning of higher-level features in turn easier. As mentioned in the examples of Table 1: learning takes sample complexity (leap- function). But if we add an intermediary monomial to our target to create , then it takes steps to learn (leap- function). If we have a full staircase, it only requires (leap- function). This thus gives an adaptive learning process that follows a curriculum learning procedure where features of increasing complexity guide the learning.
3 Related works
In parallel, several works have studied the dynamics of SGD in simpler non-convex models in high dimensions . Our analysis relies on a similar drift plus martingale decomposition of online-SGD as in . In particular, the leap complexity is related to the information-exponent introduced in . The latter considers a single-index model trained with online-SGD on a non-convex loss and the information exponent captures the scaling of the correlation between the model at a typical initialization and the global solution. showed that, with information exponent , online-SGD requires steps to converge, similarly to the scaling presented in this paper. However our analysis and the definition of the leap complexity differ from in two major ways. First, our model is not a single parameter model, so a much more involved analysis is required for the dynamics. Second, the information exponent is a coefficient that only applies at initialization, while the leap-complexity is a measure of targets that controls the entire learning trajectory (our neural networks visit several saddles during training).
Lower bounds on learning leap functions
Linear methods such as kernel methods suffer exponentially in the degree of the target function, and cannot use the “hierarchical” structure to learn faster. This was proved in for the Boolean case, and this work extends the result to the Gaussian case:
Let be a degree- polynomial over the Boolean hypercube (resp., Gaussian measure). Then there are , such that any linear method needs samples to learn to less than error, where is an unknown permutation (resp., rotation) as in (2).
Consider now the Correlational Statistical Query (CSQ) model of computation . A CSQ algorithm accesses the data via expectation queries, plus additive noise. We show that for CSQ methods the query complexity scales exponentially in the leap of the target function, which can be much less than the degree.
We note that the above lower-bounds are for CSQ models or noisy population-GD models, and not for online-SGD since the latter takes a single sample per time step. Our proof does show a correspondence between online-SGD and population-GD, but without the additional noise. It is however intriguing that the regularity in the network model for online-SGD appears to act comparatively in terms of constraints to a noisy population-GD model (on possibly non-regular architectures), and we leave potential investigations of such correspondences to future work (see also discussion in Appendix B.5). Further, we note that the correspondence to CSQ may not hold beyond the finite regime. First there is the ‘extremal case’ of learning the full parity function, which is efficiently learnable in CSQ (with 0 queries) but not necessarily with online-SGD on regular networks: shows it is efficiently learnable by a i.i.d. Rademacher initialization, but not necessarily by a Gaussian isotropic initialization. Further, the positive result of the Rademacher initialization may disappear under proper hyperparameter ‘stability’ assumptions. Beyond this extremal case, a more important nuance arises for large : the fitting of the function on the support may become costly for regular neural networks in certain cases. For example, let be a known function and consider learning which depends on unknown coordinates as . This is a leap-1 function where the linear part reveals the support and the permutation, and with a parity term on the indices such that . In this case, SGD on a regular network would first pick up the support, and then have to express a potentially large degree monomial on that support, which may be hard if is large (i.e., ). The latter part may be non trivial for SGD on a regular network, while, since is known, it would require 0 queries for a CSQ algorithm once the permutation was determined from learning the linear coefficients.
Learning leap functions with SGD on neural networks
We consider the following assumption on the activation function:This is satisfied, for example, by the shifted sigmoid for almost all shifts .
For the purposes of the analysis, we make two modifications to SGD training. First, we train layerwise: training and then , while keeping the biases frozen during the whole training. Second, during the training of the first layer weights , we project the weights in order to ensure that they remain bounded in magnitude. See Algorithm 1 for pseudocode, and see below for a detailed explanation. These modifications are not needed in practice for SGD to learn, as we demonstrate in our experiments in Figure 1 and Appendix A.
Analyzing layerwise training is a fairly standard tool in the theoretical literature to obtain rigorous analyses; it is used in a number of works, including . In our setting, layerwise training allows us to analyze the complicated dynamics of neural network training, but it also leads to a major issue. During the training of the first layer, the target function is not fully fitted because we do not train the second layer concurrently. Therefore the first-layer weights continue to evolve even after they pick up the support of . This is a challenge since we must train the first-layer weights for a large number of steps, and so they can potentially grow to a very large magnitude, leading to instability.We emphasize that this problem is due to layerwise training, since in practice if we train both layers at the same time the residual quickly goes to zero after the support is picked up, and so the first-layer weights stop evolving and remain bounded in magnitude (see Appendix A).
We correct the issue by projecting each neuron’s first-layer weights to ensure that the coordinates do not blow up. First, we keep the “small” coordinates of on the unit sphere, i.e., for some parameter , we define the “small” coordinates for neuron at time by and
We project these coordinates on the unit sphere using the operator defined by
and use the spherical gradient with respect to the sphere , i.e., for any function ,
In the second phase, the training of the second layer weights is by standard SGD (without projection) with added ridge-regularization term to encourage low-norm solutions.
2 Learning a single monomial
We first consider the case of learning a single monomial with Hermite exponents :
We assume (the case is straightforward). is a leap- function. We start by proving that, during the first phase, the first layer weights grow in the directions of which are the variables in the support of the target function.
Assume satisfy Assumption 1. Then for sufficiently small (depending on ) and the following holds. For any constant , there exist for , that only depend on and such that
and for large enough that , the following event holds with probability at least . For any neuron ,
Early stopping: for all and .
And for any neuron such that ,
On the support: \big{|}w_{j,i}^{\overline{T}_{1}}-\text{sign}(w_{j,i}^{0})\cdot\Delta\big{|}\leq C_{5}/\sqrt{d\log(d)} for .
Outside the support: for , and .
Theorem 1 shows that after the end of the first phase, the coordinates aligned with the support are all close to with the same signs as as long as has the same sign as at initialization. Furthermore, the correlation with the support only appears at the end of the dynamics, and does not appear if we stop early.
The proof of Theorem 1 follows a similar proof strategy as , namely a decomposition of the dynamics into a drift and martingale terms with information exponent . However, our problem is multi-index, and the analysis will require a tighter control of the different contributions to the dynamics as the dynamics move from saddle to saddle. An heuristic explanation for this result can be found in Appendix B.3. The complete proof of Theorem 1 is deferred to Appendix C.
Let or and assume satisfies Assumption 1. For any constants and , there exist for , that only depend on and such that taking width , bias initialization scale , and and second-layer initialization scale , and second-layer regularization , and , and , and
for we have with probability at least :
If we train the first layer weights for steps and for , then we cannot fit using the second-layer weights, i.e.,
This result suggests that the dynamics of SGD with one monomial can be decomposed into a ‘search phase’ (plateau in the learning curve) and a ‘fitting phase’ (rapid decrease of the loss) similarly to . SGD progressively aligns the first layer weights with the support, with little progress, and as soon as SGD picks up the support, the second layer weights can drive the risk quickly to . Because of the layer-wise training, we only show in Corollary 1.(b) that with early stopping on the training of the first layer weights, we cannot approximate the function at all using the second layer weights (hence, we cannot learn it even with infinite number of samples). The proof of Corollary 1 is in Appendix E.1.
3 Learning multiple monomials
We now consider with several monomials in its decomposition. In order to simplify the statement and the proofs, we will specifically consider the case of nested monomials
where and are positive integers. For , we denote , and the size of the biggest leap (such that is a leap- function), and the total degree of the polynomial . We will assume that (i.e., leap of size at least between monomials). This specific choice for allows for a more compact proof, similar to Theorem 1. However, the compositionality of is not a required structure for the sequential alignment to hold and we describe in Appendix D.2 how to modify the analysis for more generalHowever, our current proof techniques do not allow for fully general leap functions: e.g., has its two monomials pushing the ’s in two opposite directions. .
We first prove that the first-layer weights grow in the relevant directions during training.
On the support: \big{|}w_{j,i}^{\overline{T}_{1}}-\text{sign}(w_{j,i}^{0})\cdot\Delta\big{|}\leq C_{5}/\sqrt{d\log(d)} for .
Outside the support: for and .
The proof follows by showing the sequential alignment of the weights to the support: with high probability and for each neurons satisfying the sign condition at initialization, it takes between and steps to align with coordinates , after having picked up coordinates . The proof can be found in Appendix D.
While Theorem 2 captures the tight scaling in overall number of steps, it does not capture the number of steps for smaller leaps shown in Figure 1 in the case of increasing leaps. In Appendix D.2.1, we show that the scaling of steps to align to the next monomial can be obtained by varying the step size, in the case of increasing leaps. Note that in practice, neural networks with constant step size seem to achieve this optimal scaling for escaping each saddle (such as in Figure 1). Hence, there might be a mechanism in the SGD training that can implicitly control the martingale part of the dynamics, without rescaling the step sizes. However, understanding such a mechanism would require to study the joint training of both layers, which is currently out of reach of our proof techniques.
As in the single monomial case, we consider fitting the second layer weights only for a specific class of functions (where all monomials are multilinear):
We require extra assumptions on the activation function to prove that the fitting is possible. The following is an informal statement, and we leave the formal statement and proof to Appendix E.2.
Discussion
One direction for future work is to remove the modifications to vanilla SGD used in the analysis (layerwise training and the projection step). Another direction is to prove the conjecture by extending our analysis of the training dynamics to general functions, beyond those of the form (11). Another direction is to study extensions of the leap complexity measure beyond isotropic input distributions.
Acknowledgement
Part of this work was supported by the NSF-Simons Research Collaborations on the Mathematical and Scientific Foundations of Deep Learning (MoDL) Award and the EPFL PhD Exchange Fellowship. EB was also generously supported by Apple with an AI/ML fellowship. TM also acknowledges the NSF grant CCF-2006489 and the ONR grant N00014-18-1-2729.
References
Appendix A Additional numerical simulations
In Figures 2, 3, 4 and 5 we plot the risk versus number of samples for SGD training of 5-layer ResNets with fully-connected layers for various different target functions and for Boolean and Gaussian data. In these plots, the saddle-to-saddle dynamics are visible, which are caused by the neural network sequentially picking up the support using the hierarchical structure of the monomials in the function. In Figures 6 and 7, we study learning a leap-1 function (merged-staircase function), and we experiment with the effect of adding depth to see its effect on fitting. There is also an interesting edge-of-stability behavior during the “second-layer fitting” part, where the loss does not decrease monotonically . We leave understanding this to future work.
Appendix B Additional discussion from the main text
In addition to the references listed in the main text, we further review other relevant papers.
A line of work in computational learning theory studied the complexity of learning Boolean functions under the uniform input distribution. It was realized that functions with concentrated Fourier spectrum can be learned efficiently, both in sample and time complexity using the sparse Fourier algorithm . Namely, under knowledge of a set of basis elements such that for all , one can learn with error , sample complexity and polynomial time complexity if is polynomial using the sparse Fourier algorithm that estimates the coefficients in . Many interesting classes of functions fall under this setting, such as juntas, low degree functions, bounded-size or -depth decision trees . While has to be knownThe set knowledge can be relaxed under the query access model using the Kushilevitz-Mansour algorithm (based on the Goldreich-Levin algorithm) that uses a divide-and-conquer procedure to identify the coefficients to be estimated . under the random sample model, no degree constraints are imposed. In particular, the low-degree assumption (degree at most ) is just a special case that provides this knowledge (with order time complexity), monomials of degree or are equivalent in the eye of the sparse Fourier algorithm. This is not necessarily the case for SGD-trained neural networks.
A line of work has considered SGD learning on ‘unconstrained’ neural networks (besides polynomial size) and shows that we can emulate any efficient PAC or SQ algorithm . Such networks are far from the practical neural networks used in applications. Against this state of affairs, several works have attempted to derive computational lower bounds on learning with regular neural networks. For example, shows that for fully connected 2-layer networks, if the initial alignment (INAL) of a network with a Boolean target function (measured by the maximal expected correlation between target and neurons) is not significant, then noisy-GD cannot amplify the correlation to any significant level. This is achieved by showing that a low INAL implies a large minimal degree in the target function (thus a large leap) under some additional conditions. Another work uses the permutation, sign-flip, or rotational equivariance of noisy-GD training of fully-connected neural networks to show a lower bound on the number of gradient descent steps required for global convergence, when we have access to population gradients with an additive Gaussian noise. In particular, for leap- functions on the hypercube and the hypercube, steps are to shown to be required, where is the Gaussian noise variance. This roughly matches the conjecture in this paper in its exponential dependence on the leap – however, the computational model is different (noisy-GD versus online-SGD).
Finally let’s remark that a large body of work in the statistics and machine learning literature has studied the problem of learning multi-index models. These include for example phase retrieval , intersection of halfspaces and subspace juntas . We refer to and references therein for an overview of this line of work. In particular, it is well understood that in order to break the “curse of dimensionality”, the algorithm needs to estimate the low-dimensional support. In contrast with this line of work, we consider learning these multi-index functions with generic SGD on regular neural networks, with no a priori information on the target function. Surprisingly, we show that this generic algorithm can nearly match the computational complexity of the best CSQ algorithm. Note that specialized algorithms can achieve better sample and computational complexity: for example, showed an algorithm that can learn low-rank Gaussian polynomials in samples and runtime, regardless of the leap-complexity, by going beyond CSQ algorithms.
B.2 Discussion on the definition of the leap complexity
It was noted in that some “degenerate” leap-1 functions on the hypercube are not learned in SGD-steps. Take for example : by permutation symmetry on the support of , steps of SGD will learn first layer weights aligned with on the support . SGD will require many more steps to break this symmetryWe conjecture steps are required, see following discussion in the Gaussian case. and fit . circumvents this difficulty under a smoothed complexity analysis, and shows that the set of degenerate leap- functions has of Lebesgue-measure . Alternatively, a possible approach to learn these degenerate cases (for “axis-aligned” sparse functions) is to use different random learning rates for each coordinates and break the symmetry in learning.
B.3 Intuition for the proof of Theorem 1
In this section, we give some intuition behind the proof of Theorem 1. The complete proof can be found in Appendix C.
We first consider a simple SGD dynamics, with no projection step, and neglect the biases. We later discuss our choice of algorithm and how the analysis needs to be modified to control the projection step. The dynamics on the first layer weights is now simply given by
Recall that we initialize the second layer weights . By Assumption 1, we have
With high probability over a polynomial number of steps, with constant chosen sufficiently large. Hence,
and we can chose with sufficiently small, while keeping constant, so that we can neglect the interaction term between the different neurons and get:
Let us directly consider the correlation loss and track the dynamics of a unique neuron . We assume that , and . We further make the following heuristic simplification: we assume the dynamics is described by only two parameters
with SGD updates and , i.e.,
where and . We deduce that to leading term (assuming )
Let us now control the different contributions to the dynamics:
Martingale part: By Doob’s maximal inequality for martingales, we have with high probability
We choose so that we can neglect the martingale contribution during the entire dynamics by taking .
Drift part for : We now neglect the martingale term and write for all
We can study this sequence (see ) and show that
In order for , we need to take .
Drift part for : Again, by neglecting the martingale contribution and for ,
We can show that this sequence is bounded by
where we used Eq. (14) in the last inequality.
We deduce from Eq. (15) that for chosen such that , then . Hence, during the dynamics, the weights not aligned with the support of remain small, of order , while the weights aligned with the support of become of order . From the bounds in (i) and (ii), we need to choose and such that (martingale part) and (drift part), i.e., we can take
While the above heuristic derivation was useful to get intuitions, the assumption that the weights remain equal (or approximately equal) on and outside the support is not valid. Because of the statistical fluctuations over steps, different coordinates over different neurons will grow to be order 1 on the support at a stochastic time (with high probability between and for some large enough constant ). To prevent these coordinates to continue growing (because we neglected the interaction term in the dynamics, which could otherwise prevent this growth), we introduce the projection step
where is the projection step defined in Eq. (10), and we use the spherical gradient defined in Eq. (9). Note that because of the choice and the definition of the set on which we do the projection on the sphere, and commute.
Thanks to the spherical gradient, we can show that the projection steps only have a negligible impact on the dynamics (similarly to the analysis in ). By carefully arranging these additional terms, we can essentially recover the drift plus martingale analysis presented heuristically above.
B.4 Going beyond sparsity
In the regime the complexity scaling in is dominated by the ‘hard’ part of learning the low-dimensional latent space on which the function depends, and the complexity of fitting the function on the support is secondary and only results in constants. This also makes the conjecture fairly general in terms of architecture choices as long as there is enough expressivity to fit the function on the support. One could also consider functions that depend on a finite number of basis elements, without necessarily involving a finite number of coordinates. For instance the full parity function is such an example. For SQ algorithms, the class of monomials of degree 0 (more generally ) has equivalent complexity to the class of monomials of degree (more generally degree ), and the SQ-dimension is symmetrical for these dual cases. However for SGD learning on regular nets, this is not exactly the case. It is true that the full parity can be learned by regular nets under a specific setting; provides a regular 2-layer neural net that can learn the full parity if the weight measure of the first layer at initialization is i.i.d. Rademacher(1/2) and the activation is a ReLU. A constant number of step can also be sufficient in such cases, as for the 0-degree monomial. It is however conjectured that this is not achievable with a polynomial number of steps for weights that have a Gaussian initialization. Thus, for isotropic layers, it is possible that the full parity is not polytime learnable. This means that the generalized notion of leap to non-coordinate sparse may depend on more specific choices of the parameters. Further, in the non-isotropic case where the full parity is efficiently learnable, one may define the leap with basis sets that can either grow from the 0-monomial or descend from the full-monomial, with the mirror symmetry as for SQ algorithms.
Another notion to factor in when considering non-coordinate sparse function is the fitting of the function once the support is learned. First of all, there may be a non-polynomial number of coefficients to handle, although one can probably cover enough interesting cases with functions that are well-approximated by polynomially many coefficients . Further, there is the fitting of the function by the neural net that may now turn non-trivial. Consider even a function with few basis elements, , where is an arbitrary, but known function, and is large. SGD on a regular neural network would first pick up the coordinates in the support and then learn the monomial based on that support. The latter part may not be trivial for SGD on a regular net, while it would require 0 queries for an SQ algorithm (once the linear part is learned, the permutation is identified and the coefficients in front of each variable would allow us to calculate ). Thus the complexity of learning the second monomial on the detected support set is likely to factor in for such cases, and this is likely going to depend more on the model hyperparameters and architecture choice. In less contrived cases, the naive generalization of the leap applied verbatim to non-constant remains likely relevant.
B.5 Lower-bounds: beyond noisy GD
Note that the CSQ and noisy-GD models do not exactly match the SGD learning model; we do prove in this paper that the drift of the population gradient dominates the dynamic on the considered horizon, but the CSQ model also has noise added to the query outputs. It is nonetheless interesting that the regularity of the network model drives us to an achievability result that matches that of CSQ lower-bounds. Since it is known how to go beyond the CSQ/SQ lower-bounds with non-regular networks , e.g., learning dense parities by emulating matrix inversions with irregular networks, our results raise an intriguing question: may the model “regularity” act comparably to a CSQ constraint? We leave this to future work.
Appendix C Proof of Theorem 1: alignment with a single monomial
In this appendix, we prove the alignment of the first layer’s weights with the support of one monomial. The proof will follow from a similar proof strategy as in , namely decomposing the dynamics into drift and martingale terms. However, it will differ in a key aspect: while considers a single-index model, we will need to track for each neuron parameters (the first coordinates of ) and show that the other parameters remain well behaved along their whole trajectories, which requires a tighter control of the different contributions to the dynamics.
Recall that we denote by a constant that only depends on (Assumption 1) and the sub-Gaussianity of the label noise . Throughout the proofs, we will write for generic constants that only depend on and . The values of these constants are allowed to change from line to line or within the same line.
In the proof, we will consider to be small enough constants that can depend on and , but are independent of . We will track the dependency in when necessary, and otherwise use that they are bounded by (in particular, the constants in the proof will be independent of ). These constants will be fixed in Theorem 1.
We will show that we can take initialization scale of second layer weights and step size such that the dynamics of the first layer training can be approximated by a correlation dynamics, with no interactions between the neurons, so that we can analyze each neuron independently. We consider below an arbitrary neuron for . In the case that we prove that the event claimed in Theorem 1.(b) and (c) holds with probability at least for neuron . Theorem 1.(a) will follow from a similar analysis. The result for all neurons follows by a union bound.
We further consider small enough such that for (see comments below Lemma 2). Hence, the biases will not impact the training of the first layer weights and for the simplicity, we will fix in the proof.
Without loss of generality, we assume that all of the first-layer coordinates of neuron have positive sign at initialization (and therefore by our choice of ). To see why, define and consider instead initializing the network at where and for all . Then consider training the network with samples where and . The distribution of data is the same as that of , and the the training dynamics match those of up to sign flips, and .
with .
Let us introduce the following stopping times on the dynamics:
Note that \{\tau=t\}\in{\mathcal{F}}_{t}:=\sigma\big{(}{\bm{\Theta}}^{0},\{{\bm{x}}^{s},y^{s}\}_{s\leq t}\big{)} for and . For and , we have , and . We further define for all ,
where is a constant that will be chosen large enough. In particular, at time , the -th coordinate is removed from the set on which we do the projection, i.e., . We will show in the proof that with high probability.
By concentration of polynomials of Gaussian variables, we have:
Assume that . Then for any , there exists large enough that only depends on and , such that for ,
For , we must have . Using the bounds (53) in Lemma 5 and a union bound, there exists a constant such that
Note that for , we have , and therefore . Let us introduce the truncated spherical gradient defined by
where is a multiplicative factor that models the projection step ,
It is easy to check that and for all . With these notations, our dynamics are now simply given by
For and ,
We recall the following useful identities (where )
In particular, by integration by parts, we have
which gives Eq. (19) by using . Eqs. (20) and (21) are obtained similarly. ∎
From Assumption 1, we can choose small enough and depending only on and such that that for all and
We further assume that is chosen small enough such that and . With this choice of , there exist constants that only depend on such that for all , if ,
C.2 Bounding the different contributions to the dynamics
The following lemma tracks the contribution of the projection on the sphere :
Assume that , and . Then there exist constants (that only depend on and ) such that for and all , if ,
First consider the case . We have . Note that on , we have and therefore and . We therefore have
where we used that and by definition of the spherical gradient. Furthermore, for . Therefore, there exists a constant such that bound (25) holds.
In the case , we have for and the coordinates that are removed at time satisfy . Hence
We can then use Eq. (27) and that to derive Eq. (26). ∎
Let us decompose the different contributions to the dynamics. We define the martingale updates. Let us bound the change of a coordinate after one update. For , if , then by Eq. (25), we have for
(Note that for , we have .) For , we have and because . Hence, we can rearrange Eqs. (28) and obtain
On the other hand, if , by Eq. (26), we have for ,
Rearranging these equations, we obtain for and (using ),
On the other hand, if , then
Define , and
Note , so that and are still martingale updates. By induction on Eqs (29) and (30), we deduce that for ,
Let us introduce the following quantities:
The term , which is the sum of population gradients, plays the role of a drift term, while the term is a martingale and corresponds to the comparison between the stochastic and the population gradients.
We can choose a constant large enough, depending only on , such that for
In particular, this implies that for any and constant sufficiently small,
Similarly, we can choose a constant large enough, depending only on , such that for
Hence for and satisfying Eqs (32) and (34), we get the following bounds on the trajectory for : for or , ,
while for and ,
We prove the following bounds on the martingale part:
Assume and are chosen as in Lemma 3. Fix and . There exists a constant that only depends on , , and , such that if we choose
then with probability at least , we have
We will show the theorem for . The result for will follow from an union bound on all , which are also martingales (the proofs for and will follow by the same argument).
Denote . We will use a truncation argument. For some , define for all and ,
For , we have and we can use Lemma 5 to choose that only depends on such that
and for all and
Hence with probability at least ,
Let us now apply Doob’s maximal inequality on : the increments are bounded by , hence we have
Choosing and as in Eq. (37), as well as a union bound, yields the result. ∎
C.3 Proof of Theorem 1
We consider the dynamics (18) up to time where . We assume and satisfy conditions (32), (34) and (37). In particular, with probability at least , the dynamics of satisfies the bounds in Eqs (35) and (36), with satisfying the bounds (38). In the rest of the proof, we show that on this high probability event, we can choose and such that and Theorem 1.(a), (b) and (c) are satisfied.
Step 1: Controlling the coordinates at the end of the dynamics.
Let us first show that as soon as , then stays close to . Note that we can choose constant large enough independent of , such that if , then by (24)
Hence, for any , consider (in particular, ). From Eq. (36) and by Lemma 4, we have
where we used that for and therefore by Eq. (39), we have . We deduce that
Similarly, we show that for any , we have . Indeed, for any , we have
where we used that by definition of , and for all . We deduce that
Step 2: Bounding the growth of for .
Define (i.e., the minimum of that have ). Note that for by Eq. (41) and therefore . By Eq. (33) and Lemma 2, we have for ,
Combining this lower bound with Eq. (35) and the bound on the martingale in Lemma 4, we get that for all
Similarly, consider for all . By Eq. (40), we have for . Hence by Eq. (36), we get
and we deduce that if then
Combining the above bounds, we deduce that
On the other hand, consider and let us lower bound the time . For , we have and . Hence by Eq. (35),
and therefore by Lemma 6, we get for ,
Step 3: Bounding the coordinates .
From Eq. (35) and Lemma 2, we have for all and ,
Consider such that . Then by Eq. (35), we have, for any , that
Using that for and for in Eq. (48), we get that, for any ,
We deduce that for all and and
and therefore taking sufficiently small, .
Choose and that satisfy Eqs. (32), (34), (37) and (43). We have with probability at least by Lemma 1 and Lemma 4 that . And Eqs. (50) and (42) imply that . Furthermore, by Eq. (43), we have which implies Theorem 1.(b) by Eq. (40). Theorem 1.(c) follows from Eq. (50).
Step 5: Upper bound for all neurons with early stopping.
Theorem 1.(a) follows from Eqs. (45) and (47) for neurons with initialization satisfying
For neurons that do not satisfy this condition, the analysis in Section C.2 still holds and we get bounds on the dynamics similar to the ones in Eq. (35), with the difference that , and therefore the drift has a negative contribution to the dynamics, and are now defined on all coordinates instead of only .
We can upper bound the drift contribution using that satisfy for
and therefore, taking the same bounds (44) and (46), we get for for a constant sufficiently large that
Furthermore, denoting , we have
Using the analysis of Lemma 6, we get that the drift has the same upper bound as in Eq. (52),
for and therefore
The bound on coordinates follows similarly to step 3. In particular, we deduce that for and , we have for all
and therefore , which concludes the proof of Theorem 1.(a).
C.4 Technical lemmas
Assume that . Then there exist constants that only depend on and such that
Recall that |g_{i}^{t}/a^{0}|\leq|\gamma_{i}^{t}|\big{(}|yx_{i}\sigma^{\prime}(\langle{\bm{w}},{\bm{x}}\rangle)|+|yw_{i}\langle{\bm{w}},{\bm{x}}\rangle\sigma^{\prime}(\langle{\bm{w}},{\bm{x}}\rangle)|\big{)}. Conditioning on and assuming that , we obtain
Furthermore, again assuming that , we have for any ,
The following lemma provides simple upper and lower bounds on sequences satisfying some geometric bound on their evolution. The upper bound can be seen as a discrete version of Bihari–LaSalle inequality. This upper bound was proven in [6, Appendix C], and we modify their proof to obtain a lower bound.
Note that by induction, we have for any where
For , it is straightforward to get and .
For , we consider the upper bound on . First, notice that
Hence, rearranging the terms, we get for any ,
Hence, as long as , we get
Appendix D Proof of Theorem 2: sequential alignment to the support
In this appendix, we consider the sequential alignment to the support in Section 3.3. The proofs will follow from a similar argument as in the single monomial case. However, the dynamics will be now split in phases corresponding to the alignment to each of the monomials.
Recall that throughout the proofs, we will denote for simplicity generic constants that only depend on and (note that all the other constants ). The values of these constants are allowed to change from line to line or within the same line.
We will use notations and results from Appendix C and outline the main difference with the proof of Theorem 1. We can again reduce the problem to tracking one neuron, and we assume without loss of generality that and for all .
Let us introduce the following new stopping times on the dynamics: for ,
The population gradients for are now given by:
Denote for . For and , the population gradient is given by: if (i.e., ),
while if (i.e., )
For and ,
The proof follows from Lemma 2 applied to a sum of monomials. ∎
Again, by Assumption 1, we can choose small enough and depending only on and such that Eqs (22) and (23) are satisfied (with replaced by ). We can further chose small enough (only depending on and such that there exists constants such that for all , if , then if ,
and if for ,
while for and ,
and for .
We consider the dynamics up to time where . We again assume that and satisfy conditions (32), (34) and (37), so that the dynamics of satisfies the bounds in Eqs (35) and (36), with satisfying the bounds (38), with probability at least . The following steps will follow closely the proof of Theorem 1.
Step 1: Controlling the coordinates during the first phases.
Note that for , during the phase, we have for ,
Assume that and
where we used the assumption that . Using the same argument as in Step 3 of Section C.3, we deduce that as long as , then
In particular, we deduce that we must have and therefore for all .
By induction, we deduce that for all , and
Step 2: Bounding the growth of for .
The same argument as in Step 1 of Section C.3 (recalling that by the previous argument, for all ) yields
Denote (noting that for but ). Furthermore, for and , we have . Hence, for ,
Furthermore, by the previous step. We deduce by Lemma 6: for ,
Similarly, we obtain similar bounds on (see Step 2 of Section C.3).
Theorem 2.(a) follows by Step 2 and taking and that satisfy (32), (34) and (37), and the growth conditions in Step 2. Theorem 2.(b) follows by the same argument as in Step 3 of Section C.3. ∎
D.2 Extending the analysis: adaptive step size and non-nested monomials
with increasingThe case in the first phase of the dynamics can be studied easily by modifying the proof of Theorem 2 and noting that the drift is now just a sum of constant terms. leaps , so that neurons align with the support sequentially at increasing time scales. As mentioned below Theorem 2, the time complexity to escape each of these leaps is only tight for the biggest leap if we take a constant step size . Indeed, for the first phases of the dynamics, SGD requires a number of steps much smaller than to align to the -th monomial. In that case, we can take bigger step sizes and still have negligible contribution from the martingale part of the dynamics. In practice, such as in Figure 1, we can see a saddle-to-saddle dynamicsAgain, we expect this saddle-to-saddle dynamic to occur in the case of increasing leaps, otherwise we might have mixing of the different phases for different neurons and no plateaus, except at the biggest leap. to occur, with a number of steps to escape each saddle even for constant step size.
To prove these tight scalings for each plateau with constant step size, we would need to study the joint training of the two layers, which is currently out of reach of our proof techniques. Instead, we show in the next theorem that we can use a learning rate schedule , i.e.,
to get a scaling to align to each new monomial.
the following events hold with probability at least . For any neuron ,
Early stopping for : for all and .
For any neuron such that for all ,
On the support: \big{|}w_{j,i}^{T_{l}}-\text{sign}(w_{j,i}^{0})\cdot\Delta\big{|}\leq C_{4}/\sqrt{d\log(d)} for and .
Outside the support: for and . Furthermore, .
There are two key differences between Theorem 3 and Theorem 2. First we prove a tighter scaling of number of steps for the first phases of the training. Second we show that the alignment is sequential for all the neurons at the same time: at the end of each phase, we exactly picked up the support and nothing else. In particular, using a similar proof as in Corollary 1.(b), we can show that the neural network at time cannot fit the remaining monomials at all using the second layer weights. This agrees with the picture obtained in the numerical simulation in Figure 1.
The and are chosen such that the martingale term remain negligible during the whole dynamics. Furthermore, because of the separation of time scales between the different phases of the dynamics, we can show that for and step size , the contribution of the drift terms coming from the next monomials remains small. The proof follows almost identically to the proofs of Theorems 1 and 2.
D.2.2 Non-nested monomials
Below we describe how we can modify the proof of Theorem 2 for non-compositional and leave the task of proving Conjecture 1 for general leap functions to future works.
and for (each new coordinates appear in the next monomial) and with , and denote (with ). Denote which corresponds to the leap complexity of .
First note that the same formulas as in Lemma 7 hold with
however, we cannot simplify the gradient to be of order during the -th phase. Below, we outline how to modify the proof of Theorem 1 in Section D.1 to the case (61). The bounds on the martingale terms and on the dynamics from Section C.2 still hold in that case, with the difference being in the formulas of the population gradients.
By taking small enough, there exists constants such that we can upper and lower bound as follows. For , during the phase, we have for ,
We can plug these population gradients in steps 1 and 2 in Theorem 2, and control the contribution of each of these terms using and , with similar arguments as in step 2 of Section C.3.
Appendix E Fitting the second layer weights: proof of Corollaries 1 and 2
We first focus on the case and prove parts (a) and (b) separately in Sections E.1.1 and E.1.2. The case of follows from a similar argument and we outline the differences in Section E.1.3.
Recall that in this case and we can use both interchangeably. We consider the case of no biases in this part, i.e., fixing .
By Theorem 1, with probability at least , for each neuron satisfying at initialization, we get at the end of the dynamics:
For the remainder of the proof, assume the above event is true.
There exists a constant that depends only on such that the following is true. For any weights which coincide on the first coordinates where , and with biases such that , and with , there exists a constant that only depends on such that
First if we replace by and by , the error is bounded by
which is accounted for by the first two terms since we can take large enough depending on . For the last two error terms,
If satisfies , then is distributed as . So for all ,
which concludes the proof of the lemma. ∎
with a constant in the that depends only on . So if we define the coefficient
then we can approximate as follows for any such that ,
where we use that for any , we have .
Putting this together with Lemma 8 and the guarantees on the first layer weights after training (62) and (63), we obtain the following lemma.
There exists a constant depending only on such that with probability at least there exists a set of weights satisfying
(First layer weights are the trained weights) For all , we have .
(Second-layer weights are small) We have .
(Squared error is small) We have .
Consider the event that for each the set is of size . This holds with probability at least by a union bound and a Hoeffding bound, so we condition on it from now on. Consider the event that for all we have
and note that this holds with probability at least by a Hoeffding bound for a constant depending on , so we also condition on it.
Let be given by if , and 0 otherwise. From this it follows that
for a constant depending only on . By Lemma 8,
by taking a small enough choice of parameters and large enough . ∎
Now that we have constructed the certificate , we show that SGD on the second layer converges quickly to a solution with low population loss by a bias-variance analysis of SGD for ridge-regularized least-squares linear regression in Lemma 12. We train the second-layer while keeping the weights of the first layer fixed, which corresponds to linear regression with input embedding
So, plugging in Lemma 9 and taking ,
By taking and , for small enough ,
E.1.2 Converse if early stopping
We now prove the converse. The proof will follow very similarly to the proof of [34, Theorem 1]. By Theorem 1, if we train the first layer for time steps for a large enough , then with probability at least for each neuron ,
and some constant . In particular, this implies that for large enough ,
For ease of notations, denote . Let us introduce and .
Corollary 1 will follow by showing that there exist constants that only depend on such that with high probability, we have
These are proved in the following two lemmas.
Under the same setting as in Corollary 1, there exist constants such that with probability at least ,
Consider the event described in Eq. (64). By rotational invariance of the distribution of , the entries are given by
where , and .
We can do a Taylor expansion and bound the second term
where we used that by Eq. (64).
Hence, for , we have . Note that . By standard concentration, using that , there exists constants such that with probability at least , we have
Using the same computation as above, we can replace and by while only incurring an error , and show that
From the above bounds, we deduce (using ) that with high probability
For not constant, and using that , we deduce that
Under the same setting as in Corollary 1, there exists constants such that with probability at least ,
First note that for any , the correlation of with is bounded by
Indeed, as in the proof of Lemma 2, we use the formula from integration by parts:
We conclude by noting that on the high probability event (64), we have . ∎
E.1.3 Proof for a single-index Hermite monomial
Let’s now consider . In this case, we consider the biases , where is chosen sufficiently small as discussed in Theorem 1. We can use the same proof strategy as in Section E.1.1 and construct good features
for any , by considering neurons with initializations with and , and (by an easy modification of Lemma 8). We will take sufficiently many neurons (but still independent of ) so that we have a sufficiently large for any intervals of size for with high probability.
Let us now construct a certificate for based on these good features. By a Taylor approximation, for any and ,
In particular, we can rescale and sum these coefficients such that for some that has second moment bounded by ,
We can now construct a certificate by sampling from the signed measure , and for each constructing an approximate good feature, as described in Lemma 8. The proof for the low test error then follows from applying the bound on the least squares linear regression of Lemma 12.
For the lower bound with early stopping, we use that
and we can conclude using the same argument as in Section E.1.2.
E.2 Proof of Corollary 2: sequential learning of monomials
Let us formally state Corollary 2 and prove it.
we have for large enough , that with probability at least at the end of the dynamics,
In contrast to the proof of Corollary 1, we only prove this result for “diverse” enough activation functions. For the proof, we will construct a specific activation function that have this “diversity” property. This activation depends on (or upper bound on ), but otherwise is independent of . The idea is that we will use biases of different magnitudes, which will change the signs of the Hermite coefficients of the activation, in order to ensure enough neurodiversity to learn the sum of increasing monomials. This is required due to the specific choice of training of the first layer weights considered in this paper. However, we show in simulations that standard ReLus activations are enough to learn these functions.
for all . This can be achieved as follows. Let be a constant that we will take large enough. Then for any , define the “truncated Hermite function”
And we show that is invertible when viewed as a matrix. For large enough depending on , the diagonal elements are lower-bounded by a constant:
And the off-diagonal elements are small. When , for large enough we have
And similarly when but , for large enough we have
So if we take large enough the system of equations defined by is invertible, so coefficients exist such that satisfies (65).
and where has second moment bounded by . Since we can estimate to error for each , we can approximate via a linear combination
We conclude analogously to the proof of Corollary 1, using the bounded-norm certificate to obtain a generalization guarantee.
E.3 Technical result: last iterate convergence of SGD on linear models
We analyze of the last iterate for online-SGD on a linear model with ridge-regularized least-squares loss by using the well-known bias-variance decomposition . A very similar analysis also appears in the appendix of ; the key difference is that we analyze online gradient descent with one sample per iteration (as opposed to online minibatch gradient descent) with a small learning rate in order to match the setting of the theorem. Compare also to which gives final-iterate bounds for the final risk, but these hold in expectation instead of with exponentially high probability.
For a parameter , the ridge-regularized square loss is
Each iteration of the dynamics of online-SGD on the ridge-regulariezd square loss is is given by
Let be the minimizer of , which is unique by strict convexity when . We prove the following convergence to the optimum. For any iteration , define the gap to optimality
by the first-order optimality condition . So
It remains to bound . We write the evolution of as:
Inductively, one obtains the well-known “bias-variance” decomposition
To bound the variance term, define the norm squared of the variance term:
The lemma follows by plugging in the expression for and using that is optimal, so . ∎
Construct . Then is a super-martingale:
So by the Azuma-Hoeffding inequality, since ,
Appendix F Lower bounds for linear methods and CSQ methods
and estimates the target function using the linear prediction model
The takeaway of this section is that to learn any degree- functions with small support on isotropic data, linear methods must pay at least samples (and “width” ) when the support is not known. This is proved by in the case of the binary hypercube:
Let be the degree of . Consider the class of functions which depend as on some subset of coordinates
For any linear method, let be the function estimated by the linear method on (possibly noisy) samples . Then there are constants such that
We now give an analogous result for the Gaussian data distribution, where the degree also drives the complexity for linear methods. This bound is new and was not derived in .
Let be the degree of . Consider the class of functions which depend as on some subspace of coordinates
First, we we can write a degree- monomial as a linear combination of functions in .
There are semiorthogonal matrices and coefficients such that
Furthermore, for all we have , which is a constant depending only on .
Notice that , so this is a valid semi-orthogonal matrix, and so . Now let us show that we can write the monomial as a linear combination of functions of the form . Specifically, for any with we have
with a nonzero proportionality constant that only depends on . Therefore,
with a nonzero proportionality constant that only depends on . This proves the claim. ∎
We will use this claim to lower-bound the error of the linear method on . Notice that the linear method must predict , where . So the error is lower-bounded by the norm of the orthogonal projection to this subspace. For throughout,
Putting together the equations proves the lemma.
F.2 Correlational Statistical Query (CSQ) methods
First, we give a lower bound on the CSQ complexity of learning a function with leaps when is drawn uniformly from the hypercube. The below lower bound is qualitatively similar to the argument in based on the “alignment” quantity. The bounds of have tighter constants in the exponents of the bound, but they have the disadvantage that they apply only to noisy population gradient descent instead of to general CSQ algorithms.
Suppose that the CSQ algorithm knows , which can only help it. Then the problem of learning from CSQ queries is equivalent to the problem of learning from CSQ queries. However, for random permutations conditioned on we have
So by a union bound, with probability all first CSQ queries can return 0. The final output of the algorithm can also be viewed as a statistical query. So with probability at least ,
The proposition follows by letting be a small enough positive constant depending on . ∎
be its isotropic leap (as defined in Appendix B.2). Consider the class of functions which given by applying on some subspace of coordinates