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 NN neurons,

When the first layer W\boldsymbol{W} is fixed and only the second layer a\boldsymbol{a} is optimized, we arrive at a kernel model, where the kernel defined by features x→σ(W⊤x)\boldsymbol{x}\to\sigma(\boldsymbol{W}^{\top}\boldsymbol{x}) (often called the hidden representation) is referred to as the conjugate kernel (CK) [Nea95]. When W\boldsymbol{W} 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 W\boldsymbol{W} 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 x→σ(W⊤x)\boldsymbol{x}\to\sigma(\boldsymbol{W}^{\top}\boldsymbol{x}), before and after the gradient descent stepSome of our results also apply to multiple gradient steps on the first layer W\boldsymbol{W}, 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 W\boldsymbol{W} 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) f∗f^{*}, and the top eigenvector of the CK matrix aligns with the training labels y\boldsymbol{y}.

Next in Section 4 we study how the aforementioned alignment improves the kernel. We consider a more specialized setting where the teacher f∗f^{*} 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 L2L^{2}-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 W\boldsymbol{W}, 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: η=Θ(1)\eta=\Theta(1). 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 W\boldsymbol{W} with learning rate η=Θ(1)\eta=\Theta(1); 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: η=Θ(N)\eta=\Theta(\sqrt{N}). In Section 4.3, we analyze a larger learning rate that coincides with the maximal update parameterization in [YH20]. For certain target functions f∗f^{*}, 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 n=Ω(d2)n=\Omega(d^{2}) is required to go beyond this “linear” regime. As we will see in certain settings, such limitation can also be overcome (in the n≍dn\asymp d 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 f∗f^{*}, (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. n,d,N→∞n,d,N\to\infty, n/d→ψ1n/d\to\psi_{1}, N/d→ψ2N/d\to\psi_{2}, where ψ1,ψ2∈(0,∞)\psi_{1},\psi_{2}\in(0,\infty).

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 ψ1\psi_{1}, and increasing the network width corresponds to enlarging ψ2\psi_{2}. The proportional scaling of n,d,Nn,d,N (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 {xi,yi}i=1n\{\boldsymbol{x}_{i},y_{i}\}_{i=1}^{n}, 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 gg satisfies certain smoothness conditions as in [EK10]. Denote the associated RKHS as H\mathcal{H}, and [K]ij=k(xi,xj)[\boldsymbol{K}]_{ij}=k(\boldsymbol{x}_{i},\boldsymbol{x}_{j}). The kernel ridge estimator is given by

We denote the prediction risk of the above kernel estimators as RCK(λ),RNTK(λ),Rker(λ)\mathcal{R}_{\text{CK}}(\lambda),\mathcal{R}_{\text{NTK}}(\lambda),\mathcal{R}_{\text{ker}}(\lambda), 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 W1\boldsymbol{W}_{1} in the two-layer NN (1.1). We first show that the first gradient step on W\boldsymbol{W} can be approximated by a rank-1 matrix, which contains information of the training labels y\boldsymbol{y}. Based on this property, we provide a signal (spike) plus noise (bulk) decomposition of W1\boldsymbol{W}_{1}, and prove that the isolated singular vector is aligned to the linear component of the teacher f∗f^{*}.

Define G0=1ηN(W1−W0)\boldsymbol{G}_{0}=\frac{1}{\eta\sqrt{N}}(\boldsymbol{W}_{1}-\boldsymbol{W}_{0}) and a rank-1 matrix A:=μ1nNX⊤ya⊤\boldsymbol{A}:=\frac{\mu_{1}}{n\sqrt{N}}\boldsymbol{X}^{\top}\boldsymbol{y}\boldsymbol{a}^{\top}. Under Assumption 1, there exist some constants c,C>0c,C>0 such that for all large n,N,dn,N,d, with probability at least 1−ne−clog⁡2n1-ne^{-c\log^{2}n},

Proposition 2 suggests that the first-step gradient can be approximated in operator norm by a rank-1 matrix A\boldsymbol{A}; thus, when the learning rate is reasonably large, we expect a “spike” to appear in the updated weight matrix W1\boldsymbol{W}_{1}. Intuitively, since this rank-1 direction relates to the label vector y\boldsymbol{y}, the resulting W1\boldsymbol{W}_{1} may be “aligned” to the target function f∗f^{*}. This intuition is confirmed in the next subsection.

In other words, if we write η=Θ(Nα)\eta=\Theta(N^{\alpha}), then α≥0\alpha\geq 0 is required so that the change in the weight matrix is non-negligible (one may verify that for η=od(1)\eta=o_{d}(1), the test performance of kernel ridge regression remains unchanged after one GD step). On the other hand, when α>1/2\alpha>1/2, the gradient “overwhelms” the initialized parameters W0\boldsymbol{W}_{0}, and the preactivation feature ⟨x,wi⟩\langle\boldsymbol{x},\boldsymbol{w}_{i}\rangle in the NN (1.1) becomes unbounded as N→∞N\to\infty. 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 f∗f^{*},

When η=Θ(1)\eta=\Theta(1) in (3.2), we show a BBP phase transition (named after Baik, Ben Arous, Péché [BAP05]) for the leading singular value of W1\boldsymbol{W}_{1}, and quantify the alignment between the corresponding singular vector u1\boldsymbol{u}_{1} and the linear component of target function β∗\boldsymbol{\beta}_{*}. It is worth noting that in our analysis, the signal β∗\boldsymbol{\beta}_{*} is “hidden” in the rank-one perturbation A\boldsymbol{A} 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 ∣⟨u1,β∗⟩∣|\langle\boldsymbol{u}_{1},\boldsymbol{\beta}_{*}\rangle| does not directly follow from classical results on the BBP transition.

Then the leading singular value s1(W1)s_{1}(\boldsymbol{W}_{1}) and the corresponding left singular vector u1\boldsymbol{u}_{1} satisfy

if θ1>ψ21/4\theta_{1}>\psi_{2}^{1/4}; otherwise, s1(W1)→1+ψ2s_{1}(\boldsymbol{W}_{1})\to 1+\sqrt{\psi_{2}} and ∣⟨u1,β∗⟩∣→0|\langle\boldsymbol{u}_{1},\boldsymbol{\beta}_{*}\rangle|\to 0, in probability, as n,N,d→∞n,N,d\to\infty.

We make the following observations. Beyond the threshold θ1>ψ21/4\theta_{1}>\psi_{2}^{1/4}, increasing the learning rate η\eta enlarges the leading singular value (spike) s1(W1)s_{1}(\boldsymbol{W}_{1}). As for the overlap, one can numerically verify ∣⟨u1,β∗⟩∣2|\langle\boldsymbol{u}_{1},\boldsymbol{\beta}_{*}\rangle|^{2} is upper-bounded by θ24−ψ2θ22(θ22+1)<1\frac{\theta_{2}^{4}-\psi_{2}}{\theta_{2}^{2}(\theta_{2}^{2}+1)}<1 (obtained when ψ1=n/d→∞\psi_{1}=n/d\to\infty), from which we deduce that better alignment is achieved when we take a bigger step, or when the nonlinearity σ\sigma and target f∗f^{*} have larger linear components (i.e., larger μ1,μ1∗\mu_{1},\mu_{1}^{*}).

Theorem 3 is numerically verified in Figure 3. Observe that after one gradient step with η=Θ(1)\eta=\Theta(1), the bulk of the spectrum of W\boldsymbol{W} remains unchanged and is given by the Marchenko-Pastur law (red), but a spike may appear (prediction from Theorem 3 is indicated by marker “×\times”) when η\eta exceeds a certain threshold; furthermore, the corresponding singular vector u1\boldsymbol{u}_{1} aligns with the linear component β∗\boldsymbol{\beta}_{*} 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 W\boldsymbol{W}, 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 σ\sigma, the expected feature matrix (after one gradient step with η=Θ(1)\eta=\Theta(1)) satisfies

Consequently, Theorem 3 implies the same BBP transition for ΣΦ\boldsymbol{\Sigma}_{\Phi}. When the population ΣΦ\boldsymbol{\Sigma}_{\Phi} 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 η=Θ(1)\eta=\Theta(1), which we denote as CK1=ΦΦ⊤\mathbf{CK}_{1}=\boldsymbol{\Phi}\boldsymbol{\Phi}^{\top}. Observe that the bulk of the spectrum remains unchanged compared to CK0\mathbf{CK}_{0}, which can be analytically computed (red). On the other hand, similar to W1\boldsymbol{W}_{1}, an isolated eigenvalue (spike) appears in CK1\mathbf{CK}_{1}, the location of which can be predicted by the Gaussian equivalent model in Conjecture 4 (marker “×\times”).

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 x→σ(W1⊤x)\boldsymbol{x}\rightarrow\sigma(\boldsymbol{W}_{1}^{\top}\boldsymbol{x}) “adapts” to the teacher f∗f^{*}, 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 f∗f^{*} is a single-index model, and compare the CK prediction risk before and after one gradient descent step on W\boldsymbol{W}.

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 a\boldsymbol{a} on top of fixed feature map (e.g., defined by randomly initialized W0\boldsymbol{W}_{0}), and such RF models cannot learn a single-index f∗f^{*} efficiently in high dimensions [YS19].

As stated in Section 2.3, the RF ridge estimator defined by the two-layer NN (1.1) has Ω(1)\Omega(1) prediction risk unless σ∗\sigma^{*} 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 η=Θ(1)\eta=\Theta(1), the Gaussian equivalent model provides an accurate description of the prediction risk of ridge regression (on the trained CK) at any fixed time step tt, although most of our analysis deals with t=1t=1. The important observation is that even though the trained weights Wt\boldsymbol{W}_{t} are no longer i.i.d., the Gaussian equivalence property can still hold when Wt−W0\boldsymbol{W}_{t}-\boldsymbol{W}_{0} 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 P>1f∗\textsf{P}_{>1}f^{*} as follows.

Hence, when η=Θ(1)\eta=\Theta(1), even though training the first-layer W\boldsymbol{W} 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 f∗f^{*}, the trained feature map needs to violate the GET. In the case of one gradient step on W\boldsymbol{W}, 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 η=Θ(1)\eta=\Theta(1), 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 G0\boldsymbol{G}_{0} shown in Section 3.1. Therefore, in this subsection we focus on t=1t=1 and analyze how the trained features improves over the initialized RF. To quantify the discrepancy in the prediction risk (4.1), we write R0(λ)\mathcal{R}_{0}(\lambda) as the prediction risk of the initialized RF ridge regression estimator (on the feature map x→σ(W0⊤x)\boldsymbol{x}\to\sigma(\boldsymbol{W}_{0}^{\top}\boldsymbol{x})), and R1(λ)\mathcal{R}_{1}(\lambda) as the prediction risk of the ridge estimator on the new feature map x→σ(W1⊤x)\boldsymbol{x}\to\sigma(\boldsymbol{W}_{1}^{\top}\boldsymbol{x}) after one feature learning step.

Importantly, due to the alignment between the trained features and the teacher model f∗f^{*} 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 β∗\boldsymbol{\beta}_{*} 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 G0\boldsymbol{G}_{0} 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 η=Θ(1)\eta=\Theta(1), we have

where δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) is defined by (C.96) in Appendix C.3. δ\delta is a non-negative function of η,λ,ψ1,ψ2∈(0,+∞)\eta,\lambda,\psi_{1},\psi_{2}\in(0,+\infty) with parameters μ1∗,μ1,μ2\mu_{1}^{*},\mu_{1},\mu_{2}, and it vanishes if and only if (at least) one of μ1∗,μ1\mu_{1}^{*},\mu_{1} and η\eta is zero.

Performance of the initial RF ridge estimator R0(λ)\mathcal{R}_{0}(\lambda) has been characterized by many prior works (e.g., [GLK+20, MM22]); hence the precise asymptotics of δ\delta provided in Theorem 7 allows us to explicitly compute the asymptotic prediction risk of the CK model after one feature learning step R1(λ)\mathcal{R}_{1}(\lambda).

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 (δ>0\delta>0) holds for any ψ1,ψ2∈(0,∞)\psi_{1},\psi_{2}\in(0,\infty), that is, taking one gradient step (with learning rate η=Θ(1)\eta=\Theta(1)) is always beneficial, even when the training set size nn 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 μ1,μ1∗≠0\mu_{1},\mu_{1}^{*}\neq 0. On the other hand, the GET (in particular Fact 6) also implies an upper bound on the possible improvement: δ≤R0(λ)−μ2∗2\delta\leq\mathcal{R}_{0}(\lambda)-\mu_{2}^{*2} as n,d,N→∞n,d,N\to\infty; this is to say, the trained CK remains in the “linear” regime.

Now we consider the following special cases where the expression of δ\delta can be further simplified.

We first analyze the setting where the sample size nn is larger than any constant times dd, that is, we let n,d,N→∞n,d,N\to\infty proportionally, and then take the limit ψ1→∞\psi_{1}\to\infty. 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: ψ1→∞\psi_{1}\to\infty, ψ2∈(0,∞)\psi_{2}\in(0,\infty). Under the same assumptions as Theorem 5 and η=Θ(1)\eta=\Theta(1), lim⁡ψ1→∞δ(η,λ,ψ1,ψ2)\lim_{\psi_{1}\to\infty}\delta(\eta,\lambda,\psi_{1},\psi_{2}), defined in (C.120), is (i)(i) non-negative, (ii)(ii) vanishing if and only if one of μ,μ1∗,η\mu_{,}\mu_{1}^{*},\eta is zero, and (iii)(iii) increasing with respect to the learning rate η\eta.

Proposition 8 predicts that the prediction risk R1(λ)\mathcal{R}_{1}(\lambda) further decreases as we use a larger learning rate η\eta, which is empirically verified in Figure 5(a). We note that the large learning rate setting (η=Θ(N)\eta=\Theta(\sqrt{N})) in Section 4.3 cannot be covered by the above proposition by increasing η\eta, as here η\eta does not scale with NN.

We also address the highly overparameterized regime, i.e., ψ2→∞\psi_{2}\to\infty. 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: ψ1∈(0,∞)\psi_{1}\in(0,\infty), ψ2→∞\psi_{2}\to\infty. Then under the same assumptions as Theorem 5 and η=Θ(1)\eta=\Theta(1), we have lim⁡ψ2→∞δ(η,λ,ψ1,ψ2)=0\lim_{\psi_{2}\to\infty}\delta(\eta,\lambda,\psi_{1},\psi_{2})=0.

Proposition 9 agrees with Figure 5(b), where we see that the risk improvement is more prominent when the width NN is not too large. One explanation is that as ψ2=N/d\psi_{2}=N/d 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 W1\boldsymbol{W}_{1} 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 R1(λ)\mathcal{R}_{1}(\lambda), which we then compare against the kernel ridge lower bound.

It is worth noting that the definition of τ∗\tau^{*} does not involve the specific value of learning rate η\eta. This is because for any choice of η=Θ(N)\eta=\Theta(\sqrt{N}), due to the Gaussian initialization of aia_{i}, 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 τ∗\tau^{*} is a simple Gaussian integral which can be numerically or analytically computed (see Appendix D.2 for some examples). For instance, when σ=σ∗=erf\sigma=\sigma^{*}=\text{erf}, one can easily verify that κ∗=3\kappa^{*}=\sqrt{3} and τ∗=0\tau^{*}=0.

Under the same assumptions as Lemma 10, after one gradient step on W\boldsymbol{W} with η=Θ(N)\eta=\Theta(\sqrt{N}), there exist constants C,ψ1∗>0C,\psi_{1}^{*}>0 such that for any n/d>ψ1∗n/d>\psi_{1}^{*}, the ridge regression estimator (4.1) satisfies

with probability 11 as n,d,N→∞n,d,N\to\infty, if we choose the ridge penalty: nε−1<N−1λ<n−εn^{\varepsilon-1}<N^{-1}\lambda<n^{-\varepsilon} for some small ε>0\varepsilon>0.

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 1010 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 ψ1=n/d\psi_{1}=n/d 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 ψ1∗\psi_{1}^{*} such that for any ψ1>ψ1∗\psi_{1}>\psi_{1}^{*}, the following holds with probability 1 when n,d,N→∞n,d,N\to\infty 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 (σ,σ∗)(\sigma,\sigma^{*}) for which τ∗\tau^{*} is small enough. In general settings, learning a good representation likely requires more than one gradient step (even if f∗f^{*} 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 W\boldsymbol{W} 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 W1\boldsymbol{W}_{1}, and quantified the improvement in the prediction risk of conjugate kernel ridge regression under two different scalings of first-step learning rate η\eta. 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 η ⁣= ⁣Θ(1)\eta\!=\!\Theta(1) and η ⁣= ⁣Θ(N)\eta\!=\!\Theta(\sqrt{N}) 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 W\boldsymbol{W} (via large gradient step), one can also introduce sufficiently large low-rank shifts to the input X\boldsymbol{X} to enable the initial RF estimator to fit a nonlinear f∗f^{*}. Intuitively, this may be due to the “dual” relation of X\boldsymbol{X} and W\boldsymbol{W} 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 tt scales with n,d,Nn,d,N. Because of our mean-field parameterization, the first-layer weight W\boldsymbol{W} needs to travel sufficiently far away from initialization to achieve small training loss (see Figure 2). Hence in our experimental simulations (where n,d,Nn,d,N are large but finite), as the number of steps tt or learning rate η\eta 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 tt, 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 η=Θ(N)\eta=\Theta(\sqrt{N}) and specific choices of f∗f^{*}), 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 f∗f^{*}, we recall the orthogonal decomposition:

In Figure 8(a), we compute the overlap between the leading eigenvector of W1\boldsymbol{W}_{1} and the linear component of the teacher model β∗\boldsymbol{\beta}_{*}. Observe that the empirical simulations (dots) closely match the analytic predictions of Theorem 3 (solid curves). Also, note that increasing the learning rate η\eta or the sample size ψ1=n/d\psi_{1}=n/d 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 σ=σ∗=SoftPlus\sigma=\sigma^{*}=\text{SoftPlus}, for which τ∗≈0.03>0\tau^{*}\approx 0.03>0, 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 τ∗\tau^{*} 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 tt, which takes value between 0 and 1,

where CKt\mathbf{CK}_{t} denotes the CK matrix defined by Wt\boldsymbol{W}_{t}. Figure 8(c) shows the KTA for two-layer NN under our mean-field parameterization and also the NTK parameterization (which omits the 1N\frac{1}{\sqrt{N}}-prefactor in (1.1)). We optimize the first-layer weights W\boldsymbol{W} until the training loss reaches 10−210^{-2} 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: f∗(x)=μ1∗⟨x,β∗⟩+P>1f∗(x)f^{*}(\boldsymbol{x})=\mu_{1}^{*}\langle\boldsymbol{x},\boldsymbol{\beta}_{*}\rangle+\textsf{P}_{>1}f^{*}(\boldsymbol{x}). Given the ridge regression estimator on the input features: θ^Lin≜(X⊤X+λnId)−1X⊤y\hat{\boldsymbol{\theta}}_{\text{Lin}}\triangleq(\boldsymbol{X}^{\top}\boldsymbol{X}+\lambda n\boldsymbol{I}_{d})^{-1}\boldsymbol{X}^{\top}\boldsymbol{y}, 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 x\boldsymbol{x} uniform on sphere or hypercube, [GMMM21, MMM21] showed that RF and kernel ridge estimators can learn at most a degree-kk 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 kk 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 22-∞\infty 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 η\eta and prefactor N\sqrt{N}):

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 n,d,Nn,d,N needs to be tracked separately in some of our calculations.

Proof. We analyze the three matrices of interest separately.

We first upper-bound ∥A∥F2\|\boldsymbol{A}\|_{F}^{2}. 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 t≥0t\geq 0. Hence, from (B.5), we arrive at

Note that the same probability bounds also applies to ∥A∥\|\boldsymbol{A}\| and ∥A∥2,∞\|\boldsymbol{A}\|_{2,\infty}. Thus, we may take t=1Nt=\sqrt{\frac{1}{N}} to obtain the desired result. Now we provide lower bounds for ε⊤XX⊤ε\boldsymbol{\varepsilon}^{\top}\boldsymbol{X}\boldsymbol{X}^{\top}\boldsymbol{\varepsilon} and f∗(X)⊤XX⊤εf^{*}(\boldsymbol{X})^{\top}\boldsymbol{X}\boldsymbol{X}^{\top}\boldsymbol{\varepsilon}. First, we define events A1\mathcal{A}_{1}, A2\mathcal{A}_{2} and A3\mathcal{A}_{3} by

Choosing t=σε2nd/4t=\sigma_{\varepsilon}^{2}nd/4, 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 ∥A∥2,∞\|\boldsymbol{A}\|_{2,\infty}, by definition we know that

We first control the operator norm of the random feature matrix σ⊥′(XW0)\sigma^{\prime}_{\perp}(\boldsymbol{X}\boldsymbol{W}_{0}). Since σ⊥′\sigma^{\prime}_{\perp} is centered, [FW20, Lemma D.4] implies that

where event AB\mathcal{A}_{B} is defined by

Next, we estimate the failure probability of event ABc\mathcal{A}_{B}^{c}. By Bernstein’s inequality, for any t≥0,t\geq 0, we have

where we write w10\boldsymbol{w}_{1}^{0} as the first column of W0\boldsymbol{W}_{0} (and similarly for all wi0\boldsymbol{w}_{i}^{0}). Following the proof of Proposition 3.3 in [FW20], we can obtain that

for any t≥0t\geq 0. Besides, inequality (B.10) implies that for any t≥0t\geq 0,

By choosing t=c′Ndt=c^{\prime}\sqrt{\frac{N}{d}} in (B.17) and B:=c′NdB:=c^{\prime}\sqrt{\frac{N}{d}} for sufficient large c′>0c^{\prime}>0, we can claim that there exists sufficient large constant c>0c>0 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 t1,t2≥0t_{1},t_{2}\geq 0. Because f∗f^{*} is Lipschitz, f∗(X)f^{*}(\boldsymbol{X}) is a sub-Gaussian random vector with similar tail bound

Let t1=t2=log⁡nt_{1}=t_{2}=\log n. Applying all these three tail bounds (B.19) and (B.10), (B.13) gives us the first part of the probability bound in (ii)(ii). 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 ∣σ′∣|\sigma^{\prime}| being upper-bounded by λσ\lambda_{\sigma}.

To control \mathopen{}\mathclose{{}\left\|{\sigma(\boldsymbol{X}\boldsymbol{W}_{0})\boldsymbol{a}}}\right\|_{\infty}, note that since a\boldsymbol{a} is centered by Assumption 1, we can apply Bernstein inequality for a\boldsymbol{a} and W\boldsymbol{W} 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 j∈[n]j\in[n], σ(xj⊤W0)a=∑i=1Naiσ(xj⊤wi0)\sigma(\boldsymbol{x}_{j}^{\top}\boldsymbol{W}_{0})\boldsymbol{a}=\sum_{i=1}^{N}a_{i}\sigma(\boldsymbol{x}_{j}^{\top}\boldsymbol{w}_{i}^{0}) is the sum of NN 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 CC. Thus, by Bernstein inequality [Ver18, Theorem 2.8.1], for each j∈[n]j\in[n],

Then we take the union over all xj\boldsymbol{x}_{j} and obtain \mathopen{}\mathclose{{}\left\|{\sigma(\boldsymbol{X}\boldsymbol{W}_{0})\boldsymbol{a}}}\right\|_{\infty}\leq\log n with probability at least 1−2ne−c(log⁡n)21-2ne^{-c(\log n)^{2}}. Hence, by (B.10), (B.20) and (B.22), we get

Part (iii)(iii) is established by choosing t=dt=\sqrt{d}. This concludes the proof of the lemma. ∎

Proposition 2 is a direct consequence of the above norm bounds.

Proof of Proposition 2. Notice that G0−A=B+C\boldsymbol{G}_{0}-\boldsymbol{A}=\boldsymbol{B}+\boldsymbol{C}. In the proportional regime, by Lemma 14, there exist universal constants C,c>0C,c>0 such that

On the other hand, part (i)(i) in Lemma 14 implies that

for some constant c,C>0c,C>0. Here we used the fact ∥A∥=∥A∥F\|\boldsymbol{A}\|=\|\boldsymbol{A}\|_{F} because it is a rank-one matrix. Conditioning on the two events stated above, we have

As long as nn is sufficiently large such that log⁡2nn<12\frac{\log^{2}n}{\sqrt{n}}<\frac{1}{2}, we can obtain

B.1.2 Decomposition of Matrix A

Using the orthogonal decomposition (3.4), we can further decompose the rank-1 matrix A\boldsymbol{A} as follows

for some constants C,C′,c>0C,C^{\prime},c>0 that only depend on μ1\mu_{1}, σε\sigma_{\varepsilon} and f∗f^{*}.

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 (i)(i) in Lemma 14, we can further decompose \mathopen{}\mathclose{{}\left\|{\boldsymbol{A}}}\right\|_{F} into

Since P>1f∗\textsf{P}_{>1}f^{*} 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 ∥A2∥F\|\boldsymbol{A}_{2}\|_{F}.

whose squared Frobenius norms are given by

conditioned on event Aε\mathcal{A}_{\varepsilon}. Thus, the Hanson-Wright inequality (Theorem 6.2.1 [Ver18]) indicates that

Thus, by choosing t=ndt=nd and employing (B.31), we have

where we simplified the expression using the assumption that n≥dn\geq d.

We can further decompose ∥A2′∥F2\|\boldsymbol{A}_{2}^{\prime}\|_{F}^{2} into two parts:

On the other hand, Bernstein’s inequality [Ver18, Theorem 2.8.1] indicates that for all 1≤i≠j≤n1\leq i\neq j\leq n,

for all t>0t>0. As for J2J_{2}, we apply (B.35) and (B.40) to all ∥xi∥\|\boldsymbol{x}_{i}\|. Letting t=dt=d in (B.40) and taking union bounds for all 1≤i≤n1\leq i\leq n, we obtain

Hence, the above equation and (B.35) lead the following bound

Therefore, by letting t=dnNt=\frac{d}{nN} in (B.41) and combining (B.31) and (B.43), we can conclude that

Finally for A2′′′\boldsymbol{A}_{2}^{\prime\prime\prime}, we may employ a similar decomposition as A2′\boldsymbol{A}_{2}^{\prime},

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 Gt\boldsymbol{G}_{t} (we drop the learning rate η=Θ(1)\eta=\Theta(1)).

For At\boldsymbol{A}^{t}, 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 a\boldsymbol{a} in (B.8), (B.20), the norm control of y\boldsymbol{y} due to (B.9) and (B.20), the operator norm bound on X\boldsymbol{X} 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 σ(XW0)\sigma(\boldsymbol{X}\boldsymbol{W}_{0}), since σ\sigma 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 NN and constants c′,C′>0c^{\prime},C^{\prime}>0. Similarly for Bt\boldsymbol{B}^{t}, we have

Again using [OS20, Setion 6.6.1], we have

Thanks to the norm control of X\boldsymbol{X} in (B.10), the norm control of y\boldsymbol{y} and ft(X)f_{t}(\boldsymbol{X}) 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 NN. Consequently, given the induction hypothesis, we know that for the next time step (t+1)(t+1) with learning rate η=Θ(1)\eta=\Theta(1), there exist some constants c′,C′c^{\prime},C^{\prime} 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 Cp>0C_{p}>0 is a universal constant. Furthermore, if D\boldsymbol{D} 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 u⊤Du\boldsymbol{u}^{\top}\boldsymbol{D}\boldsymbol{u}. Note that X⊤=[x1,…,xn]\boldsymbol{X}^{\top}=[\boldsymbol{x}_{1},\ldots,\boldsymbol{x}_{n}] has i.i.d. columns. Hence, we can expand the first quadratic form as follows:

as n→∞n\to\infty, where (iv)(iv) is obtained by Lemma 2.2 in [BS98] because vi⊤Dvi\boldsymbol{v}_{i}^{\top}\boldsymbol{D}\boldsymbol{v}_{i} are i.i.d. for i∈[n]i\in[n], 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 β∗⊤Du\boldsymbol{\beta}_{*}^{\top}\boldsymbol{D}\boldsymbol{u}, 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 W0⊤W0\boldsymbol{W}_{0}^{\top}\boldsymbol{W}_{0} is the Marchenko–Pastur distribution μψ2MP\mu_{\psi_{2}}^{\text{MP}} with parameter ψ2\psi_{2}; let m(z)m(z) be the Stieltjes transform of μψ2MP\mu_{\psi_{2}}^{\text{MP}}. Also, we denote the limiting eigenvalue distribution for W0W0⊤\boldsymbol{W}_{0}\boldsymbol{W}_{0}^{\top} by μˉψ2MP\bar{\mu}_{\psi_{2}}^{\text{MP}} whose Stieltjes transform is mˉ(z)\bar{m}(z), which is referred to as the companion transform. The relation between m(z)m(z) and mˉ(z)\bar{m}(z) is given as

Moreover, m(z)m(z) is uniquely determined by the fixed-point equation

Following the above notions, we define the resolvent Q0(z):=(W0W0⊤−zI)−1\boldsymbol{Q}_{0}(z):=(\boldsymbol{W}_{0}\boldsymbol{W}_{0}^{\top}-z\boldsymbol{I})^{-1}, for

and any small ϵ>0\epsilon>0. Then, under the same assumptions of Theorem 3, for all sufficiently large NN, ∥Q0(z)∥≤2/ϵ\|\boldsymbol{Q}_{0}(z)\|\leq 2/\epsilon uniformly for all z∈Ωϵz\in\Omega_{\epsilon}.

Proof. Write z=x+iyz=x+iy, with x>\big{(}1+\sqrt{\psi_{2}}\big{)}^{2}+\epsilon. For any i∈[N]i\in[N], let λi\lambda_{i} be the ii-th eigenvalue of W0W0⊤\boldsymbol{W}_{0}\boldsymbol{W}_{0}^{\top}. Then we have

By Theorem 5.11 in [BS10], for sufficiently large NN, \mathopen{}\mathclose{{}\left|\lambda_{1}-\big{(}1+\sqrt{\psi_{2}}\big{)}^{2}}\right|\leq\epsilon/2. Hence, 1/∣z−λi∣≤ϵ/21/|z-\lambda_{i}|\leq\epsilon/2 for i∈[N]i\in[N] and ∥Q0(z)∥≤2/ϵ\|\boldsymbol{Q}_{0}(z)\|\leq 2/\epsilon for all z∈Ωϵz\in\Omega_{\epsilon}.

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 u=ημ1nX⊤y\boldsymbol{u}=\frac{\eta\mu_{1}}{n}\boldsymbol{X}^{\top}\boldsymbol{y}. Under the same assumptions as Theorem 3, for any ϵ>0\epsilon>0,

in probability as n/d→ψ1n/d\to\psi_{1} and N/d→ψ2N/d\to\psi_{2}, uniformly on any compact subset of Ωϵ\Omega_{\epsilon} defined in (B.89), where scalars θ1\theta_{1} and θ2\theta_{2} are defined in Theorem 3.

Proof. Firstly note that (B.90) directly follows from Lemma 19 and Hoeffding’s inequality for a\boldsymbol{a}. The remaining concentration statements will be established by applying Lemma 18 to different choices of D\boldsymbol{D} and the Hanson-Wright inequality for a\boldsymbol{a}. In particular, due to Lemma 19, we know that ∥Q0(z)∥\|\boldsymbol{Q}_{0}(z)\|, ∥W0⊤Q0(z)W0∥\|\boldsymbol{W}_{0}^{\top}\boldsymbol{Q}_{0}(z)\boldsymbol{W}_{0}\|, ∥Q0(z)2∥\|\boldsymbol{Q}_{0}(z)^{2}\| and ∥W0⊤Q0(z)2W0∥\|\boldsymbol{W}_{0}^{\top}\boldsymbol{Q}_{0}(z)^{2}\boldsymbol{W}_{0}\| are all uniformly bounded on Ωϵ\Omega_{\epsilon} for large NN. Take D=Q0(z)\boldsymbol{D}=\boldsymbol{Q}_{0}(z) in Lemma 18, we obtain that

uniformly on any compact subset of Ωϵ\Omega_{\epsilon}, 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 det⁡(AB)=det⁡(A)det⁡(B)\det(\boldsymbol{A}\boldsymbol{B})=\det(\boldsymbol{A})\det(\boldsymbol{B}), and the Sylvester’s determinant identity det⁡(I+AB)=det⁡(I+BA)\det(\boldsymbol{I}+\boldsymbol{A}\boldsymbol{B})=\det(\boldsymbol{I}+\boldsymbol{B}\boldsymbol{A}) for A,B\boldsymbol{A},\boldsymbol{B} of appropriate dimensions, the above equations has the same solution as

Now we compute the root of P(z)P(z) on \mathopen{}\mathclose{{}\left((1+\sqrt{\psi_{2}})^{2},+\infty}\right). From (B.88) we have

which implies that the root of P(z)P(z) on the real line is given as λ0:=(1+θ12)(ψ2+θ12)θ12\lambda_{0}:=\frac{(1+\theta_{1}^{2})(\psi_{2}+\theta_{1}^{2})}{\theta_{1}^{2}}. Based on the expression of m(z)m(z) and (B.88), we have

Since λ^\hat{\lambda} does not reside in the spectrum of W0W0⊤\boldsymbol{W}_{0}\boldsymbol{W}_{0}^{\top}, we can further write

By multiplying [u,W0a]⊤[\boldsymbol{u},\boldsymbol{W}_{0}\boldsymbol{a}]^{\top} from the left hand side of the above equality, we arrive at 0=Mn(λ^)[v^1v^2]\mathbf{0}=\boldsymbol{M}_{n}(\hat{\lambda})\begin{bmatrix}\hat{v}_{1}\\ \hat{v}_{2}\end{bmatrix}, where the 2-by-2 matrix is given as

In addition, multiplying β∗⊤\boldsymbol{\beta}_{*}^{\top} from the left hand side of (B.98) yields

We first consider the scenario θ1>ψ21/4\theta_{1}>\psi_{2}^{1/4}, where λ^→λ0\hat{\lambda}\to\lambda_{0} in probability and λ0\lambda_{0} is outside the support of μψ2MP\mu_{\psi_{2}}^{\text{MP}}. For sufficiently small ϵ>0\epsilon>0 and all large nn, λ^∈Ωϵ\hat{\lambda}\in\Omega_{\epsilon}, and thus Lemma 20 gives

in probability as n,d,N→∞n,d,N\to\infty proportionally. Notice that both v^1\hat{v}_{1} and v^2\hat{v}_{2} are uniformly bounded by some constants with high probability since ∥a∥→1\|\boldsymbol{a}\|\to 1, ∥W0∥→(1+ψ2)\|\boldsymbol{W}_{0}\|\to(1+\sqrt{\psi_{2}}) almost surely, and Lemma 18 implies ∥u∥→θ1\|\boldsymbol{u}\|\to\theta_{1} in probability as n,d,N→∞n,d,N\to\infty. Therefore, when n/d→ψ1n/d\to\psi_{1} and N/d→ψ2N/d\to\psi_{2}, (B.100) and (B.103) provide the limits of v^1\hat{v}_{1} and v^2\hat{v}_{2}, which we denote by v1v_{1} and v2v_{2}, respectively. More precisely,

Thus, we can apply formula (B.87), condition P(λ0)=0P(\lambda_{0})=0 in (B.96) and the following well-known facts of the Stieltjes transform (e.g., see [BS10]) with λ0=(1+θ12)(ψ2+θ12)/θ12\lambda_{0}=(1+\theta_{1}^{2})(\psi_{2}+\theta_{1}^{2})/\theta_{1}^{2}:

On the other hand, if θ14≤ψ2\theta_{1}^{4}\leq\psi_{2}, we have proved that λ^\hat{\lambda} is approaching to the right-edge of the bulk of μψ2MP\mu_{\psi_{2}}^{\text{MP}}. In fact, with probability one, λ1<λ^\lambda_{1}<\hat{\lambda}, where λ1\lambda_{1} is the largest eigenvalue of W0W0⊤\boldsymbol{W}_{0}\boldsymbol{W}_{0}^{\top}; 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 det⁡Mn(λ^)=0\det\boldsymbol{M}_{n}(\hat{\lambda})=0. Hence by the definition of Mn(z)M_{n}(z),

Also, by the Cauchy–Schwarz inequality, we have

Since λ^→(1+ψ2)2\hat{\lambda}\to(1+\sqrt{\psi_{2}})^{2}, given any small ϵ∈(0,ψ2)\epsilon\in(0,\sqrt{\psi_{2}}), for all sufficiently large NN, 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 NN, ∣a⊤W0⊤Q0(λ^)u∣|\boldsymbol{a}^{\top}\boldsymbol{W}_{0}^{\top}\boldsymbol{Q}_{0}(\hat{\lambda})\boldsymbol{u}| is bounded by some universal constant related to ψ2\psi_{2}. Then by (B.109), we can conclude a⊤W0⊤Q0(λ^)W0a⋅u⊤Q0(λ^)u\boldsymbol{a}^{\top}\boldsymbol{W}_{0}^{\top}\boldsymbol{Q}_{0}(\hat{\lambda})\boldsymbol{W}_{0}\boldsymbol{a}\cdot\boldsymbol{u}^{\top}\boldsymbol{Q}_{0}(\hat{\lambda})\boldsymbol{u} is also asymptotically bounded by some constant. On the other hand, since all eigenvalues of W0W0⊤\boldsymbol{W}_{0}\boldsymbol{W}_{0}^{\top} is smaller than λ^\hat{\lambda}, we obtain

This directly implies −a⊤W0⊤Q0(λ^)W0a-\boldsymbol{a}^{\top}\boldsymbol{W}_{0}^{\top}\boldsymbol{Q}_{0}(\hat{\lambda})\boldsymbol{W}_{0}\boldsymbol{a} has a constant upper bound for all large NN. 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 C>0C>0 and all large NN. 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 ∥R∥→0\|\boldsymbol{R}\|\to 0, by a simple adaptation of Lemma A.2 in [BGN11], one can directly verify β∗⊤Q0(λ^)Ru1→0\boldsymbol{\beta}_{*}^{\top}\boldsymbol{Q}_{0}(\hat{\lambda})\boldsymbol{R}\boldsymbol{u}_{1}\to 0, as N,n,d→∞N,n,d\to\infty proportionally. Hence we conclude that u1⊤β∗→0\boldsymbol{u}_{1}^{\top}\boldsymbol{\beta}_{*}\to 0 if θ14≤ψ2\theta_{1}^{4}\leq\psi_{2}. 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 i∈[n]i\in[n].

Define the set of weight matrices perturbed from the Gaussian initialization W0\boldsymbol{W}_{0} as

Note that for learning rate η=Θ(1)\eta=\Theta(1), we can verify that W\mathcal{W} 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 W\boldsymbol{W} remains “close” to the initialization W0\boldsymbol{W}_{0}.

where a^\hat{\boldsymbol{a}} and ˉa\bar{}\boldsymbol{a} are defined in (C.1) and (C.2).

From Proposition 22 we know that Theorem 5 holds if the optimized weight matrix W\boldsymbol{W} falls into the set W\mathcal{W} 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 λ>0\lambda>0, the RHS of the above equation is bounded in probability.

Recall the single-index teacher assumption: yi=σ∗(⟨xi,β∗⟩)+εiy_{i}=\sigma^{*}(\langle\boldsymbol{x}_{i},\boldsymbol{\beta}^{*}\rangle)+\varepsilon_{i} for i∈[n]i\in[n]. Observe that for W∈W\boldsymbol{W}\in\mathcal{W}, the following near-orthogonality condition between the neurons holds with high probability

Importantly, for W\boldsymbol{W} satisfying the near-orthogonality condition (C.4), we can utilize the following central limit theorem from [HL20] derived via Stein’s method.

where z∼N(0,1)z\sim\mathcal{N}(0,1), and K′K^{\prime} only depends on constant KK.

Following [HL20], we construct an interpolating sequence between the nonlinear and linear features model. For any 0≤k≤n0\leq k\leq n, we define

Note that when γ1=γ2=0\gamma_{1}=\gamma_{2}=0, setting k=0k=0 recovers the estimator on nonlinear features a^\hat{\boldsymbol{a}}, and similarly, setting k=nk=n gives the estimator on the linear Gaussian features aˉ\bar{\boldsymbol{a}}.

We remark that the perturbation Q(g)Q(\boldsymbol{g}) allows us to compute the prediction risk by taking the derivative of the objective w.r.t. γ1,γ2\gamma_{1},\gamma_{2} 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 λ2\frac{\lambda}{2}-strongly convex (i.e., the strongly-convex regularizer dominates the concave part of Q(g)Q(\boldsymbol{g}) when γ1<0\gamma_{1}<0).

Proof. We follow the proof of [HL20, Lemma 23] and first analyze one coordinate of gk∗\boldsymbol{g}_{k}^{*} defined by (C.6), which WLOG we select to be the last coordinate. For concise notation, we instead augment the weight matrix with an (N+1)(N+1)-th column and study the corresponding [gk∗]N+1[\boldsymbol{g}_{k}^{*}]_{N+1}. Denote the weight vector wN+1=wN+10+δN+1\boldsymbol{w}_{N+1}=\boldsymbol{w}^{0}_{N+1}+\boldsymbol{\delta}_{N+1}, where wN+10\boldsymbol{w}^{0}_{N+1} is the (N+1)(N+1)-th column of the initialized W0\boldsymbol{W}_{0}, and δ\boldsymbol{\delta} is the perturbation (i.e., gradient update for W\boldsymbol{W}).

The (N+1)(N+1)-th coordinate of interest, which we denote as u∗u^{*}, can be written as the solution to the following optimization problem,

By [HL20, Equation (249)], we know that for W∈W\boldsymbol{W}\in\mathcal{W},

We control each term on the right hand side of (C.9) separately. Note that W∈W\boldsymbol{W}\in\mathcal{W} 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 c1>0c_{1}>0 and large NN,

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 W\boldsymbol{W} to satisfy (C.4)), we have

for some c3>0c_{3}>0. The case where i≤ki\leq k (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 c6>0c_{6}>0 and all large NN. 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 i∈[N]i\in[N], we complete the proof by a union bound over the NN 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 W∈W\boldsymbol{W}\in\mathcal{W} and 1≤k≤n1\leq k\leq n, 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 W\mathcal{W} occurs with high probability. For one gradient step on the squared loss with learning rate η=Θ(1)\eta=\Theta(1), Lemma 14 together with \mathopen{}\mathclose{{}\left\|{\boldsymbol{\beta}_{*}}}\right\|=1 entail that for proportional n,d,Nn,d,N, there exists some constant c,C>0c,C>0 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 x→σ(W1⊤x)\boldsymbol{x}\to\sigma(\boldsymbol{W}_{1}^{\top}\boldsymbol{x}). 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 λ>0\lambda>0. Then for η=Θ(1)\eta=\Theta(1), we have

as n,d,N→∞n,d,N\to\infty at comparable rate, where we dropped the constant σε2\sigma_{\varepsilon}^{2} in (ii)(ii).

For the bias terms, we consider perturbation on W1\boldsymbol{W}_{1} in the operator norm. Again, Lemma 14 entails that

Based on this result, it is straightforward to show that

where in (iii)(iii) we used the fact that σ∗\sigma^{*} 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 η=Θ(1)\eta=\Theta(1), 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 TiT_{i} as n,d,N→∞n,d,N\to\infty proportionally.

In the following lemma, we show that each TiT_{i} will concentrate around some Ti0T_{i}^{0} given by

Under Assumptions 1 and 2, as n,d,N→∞n,d,N\to\infty 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 c>0c>0. Therefore, (C.58) concentrates around μ1ηnNTr⁡DX⊤Ψ\frac{\mu_{1}\eta}{n\sqrt{N}}\operatorname{Tr}\boldsymbol{D}\boldsymbol{X}^{\top}\boldsymbol{\Psi} with high probability. Thus it remains to show that μ1ηnNTr⁡DX⊤Ψ\frac{\mu_{1}\eta}{n\sqrt{N}}\operatorname{Tr}\boldsymbol{D}\boldsymbol{X}^{\top}\boldsymbol{\Psi} is vanishing in probability.

Now condition on event M\mathcal{M}, the Lipschitz property of σ∗\sigma^{*} entails that

Thus, as d→∞d\to\infty, ∥Ψ∥F/N≲(log⁡d)2/d\|\boldsymbol{\Psi}\|_{F}/\sqrt{N}\lesssim(\log d)^{2}/\sqrt{d} with high probability, which finally implies that μ1ηnNTr⁡DX⊤Ψ\frac{\mu_{1}\eta}{n\sqrt{N}}\operatorname{Tr}\boldsymbol{D}\boldsymbol{X}^{\top}\boldsymbol{\Psi} converges to zero in probability as n,d,N→∞n,d,N\to\infty proportionally. This completes the proof. ∎

Finally, we use the following simplification of quadratic forms to obtain the desired Ti0T^{0}_{i}.

C.3.2 Risk Calculation via Linear Pencils

In this section, we derive analytic expressions of the terms TiT_{i} defined in (C.49) as n,d,N→∞n,d,N\to\infty proportionally. In particular, the exact values are described by self-consistent equations defined in the following proposition.

Given Assumption 1 and λ>0\lambda>0. For each TiT_{i} defined in (C.49) and 1≤i≤12,1\leq i\leq 12, we have

in probability, as n/d→ψ1n/d\to\psi_{1} and N/d→ψ2N/d\to\psi_{2}, where τi\tau_{i}’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 m1(z)m_{1}(z) and m2(z)m_{2}(z) 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 σW2=0\sigma_{W_{2}}=0 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 T2T_{2}, (C.83) indicates

with z=ψ1ψ2λz=\frac{\psi_{1}}{\psi_{2}}\lambda. τ7\tau_{7} can also be derived in similar fashion.

For T6,T_{6}, we utilize the computations in Appendix I.6.1 of [TAP21] by setting the covariance Σ=I\Sigma=\boldsymbol{I}. More precisely, based on Equations (S370) and (S418) in [TAP21],

where (i)(i) and (ii)(ii) are both due to (C.80) and (C.81). Hence we obtain the formulae of τ6\tau_{6} and τ9\tau_{9}.

Recall the following derivative trick of the Stieltjes transform,

Once again by setting σW2=0\sigma_{W_{2}}=0 in [AP20], we can directly employ the computation of E32E_{32} in Section S4.3.4 of [AP20] to derive τ11\tau_{11} and τ12\tau_{12}. In particular, due to (C.82),

Note that Equation (S148) in [AP20] established that

whence, letting z=ψ1ψ2λz=\frac{\psi_{1}}{\psi_{2}}\lambda, we obtain τ11\tau_{11} and τ12\tau_{12}.

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., R0(λ)−R1(λ)\mathcal{R}_{0}(\lambda)-\mathcal{R}_{1}(\lambda) in Theorem 7. The following statement is the complete version of Theorem 7.

Given Assumptions 1 and 2, consider ψ1,ψ2∈(0,+∞)\psi_{1},\psi_{2}\in(0,+\infty). Fix η=Θ(1)\eta=\Theta(1) and λ>0\lambda>0. Denote R0(λ)\mathcal{R}_{0}(\lambda) and R1(λ)\mathcal{R}_{1}(\lambda) as the prediction risk of CK ridge regression in (4.1) using initial weight W0\boldsymbol{W}_{0} and first-step updated W1\boldsymbol{W}_{1}, respectively. Then the difference between these two prediction risk values satisfies

where δ\delta is a non-negative function of η,λ,ψ1\eta,\lambda,\psi_{1} and ψ2∈(0,+∞)\psi_{2}\in(0,+\infty) with parameters μ1∗,μ1,μ2\mu_{1}^{*},\mu_{1},\mu_{2} given as

Here the scalars τi\tau_{i}’s are defined in Proposition 29. Furthermore, δ(η,λ,ψ1,ψ2)=0\delta(\eta,\lambda,\psi_{1},\psi_{2})=0 if and only if at least one of μ1∗,μ1\mu_{1}^{*},\mu_{1} and η\eta is zero.

Proof. Due to Lemma 26 (or the decomposition (C.25), (C.26) and (C.27)), we can see that variance VV is unchanged after one gradient descent step with η=Θ(1)\eta=\Theta(1). 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 B\boldsymbol{B} and C\boldsymbol{C} in W1\boldsymbol{W}_{1} and take W1:=W0+ua⊤\boldsymbol{W}_{1}:=\boldsymbol{W}_{0}+\boldsymbol{u}\boldsymbol{a}^{\top}, where u=μ1ηnX⊤y\boldsymbol{u}=\frac{\mu_{1}\eta}{n}\boldsymbol{X}^{\top}\boldsymbol{y} and y=f∗(X)+ε\boldsymbol{y}=f^{*}(\boldsymbol{X})+\boldsymbol{\varepsilon}, without changing the bias terms.

First note that if μ1=0\mu_{1}=0, then u=0\boldsymbol{u}=\mathbf{0} and therefore R0(λ)=R1(λ)\mathcal{R}_{0}(\lambda)=\mathcal{R}_{1}(\lambda) as n→∞n\to\infty. In the following, we take μ1≠0\mu_{1}\neq 0 which implies that θ1\theta_{1} defined in Theorem 3 will not vanish. Now we aim to extract the low-rank perturbation ua⊤\boldsymbol{u}\boldsymbol{a}^{\top} from bias terms (C.25) and (C.26). We adhere to the notions in (C.44), (C.49) and (C.54) and define D:=T1(T2−T3)−1D:=T_{1}(T_{2}-T_{3})-1. Similar to [MM22, Lemma C.1], we use the following linearization trick to separate the gradient step ua⊤\boldsymbol{u}\boldsymbol{a}^{\top} from the matrices R,Φˉ,Σ‾Φ\boldsymbol{R},\bar{\boldsymbol{\Phi}},\overline{\boldsymbol{\Sigma}}_{\Phi} and W1\boldsymbol{W}_{1}.

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 Δua:=μ12NW0⊤ua⊤\boldsymbol{\Delta}_{ua}:=\frac{\mu_{1}^{2}}{N}\boldsymbol{W}_{0}^{\top}\boldsymbol{u}\boldsymbol{a}^{\top}, Δau:=Δua⊤\boldsymbol{\Delta}_{au}:=\boldsymbol{\Delta}_{ua}^{\top} and Δaua:=μ12T10Naa⊤\boldsymbol{\Delta}_{aua}:=\frac{\mu_{1}^{2}T_{10}}{N}\boldsymbol{a}\boldsymbol{a}^{\top}. Hence,

Analogously, we can decompose B2B_{2} in (C.26) as follows

where we repeatedly make use of Lemma 27 and the concentration for a\boldsymbol{a} 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 δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) defined in Theorem 7. Also observe the following equivalences from Proposition 29,

Hence, we can simplify δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) as follows

Finally, we validate that the function δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) is non-negative on variables η,λ,ψ1\eta,\lambda,\psi_{1} and ψ2∈(0,+∞)\psi_{2}\in(0,+\infty). Observe that the formula of δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) in (C.111) is decomposed into two parts. From Proposition 29 we know that τ1\tau_{1} and m1m_{1} are the limits of tr⁡R0(z)\operatorname{tr}\boldsymbol{R}_{0}(z) and tr⁡ˉR0(z)\operatorname{tr}\bar{}\boldsymbol{R}_{0}(z) evaluated at z=ψ1λ/ψ2z=\psi_{1}\lambda/\psi_{2}; this indicates that τ1∈(0,ψ2/λψ1]\tau_{1}\in(0,\psi_{2}/\lambda\psi_{1}] is non-negative. For the same reason, m2∈(0,ψ2/λψ1]m_{2}\in(0,\psi_{2}/\lambda\psi_{1}] and −m1′,−m2′∈(0,ψ22/λ2ψ12]-m_{1}^{\prime},-m_{2}^{\prime}\in(0,\psi_{2}^{2}/\lambda^{2}\psi_{1}^{2}]. Also due to Proposition 29, we have

Therefore, τ1(τ7−τ5)(τ4+τ12−2τ6)≤0\tau_{1}(\tau_{7}-\tau_{5})(\tau_{4}+\tau_{12}-2\tau_{6})\leq 0 and τ1(τ2−τ3)−1≤−1\tau_{1}(\tau_{2}-\tau_{3})-1\leq-1. This entails that the first part of δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) is non-negative:

As for the second part, it suffices to evaluate Δ:=τ1(τ4+τ12−2τ6)+(τ7−τ5)τ8\Delta:=\tau_{1}(\tau_{4}+\tau_{12}-2\tau_{6})+(\tau_{7}-\tau_{5})\tau_{8} since

Plugging in quantities in (C.112) with z=λψ1/ψ2z=\lambda\psi_{1}/\psi_{2}, we have

where (i)(i) and (ii)(ii) are due to (C.80) and (C.86), respectively. By Lemma A.1 in [TAP21], we know function zμ12m1(z)τ1(z)z\mu_{1}^{2}m_{1}(z)\tau_{1}(z) has non-positive derivative when z>0z>0. This implies that Δ≥0\Delta\geq 0 and hence the second part of δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) is also non-negative.

Finally, we note that when μ1∗=0\mu_{1}^{*}=0, the function δ(η,λ,ψ1,ψ2)=0\delta(\eta,\lambda,\psi_{1},\psi_{2})=0. This is because

when μ1∗=0\mu_{1}^{*}=0. Whereas when η=0\eta=0, we know that θ1=θ2=0\theta_{1}=\theta_{2}=0, which entails δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) is also vanishing. Also observe that in (C.112), m1,m2,m1′,m2′,τ1m_{1},m_{2},m_{1}^{\prime},m_{2}^{\prime},\tau_{1} are all positive. Hence we conclude that if δ(η,λ,ψ1,ψ2)=0\delta(\eta,\lambda,\psi_{1},\psi_{2})=0, then at least one of η,μ1μ1∗\eta,\mu_{1}\mu_{1}^{*} must be zero.

C.4 Analysis of Special Cases

While the previous subsection provides explicit formulae of δ\delta, 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 ψ1→∞\psi_{1}\to\infty and the large width limit ψ2→∞\psi_{2}\to\infty, where the calculation simplifies and enables us to further characterize properties of δ\delta. In both cases, we start with Theorem 30 and take one of aspect ratios (ψ1\psi_{1} or ψ2\psi_{2}) to infinity.

In this subsection we prove Proposition 8. We introduce two positive parameters

where μψ2MP\mu^{\text{MP}}_{\psi_{2}} is Marchenko–Pastur distribution with rate ψ2∈(0,∞)\psi_{2}\in(0,\infty). Now we consider the large-sample limit: ψ1→∞,ψ2∈(0,∞)\psi_{1}\to\infty,\psi_{2}\in(0,\infty). 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 s1s_{1} and s2s_{2}.

Under the same assumptions as Theorem 30 and take ψ1→∞\psi_{1}\to\infty. Then the difference between the prediction risks before and after one feature learning step R0(λ)−R1(λ)\mathcal{R}_{0}(\lambda)-\mathcal{R}_{1}(\lambda) satisfies

In this case δ(η,λ,∞,ψ2)\delta(\eta,\lambda,\infty,\psi_{2}) is a non-negative function of η,λ,ψ2∈(0,+∞)\eta,\lambda,\psi_{2}\in(0,+\infty), and δ=0\delta=0 if and only if one of μ1,μ1∗,η\mu_{1},\mu_{1}^{*},\eta is zero. Furthermore, δ(η,λ,∞,ψ2)\delta(\eta,\lambda,\infty,\psi_{2}) is increasing with respect to the learning rate η≥0\eta\geq 0.

Proof. Following Theorem 30, it suffices to consider the limit of δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) when ψ1→∞\psi_{1}\to\infty. This reduces to simplifying the asymptotics of τi\tau_{i}’s defined in Proposition 29, as δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) is determined by τi\tau_{i}’s in Theorem 30. We aim to prove the following:

as ψ1→∞\psi_{1}\to\infty, where s1s_{1} and s2s_{2} are defined in (C.119). The trivial cases when μ1,η=0\mu_{1},\eta=0 have been studied in Theorem 30. So, WLOG, we assume μ1,η>0\mu_{1},\eta>0 in the following derivations.

Recall the definitions of τ1(z),m1(z)\tau_{1}(z),m_{1}(z) and m2(z)m_{2}(z). One can easily see that m1,m2→0m_{1},m_{2}\to 0 as ψ1→∞\psi_{1}\to\infty. For any z≥0z\geq 0, (C.79) and (C.80) can be written as follows

Notice that τ1(z)=lim⁡tr⁡R0(z)\tau_{1}(z)=\lim\operatorname{tr}\boldsymbol{R}_{0}(z), and based on [FW20, Theorem 3.4], for any z≥0z\geq 0,

Therefore, when ψ1→∞\psi_{1}\to\infty, 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 s1s_{1},

which verifies the first statement in (C.124). Now recall the value of interest z=λψ1ψ2z=\frac{\lambda\psi_{1}}{\psi_{2}}. Due to the relationship between τ1\tau_{1} and m1m_{1}, it is straightforward to deduce that λψ1ψ2m1→1\lambda\frac{\psi_{1}}{\psi_{2}}m_{1}\to 1 as ψ1→∞\psi_{1}\to\infty. As for τ2\tau_{2}, in terms of (C.126), we have

Next we compute the limit of m1′m_{1}^{\prime}. Taking derivative with respect to zz at both sides of (C.128), we arrive at ψ12ψ22τ1′→−s2\frac{\psi_{1}^{2}}{\psi_{2}^{2}}\tau_{1}^{\prime}\to-s_{2}; here τ1′\tau_{1}^{\prime} represents the derivative τ1′(z)\tau_{1}^{\prime}(z) at z=λψ1ψ2z=\frac{\lambda\psi_{1}}{\psi_{2}}. Combining this relation and (C.129),(C.81), we can deduce that

Note that here we used τ1→0\tau_{1}\to 0 when ψ1→∞\psi_{1}\to\infty. Lastly, for τ11\tau_{11} and τ12\tau_{12}, by (C.126),

where we applied the previously established convergence of τ1\tau_{1} and τ1′\tau_{1}^{\prime}. Also recall that (C.130) implies that m2/m1m_{2}/m_{1} converges to 1/\mathopen{}\mathclose{{}\left(1+\psi_{2}\mu_{1}^{2}s_{1}}\right) as ψ1→∞\psi_{1}\to\infty. Together with the convergence of τ9\tau_{9} in (C.124), we get

which implies the convergence of τ11\tau_{11} in (C.124).

As a result, by replacing τi\tau_{i}’s in (C.96) with the corresponding reparameterized τi\tau_{i}’s in (C.124), we arrive at the following expression of δ\delta:

By definitions of A,B,CA,B,C in (C.121), (C.122) and (C.123), we can see that A=−μ12θ22αs1A=-\mu_{1}^{2}\theta_{2}^{2}\alpha s_{1}, B=βB=\beta and C=−μ12θ22αγC=-\mu_{1}^{2}\theta_{2}^{2}\alpha\gamma; this leads to the equivalent expression

Now we claim that A,B,CA,B,C are all non-negative, for any η,λ,ψ2≥0\eta,\lambda,\psi_{2}\geq 0. With a slight abuse of terminology, in the following we denote z=−(μ22+λ)/μ12<0z=-(\mu_{2}^{2}+\lambda)/\mu_{1}^{2}<0. Recall that s1=m(z)/μ12s_{1}=m(z)/\mu_{1}^{2} and s2=m′(z)/μ14s_{2}=m^{\prime}(z)/\mu_{1}^{4}; We can therefore simplify (C.139), (C.140) and (C.141) as follows

where (i)(i) is due to (C.132) and (ii)(ii) is obtained by taking derivative with respect to zz in (C.131). In addition,

which implies that γ>0\gamma>0. We also denote the companion Stieltjes transform of m(z)m(z) by mˉ(z)\bar{m}(z), which is the Stieltjes transform of the limiting eigenvalue distribution for W0W0⊤\boldsymbol{W}_{0}\boldsymbol{W}_{0}^{\top}. Recall the following relation between m(z)m(z) and mˉ(z)\bar{m}(z): \bar{m}(z)+\frac{1}{z}=\psi_{2}\mathopen{}\mathclose{{}\left(m(z)+\frac{1}{z}}\right). Since m(z)+zm′(z)m(z)+zm^{\prime}(z) is positive, we can deduce that

where the last equality is obtained by taking derivative of (B.87) on both sides with respect to zz. In summary, we have shown that α<0\alpha<0 and β,γ>0\beta,\gamma>0 when λ,μ1>0\lambda,\mu_{1}>0. Hence by definition, A,B,CA,B,C are all non-negative and so is δ(η,λ,∞,ψ2)\delta(\eta,\lambda,\infty,\psi_{2}).

Finally, we verify that δ(η,λ,∞,ψ2)\delta(\eta,\lambda,\infty,\psi_{2}) is an increasing function of η≥0\eta\geq 0. Observe that η\eta only appears in θ2\theta_{2} in the expression of δ(η,λ,∞,ψ2)\delta(\eta,\lambda,\infty,\psi_{2}) in (C.138). Hence, it suffices to take the derivative of δ(η,λ,∞,ψ2)\delta(\eta,\lambda,\infty,\psi_{2}) with respect to θ2\theta_{2} and verify that this partial derivative is positive. One can check that

By the definition of γ\gamma 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 λs2≤s1\lambda s_{2}\leq s_{1}. Finally, recall that α<0\alpha<0 and β,γ>0\beta,\gamma>0; this implies δ(η,λ,∞,ψ2)\delta(\eta,\lambda,\infty,\psi_{2}) is increasing with regard to η∈[0,+∞)\eta\in[0,+\infty) and completes the proof.

C.4.2 Case II: Highly overparameterized regime

Next we consider the large width limit ψ2→∞\psi_{2}\to\infty and establish Proposition 9.

Note that 0≤zm1(z)≤10\leq zm_{1}(z)\leq 1, for all z≥0z\geq 0, and thus 0≤ψ1ψ2zm1(z)≤ψ1ψ20\leq\frac{\psi_{1}}{\psi_{2}}zm_{1}(z)\leq\frac{\psi_{1}}{\psi_{2}}. On the other hand, μψ1MP\mu^{\text{MP}}_{\psi_{1}} is compactly supported. Therefore, by taking z=ψ1λ/ψ2z=\psi_{1}\lambda/\psi_{2} and letting ψ2→∞\psi_{2}\to\infty at both sides of (C.147), we arrive at

which is a finite positive value determined by ψ1,μ1,μ2\psi_{1},\mu_{1},\mu_{2}. With this in mind, we conclude that lim⁡ψ2→∞m2\lim_{\psi_{2}\to\infty}m_{2} is also finite, since m2(z)m_{2}(z) is determined by (C.126) and we can take z=ψ1λ/ψ2z=\psi_{1}\lambda/\psi_{2} with ψ2→∞\psi_{2}\to\infty. In addition, since 0≤−zm1′(z)≤m1(z)0\leq-zm_{1}^{\prime}(z)\leq m_{1}(z) for any z≥0z\geq 0, we may take the derivative with respect to zz at both sides of (C.147) to obtain m1′(z)m_{1}^{\prime}(z), and take z=ψ1λ/ψ2z=\psi_{1}\lambda/\psi_{2} and ψ2→∞\psi_{2}\to\infty to conclude that the limit of m1′m_{1}^{\prime} is finite as well. Similarly, by taking derivative with respect to zz in (C.126), one can also verify that as ψ2→∞\psi_{2}\to\infty, the limit of m2′m_{2}^{\prime} remains finite. From these estimates we know that

are vanishing as ψ2→∞\psi_{2}\to\infty, whereas

will converge to some finite values. The proposition is established based on the definition of δ(η,λ,ψ1,ψ2)\delta(\eta,\lambda,\psi_{1},\psi_{2}) 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): f∗(x)=σ∗(⟨x,β∗⟩)f^{*}(\boldsymbol{x})=\sigma^{*}(\langle\boldsymbol{x},\boldsymbol{\beta}_{*}\rangle), and study the impact of one gradient step with large learning rate η=Θ(N)\eta=\Theta(\sqrt{N}). For simplicity, we denote η=ηˉN\eta=\bar{\eta}\sqrt{N} where ηˉ>0\bar{\eta}>0 is a fixed constant not depending on NN.

Recall that W1=W0+ηNG0\boldsymbol{W}_{1}=\boldsymbol{W}_{0}+\eta\sqrt{N}\boldsymbol{G}_{0}, where G0=A1+A2+B+C\boldsymbol{G}_{0}=\boldsymbol{A}_{1}+\boldsymbol{A}_{2}+\boldsymbol{B}+\boldsymbol{C} is defined in Lemma 14 and 15, and the full-rank term B\boldsymbol{B} is given as

We first refine the estimate on the Frobenius norm of certain submatrix of B\boldsymbol{B}; the choice of such submatrices will be explained in Section D.2.

Proof. Let X⊤=(x1,x2,…,xn)\boldsymbol{X}^{\top}=(\boldsymbol{x}_{1},\boldsymbol{x}_{2},\ldots,\boldsymbol{x}_{n}) and y⊤=(y1,…,yn)\boldsymbol{y}^{\top}=(y_{1},\ldots,y_{n}). Then, matrix Br\boldsymbol{B}_{r} can be written as

Here, J1J_{1} represents the sum for distinct i≠j∈[n]i\neq j\in[n] and J2J_{2} is the sum when i=j∈[n]i=j\in[n]. Therefore,

where w∼N(0,I)\boldsymbol{w}\sim\mathcal{N}(0,\boldsymbol{I}) independent of a\boldsymbol{a} and X\boldsymbol{X}. We compute the aforementioned expectations as follows

for some constant C>0C>0. In addition, based on Lemma D.3 and G.1 in [FW20], we can show the following inequality for any t∈(0,1)t\in(0,1) under event At\mathcal{A}_{t}:

Let ζ1:=w⊤x1\zeta_{1}:=\boldsymbol{w}^{\top}\boldsymbol{x}_{1} and ζ2:=w⊤x2\zeta_{2}:=\boldsymbol{w}^{\top}\boldsymbol{x}_{2}. Conditioned on x1,x2\boldsymbol{x}_{1},\boldsymbol{x}_{2}, we know that

For i=1,2i=1,2, by Taylor expansion of σ(ζi)\sigma(\zeta_{i}) around ξi\xi_{i} (note that σ\sigma is differentiable by assumption), there exists a random variable ηi\eta_{i} between ζi\zeta_{i} and ξi\xi_{i} such that

on the event At\mathcal{A}_{t} with t∈(0,1)t\in(0,1), where C>0C>0 is a constant depending on λσ\lambda_{\sigma}. This concludes (D.6).

Also, note the probability bound for Gaussian random vector x1\boldsymbol{x}_{1} and x2\boldsymbol{x}_{2} implies that

From the above arguments, we can bound the first term I1I_{1} via the following steps:

By choosing t=C/Nd14−εt=C/Nd^{\frac{1}{4}-\varepsilon}, we conclude that NNr∣J1∣\frac{N}{N_{r}}|J_{1}| cannot exceed C/Nd14−εC/Nd^{\frac{1}{4}-\varepsilon} with probability at least 1−α2/d141-\alpha^{2}/d^{\frac{1}{4}}, for any ε∈(0,1/4)\varepsilon\in(0,1/4). As for J2J_{2}, since σ⊥′\sigma^{\prime}_{\perp} is uniformly bounded by λσ\lambda_{\sigma} and all entries of ar\boldsymbol{a}_{r} are bounded by α/N\alpha/\sqrt{N}, we have

We conclude (D.2) by combining the above estimates of J1J_{1} and J2J_{2}.

D.2 Constructing the “Oracle” Estimator

In this subsection we prove the following lemma related to Lemma 10.

where the scalar τ∗\tau^{*} is defined in (4.3).

Denote N_{r}:=\mathopen{}\mathclose{{}\left|\mathcal{A}_{r}^{\alpha}}\right| for some constant α\alpha, and ir∈[N]i_{r}\in[N] as the index such that ir∈Arαi_{r}\in\mathcal{A}_{r}^{\alpha}. We define frf_{r} as an average over neurons with indices ir∈Arαi_{r}\in\mathcal{A}_{r}^{\alpha}, and fAf_{\boldsymbol{A}} as an approximation of frf_{r} in which the first-step gradient matrix G0\boldsymbol{G}_{0} in (B.3) is replaced by the rank-1 matrix A1\boldsymbol{A}_{1} defined in (B.27):

Moreover, by definition of Arα\mathcal{A}_{r}^{\alpha}, all aira_{i_{r}}’s are close to αN\frac{\alpha}{\sqrt{N}} for ir∈Arαi_{r}\in\mathcal{A}_{r}^{\alpha}; thus Lemma 32 (in particular (D.2)) can be directly applied to Br\boldsymbol{B}_{r}. As for Cr\boldsymbol{C}_{r}, since ∥Cr∥F≤∥C∥F\|\boldsymbol{C}_{r}\|_{F}\leq\|\boldsymbol{C}\|_{F}, we use part (iii)(iii) in Lemma 14 to obtain

With these concentration estimates, we know that when n>dn>d,

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 c>0c>0; this is due to the defined step size η=Θ(N)\eta=\Theta(\sqrt{N}), (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 n,d,Nn,d,N.

By the definition of Arα\mathcal{A}_{r}^{\alpha} and the Lipschitz property of σ\sigma, one can obtain

Now define vˉ:=ημ1μ1∗Nβ∗=ηˉμ1μ1∗β∗\bar{\boldsymbol{v}}:=\frac{\eta\mu_{1}\mu_{1}^{*}}{\sqrt{N}}\boldsymbol{\beta}_{*}=\bar{\eta}\mu_{1}\mu_{1}^{*}\boldsymbol{\beta}_{*}, which corresponds to the “population” version of v\boldsymbol{v}, and denote

Combining the inequalities (D.20) and (D.22), we know that for some constant CC,

where the last inequality holds with probability at least 1-\operatorname{exp}\mathopen{}\mathclose{{}\left({-cd}}\right) for some universal constant c>0c>0, due to the operator norm bound and concentration of the sample covariance matrix 1nX⊤X\frac{1}{n}\boldsymbol{X}^{\top}\boldsymbol{X} (for instance see [Ver18, Theorem 4.6.1]).

Now we take the expectation of fˉA\bar{f}_{\boldsymbol{A}} over initial weight wir\boldsymbol{w}_{i_{r}} in (D.21) to define

Note that for fixed x\boldsymbol{x}, \langle\boldsymbol{w},\boldsymbol{x}\rangle\sim\mathcal{N}(0,\mathopen{}\mathclose{{}\left\|{\boldsymbol{x}}}\right\|^{2}/d). Since σ\sigma is λσ\lambda_{\sigma}-Lipschitz, by the Hoeffding bound on sub-Gaussian random variables, conditioned on x\boldsymbol{x}, we have

where the last inequality is due to property of the sub-Gaussian norm ∥∥x∥/d−1∥ψ2≤C/d\|\|\boldsymbol{x}\|/\sqrt{d}-1\|_{\psi_{2}}\leq C/\sqrt{d} (see e.g. [Ver18, Theorem 3.1.1]) for some universal constant C>0C>0.

as n,d,N→∞n,d,N\to\infty, for some constant C>0C>0. By the Cauchy-Schwarz inequality,

where the failure probability only relates to r,α,N,d,nr,\alpha,N,d,n and is vanishing as N,d,n→∞N,d,n\to\infty. For simplicity, we only keep the leading orders and ignore the subordinate terms in the exact probability bounds.

for some constant C>0C>0, as n,d,N→∞n,d,N\to\infty proportionally.

The above analysis illustrates that because of the Gaussian initialization of aia_{i}, for any η=Θ(N)\eta=\Theta(\sqrt{N}), we can find a subset of neurons Arα\mathcal{A}_{r}^{\alpha} that receive a “good” learning rate, in the sense that the corresponding (sub-) network defined by frf_{r} can achieve the prediction risk close to τε∗\tau_{\varepsilon}^{*} when n≫dn\gg d.

Equation (D.10) reduces the prediction risk of our constructed frf_{r} to a one-dimensional Gaussian integral, which can be numerically evaluated for pairs of (σ,σ∗)(\sigma,\sigma^{*}). Denote κ∗=α∗ηˉμ1μ1∗\kappa^{*}=\alpha^{*}\bar{\eta}\mu_{1}\mu_{1}^{*}, we give a few examples in which we set ε=0\varepsilon=0 and the corresponding τ∗\tau^{*} is small. Note that due to Assumptions 1 and 2, choices of σ\sigma and σ∗\sigma^{*} considered below are centered with respect to standard Gaussian measure Γ\Gamma.

σ=σ∗=tanh\sigma=\sigma^{*}=\text{tanh}. Numerical integration yields τ∗≈3×10−4\tau^{*}\approx 3\times 10^{-4}, κ∗≈1.6\kappa^{*}\approx 1.6.

σ=σ∗=SoftPlus\sigma=\sigma^{*}=\text{SoftPlus}. Numerical integration yields τ∗≈0.03\tau^{*}\approx 0.03, κ∗≈0.96\kappa^{*}\approx 0.96.

σ=ReLU,σ∗=SoftPlus\sigma=\text{ReLU},\sigma^{*}=\text{SoftPlus}. Numerical integration yields τ∗≈0.09\tau^{*}\approx 0.09, κ∗≈0.94\kappa^{*}\approx 0.94.

Observe that in all the above examples, τ∗\tau^{*} can be obtained by some finite α∗\alpha^{*} (or equivalently κ∗\kappa^{*}). In the following analysis of kernel ridge regression, we drop the small constant ε\varepsilon in Lemma 33 and directly apply the asymptotic statement given in (D.39).

We make the following remarks on the calculation of τ∗\tau^{*} in (4.3).

When σ=σ∗\sigma=\sigma^{*}, we intuitively expect τ∗\tau^{*} to be small when the nonlinearity is smooth such that it is to some extent unchanged under Gaussian convolution (when κ\kappa is chosen appropriately).

Adding weight decay with strength λ<1\lambda<1 to the first-layer parameters W0\boldsymbol{W}_{0} simply corresponds to multiplying ξ2\xi_{2} in the definition of τ\tau (D.30) by a factor of (1−λ)(1-\lambda).

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 R1(λ)\mathcal{R}_{1}(\lambda). We first define the following quantities which R1(λ)\mathcal{R}_{1}(\lambda) can be decomposed into (see Lemma 35):

We begin by defining a concentration event A\mathcal{A} on the empirical feature matrix Σ^Φ\widehat{\boldsymbol{\Sigma}}_{\Phi}, 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 c>0c>0 such that the following holdsNote that for λ=0\lambda=0, the LHS of the inequality may be interpreted as a pseudo-inverse.

for all large n>Nn>N, where K:=λσN∥W1∥FK:=\frac{\lambda_{\sigma}}{\sqrt{N}}\|\boldsymbol{W}_{1}\|_{F}.

Proof. First observe that the null space of ΣΦ\boldsymbol{\Sigma}_{\Phi} contains the null space of Σ^Φ\widehat{\boldsymbol{\Sigma}}_{\Phi}. Also, notice that Σ^Φ\widehat{\boldsymbol{\Sigma}}_{\Phi} is a sample covariance matrix taking the form of

for all large n,Nn,N. This proposition is proved by setting t=Nt=\sqrt{N} and noting that

Similarly, for the “ridgeless” case λ=0\lambda=0, we define

Lemma 34 entails that both Aλ\mathcal{A}_{\lambda} and A0\mathcal{A}_{0} hold with probability at least 1−2e−cN1-2e^{-c\sqrt{N}}. Following the remark on [Bac23, Lemma 7.1], under events Aλ\mathcal{A}_{\lambda} and A0\mathcal{A}_{0}, we can obtain that

and (1−t)(ΣΦ−Σ^Φ)≼t(Σ^Φ+λI)(1-t)(\boldsymbol{\Sigma}_{\Phi}-\widehat{\boldsymbol{\Sigma}}_{\Phi})\preccurlyeq t(\widehat{\boldsymbol{\Sigma}}_{\Phi}+\lambda\boldsymbol{I}), 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 Aλ\mathcal{A}_{\lambda} and A0\mathcal{A}_{0}, we know that

We now control B1,B2,V1,V2B_{1},B_{2},V_{1},V_{2} under the high probability events Aλ\mathcal{A}_{\lambda} and A0\mathcal{A}_{0}.

By the definition of fˇ\check{f}, 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 Aλ\mathcal{A}_{\lambda} and A0\mathcal{A}_{0},

where (i)(i) follows from the definition of the concentration events A\mathcal{A}, (D.44) and (D.46).

Finally, from [Bac23, Lemma 7.2], we have

where the last step is a triangle inequality due to ∥f∗−fˇ∥L22≤∥f∗−fr∥L22\|f^{*}-\check{f}\|^{2}_{L^{2}}\leq\|f^{*}-f_{r}\|^{2}_{L^{2}}.

For V1V_{1}, note that under event Aλ\mathcal{A}_{\lambda},

Similarly for V2V_{2}, under event Aλ\mathcal{A}_{\lambda}, we have

where (iii)(iii) is due to the boundedness of σ\sigma and ∥f⊥∥L2\|f_{\perp}\|_{L^{2}}. Combining V1V_{1} and V2V_{2}, and taking x=Cnε−1x=Cn^{\varepsilon-1} in (D.62) and (D.65), for some C>0C>0 and any small ε>0\varepsilon>0, we arrive at

with probability at least 1−n−ε1-n^{-\varepsilon}.

The following lemma provides a decomposition of the prediction risk R1(λ)\mathcal{R}_{1}(\lambda) in terms of B1,B2,V1,V2B_{1},B_{2},V_{1},V_{2} analyzed above.

Under the same assumptions as Lemma 10, if we choose λ=Ω(nε−1)\lambda=\Omega(n^{\varepsilon-1}) for small ε>0\varepsilon>0, then the prediction risk of the CK ridge estimator admits the following upper bound

where B1,B2B_{1},B_{2} 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 Aλ\mathcal{A}_{\lambda} and A0\mathcal{A}_{0} happens with high probability for fixed t∈(0,1)t\in(0,1), if we set λ=Ω(nε−1)\lambda=\Omega(n^{\varepsilon-1}) for some small ε>0\varepsilon>0, then Lemma 35 entails

Finally, due to the upper-bound (D.76), we conclude that

with probability one as n,d,N→∞n,d,N\to\infty proportionally and n/d>ψ1∗n/d>\psi_{1}^{*}, where τ∗\tau^{*} is defined in (4.3).