On Feature Learning in Neural Networks with Global Convergence Guarantees

Zhengdao Chen, Eric Vanden-Eijnden, Joan Bruna

Introduction

The training of neural networks (NNs) is typically a non-convex optimization problem, but remarkably, simple algorithms like gradient descent (GD) or its variants can usually succeed in finding solutions with low training losses. To understand this phenomenon, a promising idea is to focus on NNs with large widths (a.k.a. under over-parameterization), for which we can derive infinite-width limits under suitable ways to scale the parameters by the widths. For example, under a “1/width1/\sqrt{\text{width}}” scaling of the weights, the GD dynamics of wide NNs can be approximated by the linearized dynamics around initialization, and as the widths tend to infinity, we obtain the Neural Tangent Kernel (NTK) limit of NNs, where the solution obtained by GD coincides with a kernel method . Importantly, theoretical guarantees for optimization and generalization can be obtained for wide NNs under this scaling . Nonetheless, it was pointed out that this NTK analysis replies on a form of lazy training that excludes the learning of features or representations , which is a crucial ingredient to the success of deep learning, and is therefore not adequate for explaining the success of NNs .

Meanwhile, for shallow (i.e., one-hidden-layer) NNs, if we choose a “1 / width” scaling, we can derive an alternative mean-field (MF) limit as the widths tend to infinity. Under this scaling, feature learning occurs even in the infinite-width limit, and the training dynamics can be described by the Wasserstein gradient flow of a probability measure on the space of the parameters, which converges to a global minimizer of the loss function under certain conditions . Generalization guarantees have also been proved for learning with shallow NNs under the MF scaling by identifying a corresponding function space . However, currently there are three limitations to this model of over-parameterized NNs. First, the global convergence guarantees for shallow NNs only hold in the infinite-width limit (i.e. they are asymptotic). While studies the deviation between finite-width NNs and their infinite-width limits during training, the analysis is done only asymptotically to the next order in width. Second, a convergence rate has yet to be established except under special assumptions or with modifications to the GD algorithm . Third, while several works have proposed to extend the MF formulation to deep (i.e., multi-layer) NNs , there is less concensus on what the right model should be than for the shallow case. In summary, we still lack a model for the GD optimization of shallow and multi-layer NNs that goes beyond lazy training while admitting fast global convergence.

In this work, we study the optimization of both shallow NNs under the MF scaling and a type of partially-trained multi-layer NNs, and obtain theoretical guarantees of linear-rate global convergence.

We consider the scenario of training NN models to fit a training set of nn data points in dimension dd, where the model parameters are optimized by gradient flow (GF, which is the continuous-time limit of GD) with respect to the squared loss. Allowing most choices of the activation function, we prove that:

For a shallow NN, if the hidden layer is sufficiently wide and the input data are linearly independent (requiring n≤dn\leq d), then with high probability, the training loss converges to zero at a linear rate.

For a multi-layer NN where we only train the second-to-last layer, if the hidden layers are both sufficiently wide, then with high probability, the training loss converges to zero at a linear rate. Unlike for shallow NNs, here we no longer need the requirement on input dimension, demonstrating a benefit of jointly having depth and width.

We also run numerical experiments to demonstrate that our model exhibits feature learning and can achieve better generalization performance than its NTK counterpart.

2 Related works

Many recent works have studied the optimization landscape of NNs and the benefits of over-parameterization . One influential idea is the Neural Tangent Kernel (NTK) , which characterizes the behavior of GD on the infinite-width limit of NNs under a particular scaling of the parameters (e.g. for shallow NNs, replacing 1/m1/m with 1/m1/\sqrt{m} in (1)). In particular, when the network width is polynomially large in the size of the training set, the training loss converges to a global minimum at a linear rate under GD . Nonetheless, in the NTK limit, due to a relatively large scaling of the parameters at initialization, the hidden-layer features do not move significantly . For this reason, the NTK scaling has been called the lazy-training regime, as opposed to a feature-learning or rich regime . Several works have investigated the differences between the two regimes both in theory and in practice . In addition, several works have generalized the NTK analysis by considering higher-order Taylor approximations of the GD dynamics or finite-width corrections to the NTK .

An alternative path has been taken to study shallow NNs in the mean-field scaling (as in (1)), where the infinite-width limit is analogous to the thermodynamic or hydrodynamic limit of interacting particle systems . Thanks to the interchangeability of the parameters, the neural network is equivalently characterized by a probability measure on the space of its parameters, and the training can then be described by a Wasserstein gradient flow followed by this probability measure, which, in the infinite-width limit, converges to global mimima under mild conditions. Regarding convergence rate, ref. proves that if we train a shallow NN to fit a Lipschitz target function under population loss, the convergence rate cannot beat the curse of dimensionality. In contrast, we will study the setting of empirical risk minimization, where there are finitely many training data. Ref. shows that mean field Langevin dynamics on shallow NNs can converge exponentially to global minimizers in over-regularized scenarios, but we focus on GF without entropic regularization. Besides the question of optimization, shallow NNs under this scaling represent functions in the Barron space or variation-norm function space , which provide theoretical guarantees on generalization as well as fluctuation in training . Several works have proposed different mean-field limits of wide multi-layer NNs and proved convergence guarantees , but questions remain. First, due to the presence of different symmetries in a multi-layer network compared to a shallow network , the limiting object at the infinite-width limit is often quite complicated. Second, it has been pointed out that under the MF scaling of a multi-layer network, an i.i.d. initialization of the weights would lead to a collapse of the diversity of neurons in the middle layers, diminishing the effect of having large widths . In addition, while another line of work develops MF models of residual models , we are interested in multi-layer NN models with a large width in every layer.

Ref. demonstrates the importance of hierarchical learning by proving the existence of concept classes that can be learned efficiently by a deep NN with quadratic activations but not by non-hierarchical models. Ref. studies the optimization landscape and generalization properties of a hierarchical model that is similar to ours in spirit, where an untrained embedding of the input is passed into a trainable shallow model, and prove an improvement in sample complexity in learning polynomials by having neural network outputs as the embedding. However, the trainable models they consider are not shallow NNs but their linearized and quadratic-Taylor approximations, and furthermore the convergence rate of the training is not known. Ref. proposes a novel parameterization under which there exists an infinite-width limit of deep NNs that exhibits feature learning, but properties of its training dynamics is not well-understood. Our multi-layer NN models adopt an equivalent scaling (see Appendix C), and our focus is on proving non-asymptotic convergence guarantees for its partial training under GF.

Problem setup

Thus, we obtain a 33-layer feed-forward NN whose first-layer weights are random and fixed, and we call it a partially-trained 33-layer (P-33L) NN. Note that the scaling in this model is different from both the NTK scaling (1/m1/\sqrt{m} instead of 1/m1/m in (4)) and the MF scaling for multi-layer NNs adopted in (1/D1/{D} instead of 1/D1/\sqrt{D} in (5)). We show in Appendix B that when σ\sigma is homogeneous, this scaling is consistent with the Xavier initialization of neural network parameters up to a reparameterization . We also show in Appendix C that in certain cases this scaling is equivalent to the maximum-update parameterization proposed in . Numerical experiments that compare different scalings are described in Section 4.

2 Training with gradient flow

πc=12δc^(dc)+12δ−c^(dc)\pi_{c}=\frac{1}{2}\delta_{\hat{c}}(dc)+\frac{1}{2}\delta_{-\hat{c}}(dc) for some c^>0\hat{c}>0 independent from mm, which is the law of a scaled Rademacher random variable.

If σ\sigma is Lipschitz, it is differentiable almost everywhere, and we write σ′(x)\sigma^{\prime}(\bm{x}) to denote the derivative of σ\sigma when it is differentiable at x\bm{x} and otherwise. When σ\sigma is differentiable at hi(x)h_{i}(\bm{x}), there is

and the gradient of the loss function with respect to WijW_{ij} is given by

Thus, we can perform GD updates on WW according to the following rule: ∀i∈[m]\forall i\in[m] and ∀j∈[D]\forall j\in[D],

where ftf^{t} denotes the output function and h1t,...,hmth_{1}^{t},...,h_{m}^{t} denote the hidden-layer feature maps determined by the parameters at time tt. Then, induced by the evolution of WtW^{t}, each hith_{i}^{t} evolves according to

Accordingly, the output function ftf^{t} satisfies

Thus, the loss function Lt:=L[ft]\mathcal{L}^{t}:=\mathcal{L}[f^{t}] evolves according to

Compared to the NTK scaling of neural networks, the crucial difference is the 1/m1/m factor in (2), instead of 1/m1/{\sqrt{m}}. It is known that under the NTK scaling, due to the 1/m1/{\sqrt{m}} factor, the movement of the feature maps, h1,...,hmh_{1},...,h_{m}, is only of order O(1/m)O({1}/{\sqrt{m}}) while the function value changes by an amount of order Ω(1)\Omega(1). While this greatly simplifies the convergence analysis, it also implies that the hidden-layer representations are not being learned. In contrast, with the 1/m1/m factor in (2), if σ\sigma is Lipschitz with Lipschitz constant LσL_{\sigma}, there is ∣ft2(x)−ft1(x)∣≤Lσc^m∑i=1m∣hit2(x)−hit1(x)∣|f^{t_{2}}(\bm{x})-f^{t_{1}}(\bm{x})|\leq\tfrac{L_{\sigma}\hat{c}}{m}\sum_{i=1}^{m}|h_{i}^{t_{2}}(\bm{x})-h_{i}^{t_{1}}(\bm{x})|, ∀t1,t2≥0\forall t_{1},t_{2}\geq 0. Therefore, regardless of mm and DD,

which implies that the average movement of the feature maps is on the same order as the change in function value, and thus the hidden-layer representations as well as the NTK undergoes nontrivial movement during training. In Appendix C, we further justify the occurrence of feature learning using the framework developed in .

Convergence analysis

To prove that the training loss converges to zero, we need a lower bound on the absolute value of L˙t\dot{\mathcal{L}}_{t}. Indeed, if GG is positive definite, which depends on Φ\Phi and the training data, we can establish one in the following way. First, as a simple case, if we use an activation function whose derivative’s absolute value is uniformly bounded from below by a constant Kσ′>0K_{\sigma^{\prime}}>0, such as linear, cubic or (smoothed) Leaky ReLU activations, we can derive a Polyak-Lojasiewicz (PL) condition from (14) directly,

which implies Lt≤L0e−2c^2λmin⁡(G)(Kσ′)2t\mathcal{L}_{t}\leq\mathcal{L}_{0}e^{-2\hat{c}^{2}{\lambda_{\min}(G)}\left(K_{\sigma^{\prime}}\right)^{2}t}, indicating that the training loss decays to at a linear rate.

For more general choices of the activation function, a challenge is to guarantee that, heuristically speaking, for each a∈[n]a\in[n], \sigma^{\prime}\big{(}h_{i}(\bm{x}_{a})\big{)} does not become near zero for too many i∈[m]i\in[m] before the loss vanishes. To facilitate a finer-grained analysis, we need the following mild assumption on σ\sigma:

Intuitively, II is an active region of σ\sigma, within which the derivative has a magnitude bounded away from zero. This assumption is satisfied by the majority of activation functions in practice, including smooth ones such as tanh⁡\tanh and sigmoid as well as non-smooth ones such as ReLU. Then, under the following initialization scheme, we prove a general result for models with a fixed embedding.

πw\pi_{\bm{w}} is the DD-dimensional standard Gaussian distribution, i.e., each WijW_{ij} is sampled independently from a standard Gaussian distribution.

Suppose that Assumptions 1, 2 and 3 are satisfied, and λmin⁡(G)>0\lambda_{\min}(G)>0. Then ∃c^0\exists\hat{c}_{0}, rr and C>0C>0 such that ∀δ>0\forall\delta>0, if c^≥c^0λmax⁡(G)/λmin⁡(G)\hat{c}\geq\hat{c}_{0}\lambda_{\max}(G)/\lambda_{\min}(G) and m≥C(1+c^2)log⁡(n/δ)m\geq C(1+\hat{c}^{2})\log\left(n/\delta\right), then with probability at least 1−δ1-\delta, it holds that ∀t≥0\forall t\geq 0,

Here, c^0\hat{c}_{0}, rr and CC depend on I,Gmin⁡,Gmax⁡,∥y∥,LσI,G_{\min},G_{\max},\|\bm{y}\|,L_{\sigma} and Kσ′K_{\sigma^{\prime}} (but not on mm, nn, dd, DD, δ\delta, or λmin⁡(G)\lambda_{\min}(G)).

The result is proved in Appendix E, and below we briefly describe the intuition. A key to the proof is to guarantee that enough neurons remain in the active region throughout training. Specifically, with respect to each training data point (i.e. for each a∈[n]a\in[n]), we can keep track of the proportion of neurons (among all i∈[m]i\in[m]) for which hit(xa)∈Ih_{i}^{t}(\bm{x}_{a})\in I. We show that if the proportion is large enough at initialization (shown by Lemma 3 in Appendix E.2 under Assumption 3), then it cannot drop dramatically without a simultaneous decrease of the loss value, as long as the cic_{i}’s are not too small in absolute value. This property of the dynamics is formalized in the following lemma:

Consider the dynamics of Lt\mathcal{L}^{t} and \big{\{}h_{i}^{t}(\bm{x}_{a})\big{\}}_{i\in[m],a\in[n]} governed by (11) and (14). Assume that λmin⁡(G)>0\lambda_{\min}(G)>0, and ∀i∈[m]\forall i\in[m], ∣ci∣=c^>0|c_{i}|=\hat{c}>0. Under Assumption 2, define

where κ(λ1,λ2)=9λ2(Ir−Il)2λ1Kσ′\kappa(\lambda_{1},\lambda_{2})=\frac{9\lambda_{2}(I_{r}-I_{l})}{2\lambda_{1}K_{\sigma^{\prime}}}.

Suppose that Assumptions 1, 2 and 3 are satisfied. If the training data are linearly-independent vectors, then under GF (10) on the first-layer weights of the shallow NN, the training loss converges to zero at a linear rate.

While the assumption that n≤dn\leq d is restrictive, we note that existing convergence rate guarantees for the GD-type training of shallow NNs in the MF scaling need strong additional assumptions , modifications to the GD algorithm , or restrictions to certain special tasks .

2 Models with a high-dimensional random embedding

A clear limitation of Corollary 1 is that it is only applicable when n≤dn\leq d, since otherwise the Gram matrix G(0)G^{(0)} cannot be positive definite. This motivates us to consider the use of a high-dimensional embedding Φ\Phi to lift the effective input dimension. In particular, we focus on the scenario where DD is large and Φ\Phi is random. While the Gram matrix GG in this case is also random, we only need that it concentrates around a deterministic and positive definite limit as DD tends to infinity:

Condition 1 is sufficient for us to apply Lemma 1 and obtain the following global convergence guarantee, which extends Theorem 1 to models with a high-dimensional random embedding. The proof is given in Appendix F.

Under Assumptions 1, 2, 3 and Condition 1, ∃c^0\exists\hat{c}_{0}, rr and C>0C>0 such that ∀δ>0\forall\delta>0, if c^≥c^0λmax⁡(Gˉ)/λmin⁡(Gˉ)\hat{c}\geq\hat{c}_{0}\lambda_{\max}(\bar{G})/\lambda_{\min}(\bar{G}), m≥C(1+c^2)log⁡(n/δ)m\geq C(1+\hat{c}^{2})\log\left({n}/{\delta}\right) and D≥Dmin⁡(12δ,12λmin⁡(Gˉ))D\geq D_{\min}(\frac{1}{2}\delta,\frac{1}{2}\lambda_{\min}(\bar{G})), then with probability at least 1−δ1-\delta, it holds that ∀t≥0\forall t\geq 0,

Here, c^0\hat{c}_{0}, rr and CC depend on I,Gˉmin⁡,Gˉmax⁡,∥y∥,LσI,\bar{G}_{\min},\bar{G}_{\max},\|\bm{y}\|,L_{\sigma} and Kσ′K_{\sigma^{\prime}} (but not mm, nn, dd, DD, δ\delta, or λmin⁡(Gˉ)\lambda_{\min}(\bar{G})).

Consider the P-33L NN model defined in (4). In this case, the Gram matrix is G(1)G^{(1)}, defined by

Thus, for the convergence result, the assumption we need on the limiting Gram matrix is

πz\pi_{\bm{z}} is sub-Gaussian and the matrix Gˉ(1)\bar{G}^{(1)}, which depends on the choice of σ\sigma and the training set, is positive definite with \lambda_{\min}\big{(}\bar{G}^{(1)}\big{)}>0 and (Gˉ(1))max⁡<∞(\bar{G}^{(1)})_{\max}<\infty.

This assumption also plays an important role in the NTK analysis, and it is satisfied if, for example, πz\pi_{\bm{z}} is the dd-dimensional standard Gaussian distribution, no two data points are parallel, and σ\sigma is either the ReLU function or analytic and not a polynomial . When Assumption 4 is satisfied, as long as σ\sigma is Lipschitz, we can use standard concentration techniques to verify Condition 1. Thus, Theorem 2 implies that

Under Assumptions 1, 2, 3 and 4, ∃c^0\exists\hat{c}_{0}, rr, C1C_{1} and C2>0C_{2}>0 such that ∀δ>0\forall\delta>0, if c^≥c^0λmax⁡(Gˉ(1))/λmin⁡(Gˉ(1))\hat{c}\geq\hat{c}_{0}\lambda_{\max}(\bar{G}^{(1)})/\lambda_{\min}(\bar{G}^{(1)}), m≥C1(1+c^2)log⁡(n/δ)m\geq C_{1}(1+\hat{c}^{2})\log\left({n}/{\delta}\right) and D≥C2n2log⁡(n/δ)/λmin⁡(Gˉ(1))2D\geq C_{2}n^{2}\log(n/\delta)/\lambda_{\min}(\bar{G}^{(1)})^{2}, then with probability at least 1−δ1-\delta, it holds that ∀t≥0\forall t\geq 0,

Here, c^0\hat{c}_{0}, rr, C1C_{1} and C2C_{2} depend on I,Gˉmin⁡(1),Gˉmax⁡(1),∥y∥,Kσ′I,\bar{G}^{(1)}_{\min},\bar{G}^{(1)}_{\max},\|\bm{y}\|,K_{\sigma^{\prime}} as well as the sub-Gaussian norm of μz\mu_{\bm{z}} (but not on mm, nn, dd, DD, δ\delta or λmin⁡(Gˉ(1))\lambda_{\min}(\bar{G}^{(1)})).

The proof is given in Appendix G. Compared to Corollary 1 for shallow NNs, a highlight of Theorem 3 is that the requirement of n≤dn\leq d is no longer needed. This demonstrates an advantage of the high-dimensional random embedding realized by the first hidden layer in the P-33L NN, thus illustrating a benefit of having both depth and width in NNs from the viewpoint of optimization. Compared to the NTK result , our analysis assumes the same level of over-parameterization, but crucially allows feature training to occur, which we discuss in Section 2.2 and support empirically in Section 4.3.

Furthermore, by using a multi-layer NN with random and fixed weights as the high-dimensional random embedding, we extend the P-33L NN to a partially-trained LL-layer NN model in Appendix H, for which similar convergence results can be proved for training its second-to-last layer via GF.

Numerical experiments

Additional results and details of the experiments are provided in Appendix I.

We train shallow NNs to fit a randomly labeled data set {(x1,y1),...,(xn,yn)}\{(\bm{x}_{1},y_{1}),...,(\bm{x}_{n},y_{n})\} with d=20d=20. Specifically, we sample each xa\bm{x}_{a} i.i.d. with every entry sampled independently from a standard Gaussian distribution, and each yay_{a} i.i.d. uniformly on [−12,12][-\tfrac{1}{2},\tfrac{1}{2}] and independently from the xa\bm{x}_{a}’s. We see from Figure 3 that the convergence happens at a nearly linear rate when n=20n=20 and 4040, and the rate decreases as nn becomes larger. This is coherent with our theoretical result (Corollary 1), and interestingly also echoes a prior result that the convergence rate of optimizing a shallow NN using population loss can suffer from the curse of dimensionality , which implies a worsening of the convergence rate as the number of data points increases.

2 Experiment 2: Benefit of input embedding

3 Experiment 3: Feature learning v.s. lazy training

We consider the P-33L NN model defined in (4) and (5) with D=mD=m (i.e. both hidden layers having the same width), and compare it with 33-layer NN models under NTK and MF scalings, as we define in Table 1 based on prior literature , which undergo partial training in the same fashion. We adopt the data set used in (more details in Appendix I.3), and train the models by minimizing the unregularized squared loss for varying nn’s and mm’s.

First, we see from the top-left plot in Figure 4 that, consistently across different mm, the training loss converges at a linear rate for the model under our scaling, which is coherent with Theorem 3. Second, we see from the second row that feature learning occurs in the model under our scaling but negligibly in the model under the NTK scaling, as expected . Note also that under the MF scaling, the feature maps h1(x),...,hm(x)h_{1}(\bm{x}),...,h_{m}(\bm{x}) concentrate near at initialization due to the small scaling, but gains diversity during training. Third, we see from Figure 3 that our model yields the smallest test errors out of all three, and in addition, as nn grows the test error decreases faster under the MF scaling than under the NTK scaling, both indicating an advantage of feature learning compared to lazy training.

Conclusions and limitations

We consider a general type of models that includes shallow and partially-trained multi-layer NNs, which exhibits feature learning when trained via GF, and prove non-asymptotic global convergence guarantees that accommodates a general class of activation functions. For a randomly-initialized shallow NN in the MF scaling that is wide enough, we prove that by performing GF on the input-layer weights, the training loss converges to zero at a linear rate if the number of training data does not exceed the input dimension. For a randomly-initialized multi-layer NN with large widths, we prove that by performing GF on the weights in the second-to-last layer, the same result holds except there is no requirement on the input dimension. We also perform numerical experiments to demonstrate the advantage of feature learning in our partially-trained multi-layer NNs relative to their counterparts under the NTK scaling.

Our work focuses on the optimization rather than the approximation or generalization properties of NNs, which are also crucial to understand. In addition, as our current theoretical results on global convergence neglect the bias terms and assume that the last-layer weights are untrained, a more general version is left for future work.

Acknowledgments

The authors acknowledge support from the Henry MacCracken Fellowship, NSF RI-1816753, NSF CAREER CIF 1845360, NSF CHS-1901091 and NSF DMS-MoDL 2134216.

References

Appendix A Additional notations

For a positive integer nn, we let [n][n] denote the set {1,...,n}\{1,...,n\}.

We write ∑a‾\overline{\sum_{a}} for 1n∑a=1n\frac{1}{n}\sum_{a=1}^{n}.

We use bold letters (e.g. x\bm{x}, z\bm{z}, c\bm{c}, y\bm{y}) to denote vectors.

We use WW and {Wij}i∈[m],j∈[D]\{W_{ij}\}_{i\in[m],j\in[D]} interchangeably to refer to the same set of parameters.

Appendix B Consistency of the scaling and GD update rule with Xavier initialization

Consider a three-layer network defined by

with weight parameters \big{\{}\theta^{(1)}_{jk}\big{\}}_{j,k\in[m]}, \big{\{}\theta^{(2)}_{ij}\big{\}}_{i,j\in[m]} and \big{\{}\theta^{(3)}_{i}\big{\}}_{i\in[m]} are initialized according to Xavier initialization, which means that we sample each θjk(1)\theta^{(1)}_{jk} i.i.d. from N(0,1m+d)\mathcal{N}(0,\frac{1}{m+d}), each θij(2)\theta^{(2)}_{ij} i.i.d. from N(0,12m)\mathcal{N}(0,\frac{1}{2m}), and each θi(3)\theta^{(3)}_{i} i.i.d. from N(0,1m+1)\mathcal{N}(0,\frac{1}{m+1}). If m≫dm\gg d, both N(0,1m+d)\mathcal{N}(0,\frac{1}{m+d}) and N(0,1m+1)\mathcal{N}(0,\frac{1}{m+1}) can be approximated by N(0,1m)\mathcal{N}(0,\frac{1}{m}). Then, up to this approximation, by redefining ci=mθi(3)c_{i}=\sqrt{m}\theta^{(3)}_{i}, Wij=mθij(2)W_{ij}=\sqrt{m}\theta^{(2)}_{ij} and zjk=mθjk(1)z_{jk}=\sqrt{m}\theta^{(1)}_{jk}, we can write

and note that ci,Wijc_{i},W_{ij} and zjkz_{jk} are all initialized i.i.d. of order O(1)O(1). In addition, if σ\sigma is homogeneous, this is then equivalent to (4) and (5) when D=mD=m.

Moreover, there is ∂f∂Wij(x)=1m∂f∂θij(2)(x)\frac{\partial f}{\partial W_{ij}}(\bm{x})=\frac{1}{\sqrt{m}}\frac{\partial f}{\partial\theta^{(2)}_{ij}}(\bm{x}). Then, since performing GD on θij(2)\theta^{(2)}_{ij} with step size δ\delta means updating θij(2)\theta^{(2)}_{ij} according to

this is equivalent to updating WijW_{ij} according to

which justifies the mm factor on the right-hand-side of (9).

Appendix C Relationship to the maximum-update parameterization and feature learning

Consider the partially-trained LL-layer NN model defined in Section H in the case where D=m≫dD=m\gg d. In the framework of abc-parameterization introduced in , our model corresponds to setting

Furthermore, as we explain in Appendix B, the appropriate learning rate scales linearly with mm (as in (9)), which corresponds to having

Meanwhile, the maximum-update (μ\muP) parameterization is characterized by setting

Recall the symmetry of abc-parameterization derived in , which states that one gets a different but equivalent abc-parameterization by setting

Since our parameterization can be obtained from the maximum-update parameterization by applying the transformation above with θ=12\theta=\frac{1}{2}, they are equivalent in the function space. In particular, for our parameterization, the rr parameter defined in can be computed as

Hence, according to , our parameterization exhibits feature learning.

Appendix D Proof of Lemma 1

Since we assume that GG is positive definite and ∣ci∣=c^>0|c_{i}|=\hat{c}>0, ∀i∈[m]\forall i\in[m], we can derive from (14) that

Since II is an open interval, ∃ξ>0\exists\xi>0 such that we can find a subinterval I0⊆II_{0}\subseteq I such that the distance between I0I_{0} and the boundaries of II (if II is bounded on either side) is no less than ξ\xi, i.e.,

In particular, we can choose ξ=13(Ir−Il)\xi=\frac{1}{3}(I_{r}-I_{l}) and I0=(Il+ξ,Ir−ξ)I_{0}=(I_{l}+\xi,I_{r}-\xi). Then there is

where we set C1=λmax⁡(G)(λmin⁡(G))−12ξ−1>0C_{1}=\lambda_{\max}(G)\left(\lambda_{\min}\left(G\right)\right)^{-\frac{1}{2}}\xi^{-1}>0 for simplicity. As a consequence,

where we set C2=2−12C1(λmin⁡(G))−12(Kσ′)−1=λmax⁡(G)2ξKσ′λmin⁡(G)≤nGmax⁡2ξKσ′λmin⁡(G)C_{2}=2^{-\frac{1}{2}}C_{1}\left(\lambda_{\min}(G)\right)^{-\frac{1}{2}}(K_{\sigma^{\prime}})^{-1}=\frac{\lambda_{\max}(G)}{\sqrt{2}\xi K_{\sigma^{\prime}}\lambda_{\min}(G)}\leq\frac{nG_{\max}}{\sqrt{2}\xi K_{\sigma^{\prime}}\lambda_{\min}(G)} for simplicity. Therefore, when ηt>0\eta^{t}>0,

Appendix E Proof of Theorem 1

To apply Lemma 1, we need two additional lemmas, which we will prove in Appendix E.1 and E.2. The first one guarantees that the loss value at initialization, L0\mathcal{L}^{0}, is upper-bounded with high probability:

∀δ>0\forall\delta>0, if m≥Ω(c^2log⁡(nδ−1)Gmax⁡/∥y∥2)m\geq\Omega\left(\hat{c}^{2}\log\left(n\delta^{-1}\right)G_{\max}/\|\bm{y}\|^{2}\right), then with probability at least 1−δ1-\delta, there is

∀δ>0\forall\delta>0, if m≥log⁡(nδ−1)2(K(I,Gmin⁡,Gmax⁡))2m\geq\frac{\log(n{\delta}^{-1})}{2\left(K(I,G_{\min},G_{\max})\right)^{2}}, then with probability at least 1−δ1-\delta, there is

where K(I,λ1,λ2)=162πλ2(Ir−Il)exp⁡{−max⁡{∣Il∣,∣Ir∣}(λ1)2}K(I,\lambda_{1},\lambda_{2})=\frac{1}{6\sqrt{2\pi}\lambda_{2}}(I_{r}-I_{l})\exp\left\{-\frac{\max\{|I_{l}|,|I_{r}|\}}{(\lambda_{1})^{2}}\right\} is a positive number that depends on II, λ1\lambda_{1} and λ2\lambda_{2}.

With these two lemmas, we deduce that ∀δ>0\forall\delta>0, if m≥Ω((1+c^2/∥y∥2)log⁡(nδ−1))m\geq\Omega\left((1+\hat{c}^{2}/\|\bm{y}\|^{2})\log\left(n\delta^{-1}\right)\right), then with probability at least δ\delta, there is ∀t≥0\forall t\geq 0,

where K(I,Gmin⁡,Gmax⁡)K(I,G_{\min},G_{\max}) is defined as in Lemma 3. Therefore, if our choice of c^\hat{c} satisfies

which will allow us to finally conclude that

Note that (60) establishes a PL condition. Several other convergence analyses of NNs have also relied on variants of the PL condition .

Since at initialization, \big{\{}c_{i}\big{\}}_{i\in[m]} and \big{\{}W_{ij}^{0}\big{\}}_{i\in[m],j\in[D]} are both sampled i.i.d. and \big{\{}c_{i}\big{\}}_{i\in[m]} has mean zero, we know that ∀a∈[n]\forall a\in[n], f^{0}(\bm{x}_{a})=\frac{1}{m}\sum_{i=1}^{m}c_{i}\sigma\big{(}h_{i}^{t}(\bm{x}_{a})\big{)} is the sample mean of i.i.d. random variables with zero-mean. Moreover, since \big{\{}W_{ij}^{0}\big{\}}_{i\in[m],j\in[D]} is sampled from N(0,1)\mathcal{N}(0,1), we know that ∀i∈[m]\forall i\in[m], the random variable c_{i}\sigma\big{(}h_{i}^{t}(\bm{x}_{a})\big{)} is sub-Gaussian , with sub-Gaussian norm

where MSG>0M_{SG}>0 is some absolute constant. Thus, by Hoeffding’s inequality , ∀a∈[n]\forall a\in[n], ∀r>0\forall r>0,

where KK is some absolute constant. Hence, by union bound,

then with probability at least 1−δ1-\delta, there is

E.2 Proof of Lemma 3

Since each Wij0W_{ij}^{0} are sampled i.i.d. from N(0,1)\mathcal{N}(0,1), we know that ∀a∈[n]\forall a\in[n], independently for each i∈[m]i\in[m], hi0(xa)h_{i}^{0}(\bm{x}_{a}) follows a Gaussian distribution with mean 0 and variance GaaG_{aa}. Therefore,

Hence, by Hoeffding’s inequality, ∀a∈[n]\forall a\in[n], ∀r>0\forall r>0,

∀a∈[n]\forall a\in[n], choosing r=12π(I0;Gaa)r=\frac{1}{2}\pi\left(I_{0};G_{aa}\right), we then get

Since ∀b∈[n]\forall b\in[n], there is Gmin⁡≤Gbb≤Gmax⁡G_{\min}\leq G_{bb}\leq G_{\max},

Letting K(I,λ1,λ2)=162πλ2(Ir−Il)exp⁡{−max⁡{∣Il∣,∣Ir∣}(λ1)2}>0K(I,\lambda_{1},\lambda_{2})=\frac{1}{6\sqrt{2\pi}\lambda_{2}}(I_{r}-I_{l})\exp\left\{-\frac{\max\{|I_{l}|,|I_{r}|\}}{(\lambda_{1})^{2}}\right\}>0, we can then write

Appendix F Proof of Theorem 2

By Condition 1, we know that ∀δ>0\forall\delta>0, if D≥Dmin⁡(12δ,12λmin⁡(Gˉ))D\geq D_{\min}(\frac{1}{2}\delta,\frac{1}{2}\lambda_{\min}(\bar{G})), then with probability at least 1−12δ1-\frac{1}{2}\delta, there is ∥G−Gˉ∥2≤12λmin⁡(Gˉ)\|G-\bar{G}\|_{2}\leq\frac{1}{2}\lambda_{\min}(\bar{G}), and hence λmin⁡(G)≥12λmin⁡(Gˉ)\lambda_{\min}(G)\geq\frac{1}{2}\lambda_{\min}(\bar{G}), Gmin⁡≥12Gˉmin⁡G_{\min}\geq\frac{1}{2}\bar{G}_{\min}, λmax⁡(G)≤λmax⁡(Gˉ)+12λmin⁡(Gˉ)≤2λmax⁡(Gˉ)\lambda_{\max}(G)\leq\lambda_{\max}(\bar{G})+\frac{1}{2}\lambda_{\min}(\bar{G})\leq 2\lambda_{\max}(\bar{G}), and Gmax⁡≤2Gˉmax⁡G_{\max}\leq 2\bar{G}_{\max}. We then perform the following analysis conditioned on the event that ∥G−Gˉ∥2≤12λmin⁡(Gˉ)\|G-\bar{G}\|_{2}\leq\frac{1}{2}\lambda_{\min}(\bar{G}).

Since the sampling of \big{\{}c_{i}\big{\}}_{i\in[m]} and \big{\{}W_{ij}^{0}\big{\}}_{i\in[m],j\in[D]} is independent from the realization of GG, we know from Lemma 3 that if m≥log⁡(4nδ−1)2(K(I,12λmin⁡(Gˉ),2λmax⁡(Gˉ)))2≥log⁡(4nδ−1)2(K(I,Gmin⁡,Gmax⁡))2m\geq\frac{\log(4n{\delta}^{-1})}{2\left(K(I,\frac{1}{2}\lambda_{\min}(\bar{G}),2\lambda_{\max}(\bar{G}))\right)^{2}}\geq\frac{\log(4n{\delta}^{-1})}{2\left(K(I,G_{\min},G_{\max})\right)^{2}}, then with probability at least 1−14δ1-\frac{1}{4}\delta, there is

From Lemma 2, we also know that if m≥Ω(c^2log⁡(nδ−1)λmax⁡(Gˉ)/∥y∥2)≥Ω(c^2log⁡(nδ−1)λmax⁡(G)/∥y∥2)m\geq\Omega\left(\hat{c}^{2}\log\left(n\delta^{-1}\right)\lambda_{\max}(\bar{G})/\|\bm{y}\|^{2}\right)\geq\Omega\left(\hat{c}^{2}\log\left(n\delta^{-1}\right)\lambda_{\max}(G)/\|\bm{y}\|^{2}\right), then with probability at least 1−14δ1-\frac{1}{4}\delta, there is L0≤∥y∥2\mathcal{L}^{0}\leq\|\bm{y}\|^{2}. Therefore, in total, we know that with probability at least 1−δ1-\delta, the following conditions all hold:

in which case, by applying Lemma 1 with G=G(1)G=G^{(1)}, we get

where K1(λ1,λ2)=92λ1−1λ2Kσ′−1(Ir−Il)>0K_{1}(\lambda_{1},\lambda_{2})=\frac{9}{2}\lambda_{1}^{-1}\lambda_{2}K_{\sigma^{\prime}}^{-1}(I_{r}-I_{l})>0. Thus, by the definition of K1(⋅,⋅)K_{1}(\cdot,\cdot), we know that

Therefore, if our choice of c^\hat{c} satisfies

Hence, (50) implies that ∀t≥0\forall t\geq 0,

Appendix G Proof of Theorem 3

In view of Theorem 2, it is sufficient to verify that Condition 1 holds for Dmin⁡(δ,u)=Ω(n2u−2log⁡(nδ−1))D_{\min}(\delta,u)=\Omega\left(n^{2}u^{-2}\log(n\delta^{-1})\right), which is given by the following lemma:

∀δ≥0\forall\delta\geq 0, if D≥Ω(n2u−2log⁡(nδ−1))D\geq\Omega\left(n^{2}u^{-2}\log(n\delta^{-1})\right) , then with probability at least 1−δ1-\delta,

Hence, by Lemma 2.7.7 in , we know that ∀a,b∈[n]\forall a,b\in[n], σ(xa⊺Z)σ(xb⊺Z)\sigma(x_{a}^{\intercal}Z)\sigma(x_{b}^{\intercal}Z) is a sub-exponential random variable with sub-exponential norm

Then, by Bernstein’s inequality (Theorem 2.8.1 in ), since each zj\bm{z}_{j} is sampled i.i.d. from πz\pi_{\bm{z}}, we have that ∀a,b∈[n]\forall a,b\in[n] and ∀u>0\forall u>0,

where K>0K>0 is some absolute constant. In other words, for any δ′>0\delta^{\prime}>0, if

then we have ∣Gab(1)−Gˉab(1)∣≥u\left|G^{(1)}_{ab}-\bar{G}^{(1)}_{ab}\right|\geq u with probability at least 1−δ1-\delta. If we choose u=u′nu=\frac{u^{\prime}}{n} and δ′=δn2\delta^{\prime}=\frac{\delta}{n^{2}}, then we get, if

Hence, with probability at least 1−δ1-\delta, we have

Appendix H Generalization to deeper models

By setting Φ\Phi to be the activations of the second-to-last hidden-layer of a multi-layer NN, we can obtain generalizations of the P-33L NN to deeper architectures. For example, in the feed-forward case, we can obtain the following partially-trained LL-layer NN:

Appendix I Further details of the numerical experiments

In our models, \big{\{}c_{i}\big{\}}_{i\in[m]} is sampled i.i.d. from the Rademacher distribution μc=12δ1+12δ−1\mu_{c}=\frac{1}{2}\bm{\delta}_{1}+\frac{1}{2}\bm{\delta}_{-1}, \big{\{}\bm{z}_{j}\big{\}}_{j\in[D]} is sampled i.i.d. from N(0,Id)\mathcal{N}(0,I_{d}), and \big{\{}W_{ij}\big{\}}_{i\in[m],j\in[D]} is initialized by sampling i.i.d. from N(0,1)\mathcal{N}(0,1). In the model under NTK scaling, we additionally symmetrize the model at initialization according to the strategy used in to ensure that the function value at initialization does not blow up when the width is large. We choose to train the models using 5000050000 steps of (full-batch) GD with step size δ=1\delta=1. When the test error is computed, we use a test set of size 500500 generated by sampling i.i.d. from the same distribution as the training set.

The experiments are run with NVIDIA GPUs (1080ti and Titan RTX).

We choose σ\sigma to be tanh⁡\tanh. For each choice of nn, we run the experiment with 55 different random seeds, and Figure 3 plots the evolution of the training loss during GD averaged over the 55 runs with m=8192m=8192.

Figure 5 is the same as Figure 3 except for having m=4096m=4096. We see that the two two plots agree well.

I.2 Experiment 222

We choose σ\sigma to be ReLU. For each choice of nn and each of the two models, we experiment with 55 different random seeds, and Figure 3 plots the test error at the 5000050000 GD step averaged over the 55 runs ±\pm its standard deviation.

In Figure 6, we plot the evolution of the training loss and test error during GD for the two different models, with m=2048m=2048 or 81928192 and different choices of nn, averaged over 55 runs with different random seeds. We see in particular that the difference between the two choices of mm is negligible, suggesting that it is unlikely to obtain performance improvements with further over-parameterization.

I.3 Experiment 333

and x3,...,xdx_{3},...,x_{d} each follow the uniform distribution in $,independentlyfromeachotheraswellas, independently from each other as well asx_{1},x_{2}andandy$.

Figures 8 and 8 are the same as Figure 4 except for having n=400n=400 and 800800, respectively. We see that as nn increases, test error improves for all three models, while our P-33L NN model remains the one achieving the lowest test error.