High-dimensional Asymptotics of Feature Learning: How One Gradient Step Improves the Representation
Jimmy Ba, Murat A. Erdogdu, Taiji Suzuki, Zhichao Wang, Denny Wu, Greg Yang
Introduction
We consider the training of a fully-connected two-layer neural network (NN) with neurons,
When the first layer is fixed and only the second layer is optimized, we arrive at a kernel model, where the kernel defined by features (often called the hidden representation) is referred to as the conjugate kernel (CK) [Nea95]. When is randomly initialized, this model is an example of the random features (RF) model [RR08]. The training and test performance of RF regression has been extensively studied in the proportional limit [LLC18, MM22]. These precise characterizations reveal interesting phenomena also present in practical deep learning, such as the non-monotonic risk curve [BHMM19].
However, RF models do not fully explain the empirical success of neural networks: one crucial advantage of deep learning is the ability to learn useful features [GDDM14, DCLT18] that “adapt” to the learning problem [Suz18]. In fact, recent works have shown that such adaptivity enables NNs optimized by gradient descent to outperform a wide range of linear/kernel estimators [AZL19, GMMM19]. While many explanations of this separation between NNs and kernel models have been proposed, our starting point is the empirical finding that “non-kernel” behavior often occurs in the early phase of NN optimization, especially under large learning rate [JSF+20, FDP+20]. The goal of this work is to answer the following question:
Can we precisely capture the presence of feature learning in the early phase of gradient descent training, and demonstrate its improvement over the initial (fixed) kernel in the proportional limit?
Motivated by the above observations, we investigate a simplified scenario of the “early phase” of learning: how the first gradient step on the first-layer parameters impacts the representation of the two-layer NN (1.1). Specifically, we consider the regression setting with the squared (MSE) loss, and a student-teacher model in the proportional asymptotic limit; we characterize the prediction risk of the kernel ridge regression estimator on top of the first-layer CK feature , before and after the gradient descent stepSome of our results also apply to multiple gradient steps on the first layer , which we specify in the sequel. on the empirical risk (starting from Gaussian initialization). Our findings can be summarized as follows.
In Section 3, we show that the first gradient step on is approximately rank-1; hence under appropriate learning rate, the updated weight matrix exhibits a information (spike) plus noise (bulk) structure.
As a result, the isolated singular vector of the weight matrix aligns with the linear component of target function (teacher) , and the top eigenvector of the CK matrix aligns with the training labels .
Next in Section 4 we study how the aforementioned alignment improves the kernel. We consider a more specialized setting where the teacher is a single-index model, in which case the prediction risk of a large class of RF/kernel ridge regression estimators is lower-bounded by the -norm of the “nonlinear” component the teacher \mathopen{}\mathclose{{}\left\|{\textsf{P}_{>1}f^{*}}}\right\|_{L^{2}}^{2}, i.e., they can only learn linear functions on the input. After taking one gradient step on , we compute the CK ridge estimator using separate training data, and compare its prediction risk against this linear lower bound. Our analysis will be made under two choices of learning rate scalings (see Figure 1):
Small lr: . In Section 4.2, we extend the Gaussian Equivalence Theorem (GET) in [HL20] to the updated feature map trained via multiple gradient descent steps on with learning rate ; this allows us to precisely characterize the prediction risk using random matrix theoretical tools. We prove that after one gradient step, the ridge regression estimator on the learned CK features already exhibits nontrivial improvement over the initial RF ridge model, but it remains in the “linear regime” and cannot outperform the best linear estimator on the input.
Large lr: . In Section 4.3, we analyze a larger learning rate that coincides with the maximal update parameterization in [YH20]. For certain target functions , we prove that kernel ridge regression after one feature learning step can achieve lower risk than the lower bound \mathopen{}\mathclose{{}\left\|{\textsf{P}_{>1}f^{*}}}\right\|_{L^{2}}^{2}, and thus outperform a wide range of kernel ridge estimators (including the neural tangent kernel of (1.1)).
2 Related Works
A plethora of recent works provided precise performance analysis of RF and kernel models in the proportional limit [MM22, GLK+20, DL20, LCM20, AP20]. These results typically build upon analyses of the spectrum of kernel matrices, a key ingredient in which is the “linearization” of nonlinear random matrices via Taylor expansion [EK10] or orthogonal polynomials [CS13, PW17].
Consequently, a large class of kernel models are essentially linear in the proportional limit [LR20, BMR21]. In the case of RF models, similar property is captured by the Gaussian equivalence theorem [GMKZ20, HL20, GLR+21], which roughly states that RF estimators achieve the same prediction risk as a (noisy) linear model. For input on unit sphere, [GMMM21, MMM21] showed that sample size is required to go beyond this “linear” regime. As we will see in certain settings, such limitation can also be overcome (in the scaling) by training the feature map for one gradient step with sufficiently large learning rate.
It is well-known that under certain initialization, the learning dynamics of overparameterized NNs can be described by the neural tangent kernel (NTK) [JGH18]. However, the NTK description essentially “freezes” the model around its initialization [COB19], and thus does not explain the presence of feature learning in NNs [YH20].
In fact, various works have shown that deep learning is more powerful than kernel methods in terms of approximation and estimation ability [Bac17, Suz18, IF19, SH20, GMMM20]. Moreover, in some specialized settings, NNs optimized with gradient-based methods can outperform the NTK (or more generally any kernel estimators) in terms of generalization error [AZL19, WLLM19, GMMM19, LMZ20, DM20, SA20, AZL20, RGKZ21, KWLS21, ABAB+21] (see [MKAS21, Table 2] for survey). These results often require careful analysis of the landscape (e.g., properties of global optimum) or optimization dynamics; in contrast, our goal is to precisely characterize the first gradient step and demonstrate a similar separation.
Recent empirical studies suggest that properties of the final trained model is strongly influenced by the early stage of optimization [GAS19, LM20, PPVF21], and the NTK evolves most rapidly in the first few epochs [FDP+20]. Large learning rate in the initial steps can impact the conditioning of loss surface [JSF+20, CKL+21] and potentially improve the generalization performance [LWM19, LBD+20]. Under structural assumptions on the data, it has been proved that one gradient step with sufficiently large learning rate can drastically decrease the training loss [CLB21], extract task-relevant features [DM20, FCB22], or escape the trivial stationary point at initialization [HCG21]. While these works also highlight the benefit of one feature learning stepWe however note that the “early phase” is not always sufficient: for certain teacher model , (stochastic) gradient descent may exhibit a long initial “search” stage before nontrivial alignment can be achieved, see [AGJ21, VSL+22]. , to our knowledge this advantage has not been precisely characterized in the proportional regime (where the performance of RF models has been extensively studied).
Problem Setup and Basic Assumptions
1 Training Procedure
2 Main Assumptions
Proportional Limit. , , , where .
Following [HL20], we assume smooth and centered activation to simplify the computation; Section 4 provides empirical evidence that our results hold beyond this condition (see also [LGC+21]). We expect that the Gaussian input assumption may be replaced by weaker orthogonality conditions as in [FW20].
Under Assumption 1, increasing the sample size corresponds to enlarging , and increasing the network width corresponds to enlarging . The proportional scaling of (also referred to as the “linear-width” regime) implies that the model width is not significantly larger than the training set size, in contrast to the polynomial overparameterization often required in NTK analyses [DZPS19], which may be less realistic for practical settings.
3 Lower Bound for Kernel Ridge Regression
To illustrate the benefit of feature learning, we compare the prediction risk of ridge regression on the trained CK (after one gradient step) against the ridge estimator on the initial RF kernels. Specifically, given training data , we consider the following class of kernel models for comparison.
Rotationally Invariant Kernel Model. Consider the inner-product kernel: k(\boldsymbol{x},\boldsymbol{y})=g\mathopen{}\mathclose{{}\left(\frac{\langle\boldsymbol{x},\boldsymbol{y}\rangle}{d}}\right), and Euclidean distance kernel: k(\boldsymbol{x},\boldsymbol{y})=g\mathopen{}\mathclose{{}\left(\frac{\mathopen{}\mathclose{{}\left\|{\boldsymbol{x}-\boldsymbol{y}}}\right\|^{2}}{d}}\right), where satisfies certain smoothness conditions as in [EK10]. Denote the associated RKHS as , and . The kernel ridge estimator is given by
We denote the prediction risk of the above kernel estimators as , respectively. The following lower bound is a simple combination of known results from [EK10, HL20, MZ20, BMR21].
This proposition implies that in the proportional limit, ridge regression on the RF or rotationally invariant kernels defined above does not outperform the best linear estimator on the input data — it cannot achieve negligible prediction risk unless the target function is linear (i.e., \mathopen{}\mathclose{{}\left\|{\textsf{P}_{>1}f^{*}}}\right\|_{L^{2}}=0). In Section 4, we compare the prediction risk of the ridge estimator on trained features against this lower bound.
How Does One Gradient Step Change the Weights?
In this section, we study the properties of the updated weight matrix in the two-layer NN (1.1). We first show that the first gradient step on can be approximated by a rank-1 matrix, which contains information of the training labels . Based on this property, we provide a signal (spike) plus noise (bulk) decomposition of , and prove that the isolated singular vector is aligned to the linear component of the teacher .
Define and a rank-1 matrix . Under Assumption 1, there exist some constants such that for all large , with probability at least ,
Proposition 2 suggests that the first-step gradient can be approximated in operator norm by a rank-1 matrix ; thus, when the learning rate is reasonably large, we expect a “spike” to appear in the updated weight matrix . Intuitively, since this rank-1 direction relates to the label vector , the resulting may be “aligned” to the target function . This intuition is confirmed in the next subsection.
In other words, if we write , then is required so that the change in the weight matrix is non-negligible (one may verify that for , the test performance of kernel ridge regression remains unchanged after one GD step). On the other hand, when , the gradient “overwhelms” the initialized parameters , and the preactivation feature in the NN (1.1) becomes unbounded as . This motivates us to consider the following two regimes of learning rate scaling.
2 Alignment with the Target Function
Under Assumption 1, we may utilize the following orthogonal decomposition of the target function ,
When in (3.2), we show a BBP phase transition (named after Baik, Ben Arous, Péché [BAP05]) for the leading singular value of , and quantify the alignment between the corresponding singular vector and the linear component of target function . It is worth noting that in our analysis, the signal is “hidden” in the rank-one perturbation defined in Proposition 2; thus our setting is different from the usual low-rank signal-plus-noise models (e.g. [BGN11, BGN12, Cap18]), and the alignment we aim to quantify does not directly follow from classical results on the BBP transition.
Then the leading singular value and the corresponding left singular vector satisfy
if ; otherwise, and , in probability, as .
We make the following observations. Beyond the threshold , increasing the learning rate enlarges the leading singular value (spike) . As for the overlap, one can numerically verify is upper-bounded by (obtained when ), from which we deduce that better alignment is achieved when we take a bigger step, or when the nonlinearity and target have larger linear components (i.e., larger ).
Theorem 3 is numerically verified in Figure 3. Observe that after one gradient step with , the bulk of the spectrum of remains unchanged and is given by the Marchenko-Pastur law (red), but a spike may appear (prediction from Theorem 3 is indicated by marker “”) when exceeds a certain threshold; furthermore, the corresponding singular vector aligns with the linear component of the target function, as shown in the subfigure (see also Figure 8(a)). We investigate the impact of this alignment on the performance of kernel ridge regression in Section 4.
While our result only characterizes the weight matrix , it may also reveal interesting properties of the CK matrix. In particular, [HL20, Lemma 5] in combination with Lemma 14 imply that for odd activation , the expected feature matrix (after one gradient step with ) satisfies
Consequently, Theorem 3 implies the same BBP transition for . When the population contains a spike, it is natural to expect the empirical CK matrix to exhibit a similar transition, which we conjecture that the Gaussian equivalence property (see Section 4.1) can precisely capture.
The conjecture predicts both the eigenvalues of the CK matrix and the overlap between its spike eigenvector and the training labels. In Figure 4 we plot the eigenvalue histogram of the CK matrix after one gradient step with , which we denote as . Observe that the bulk of the spectrum remains unchanged compared to , which can be analytically computed (red). On the other hand, similar to , an isolated eigenvalue (spike) appears in , the location of which can be predicted by the Gaussian equivalent model in Conjecture 4 (marker “”).
Do the Learned Features Improve Generalization?
Thus far we have shown that after one gradient step, the first-layer weights align with the linear component of the teacher model. Intuitively, since the learned feature map “adapts” to the teacher , we may expect the ridge regression estimator on the trained CK to achieve better performance. In this section we confirm this intuition in a concrete example: we consider the setting where is a single-index model, and compare the CK prediction risk before and after one gradient descent step on .
The single-index setting has been extensively studied in the proportional regime [GLK+20, DL20, HL20], and it is an instance of the “hidden manifold model” [GMKZ20]. However, most prior works only considered training the coefficients on top of fixed feature map (e.g., defined by randomly initialized ), and such RF models cannot learn a single-index efficiently in high dimensions [YS19].
As stated in Section 2.3, the RF ridge estimator defined by the two-layer NN (1.1) has prediction risk unless is a linear function. Here our goal is to demonstrate that the trained CK model can outperform the initial RF and potentially the kernel lower bound (2.4). We first introduce the Gaussian equivalence property which will be useful in the computation of prediction risk.
The Gaussian equivalence theorem (GET) implies that the prediction risk of a nonlinear kernel model can be the same as that of a noisy linear model. Specifically, recall the prediction risk of the ridge estimator:
This is to say, for learning rate , the Gaussian equivalent model provides an accurate description of the prediction risk of ridge regression (on the trained CK) at any fixed time step , although most of our analysis deals with . The important observation is that even though the trained weights are no longer i.i.d., the Gaussian equivalence property can still hold when remains “small” (in some norm, see (C.3) for details), which entails that the neurons remain nearly orthogonal to one another.
On the other hand, the GET also implies that the kernel estimator is essentially “linear” in high dimensions. For the squared loss, it is straightforward to verify that the Gaussian equivalent model cannot learn the nonlinear component of the target function as follows.
Hence, when , even though training the first-layer for just one step leads to non-trivial improvement over the initial RF ridge estimator (which we precisely quantify in Section 4.2), the learned CK cannot outperform the best linear model on the input features. In other words, to (possibly) learn a nonlinear , the trained feature map needs to violate the GET. In the case of one gradient step on , this amounts to using a sufficiently large step size, which we analyze in Section 4.3.
2 η=Θ(1)𝜂Θ1\eta=\Theta(1): Improvement Over the Initial CK
While the Gaussian equivalence property allows us to compute the asymptotic prediction risk after multiple gradient steps with , the precise expressions can be opaque and not amenable to interpretation or quantitative characterization. Fortunately for the first gradient step, the risk calculation can be simplified by the rank-1 approximation of the gradient matrix shown in Section 3.1. Therefore, in this subsection we focus on and analyze how the trained features improves over the initialized RF. To quantify the discrepancy in the prediction risk (4.1), we write as the prediction risk of the initialized RF ridge regression estimator (on the feature map ), and as the prediction risk of the ridge estimator on the new feature map after one feature learning step.
Importantly, due to the alignment between the trained features and the teacher model demonstrated in Section 3, we cannot simply apply a rotation invariance argument (e.g., [MM22, Lemma 9.2]) to remove the dependency on the true parameters and reduce the prediction risk to trace of certain rational functions of the kernel matrix; in other words, knowing the spectrum (or the Stieltjes transform) of the CK is not sufficient. Instead, we utilize the GET and the almost rank-1 property of in Proposition 2, which, in combination with techniques from operator-valued free probability theory [MS17], enables us to obtain the asymptotic expression of the difference in the prediction risk before and after one gradient step.
Under the same assumptions as Theorem 5 and , we have
where is defined by (C.96) in Appendix C.3. is a non-negative function of with parameters , and it vanishes if and only if (at least) one of and is zero.
Performance of the initial RF ridge estimator has been characterized by many prior works (e.g., [GLK+20, MM22]); hence the precise asymptotics of provided in Theorem 7 allows us to explicitly compute the asymptotic prediction risk of the CK model after one feature learning step .
Theorem 7 confirms our intuition that training the first-layer parameters improves the CK model, as shown in Figure 5(a)(b). Remarkably, this improvement () holds for any , that is, taking one gradient step (with learning rate ) is always beneficial, even when the training set size is small. Moreover, we do not require the student and teacher models to have the same nonlinearity — a non-vanishing decrease in the prediction risk of CK ridge regression is present as long as . On the other hand, the GET (in particular Fact 6) also implies an upper bound on the possible improvement: as ; this is to say, the trained CK remains in the “linear” regime.
Now we consider the following special cases where the expression of can be further simplified.
We first analyze the setting where the sample size is larger than any constant times , that is, we let proportionally, and then take the limit . In this regime, since a large number of training data is used to compute the gradient for the first-layer parameters, we intuitively expect the benefit of feature learning to be more pronounced, and a larger step may be more beneficial.
Consider the large sample regime: , . Under the same assumptions as Theorem 5 and , , defined in (C.120), is non-negative, vanishing if and only if one of is zero, and increasing with respect to the learning rate .
Proposition 8 predicts that the prediction risk further decreases as we use a larger learning rate , which is empirically verified in Figure 5(a). We note that the large learning rate setting () in Section 4.3 cannot be covered by the above proposition by increasing , as here does not scale with .
We also address the highly overparameterized regime, i.e., . In this limit, the initialized CK model approaches the kernel ridge regression estimator, the prediction risk of which is still lower bounded by \mathopen{}\mathclose{{}\left\|{\textsf{P}_{>1}f^{*}}}\right\|_{L^{2}}^{2} due to Proposition 1. The following proposition indicates that the advantage of one-step feature learning becomes negligible in this large width setting.
Consider the large width regime: , . Then under the same assumptions as Theorem 5 and , we have .
Proposition 9 agrees with Figure 5(b), where we see that the risk improvement is more prominent when the width is not too large. One explanation is that as increases, the initial CK already achieves lower prediction risk (e.g., see [MM22, Figure 4]), so the benefit of feature learning becomes less significant.
3 η=Θ(N)𝜂Θ𝑁\eta=\Theta(\sqrt{N}): Improvement Over the Kernel Lower Bound
Due to the large step size, the columns of the updated weight matrix are no longer near-orthogonal, which is an important property used in existing analyses of the Gaussian equivalence (e.g., see Proposition 22 or [HL20, Equation (66)]). Indeed, we will see that in this regime, the ridge regression estimator on the trained CK features is no longer “linear” and can potentially outperform the kernel lower bound (2.4) in the proportional limit. However, in the absence of GET, it is difficult to derive the precise asymptotics of the CK model. As an alternative, in this subsection we establish an upper bound on the prediction risk , which we then compare against the kernel ridge lower bound.
It is worth noting that the definition of does not involve the specific value of learning rate . This is because for any choice of , due to the Gaussian initialization of , we can find a subset of weights that receive a “good” learning rate (with high probability) such that the corresponding neurons are useful in learning the teacher model. In addition, observe that is a simple Gaussian integral which can be numerically or analytically computed (see Appendix D.2 for some examples). For instance, when , one can easily verify that and .
Under the same assumptions as Lemma 10, after one gradient step on with , there exist constants such that for any , the ridge regression estimator (4.1) satisfies
with probability as , if we choose the ridge penalty: for some small .
While Theorem 11 does not provide exact expression of the prediction risk, the upper bound still allows us to compare the prediction risk of CK ridge regression before and after one large gradient step. In particular, if \mathopen{}\mathclose{{}\left\|{\mathsf{P}_{>1}f^{*}}}\right\|_{L^{2}}^{2}\geq 10\tau^{*} (the constant is not optimized), we know that the trained CK can outperform the kernel lower bound (2.4) (hence also the initialized CK) in the proportional limit, when the ratio is sufficiently large. The following corollary provides two examples of this separation (see Figure 6).
Under the same conditions as Theorem 11, there exists some constant such that for any , the following holds with probability 1 when proportionally:
In the two examples outlined above, training the features by taking one large gradient step on the first-layer parameters can lead to substantial improvement in the performance of the CK model. In fact, the new ridge regression estimator may outperform a wide range of kernel models outlined in Section 2.3. However, we emphasize that this separation is only present in specific pairs of for which is small enough. In general settings, learning a good representation likely requires more than one gradient step (even if is a simple single-index model).
Discussion and Conclusion
We investigated how the conjugate kernel of a two-layer neural network (1.1) benefits from feature learning in an idealized student-teacher setting, where the first-layer parameters are updated by one gradient descent step on the empirical risk. Based on the approximate low-rank property of the gradient matrix, we established a signal-plus-noise decomposition for the updated weight matrix , and quantified the improvement in the prediction risk of conjugate kernel ridge regression under two different scalings of first-step learning rate . To the best of our knowledge, this is the first work that rigorously characterizes the precise asymptotics of kernel models (defined by neural networks) in the presence of feature learning.
We outline a few limitations of our current analysis as well as future directions.
Scaling of Learning Rate. Our findings in Section 4 illustrate that and result in drastically different behavior. One natural question to ask is whether there exists a “phase transition” in between the two regimes (see Figure 6) that dictates whether the GET holds. Interestingly, [RGKZ21] showed that instead of breaking the near-orthogonality of weight matrix (via large gradient step), one can also introduce sufficiently large low-rank shifts to the input to enable the initial RF estimator to fit a nonlinear . Intuitively, this may be due to the “dual” relation of and in the CK model.
Rigorous Analysis of CK Spike. In Section 3.2 we put forward a Gaussian equivalence hypothesis on the isolated eigenvalue/eigenvector of the trained CK matrix (see Figure 4); understanding whether and when such property holds is an important research direction.
The authors would like to thank (in alphabetical order) Konstantin Donhauser, Zhou Fan, Hong Hu, Masaaki Imaizumi, Ryo Karakida, Bruno Loureiro, Yue M. Lu, Atsushi Nitanda, Sejun Park, Ji Xu, Yiqiao Zhong for discussions and feedback on the manuscript.
JB was supported by NSERC Grant , CIFAR AI Chairs program, Google Research Scholar Program and Amazon Research Award. MAE was supported by NSERC Grant , Connaught New Researcher Award, CIFAR AI Chairs program, and CIFAR AI Catalyst grant. TS was partially supported by JSPS KAKENHI (20H00576) and JST CREST. ZW was supported by NSF Grant DMS-2055340. Part of this work was completed when DW interned at Microsoft Research (hosted by GY).
References
Appendix A Background and Additional Results
It is worth noting that Theorem 5 does not apply to the setting where scales with . Because of our mean-field parameterization, the first-layer weight needs to travel sufficiently far away from initialization to achieve small training loss (see Figure 2). Hence in our experimental simulations (where are large but finite), as the number of steps or learning rate increases, we expect the Gaussian equivalence predictions to become inaccurate at some point. This transition is empirically demonstrated in Figure 7(a). Observe that for larger , the GET predictions overestimate the test loss; one possible explanation is that the trained kernel can learn nonlinear functions (which we show in Section 4.3 for one gradient step with and specific choices of ), which the GET cannot capture.
We provide additional empirical evidence on this explanation in Figure 7(b). To track the learning of the linear and nonlinear components of , we recall the orthogonal decomposition:
In Figure 8(a), we compute the overlap between the leading eigenvector of and the linear component of the teacher model . Observe that the empirical simulations (dots) closely match the analytic predictions of Theorem 3 (solid curves). Also, note that increasing the learning rate or the sample size both lead to greater alignment with the teacher model.
In Figure 8(b) we repeat the large learning rate experiment in Section 4.3 for a different nonlinearity , for which , and hence the upper bound in Theorem 11 is non-vanishing. In this case, we observe that the prediction risk of the CK ridge regression model (after one feature learning step) is also non-vanishing even when the step size is large; this indicates that although we do not provide precise asymptotic characterization in Theorem 11, the upper-bounding quantity in (4.3) has predictive power on the actual prediction risk.
In Section 3.2, we observed that the trained CK aligns with training labels. Here we provide additional empirical evidence by tracking the Kernel Target Alignment (KTA) [CSTEK01] between the CK and training labels during training. Specifically, we compute the following quantity at each gradient step , which takes value between 0 and 1,
where denotes the CK matrix defined by . Figure 8(c) shows the KTA for two-layer NN under our mean-field parameterization and also the NTK parameterization (which omits the -prefactor in (1.1)). We optimize the first-layer weights until the training loss reaches for both settings, and compute the KTA on the training and test data at gradient step. Observe that the trained CK in the mean-field model aligns with both the training and test labels (purple), whereas the NN in the kernel regime does not exhibit such alignment (orange).
A.2 Additional Related Works
The neural tangent kernel (NTK) [JGH18] describes the learning dynamics of wide neural network under specific parameter scaling. Such description is based on linearizing the NN around its initialization, and the limiting kernel can be computed for various architectures [ADH+19, Yan20]. Thanks to strong convexity of the kernel objective, global convergence rate guarantees of gradient descent can be established [DZPS19, JT20]. As mentioned in Section 1.2, this first-order Taylor expansion fails to explain the adaptivity of NNs; therefore, recent works also analyzed higher-order approximations of the training dynamics [DGA20, HY20]. Noticeably, a quadratic model (i.e., second-order approximation) can outperform kernel (NTK) estimators in certain settings [AZLL19, BL20].
In contrast to the aforementioned local approximations (via Taylor expansion and truncation), the mean-field regime (e.g., [NS17, MMN18, CB18]) deals with a different scaling limit under which the evolution of parameters can be described by some partial differential equation (for comparison between regimes see [WGL+20, GSJW20]). While the mean-field limit can capture the presence of feature learning [CB20, Ngu21], quantitative guarantees often require additional conditions such as KL regularization [NWS22, Chi22]. Note that our parameterization (1.1) mirrors the mean-field scaling, but we circumvent the difficulty of analyzing the nonlinear PDE because only the “early phase” (one gradient step) is considered.
Finally, we highlight two concurrent papers that studied the mean-field dynamics of two-layer NNs (under one-pass SGD) in the high-dimensional asymptotic regime, and showed learnability results for certain target functions. [ABAM22] established a separation between NNs and kernel methods in learning “staircase-like” functions on hypercube; [VSL+22] analyzed how the model width and step size impact the learning of a well-specified two-layer NN teacher model.
A.3 Linearity of Kernel Ridge Regression
As previously mentioned, our kernel ridge regression lower bound (Proposition 1) is a simple combination of existing results, which we briefly outline below.
We first discuss the prediction risk of the ridge regression estimator on the input features. Recall that under Assumptions 1 and 2, we may write: . Given the ridge regression estimator on the input features: , we have the following bias-variance decomposition,
Following a similar computation as [BMR21, Theorem 4.13] and using the asymptotic formulae in [DW18, WX20], we can derive the following expression,
The error of this linear approximation has been studied in [MZ20, Lemma B.8] and [WZ21, Theorem 2.7], which, together with [BMR21, Theorem 4.13], entail the following equivalence under Assumption 1,
For high-dimensional input uniform on sphere or hypercube, [GMMM21, MMM21] showed that RF and kernel ridge estimators can learn at most a degree- polynomial when n=\mathcal{O}\mathopen{}\mathclose{{}\left(d^{k+1-\varepsilon}}\right); for the proportional scaling, this implies our lower bound \mathopen{}\mathclose{{}\left\|{\mathsf{P}_{>1}f^{*}}}\right\|_{L^{2}}^{2} (but under different input assumptions). [DWY21] provided a similar result for more general data distributions and a class of rotation invariant kernels based on power series expansion, but the dependence on is not sharp enough to cover the linear lower bound in Proposition 1.
Appendix B Proof for the Weight Matrix
For our later analysis, a key quantity to control is the entry-wise - matrix norm defined as
In addition, for the Hadamard product with rank-1 matrix, we have the following property.
We begin with the first gradient step. Recall the definition of the gradient matrix under the squared loss (we omit the learning rate and prefactor ):
Furthermore, we have the following probability bounds.
In Lemma 14, we do not use the proportional scaling in Assumption 1 to simplify the expressions. This is because the dependence on needs to be tracked separately in some of our calculations.
Proof. We analyze the three matrices of interest separately.
We first upper-bound . Notice that
We know that Gaussian random matrices and vectors satisfy
where the last inequality is from [Ver18, Exercise 4.6.2]. Based on Cauchy-Schwarz inequality, we can employ (B.6) and (B.7) to obtain
for any . Hence, from (B.5), we arrive at
Note that the same probability bounds also applies to and . Thus, we may take to obtain the desired result. Now we provide lower bounds for and . First, we define events , and by
Choosing , we have
Similarly, by the general Hoeffding inequality, one can easily see that
Also, since the operator norm has the following lower bound
by (B.8), (B.11) and (B.12), we arrive at
As for the last inequality on , by definition we know that
We first control the operator norm of the random feature matrix . Since is centered, [FW20, Lemma D.4] implies that
where event is defined by
Next, we estimate the failure probability of event . By Bernstein’s inequality, for any we have
where we write as the first column of (and similarly for all ). Following the proof of Proposition 3.3 in [FW20], we can obtain that
for any . Besides, inequality (B.10) implies that for any ,
By choosing in (B.17) and for sufficient large , we can claim that there exists sufficient large constant such that
Combining (B.15) and the above inequality, we have
In addition, the following tail bound is due to property of (sub-)Gaussian random variables:
for any . Because is Lipschitz, is a sub-Gaussian random vector with similar tail bound
Let . Applying all these three tail bounds (B.19) and (B.10), (B.13) gives us the first part of the probability bound in . As for the second part, following the observation
we can adopt (B.8), (B.9) and (B.10) to conclude the second probability bound.
where the last inequality can be deduced by
which is uniformly bounded by a constant. Therefore, by (B.7) and (B.21), we get
As for the tail control, because of Fact 13, we consider the following upper-bound,
where the last inequality is due to being upper-bounded by .
To control \mathopen{}\mathclose{{}\left\|{\sigma(\boldsymbol{X}\boldsymbol{W}_{0})\boldsymbol{a}}}\right\|_{\infty}, note that since is centered by Assumption 1, we can apply Bernstein inequality for and conditioned on the event \mathcal{M}:=\mathopen{}\mathclose{{}\left\{\mathopen{}\mathclose{{}\left|\|\boldsymbol{x}_{i}\|/\sqrt{d}-1}\right|\leq\nicefrac{{1}}{{2}},~{}i\in[n]}\right\}. Conventionally, we denote \mathopen{}\mathclose{{}\left\|{\cdot}}\right\|_{\psi_{2}} as the sub-Gaussian norm. Since \big{\|}\|\boldsymbol{x}_{i}\|-\sqrt{d}\big{\|}_{\psi_{2}} is bounded by some absolute constant ([Ver18, Theorem 3.1.1]), we know that
Notice that for any , is the sum of independent and centered sub-Exponential random variables, where, in terms of [FW20, Lemma D.5], the sub-Exponential norm \mathopen{}\mathclose{{}\left\|{\cdot}}\right\|_{\psi_{1}} of each term is bounded by the sub-Gaussian norm of the entries as follows,
for some absolute constant . Thus, by Bernstein inequality [Ver18, Theorem 2.8.1], for each ,
Then we take the union over all and obtain \mathopen{}\mathclose{{}\left\|{\sigma(\boldsymbol{X}\boldsymbol{W}_{0})\boldsymbol{a}}}\right\|_{\infty}\leq\log n with probability at least . Hence, by (B.10), (B.20) and (B.22), we get
Part is established by choosing . This concludes the proof of the lemma. ∎
Proposition 2 is a direct consequence of the above norm bounds.
Proof of Proposition 2. Notice that . In the proportional regime, by Lemma 14, there exist universal constants such that
On the other hand, part in Lemma 14 implies that
for some constant . Here we used the fact because it is a rank-one matrix. Conditioning on the two events stated above, we have
As long as is sufficiently large such that , we can obtain
B.1.2 Decomposition of Matrix A
Using the orthogonal decomposition (3.4), we can further decompose the rank-1 matrix as follows
for some constants that only depend on , and .
The expectation follows from (B.7) and the following inequality,
The probability bound also follows from the same argument as Lemma 14.
Following the proof of part in Lemma 14, we can further decompose \mathopen{}\mathclose{{}\left\|{\boldsymbol{A}}}\right\|_{F} into
Since is a Lipschitz function as well, we can again apply the Lipschitz concentration (B.9). Hence, combining (B.10), (B.8) and (B.9), one can conclude the bound on the expectation of .
whose squared Frobenius norms are given by
conditioned on event . Thus, the Hanson-Wright inequality (Theorem 6.2.1 [Ver18]) indicates that
Thus, by choosing and employing (B.31), we have
where we simplified the expression using the assumption that .
We can further decompose into two parts:
On the other hand, Bernstein’s inequality [Ver18, Theorem 2.8.1] indicates that for all ,
for all . As for , we apply (B.35) and (B.40) to all . Letting in (B.40) and taking union bounds for all , we obtain
Hence, the above equation and (B.35) lead the following bound
Therefore, by letting in (B.41) and combining (B.31) and (B.43), we can conclude that
Finally for , we may employ a similar decomposition as ,
B.1.3 Multiple Gradient Steps
We first control the difference in the prediction of the trained neural network compared to the initialized model. Following the same argument as [OS20, Setion 6.6.1], we know that
Note that \mathopen{}\mathclose{{}\left\|{\boldsymbol{W}_{t}-\boldsymbol{W}_{0}}}\right\|_{F}=\mathcal{O}(1) with high probability due to the induction hypothesis. We now compute the next gradient update (we drop the learning rate ).
For , following the same argument as Lemma 14, we have
Now recall that \mathopen{}\mathclose{{}\left\|{f_{0}(\boldsymbol{X})}}\right\|\leq\frac{1}{\sqrt{N}}\mathopen{}\mathclose{{}\left\|{\boldsymbol{a}}}\right\|\mathopen{}\mathclose{{}\left\|{\sigma(\boldsymbol{X}\boldsymbol{W}_{0})}}\right\|. Combining the norm control of in (B.8), (B.20), the norm control of due to (B.9) and (B.20), the operator norm bound on and \mathopen{}\mathclose{{}\left\|{\sigma(\boldsymbol{X}\boldsymbol{W}_{0})}}\right\| given in (B.10) and (B.19) (where we applied [FW20, Lemma D.4] to the matrix , since is centered), and the upper bound on \mathopen{}\mathclose{{}\left\|{f_{t}(\boldsymbol{X})}}\right\| given in (B.49), we arrive at
for large enough and constants . Similarly for , we have
Again using [OS20, Setion 6.6.1], we have
Thanks to the norm control of in (B.10), the norm control of and from (B.8), (B.9), (B.20), and (B.49), and the operator norm of the CK matrix given in (B.19), we get
for large enough . Consequently, given the induction hypothesis, we know that for the next time step with learning rate , there exist some constants such that
B.2 Calculation of Alignment with Target Function
In this section we prove Theorem 3. We first characterize certain quadratic forms which will appear in many parts of our analysis.
The following lemma is a direct adaptation from Lemma 2.7 and Lemma A.1 in [BS98]. We also refer readers to section B.5 in [BS10] for more details.
where is a universal constant. Furthermore, if is a non-negative definite matrix, then we have
Equipped with Lemma 17, we introduce a quadratic concentration lemma specialized to our setting.
Proof. We first consider the concentration for . Note that has i.i.d. columns. Hence, we can expand the first quadratic form as follows:
as , where is obtained by Lemma 2.2 in [BS98] because are i.i.d. for , and the last inequality (B.82) follows from (B.73), (B.76) and (B.80). This yields the convergence in probability.
For the second part , notice that
is a sample mean of i.i.d. centered random variables. Therefore by Lemma 17, we have
B.2.2 Analysis of Spike in Weight Matrix
Observe that the limiting eigenvalue distribution of is the Marchenko–Pastur distribution with parameter ; let be the Stieltjes transform of . Also, we denote the limiting eigenvalue distribution for by whose Stieltjes transform is , which is referred to as the companion transform. The relation between and is given as
Moreover, is uniquely determined by the fixed-point equation
Following the above notions, we define the resolvent , for
and any small . Then, under the same assumptions of Theorem 3, for all sufficiently large , uniformly for all .
Proof. Write , with x>\big{(}1+\sqrt{\psi_{2}}\big{)}^{2}+\epsilon. For any , let be the -th eigenvalue of . Then we have
By Theorem 5.11 in [BS10], for sufficiently large , \mathopen{}\mathclose{{}\left|\lambda_{1}-\big{(}1+\sqrt{\psi_{2}}\big{)}^{2}}\right|\leq\epsilon/2. Hence, for and for all .
The following lemma characterizes the asymptotics of certain quantities in terms of the Stieltjes transform which will be useful in the subsequent analysis.
Recall the definition . Under the same assumptions as Theorem 3, for any ,
in probability as and , uniformly on any compact subset of defined in (B.89), where scalars and are defined in Theorem 3.
Proof. Firstly note that (B.90) directly follows from Lemma 19 and Hoeffding’s inequality for . The remaining concentration statements will be established by applying Lemma 18 to different choices of and the Hanson-Wright inequality for . In particular, due to Lemma 19, we know that , , and are all uniformly bounded on for large . Take in Lemma 18, we obtain that
uniformly on any compact subset of , due to (B.87), Lemma 7.4 of [DW18] and Lemma 2.14 of [BS10].
Finally, we recall the following control of singular values (also referred to as Weyl’s inequality).
By (B.95), the equality , and the Sylvester’s determinant identity for of appropriate dimensions, the above equations has the same solution as
Now we compute the root of on \mathopen{}\mathclose{{}\left((1+\sqrt{\psi_{2}})^{2},+\infty}\right). From (B.88) we have
which implies that the root of on the real line is given as . Based on the expression of and (B.88), we have
Since does not reside in the spectrum of , we can further write
By multiplying from the left hand side of the above equality, we arrive at , where the 2-by-2 matrix is given as
In addition, multiplying from the left hand side of (B.98) yields
We first consider the scenario , where in probability and is outside the support of . For sufficiently small and all large , , and thus Lemma 20 gives
in probability as proportionally. Notice that both and are uniformly bounded by some constants with high probability since , almost surely, and Lemma 18 implies in probability as . Therefore, when and , (B.100) and (B.103) provide the limits of and , which we denote by and , respectively. More precisely,
Thus, we can apply formula (B.87), condition in (B.96) and the following well-known facts of the Stieltjes transform (e.g., see [BS10]) with :
On the other hand, if , we have proved that is approaching to the right-edge of the bulk of . In fact, with probability one, , where is the largest eigenvalue of ; this is because \det\boldsymbol{M}_{n}(z)=\mathopen{}\mathclose{{}\left(\boldsymbol{u}^{\top}\boldsymbol{Q}_{0}(z)\boldsymbol{W}_{0}\boldsymbol{a}+1}\right)^{2}+\boldsymbol{u}^{\top}\boldsymbol{Q}_{0}(z)\boldsymbol{u}\mathopen{}\mathclose{{}\left(\|\boldsymbol{a}\|^{2}-\boldsymbol{a}^{\top}\boldsymbol{W}_{0}^{\top}\boldsymbol{Q}_{0}(z)\boldsymbol{W}_{0}\boldsymbol{a}}\right) satisfies
while . Hence by the definition of ,
Also, by the Cauchy–Schwarz inequality, we have
Since , given any small , for all sufficiently large , we have
which indicates that \boldsymbol{a}^{\top}\boldsymbol{W}_{0}^{\top}\boldsymbol{Q}_{0}(\hat{\lambda})\boldsymbol{u}\in\mathopen{}\mathclose{{}\left(\frac{1}{-2-\sqrt{\psi_{2}}+\epsilon},\frac{1}{\sqrt{\psi_{2}}-\epsilon}}\right). Therefore, for all large , is bounded by some universal constant related to . Then by (B.109), we can conclude is also asymptotically bounded by some constant. On the other hand, since all eigenvalues of is smaller than , we obtain
This directly implies has a constant upper bound for all large . Following the proofs of [BGN11, Theorem 2.3] and [BGN12, Theorem 2.10] (with slight modifications of Lemma A.2 and Proposition A.3 in [BGN11]), it is straightforward to control the following quadratic forms by verifying the weak convergence of certain weighted spectral measures in combination with the Portmanteau theorem:
Consequently, we know that -\boldsymbol{a}^{\top}\boldsymbol{W}_{0}^{\top}\boldsymbol{Q}_{0}(z)\boldsymbol{W}_{0}\boldsymbol{a}\in\mathopen{}\mathclose{{}\left(1/\sqrt{\psi_{2}},C(1+\sqrt{\psi_{2}})^{2}}\right) for some constant and all large . Hence -\boldsymbol{u}^{\top}\boldsymbol{Q}_{0}(\hat{\lambda})\boldsymbol{u}\in\mathopen{}\mathclose{{}\left(1/(1+\sqrt{\psi_{2}})^{2},C\sqrt{\psi_{2}}}\right). Now notice that
Since , by a simple adaptation of Lemma A.2 in [BGN11], one can directly verify , as proportionally. Hence we conclude that if . The theorem is established by combining the above cases.
Appendix C Proof for Small Learning Rate (η=Θ(1)𝜂Θ1\eta=\Theta(1))
To validate Theorem 5, we follow the proof strategy of [HL20], which established the GET for RF models using the Lindeberg approach and leave-one-out arguments [EK18]. We remark that concurrent to our work, [MS22] proved the Gaussian equivalence property for a larger model class under an assumed central limit theorem, which is verified for two-layer RF or NTK models, and thus cannot directly imply our results on the trained features.
where we abbreviated \boldsymbol{\phi}_{i}=\boldsymbol{\phi}_{\boldsymbol{x}_{i}}=\frac{1}{\sqrt{N}}\sigma(\boldsymbol{W}^{\top}\boldsymbol{x}_{i}),\bar{\boldsymbol{\phi}}_{i}=\bar{\boldsymbol{\phi}}_{\boldsymbol{x}_{i}}=\frac{1}{\sqrt{N}}\mathopen{}\mathclose{{}\left(\mu_{1}\boldsymbol{W}^{\top}\boldsymbol{x}_{i}+\mu_{2}\boldsymbol{z}_{i}}\right) for .
Define the set of weight matrices perturbed from the Gaussian initialization as
Note that for learning rate , we can verify that is a high-probability event after any finite number of gradient steps, as characterized in Lemma 14 and 16. The following proposition is a reformulation and extension of [HL20, Theorem 1], stating that the Gaussian equivalence property holds as long as remains “close” to the initialization .
where and are defined in (C.1) and (C.2).
From Proposition 22 we know that Theorem 5 holds if the optimized weight matrix falls into the set with sufficiently high probability. This condition is in turn verified by Lemma 14 and 16. Also note that in our setting of MSE loss and , the RHS of the above equation is bounded in probability.
Recall the single-index teacher assumption: for . Observe that for , the following near-orthogonality condition between the neurons holds with high probability
Importantly, for satisfying the near-orthogonality condition (C.4), we can utilize the following central limit theorem from [HL20] derived via Stein’s method.
where , and only depends on constant .
Following [HL20], we construct an interpolating sequence between the nonlinear and linear features model. For any , we define
Note that when , setting recovers the estimator on nonlinear features , and similarly, setting gives the estimator on the linear Gaussian features .
We remark that the perturbation allows us to compute the prediction risk by taking the derivative of the objective w.r.t. around 0 — see [HL20, Proposition 1] for details. Note that when \mathopen{}\mathclose{{}\left\|{\boldsymbol{W}}}\right\|=\Theta(1), we may choose \gamma^{*}=\frac{N}{n}\cdot\frac{\lambda/4}{\mu_{1}^{2}\mathopen{}\mathclose{{}\left\|{\boldsymbol{W}}}\right\|^{2}+\mu_{2}^{2}}>0 such that for \mathopen{}\mathclose{{}\left|\gamma_{1}}\right|\leq\gamma^{*},\mathopen{}\mathclose{{}\left|\gamma_{2}}\right|\leq 1, the overall objective (C.6) is -strongly convex (i.e., the strongly-convex regularizer dominates the concave part of when ).
Proof. We follow the proof of [HL20, Lemma 23] and first analyze one coordinate of defined by (C.6), which WLOG we select to be the last coordinate. For concise notation, we instead augment the weight matrix with an -th column and study the corresponding . Denote the weight vector , where is the -th column of the initialized , and is the perturbation (i.e., gradient update for ).
The -th coordinate of interest, which we denote as , can be written as the solution to the following optimization problem,
By [HL20, Equation (249)], we know that for ,
We control each term on the right hand side of (C.9) separately. Note that implies that \mathopen{}\mathclose{{}\left\|{\boldsymbol{\delta}_{N+1}}}\right\|=\mathcal{O}\mathopen{}\mathclose{{}\left(\frac{\text{polylog}d}{\sqrt{d}}}\right) due to the definition (C.3). Since \mathopen{}\mathclose{{}\left|\boldsymbol{\beta}_{*}^{\top}\boldsymbol{w}_{N+1}}\right|\leq\mathopen{}\mathclose{{}\left|\boldsymbol{\beta}_{*}^{\top}\boldsymbol{w}^{0}_{N+1}}\right|+\mathopen{}\mathclose{{}\left\|{\boldsymbol{\delta}_{N+1}}}\right\|\mathopen{}\mathclose{{}\left\|{\boldsymbol{\beta}_{*}}}\right\|, by combining [HL20, Equation (252)] and our assumption that \mathopen{}\mathclose{{}\left\|{\boldsymbol{\beta}_{*}}}\right\|=1, we know that for some constant and large ,
Similarly, \mathopen{}\mathclose{{}\left|\boldsymbol{w}_{N+1}^{\top}\boldsymbol{W}\boldsymbol{g}_{k}^{*}}\right|\leq\mathopen{}\mathclose{{}\left|\boldsymbol{w}^{0^{\top}}_{N+1}\boldsymbol{W}\boldsymbol{g}_{k}^{*}}\right|+\mathopen{}\mathclose{{}\left\|{\boldsymbol{\delta}_{N+1}}}\right\|\mathopen{}\mathclose{{}\left\|{\boldsymbol{W}\boldsymbol{g}_{k}^{*}}}\right\|, and therefore by [HL20, Lemma 17] (note that the lemma only requires to satisfy (C.4)), we have
for some . The case where (i.e., the features are linear) follows from the exact same argument. Also, because of \mathopen{}\mathclose{{}\left|\boldsymbol{r}_{i}^{\top}\boldsymbol{g}_{k}^{*}}\right|\leq\mathopen{}\mathclose{{}\left|\boldsymbol{r}_{i}^{0\top}\boldsymbol{g}_{k}^{*}}\right|+\mathopen{}\mathclose{{}\left\|{\boldsymbol{r}_{i}-\boldsymbol{r}_{i}^{0}}}\right\|\mathopen{}\mathclose{{}\left\|{\boldsymbol{g}_{k}^{*}}}\right\|, we know that [HL20, Equation (257)], [HL20, Lemma 17], and (C.12) together ensure that
for some constant and all large . Finally, since the assumption on \mathopen{}\mathclose{{}\left\|{\boldsymbol{\Delta}}}\right\|_{2,\infty} implies control of \mathopen{}\mathclose{{}\left\|{\boldsymbol{\delta}_{i}}}\right\| for all , we complete the proof by a union bound over the coordinates.
Denote the optimal value of objective (C.6) by
From [HL20, Section 2.3], we know that Proposition 23 and Lemma 24 imply that for any and , the discrepancy due to one swap can be bounded as
Proof of Theorem 5. Finally, we establish Theorem 5 by verifying that in our setting the event occurs with high probability. For one gradient step on the squared loss with learning rate , Lemma 14 together with \mathopen{}\mathclose{{}\left\|{\boldsymbol{\beta}_{*}}}\right\|=1 entail that for proportional , there exists some constant such that
C.2 Prediction Risk of the Gaussian Equivalent Model
Now we compute the prediction risk of the CK ridge estimator on the feature map after one gradient step . We restrict ourselves to the squared loss, the optimal solution of which is given by:
where the bias and variance terms are given as
Also, the risk lower bound for the Gaussian equivalent model is a direct consequence of (C.29).
In the following sections, we compare the bias and variance terms given in (C.25), (C.26) and (C.27) before and after one feature learning step. We first simplify the calculation by showing that the values of these equations remain asymptotically unchanged if we remove certain low-order terms.
We now control the errors in the bias and variance terms after ignoring the lower-order terms in the weight matrix.
Given Assumptions 1, 2 and . Then for , we have
as at comparable rate, where we dropped the constant in .
For the bias terms, we consider perturbation on in the operator norm. Again, Lemma 14 entails that
Based on this result, it is straightforward to show that
where in we used the fact that is Lipschitz and \mathopen{}\mathclose{{}\left\|{\boldsymbol{\beta}_{*}}}\right\|=1 (for example see [BMR21, Lemma A.12]). The statement is proved by combining all the above calculations.
Lemma 26 entails that the variance term in the risk does not change after one gradient step with , and for the bias terms, we may consider the rank-1 approximation of the gradient matrix studied in Proposition 2 instead. In the following section, we use this property to simplify the risk expressions.
C.3 Precise Characterization of Prediction Risk
In the following subsections we will characterize the limiting value of each as proportionally.
In the following lemma, we show that each will concentrate around some given by
Under Assumptions 1 and 2, as proportionally, we have
which implies that the second term in (C.57) converges to zero in probability. Hence we only need to control
for some constant . Therefore, (C.58) concentrates around with high probability. Thus it remains to show that is vanishing in probability.
Now condition on event , the Lipschitz property of entails that
Thus, as , with high probability, which finally implies that converges to zero in probability as proportionally. This completes the proof. ∎
Finally, we use the following simplification of quadratic forms to obtain the desired .
C.3.2 Risk Calculation via Linear Pencils
In this section, we derive analytic expressions of the terms defined in (C.49) as proportionally. In particular, the exact values are described by self-consistent equations defined in the following proposition.
Given Assumption 1 and . For each defined in (C.49) and we have
in probability, as and , where ’s are defined as follows
It is straightforward to verify all the above limits exist and are finite. Finally, in the following analysis we will repeatedly make use of the following identities:
Note that and defined in (C.79) and (C.80) have been characterized in prior works, such as Proposition 1 of [AP20]; in particular, since we are only interested in the CK, we can simply set in [AP20] (which considered the sum of the CK and the first-layer NTK). Therefore, from (C.81) we obtain \tau_{1}=\tau_{1}\mathopen{}\mathclose{{}\left(\frac{\psi_{1}}{\psi_{2}}\lambda}\right) from m_{1}:=m_{1}\mathopen{}\mathclose{{}\left(\frac{\psi_{1}}{\psi_{2}}\lambda}\right). As for the limit of , (C.83) indicates
with . can also be derived in similar fashion.
For we utilize the computations in Appendix I.6.1 of [TAP21] by setting the covariance . More precisely, based on Equations (S370) and (S418) in [TAP21],
where and are both due to (C.80) and (C.81). Hence we obtain the formulae of and .
Recall the following derivative trick of the Stieltjes transform,
Once again by setting in [AP20], we can directly employ the computation of in Section S4.3.4 of [AP20] to derive and . In particular, due to (C.82),
Note that Equation (S148) in [AP20] established that
whence, letting , we obtain and .
Having obtained the asymptotic expressions of each term in the decomposition of the prediction risk, we can now compute the difference in the prediction risk of CK ridge regression before and after one gradient descent step, i.e., in Theorem 7. The following statement is the complete version of Theorem 7.
Given Assumptions 1 and 2, consider . Fix and . Denote and as the prediction risk of CK ridge regression in (4.1) using initial weight and first-step updated , respectively. Then the difference between these two prediction risk values satisfies
where is a non-negative function of and with parameters given as
Here the scalars ’s are defined in Proposition 29. Furthermore, if and only if at least one of and is zero.
Proof. Due to Lemma 26 (or the decomposition (C.25), (C.26) and (C.27)), we can see that variance is unchanged after one gradient descent step with . Hence we only need to analyze the changes in (C.25) and (C.26). Also, due to Lemma 26 and the proof of Theorem 3, we can ignore and in and take , where and , without changing the bias terms.
First note that if , then and therefore as . In the following, we take which implies that defined in Theorem 3 will not vanish. Now we aim to extract the low-rank perturbation from bias terms (C.25) and (C.26). We adhere to the notions in (C.44), (C.49) and (C.54) and define . Similar to [MM22, Lemma C.1], we use the following linearization trick to separate the gradient step from the matrices and .
Therefore, by the Sherman-Morrison-Woodbury formula and Hanson-Wright inequality, we have
We decompose the subtracted term in (C.25) into
Now we denote , and . Hence,
Analogously, we can decompose in (C.26) as follows
where we repeatedly make use of Lemma 27 and the concentration for to simplify the computations. Therefore, one can obtain
On the other hand, from Proposition 29 we know that
where the right hand side is the quantity of interest defined in Theorem 7. Also observe the following equivalences from Proposition 29,
Hence, we can simplify as follows
Finally, we validate that the function is non-negative on variables and . Observe that the formula of in (C.111) is decomposed into two parts. From Proposition 29 we know that and are the limits of and evaluated at ; this indicates that is non-negative. For the same reason, and . Also due to Proposition 29, we have
Therefore, and . This entails that the first part of is non-negative:
As for the second part, it suffices to evaluate since
Plugging in quantities in (C.112) with , we have
where and are due to (C.80) and (C.86), respectively. By Lemma A.1 in [TAP21], we know function has non-positive derivative when . This implies that and hence the second part of is also non-negative.
Finally, we note that when , the function . This is because
when . Whereas when , we know that , which entails is also vanishing. Also observe that in (C.112), are all positive. Hence we conclude that if , then at least one of must be zero.
C.4 Analysis of Special Cases
While the previous subsection provides explicit formulae of , the expressions are rather complicated due to the self-consistent equations (C.79) and (C.80). In this section we consider two special cases: the large sample limit and the large width limit , where the calculation simplifies and enables us to further characterize properties of . In both cases, we start with Theorem 30 and take one of aspect ratios ( or ) to infinity.
In this subsection we prove Proposition 8. We introduce two positive parameters
where is Marchenko–Pastur distribution with rate . Now we consider the large-sample limit: . The following statement is the formal version of Proposition 8, and compared to the general result (Theorem 30), this special case admits a more explicit formula only determined by and .
Under the same assumptions as Theorem 30 and take . Then the difference between the prediction risks before and after one feature learning step satisfies
In this case is a non-negative function of , and if and only if one of is zero. Furthermore, is increasing with respect to the learning rate .
Proof. Following Theorem 30, it suffices to consider the limit of when . This reduces to simplifying the asymptotics of ’s defined in Proposition 29, as is determined by ’s in Theorem 30. We aim to prove the following:
as , where and are defined in (C.119). The trivial cases when have been studied in Theorem 30. So, WLOG, we assume in the following derivations.
Recall the definitions of and . One can easily see that as . For any , (C.79) and (C.80) can be written as follows
Notice that , and based on [FW20, Theorem 3.4], for any ,
Therefore, when , the measure \mu^{\text{MP}}_{\psi_{2}/\psi_{1}}\boxtimes\mathopen{}\mathclose{{}\left(\mu_{2}^{2}+\mu_{1}^{2}\cdot\mu^{\text{MP}}_{\psi_{2}}}\right) reduces to a deformed Marchenko–Pastur law \mathopen{}\mathclose{{}\left(\mu_{2}^{2}+\mu_{1}^{2}\cdot\mu^{\text{MP}}_{\psi_{2}}}\right); hence by the definition of ,
which verifies the first statement in (C.124). Now recall the value of interest . Due to the relationship between and , it is straightforward to deduce that as . As for , in terms of (C.126), we have
Next we compute the limit of . Taking derivative with respect to at both sides of (C.128), we arrive at ; here represents the derivative at . Combining this relation and (C.129),(C.81), we can deduce that
Note that here we used when . Lastly, for and , by (C.126),
where we applied the previously established convergence of and . Also recall that (C.130) implies that converges to 1/\mathopen{}\mathclose{{}\left(1+\psi_{2}\mu_{1}^{2}s_{1}}\right) as . Together with the convergence of in (C.124), we get
which implies the convergence of in (C.124).
As a result, by replacing ’s in (C.96) with the corresponding reparameterized ’s in (C.124), we arrive at the following expression of :
By definitions of in (C.121), (C.122) and (C.123), we can see that , and ; this leads to the equivalent expression
Now we claim that are all non-negative, for any . With a slight abuse of terminology, in the following we denote . Recall that and ; We can therefore simplify (C.139), (C.140) and (C.141) as follows
where is due to (C.132) and is obtained by taking derivative with respect to in (C.131). In addition,
which implies that . We also denote the companion Stieltjes transform of by , which is the Stieltjes transform of the limiting eigenvalue distribution for . Recall the following relation between and : \bar{m}(z)+\frac{1}{z}=\psi_{2}\mathopen{}\mathclose{{}\left(m(z)+\frac{1}{z}}\right). Since is positive, we can deduce that
where the last equality is obtained by taking derivative of (B.87) on both sides with respect to . In summary, we have shown that and when . Hence by definition, are all non-negative and so is .
Finally, we verify that is an increasing function of . Observe that only appears in in the expression of in (C.138). Hence, it suffices to take the derivative of with respect to and verify that this partial derivative is positive. One can check that
By the definition of in (C.141), we have \mathopen{}\mathclose{{}\left(\gamma-s_{1}\beta}\right)=\alpha(s_{1}-\lambda s_{2}). Also from (C.119) we know that . Finally, recall that and ; this implies is increasing with regard to and completes the proof.
C.4.2 Case II: Highly overparameterized regime
Next we consider the large width limit and establish Proposition 9.
Note that , for all , and thus . On the other hand, is compactly supported. Therefore, by taking and letting at both sides of (C.147), we arrive at
which is a finite positive value determined by . With this in mind, we conclude that is also finite, since is determined by (C.126) and we can take with . In addition, since for any , we may take the derivative with respect to at both sides of (C.147) to obtain , and take and to conclude that the limit of is finite as well. Similarly, by taking derivative with respect to in (C.126), one can also verify that as , the limit of remains finite. From these estimates we know that
are vanishing as , whereas
will converge to some finite values. The proposition is established based on the definition of with the help of the above statements.
Appendix D Proof for Large Learning Rate (η=Θ(N)𝜂Θ𝑁\eta=\Theta(\sqrt{N}))
In this section we restrict ourselves to a single-index target function (generalized linear model): , and study the impact of one gradient step with large learning rate . For simplicity, we denote where is a fixed constant not depending on .
Recall that , where is defined in Lemma 14 and 15, and the full-rank term is given as
We first refine the estimate on the Frobenius norm of certain submatrix of ; the choice of such submatrices will be explained in Section D.2.
Proof. Let and . Then, matrix can be written as
Here, represents the sum for distinct and is the sum when . Therefore,
where independent of and . We compute the aforementioned expectations as follows
for some constant . In addition, based on Lemma D.3 and G.1 in [FW20], we can show the following inequality for any under event :
Let and . Conditioned on , we know that
For , by Taylor expansion of around (note that is differentiable by assumption), there exists a random variable between and such that
on the event with , where is a constant depending on . This concludes (D.6).
Also, note the probability bound for Gaussian random vector and implies that
From the above arguments, we can bound the first term via the following steps:
By choosing , we conclude that cannot exceed with probability at least , for any . As for , since is uniformly bounded by and all entries of are bounded by , we have
We conclude (D.2) by combining the above estimates of and .
D.2 Constructing the “Oracle” Estimator
In this subsection we prove the following lemma related to Lemma 10.
where the scalar is defined in (4.3).
Denote N_{r}:=\mathopen{}\mathclose{{}\left|\mathcal{A}_{r}^{\alpha}}\right| for some constant , and as the index such that . We define as an average over neurons with indices , and as an approximation of in which the first-step gradient matrix in (B.3) is replaced by the rank-1 matrix defined in (B.27):
Moreover, by definition of , all ’s are close to for ; thus Lemma 32 (in particular (D.2)) can be directly applied to . As for , since , we use part in Lemma 14 to obtain
With these concentration estimates, we know that when ,
with probability at least 1-c\mathopen{}\mathclose{{}\left(\frac{\alpha^{2}}{d^{\frac{1}{4}}}+\frac{\alpha^{4}}{n}+\frac{1}{\sqrt{N}}+ne^{-c\log^{2}n}+Ne^{-\log^{2}N}}\right) for some constant ; this is due to the defined step size , (D.15), (D.16) in (D.2) of Lemma 32 outlined above. In (D.19), we ignore the constants in the upper bound since we are only interested in the rate with respect to .
By the definition of and the Lipschitz property of , one can obtain
Now define , which corresponds to the “population” version of , and denote
Combining the inequalities (D.20) and (D.22), we know that for some constant ,
where the last inequality holds with probability at least 1-\operatorname{exp}\mathopen{}\mathclose{{}\left({-cd}}\right) for some universal constant , due to the operator norm bound and concentration of the sample covariance matrix (for instance see [Ver18, Theorem 4.6.1]).
Now we take the expectation of over initial weight in (D.21) to define
Note that for fixed , \langle\boldsymbol{w},\boldsymbol{x}\rangle\sim\mathcal{N}(0,\mathopen{}\mathclose{{}\left\|{\boldsymbol{x}}}\right\|^{2}/d). Since is -Lipschitz, by the Hoeffding bound on sub-Gaussian random variables, conditioned on , we have
where the last inequality is due to property of the sub-Gaussian norm (see e.g. [Ver18, Theorem 3.1.1]) for some universal constant .
as , for some constant . By the Cauchy-Schwarz inequality,
where the failure probability only relates to and is vanishing as . For simplicity, we only keep the leading orders and ignore the subordinate terms in the exact probability bounds.
for some constant , as proportionally.
The above analysis illustrates that because of the Gaussian initialization of , for any , we can find a subset of neurons that receive a “good” learning rate, in the sense that the corresponding (sub-) network defined by can achieve the prediction risk close to when .
Equation (D.10) reduces the prediction risk of our constructed to a one-dimensional Gaussian integral, which can be numerically evaluated for pairs of . Denote , we give a few examples in which we set and the corresponding is small. Note that due to Assumptions 1 and 2, choices of and considered below are centered with respect to standard Gaussian measure .
. Numerical integration yields , .
. Numerical integration yields , .
. Numerical integration yields , .
Observe that in all the above examples, can be obtained by some finite (or equivalently ). In the following analysis of kernel ridge regression, we drop the small constant in Lemma 33 and directly apply the asymptotic statement given in (D.39).
We make the following remarks on the calculation of in (4.3).
When , we intuitively expect to be small when the nonlinearity is smooth such that it is to some extent unchanged under Gaussian convolution (when is chosen appropriately).
Adding weight decay with strength to the first-layer parameters simply corresponds to multiplying in the definition of (D.30) by a factor of .
D.3 Prediction Risk of Ridge Regression
In this section we prove Theorem 11. Recall that we aim to upper-bound the prediction risk of the CK ridge regression estimator defined as
We are interested in the prediction risk of the CK ridge regression estimator denoted as . We first define the following quantities which can be decomposed into (see Lemma 35):
We begin by defining a concentration event on the empirical feature matrix , under which the prediction risk can be controlled. We modify the proof of [Ver18, Theorem 4.7.1] to obtain a normalized version of the concentration for CK matrix as follows.
Under Assumptions 1, 2 and using the above notations, there exists some constant such that the following holdsNote that for , the LHS of the inequality may be interpreted as a pseudo-inverse.
for all large , where .
Proof. First observe that the null space of contains the null space of . Also, notice that is a sample covariance matrix taking the form of
for all large . This proposition is proved by setting and noting that
Similarly, for the “ridgeless” case , we define
Lemma 34 entails that both and hold with probability at least . Following the remark on [Bac23, Lemma 7.1], under events and , we can obtain that
and , which implies that
since \mathopen{}\mathclose{{}\left\|{\widehat{\boldsymbol{\Sigma}}_{\Phi}\mathopen{}\mathclose{{}\left(\widehat{\boldsymbol{\Sigma}}_{\Phi}+\lambda\boldsymbol{I}}\right)^{-1}}}\right\|\leq 1. Analogously, we claim that \mathopen{}\mathclose{{}\left(\widehat{\boldsymbol{\Sigma}}_{\Phi}+\lambda\boldsymbol{I}}\right)^{-1/2}\boldsymbol{\Sigma}_{\Phi}\mathopen{}\mathclose{{}\left(\widehat{\boldsymbol{\Sigma}}_{\Phi}+\lambda\boldsymbol{I}}\right)^{-1/2}\preccurlyeq\frac{1}{1-t}\boldsymbol{I}. Thus, under events and , we know that
We now control under the high probability events and .
By the definition of , we have
Following [Bac23, Proposition 7.2], we define \boldsymbol{a}_{\lambda}=\boldsymbol{\Sigma}_{\Phi}\mathopen{}\mathclose{{}\left(\boldsymbol{\Sigma}_{\Phi}+\lambda\boldsymbol{I}}\right)^{-1}\check{\boldsymbol{a}} and obtain
Therefore, we know that under events and ,
where follows from the definition of the concentration events , (D.44) and (D.46).
Finally, from [Bac23, Lemma 7.2], we have
where the last step is a triangle inequality due to .
For , note that under event ,
Similarly for , under event , we have
where is due to the boundedness of and . Combining and , and taking in (D.62) and (D.65), for some and any small , we arrive at
with probability at least .
The following lemma provides a decomposition of the prediction risk in terms of analyzed above.
Under the same assumptions as Lemma 10, if we choose for small , then the prediction risk of the CK ridge estimator admits the following upper bound
where are defined in (D.41).
Proof. Based on the definition of prediction risk, we have
Proof of Theorem 11. Since Lemma 34 ensures that events and happens with high probability for fixed , if we set for some small , then Lemma 35 entails
Finally, due to the upper-bound (D.76), we conclude that
with probability one as proportionally and , where is defined in (4.3).