Dynamics of Finite Width Kernel and Prediction Fluctuations in Mean Field Neural Networks

Blake Bordelon, Cengiz Pehlevan

Introduction

Learning dynamics of deep neural networks are challenging to analyze and understand theoretically, but recent progress has been made by studying the idealization of infinite-width networks. Two types of infinite-width limits have been especially fruitful. First, the kernel or lazy infinite-width limit, which arises in the standard or neural tangent kernel (NTK) parameterization, gives prediction dynamics which correspond to a linear model . This limit is theoretically tractable but fails to capture adaptation of internal features in the neural network, which are thought to be crucial to the success of deep learning in practice. Alternatively, the mean field or μ\mu-parameterization allows feature learning at infinite width .

With a set of well-defined infinite-width limits, prior theoretical works have analyzed finite networks in the NTK parameterization perturbatively, revealing that finite width both enhances the amount of feature evolution (which is still small in this limit) but also introduces variance in the kernels and the predictions over random initializations . Because of these competing effects, in some situations wider networks are better, and in others wider networks perform worse .

In this paper, we analyze finite-width network learning dynamics in the mean field parameterization. In this parameterization, wide networks are empirically observed to outperform narrow networks . Our results and framework provide a methodology for reasoning about detrimental finite-size effects in such feature-learning neural networks. We show that observable averages involving kernels and predictions obey a well-defined power series in inverse width even in rich training regimes. We generally observe that the leading finite-size corrections to both the bias and variance components of the square loss are increased for narrower networks, and diminish performance. Further, we show that richer networks are closer to their corresponding infinite-width mean field limit. For simple tasks and architectures the leading O(1/width)\mathcal{O}(1/\text{width}) corrections to the error can be descriptive, while for large sample size or more realistic tasks, higher order corrections appear to become relevant. Concretely, our contributions are listed below:

Starting from a dynamical mean field theory (DMFT) description of infinite-width nonlinear deep neural network training dynamics, we provide a complete recipe for computing fluctuation dynamics of DMFT order parameters over random network initializations during training. These include the variance of the training and test predictions and the O(1/width)\mathcal{O}(1/\text{width}) variance of feature and gradient kernels throughout training.

We first solve these equations for the lazy limit, where no feature learning occurs, recovering a simple differential equation which describes how prediction variance evolves during learning.

We solve for variance in the rich feature learning regime in two-layer networks and deep linear networks. We show richer nonlinear dynamics improve the signal-to-noise ratio (SNR) of kernels and predictions, leading to closer agreement with infinite-width mean field behavior.

We analyze in a two-layer model why larger training set sizes in the overparameterized regime enhance finite-width effects and how richer training can reduce this effect.

We show that large learning rate effects such as edge-of-stability dynamics can be well captured by infinite width theory, with finite size variance accurately predicted by our theory.

We test our predictions in Convolutional Neural Networks (CNNs) trained on CIFAR-10 . We observe that wider networks and richly trained networks have lower logit variance as predicted. However, the timescale of training dynamics is significantly altered by finite width even after ensembling. We argue that this is due to a detrimental correction to the mean dynamical NTK.

Infinite-width networks at initialization converge to a Gaussian process with a covariance kernel that is computed with a layerwise recursion . In the large but finite width limit, these kernels do not concentrate at each layer, but rather propagate finite-size corrections forward through the network . During gradient-based training with the NTK parameterization, a hierarchy of differential equations have been utilized to compute small feature learning corrections to the kernel through training . However the higher order tensors required to compute the theory are initialization dependent, and the theory breaks down for sufficiently rich feature learning dynamics. Various works on Bayesian deep networks have also considered fluctuations and perturbations in the kernels at finite width during inference . Other relevant work in this domain are .

Problem Setup

Review of Dynamical Mean Field Theory

Dynamical Fluctuations Around Mean Field Theory

We are interested in going beyond the infinite-width limit to study more realistic finite-width networks. In this regime, the order parameters q\bm{q} fluctuate in a O(1/N)\mathcal{O}(1/\sqrt{N}) neighborhood of q∞\bm{q}_{\infty} . Statistics of these fluctuations can be calculated from a general cumulant expansion (see App. D) . We will focus on the leading-order corrections to the infinite-width limit in this expansion.

The finite-width NN average of observable O(q)O(\bm{q}) across initializations, which we denote by <O(q)>N\left<O(\bm{q})\right>_{N}, admits an expansion of the form whose leading terms are

where <>∞\left<\right>_{\infty} denotes an average over the Gaussian distribution q∼N(q∞,−1N(∇q2S[q∞])−1)\bm{q}\sim\mathcal{N}\left(\bm{q}_{\infty},-\frac{1}{N}\left(\nabla^{2}_{\bm{q}}S[\bm{q}_{\infty}]\right)^{-1}\right) and the function V(q)≡S(q)−S(q∞)−12(q−q∞)⊤∇q2S(q∞)(q−q∞)V(\bm{q})\equiv S(\bm{q})-S(\bm{q}_{\infty})-\frac{1}{2}(\bm{q}-\bm{q}_{\infty})^{\top}\nabla^{2}_{\bm{q}}S(\bm{q}_{\infty})(\bm{q}-\bm{q}_{\infty}) contains cubic and higher terms in the Taylor expansion of SS around q∞\bm{q}_{\infty}. The terms shown include all the leading and sub-leading terms in the series in powers of 1/N1/N. The terms in ellipses are at least O(N−1)\mathcal{O}(N^{-1}) suppressed compared to the terms provided.

The proof of this statement is given in App. D. The central object to characterize finite size effects is the unperturbed covariance (the propagator): Σ=−[∇2S(q∞)]−1\bm{\Sigma}=-\left[\nabla^{2}S(\bm{q}_{\infty})\right]^{-1}. This object can be shown to capture leading order fluctuation statistics <(q−q∞)(q−q∞)⊤>N=1NΣ+O(N−2)\left<\left(\bm{q}-\bm{q}_{\infty}\right)\left(\bm{q}-\bm{q}_{\infty}\right)^{\top}\right>_{N}=\frac{1}{N}\bm{\Sigma}+\mathcal{O}(N^{-2}) (App. D.1), which can be used to reason about, for example, expected square error over random initializations. Correction terms at finite width may give a possible explanation of the superior performance of wide networks at fixed γ\gamma . To calculate such corrections, in App. E, we provide a complete description of Hessian ∇q2S(q)\nabla^{2}_{\bm{q}}S(\bm{q}) and its inverse (the propagator) for a depth-LL network. This description constitutes one of our main results. The resulting expressions are lengthy and are left to App. E. Here, we discuss them at a high level. Conceptually there are two primary ingredients for obtaining the full propagator:

Hessian sub-blocks κ\kappa which describe the uncoupled variances of the kernels, such as

Similar terms also appear in other studies on finite width Bayesian inference and in studies on kernel variance at initialization .

Blocks which capture the sensitivity of field averages to pertubations of order parameters, such as

In App. E, we calculate κ\bm{\kappa} and D\bm{D} tensors, and show how to use them to calculate the propagator. As an example of our results:

The necessary order parameters for calculating the fluctuations are obtained by solving the DMFT using numerical methods introduced in . We provide a pseudocode for this procedure in App. F. We proceed to solve the equations defining Σ\bm{\Sigma} in special cases which are illuminating and numerically feasible including lazy training, two layer networks and deep linear NNs.

Lazy Training Limit

where averages are computed over the training distribution D\mathcal{D}.

For MSE loss, the prediction error covariance ΣΔ(t,s)=NCov0(Δ(t),Δ(s))\bm{\Sigma}^{\Delta}(t,s)=N\text{Cov}_{0}(\bm{\Delta}(t),\bm{\Delta}(s)) satisfies a differential equation (App. H)

where Δk∞(t)≡exp⁡(−λkt)<ψk(x)y(x)>x{\Delta}_{k}^{\infty}(t)\equiv\exp\left(-\lambda_{k}t\right)\left<\psi_{k}(\bm{x})y(\bm{x})\right>_{\bm{x}} are the errors at infinite width for eigenmode kk.

Rich Regime in Two-Layer Networks

In this section, we analyze how feature learning alters the variance through training. We show a denoising effect where the signal to noise ratios of the order parameters improve with feature learning.

In the rich regime, the kernel evolves over time but inherits fluctuations from the training errors Δ\bm{\Delta}. To gain insight, we first study a simplified setting where the data distribution is a single training example x\bm{x} and single test point x⋆\bm{x}_{\star} in a two layer network. We will track Δ(t)=y−f(x,t)\Delta(t)=y-f(\bm{x},t) and the test prediction f⋆(t)=f(x⋆,t)f_{\star}(t)=f(\bm{x}_{\star},t). To identify the dynamics of these predictions we need the NTK K(t)K(t) on the train point, as well as the train-test NTK K⋆(t)K_{\star}(t). In this case, all order parameters can be viewed as scalar functions of a single time index (unlike the deep network case, see App. E).

where [ΘK](t,s)=Θ(t−s)K(s)[\bm{\Theta}_{K}](t,s)=\Theta(t-s)K(s), [ΘΔ](t,s)=Θ(t−s)Δ(s)[\bm{\Theta}_{\Delta}](t,s)=\Theta(t-s)\Delta(s) are Heaviside step functions and D(t,s)=<∂∂Δ(s)(ϕ(h(t))2+g(t)2)>D(t,s)=\left<\frac{\partial}{\partial\Delta(s)}(\phi(h(t))^{2}+g(t)^{2})\right> and D⋆(t,s)=<∂∂Δ(s)(ϕ(h(t))ϕ(h⋆(t))+g(t)g⋆(t))>D_{\star}(t,s)=\left<\frac{\partial}{\partial\Delta(s)}(\phi(h(t))\phi(h_{\star}(t))+g(t)g_{\star}(t))\right> quantify sensitivity of the kernel to perturbations in the error signal Δ(s)\Delta(s). Lastly κ\kappa and κ⋆\kappa_{\star} are the uncoupled variances of K(t)K(t) and K⋆⋆(t)K_{\star\star}(t) and κ⋆\kappa_{\star} is the uncoupled covariance of K(t),K⋆(t)K(t),K_{\star}(t).

In Fig. 3, we plot the resulting theory (diagonal blocks of Σq1\bm{\Sigma}_{\bm{q}_{1}} from Equation 8) for two layer neural networks. As predicted by theory, all average squared deviations from the infinite width DMFT scale as O(N−1)\mathcal{O}(N^{-1}). Similarly, the average kernels <K>\left<K\right> and test predictions <f⋆>\left<f_{\star}\right> change by a larger amount for larger γ\gamma (equation (I.1)). The experimental variances also match the theory quite accurately. The variance of the train error Δ(t)\Delta(t) peaks earlier and at a lower value for richer training, but all variances go to zero at late time as the model approaches the interpolation condition Δ=0\Delta=0. As γ→0\gamma\to 0 the curve approaches N Var(Δ(t))∼κ y2 t2 e−2tN\ \text{Var}(\Delta(t))\sim\kappa\ y^{2}\ t^{2}\ e^{-2t}, where κ\kappa is the initial NTK variance (see Section 5). While the train prediction variance goes to zero, the test point prediction does not, with richer networks reaching a lower asymptotic variance. We suspect this dynamical effect could explain lower variance observed in feature learning networks compared to lazy networks . In Fig. A.1, we show that the reduction in variance is not due to a reduction in the uncoupled variance κ(t,s)\kappa(t,s), which increases in γ\gamma. Rather the reduction in variance is driven by the coupling of perturbations across time.

2 Offline Training with Multiple Samples or Online Training in High Dimension

However, at finite width, both the Δy(t)\Delta_{y}(t) and the P−1P-1 orthogonal variables Δ⊥\bm{\Delta}_{\perp} inherit initialization variance, which we represent as ΣΔy(t,s)\Sigma_{\Delta_{y}}(t,s) and Σ⊥(t,s)\Sigma_{\perp}(t,s). In Fig. 4 (a)-(b) we show this approximate solution <∣Δ(t)∣2>∼Δy(t)2+2NΔy1(t)Δy(t)+1NΣΔy(t,t)+(P−1)NΣ⊥(t,t)+O(N−2)\left<|\bm{\Delta}(t)|^{2}\right>\sim\Delta_{y}(t)^{2}+\frac{2}{N}\Delta^{1}_{y}(t)\Delta_{y}(t)+\frac{1}{N}\Sigma_{\Delta_{y}}(t,t)+\frac{(P-1)}{N}\Sigma_{\perp}(t,t)+\mathcal{O}(N^{-2}) across varying γ\gamma and varying PP (see Appendix J for ΣΔy\Sigma_{\Delta_{y}} and Σ⊥\Sigma_{\perp} formulas). We see that variance of train point predictions fμ(t)f_{\mu}(t) increases with the total number of points despite the signal of the target vector ∑μyμ2\sum_{\mu}y_{\mu}^{2} being fixed. In this model, the bias correction 2NΔy1(t)Δy(t)\frac{2}{N}\Delta^{1}_{y}(t)\Delta_{y}(t) is always O(1/N)\mathcal{O}(1/N) but the variance correction is O(P/N)\mathcal{O}(P/N). The fluctuations along the P−1P-1 orthogonal directions begin to dominate the variance at large PP. Fig. 4 (b) shows that as PP increases, the leading order approximation breaks down as higher order terms become relevant. Analysis for online training reveals identical fluctuation statistics, but with variance that scales as ∼D/N\sim D/N (Appendix K) as we verify in Figure 4 (e)-(f).

Deep Networks

Variance can be Small Near Edge of Stability

In this section, we move beyond the gradient flow formalism and ask what large step sizes do to finite size effects. Recent studies have identified that networks trained at large learning rates can be qualitatively different than networks in the gradient flow regime, including the catapult and edge of stability (EOS) phenomena . In these settings, the kernel undergoes an initial scale growth before exhibiting either a recovery or a clipping effect. In this section, we explore whether these dynamics are highly sensitive to initialization variance or if finite networks are well captured by mean field theory. Following , we consider two layer networks trained on a single example ∣x∣2=D|\bm{x}|^{2}=D and y=1y=1. We use learning rate η\eta and feature learning strength γ\gamma. The infinite width mean field equations for the prediction ftf_{t} and the kernel KtK_{t} are (App. M)

For small η\eta, the equations are well approximated by the gradient flow limit and for small γ\gamma corresponds to a discrete time linear model. For large ηγ>1\eta\gamma>1, the kernel KK progressively sharpens (increases in scale) until it reaches 2/η2/\eta and then oscillates around this value. It may be expected that near the EOS, the large oscillations in the kernels and predictions could lead to amplified finite size effects, however, we show in Fig. 6 that the leading order propagator elements decrease even after reaching the EOS threshold, indicating reduced disagreement between finite and infinite width dynamics.

Finite Width Alters Bias, Training Rate, and Variance in Realistic Tasks

To analyze the effect of finite width on neural network dynamics during realistic learning tasks, we studied a vanilla depth-66 ReLU CNN trained on CIFAR-10 (experimental details in App. B, G.2) In Fig. 7, we train an ensemble of E=8E=8 independently initialized CNNs of each width NN. Wider networks not only have better performance for a single model (solid), but also have lower bias (dashed), measured with ensemble averaging of the logits. Because of faster convergence of wide networks, we observe wider networks have higher variance, but if we plot variance at fixed ensembled training accuracy, wider networks have consistently lower variance (Fig. 7(d)).

We next seek an explanation for why wider networks after ensembling trains at a faster rate. Theoretically, this can be rationalized by a finite-width alteration to the ensemble averaged NTK, which governs the convergence timescale of the ensembled predictions (App. G.1). Our analysis in App. G.1 suggests that the rate of convergence receives a finite size correction with leading correction O(N−1)\mathcal{O}(N^{-1}) G.2. To test this hypothesis, we fit the ensemble training loss curve to exponential function L≈Cexp⁡(−RNt)\mathcal{L}\approx C\exp\left(-R_{N}t\right) where CC is a constant. We plot the fit RNR_{N} as a function of N−1N^{-1} result in Fig. 7(e). For large NN, we see the leading behavior is linear in N−1N^{-1}, but begins to deviate at small NN as a quadratic function of N−1N^{-1}, suggesting that second order effects become relevant around N≲100N\lesssim 100.

In App. Fig. A.4, we train a smaller subset of CIFAR-10 where we find that RNR_{N} is well approximated by a O(N−1)\mathcal{O}(N^{-1}) correction, consistent with the idea that higher sample size drives the dynamics out of the leading order picture. We also analyze the effect of γ\gamma on variance in this task. In App. Fig. A.5, we train N=64N=64 models with varying γ\gamma. Increased γ\gamma reduces variance of the logits and alters the representation (measured with kernel-task alignment), the training and test accuracy are roughly insensitive to the richness γ\gamma in the range we considered.

Discussion

We studied the leading order fluctuations of kernels and predictions in mean field neural networks. Feature learning dynamics can reduce undesirable finite size variance, making finite networks order parameters closer to the infinite width limit. In several toy models, we revealed some interesting connections between the influence of feature learning, depth, sample size, and large learning rate and the variance of various DMFT order parameters. Lastly, in realistic tasks, we illustrated that bias corrections can be significant as rates of learning can be modified by width. Though our full set of equations for the leading finite size fluctuations are quite general in terms of network architecture and data structure, they are only derived at the level of rigor of physics rather than a formally rigorous proof which would need several additional assumptions to make the perturbation expansion properly defined. Further, the leading terms in our perturbation series involving only Σ\bm{\Sigma} does not capture the complete finite size distribution defined in Eq. (3), especially as the sample size becomes comparable to the width. It would be interesting to see if proportional limits of the rich training regime where samples and width scale linearly can be examined dynamically . Future work could explore in greater detail the higher order contributions from averages involving powers of V(q)V(\bm{q}) by examining cubic and higher derivatives of SS in Eq. (3). It could also be worth examining in future work how finite size impacts other biologically plausible learning rules, where the effective NTK can have asymmetric (over sample index) fluctuations . Also of interest would be computing the finite width effects in other types of architectures, including residual networks with various branch scalings . Further, even though we expect our perturbative expressions to give a precise asymptotic description of finite networks in mean field/μ\muP, the resulting expressions are not realistically computable in deep networks trained on large dataset size PP for long times TT since the number of Hessian entries scales as O(T4P4)\mathcal{O}(T^{4}P^{4}) and a matrix of this size must be stored in memory and inverted in the general case. Future work could explore solveable special cases such as high dimensional limits.

Code to reproduce the experiments in this paper is provided at https://github.com/Pehlevan-Group/dmft_fluctuations. Details about numerical methods and computational implementation can be found in Appendices F and N.

Acknowledgements

CP is supported by NSF Award DMS-2134157, NSF CAREER Award IIS-2239780, and a Sloan Research Fellowship. BB is supported by a Google PhD research fellowship and NSF Award DMS-2134157. This work has been made possible in part by a gift from the Chan Zuckerberg Initiative Foundation to establish the Kempner Institute for the Study of Natural and Artificial Intelligence. The computations in this paper were run on the FASRC cluster supported by the FAS Division of Science Research Computing Group at Harvard University. BB thanks Alex Atanasov, Jacob Zavatone-Veth for their comments on this manuscript and Boris Hanin, Greg Yang, Mufan Bill Li and Jeremy Cohen for helpful discussions.

References

Appendix

Appendix A Additional Figures

Appendix B CIFAR-10 Experimental Details

All models were trained with standard SGD with a batch size of 256256. Each element in the ensemble of EE networks is trained on identical batches presented in identical order. For the Figure 7 experiments, the raw learning rate is scaled as η=10Nγ\eta=10N\sqrt{\gamma} with γ=0.2\gamma=0.2 (note that mean field theory requires scaling the raw learning rate linearly with NN since the raw NTK is O(N−1)\mathcal{O}(N^{-1}) ). For Figure A.5, the learning rate is η=5Nγ\eta=5N\sqrt{\gamma}. We find that choosing η∝γ\eta\propto\sqrt{\gamma} gives approximately conserved training times across γ\gamma (though distinct representation dynamics). The Figure A.4 shows the dynamics of fitting P=64P=64 training points with full batch gradient descent and γ=0.1\gamma=0.1.

Appendix C Review of DMFT: Deriving the Action

In this section we derive the DMFT action which contains all of the necessary statistical information about randomly initialized finite width NN networks. From the action SS the DMFT saddle point and the propagator can be computed. This derivation follows closely the original derivation by Bordelon & Pehlevan . We start by writing the gradient flow dynamics on weight matrices

where we introduced the feature and gradient kernels

Moments of these fields can be computed through differentiation with respect to the sources j,v\bm{j},\bm{v} near zero-source (j=v=0\bm{j}=\bm{v}=0)

To average over the initial weights, we introduce a Fourier representation of the Dirac-Delta function 1=∫dzδ(z)=∫dzdz^2πexp⁡(iz^z)1=\int dz\delta(z)=\int\frac{dzd\hat{z}}{2\pi}\exp(i\hat{z}z). We perform this transformation for each of the fields to enforce their definition

We insert these Dirac delta functions so that we can directly average over the weights

To enforce the definitions of the new order parameters {Φ,G,A}\{\Phi,G,A\} we again introduce Dirac-delta functions

Analogous constraints for GG and AA are enforced with conjugate variables G^,B\hat{G},B. After introducing these variables, we find that the moment generating functional has the form

where SS is the O(1)\mathcal{O}(1) DMFT action which defines the statistical distribution over the dynamics. The action takes the form

Appendix D Cumulant Expansion of Observables

We are interested in a principled power series expansion (in 1/N1/N) of any observable average <O(q)>\left<O(\bm{q})\right> that depends on DMFT order parameters q\bm{q}. At any width NN the observable average takes the form

As discussed in the main text, the N→∞N\to\infty limit gives <O(q)>N∼O(q∞)\left<O(\bm{q})\right>_{N}\sim O(\bm{q}_{\infty}) where ∂S∂q∣q∞=0\frac{\partial S}{\partial\bm{q}}|_{\bm{q}_{\infty}}=0 by a steepest descent argument . We assume that SS’s Hessian is negative semidefinite so that Σ≡−[∇2S(q)∣q∞]−1⪰0\bm{\Sigma}\equiv-\left[\nabla^{2}S(\bm{q})|_{\bm{q}_{\infty}}\right]^{-1}\succeq 0 and Taylor expand S(q)S(\bm{q}) around the saddle point q∞\bm{q}_{\infty} giving S(q)=S(q∞)+12(q−q∞)⊤∇2S(q)(q−q∞)+V(q−q∞)S(\bm{q})=S(\bm{q}_{\infty})+\frac{1}{2}(\bm{q}-\bm{q}_{\infty})^{\top}\nabla^{2}S(\bm{q})(\bm{q}-\bm{q}_{\infty})+V(\bm{q}-\bm{q}_{\infty}). We note that the remainder function VV contains only cubic and higher powers of q−q∞≡δ/N\bm{q}-\bm{q}_{\infty}\equiv\bm{\delta}/\sqrt{N}. The variable δ\bm{\delta} will be order O(1)\mathcal{O}(1). This will allow us to verify that additional terms are suppressed in powers of 1/N1/N. Expanding both the numerator and denominator’s integrands in powers of VV, we find

where <>∞\left<\right>_{\infty} represents an average over the Gaussian fluctuation N(q∞,−1N[∇q2S(q∞)]−1)\mathcal{N}\left(\bm{q}_{\infty},-\frac{1}{N}\left[\nabla^{2}_{\bm{q}}S(\bm{\bm{q}}_{\infty})\right]^{-1}\right). We see that the series in the denominator contains terms of the form Nkk!<Vk>∞\frac{N^{k}}{k!}\left<V^{k}\right>_{\infty} while the numerator depends on terms of the form Nkk!<VkO>∞/<O>∞\frac{N^{k}}{k!}\left<V^{k}O\right>_{\infty}/\left<O\right>_{\infty}. In either of these power series, the kk-th term can contribute at most

since VV contributes only cubic and higher terms. Thus each term in the numerator and denominator’s series contains increasing powers of 1/N1/N. Concretely, each of the two series have terms of order {N0,N−1,N−1,N−2,N−2,...}\{N^{0},N^{-1},N^{-1},N^{-2},N^{-2},...\}. Thus any quantity of the form <O><O>∞\frac{\left<O\right>}{\left<O\right>_{\infty}} admits a ratio of power series in powers of 1/N1/N. One could truncate each of the series in the numerator and denominator to a desired order in NN. Alternatively, the denominator could be expanded giving a single series (the cumulant expansion ). The first few terms in the cumulant expansion have the form

In this work, we mainly are interested in the leading order correction to <O>\left<O\right> which can always be obtained with the truncation after the terms linear in VV for any observable OO.

We will now analyze the fluctuation statistics of our order parameters around the saddle point <(q−q∞)(q−q∞)⊤>N\left<(\bm{q}-\bm{q}_{\infty})(\bm{q}-\bm{q}_{\infty})^{\top}\right>_{N} which has the form

as stated in the main text and verified empirically in Figure 3 (a). The reason that the terms in the numerator involving VV can be no larger than O(N−2)\mathcal{O}(N^{-2}) comes from vanishing of odd moments for q−q∞\bm{q}-\bm{q}_{\infty} in the unperturbed distribution. Thus the leading expression for <(q−q∞)(q−q∞)⊤>\left<(\bm{q}-\bm{q}_{\infty})(\bm{q}-\bm{q}_{\infty})^{\top}\right> only depends on Σ\bm{\Sigma} and not on VV.

D.2 Mean Deviation from DMFT

Although the square displacement from DMFT only depended on Σ\Sigma and not on VV, we note that the average order parameter displacement <q−q∞>\left<\bm{q}-\bm{q}_{\infty}\right> does receive a O(1/N)\mathcal{O}(1/N) correction that depends on the perturbed potential VV

where in the last line we used Stein’s lemma (Gaussian integration by parts) for the Gaussian distribution over q\bm{q}. Note that <∂V∂q>∞∼O(1N)\left<\frac{\partial V}{\partial\bm{q}}\right>_{\infty}\sim\mathcal{O}\left(\frac{1}{N}\right) since the derivative of the cubic term in VV gives a quadratic function of q−q∞\bm{q}-\bm{q}_{\infty}, whose average must be O(N−1)\mathcal{O}(N^{-1}). In this work, we focus primarily on the structure of the propagator, but outline a general recipe for getting the leading mean correction in Appendix G and H.2.

D.3 Covariance of Order Parameters

Lastly, we combine the previous two observations to reason about the scaling of the order parameter covariance over initializations. We note that the leading covariance of the order parameters over random initializations is also given by the propagator: Cov(q)∼1NΣ+O(N−2)\text{Cov}(\bm{q})\sim\frac{1}{N}\bm{\Sigma}+\mathcal{O}(N^{-2}), since

due to the arguments above which showed that <(q−q∞)(q−q∞)⊤>∼1NΣ+O(N−2)\left<(\bm{q}-\bm{q}_{\infty})(\bm{q}-\bm{q}_{\infty})^{\top}\right>\sim\frac{1}{N}\bm{\Sigma}+\mathcal{O}(N^{-2}) and that q∞−<q>N∼O(N−1)\bm{q}_{\infty}-\left<\bm{q}\right>_{N}\sim\mathcal{O}(N^{-1}). Therefore, in the leading order picture, it is safe to associate Σ\bm{\Sigma} with the covariance of order parameters over random initializations of the network weights.

Appendix E Propagator Structure for the full DMFT Action

This enumerates all possible non-vanishing terms in the Hessian. We can now construct a block matrix of these Hessians by partitioning our order parameters q=[q1,q2]⊤\bm{q}=[\bm{q}_{1},\bm{q}_{2}]^{\top} where

This choice will become apparent shortly.

To calculate the full propagator Σ=−[∇q2S]−1\bm{\Sigma}=-\left[\nabla^{2}_{\bm{q}}S\right]^{-1}, we will assume invertibility of the upper block Σ0=−[∇q12S]−1\bm{\Sigma}^{0}=-\left[\nabla^{2}_{\bm{q}_{1}}S\right]^{-1} and use this in the Schur complement

We seek a physically sensible inverse where the variance of q12\bm{q}_{1}^{2} is vanishing . This leads to the following sub-propagator Σ0\bm{\Sigma}^{0}

Thus given κ,U\bm{\kappa},\bm{U}, we can solve for Σ0\bm{\Sigma}^{0} and ultimately for the full propagator Σ\bm{\Sigma}. The relevant entries in κ\bm{\kappa} and U\bm{U} are given by those second derivatives calculated above. We note that each of the field derivatives needed for U\bm{U} can be computed implicitly from the field dynamics. For example, for the Δμ(t)\Delta_{\mu}(t) derivatives we have

Appendix F Solving for the Propagator

In this section we sketch out the required steps to obtain the propagator Σ\bm{\Sigma}.

Step 3: After populating the entries of the block matrix for the Hesssian ∇2S\nabla^{2}S, we then calculate the propagator Σ\Sigma with a matrix inversion. Since we discretized time, this is a finite dimensional matrix.

The step 1 above demands a solution to the infinite width DMFT equations (solving for the saddle point q∞\bm{q}_{\infty}). We will now give a detailed set of instructions about how the infinite width limit for q∞q_{\infty} is solved (step 1 above). This corresponds to the algorithm of Bordelon & Pehlevan 2022 to solve the saddle point equations ∂∂qS(q)∣q∞=0\frac{\partial}{\partial\bm{q}}S(\bm{q})|_{\bm{q}_{\infty}}=0 .

Step 3: For each sample, solve integral equations for h(t)h(t) and z(t)z(t).

These will be samples from the single site distribution for h,zh,z

Repeat steps 2-5 until the order parameters converge.

Below we provide a pseudocode algorithm to solve for the propagator elements.

The above propagator solver builds on the solution to the DMFT equations which is provided below.

Appendix G Leading Correction to the Mean Order Parameters

In this section we use the propagator structure derived in the last section to reason about the leading finite size correction to <q>\left<\bm{q}\right> at width NN. Letting the indices i,j,k,ni,j,k,n enumerate all entries of the order parameters in q\bm{q} (technically this is a sum over samples and an integral over time for gradient flow), we find the leading Pade Approximant for the mean has the form (App D)

where δj=N(qj−qj∞)\delta_{j}=\sqrt{N}(q_{j}-q_{j}^{\infty}) and the derivatives are computed at the saddle point. In the last line, we utilized Wick’s theorem and the permutation symmetry of the third derivative ∂3S∂qi∂qj∂qk\frac{\partial^{3}S}{\partial q_{i}\partial q_{j}\partial q_{k}} to evaluate the four point averages in terms of the propagator Σij\Sigma_{ij}, which was provided in the preceding section E. In practice computing even the full set of second derivatives for the DMFT action to get Σ\Sigma is quite challenging. Despite the challenge of computing the mean order parameter correction, these corrections are relevant in practice and crucially distinguish the training timescales of deep networks at different widths as we show in Figures 7 and A.4.

Supposing that we solved for the propagator Σ\bm{\Sigma}, using the formalism in the preceeding section, we can compute the O(N−1)\mathcal{O}(N^{-1}) correction to the average network prediction error due to finite size. We let <Δ(t)>\left<\bm{\Delta}(t)\right> represent the average of errors over an ensemble of width NN networks.

where ΣμννKΔ(t,t)\Sigma^{K\Delta}_{\mu\nu\nu}(t,t) is the leading covariance (propagator element) between the kernel Kμν(t)K_{\mu\nu}(t) and prediction error Δν(t)\Delta_{\nu}(t). We see that the average kernel <Kμν(t)>\left<K_{\mu\nu}(t)\right> (which depends on the finite width NN) plays an important role in characterizing the timescales of the average prediction dynamics. Once this equation is solved for <Δμ(t)>\left<\Delta_{\mu}(t)\right>, the square loss at width NN and time tt has the form

We will now comment on the structure of the cross term in this above solution. First, if <K>⪰K∞\left<\bm{K}\right>\succeq\bm{K}^{\infty} and ΣKΔ\Sigma^{K\Delta} is negligible then the average errors at finite width will decay more rapidly than the infinite width model. However, we suspect that in general, <K>−K∞\left<\bm{K}\right>-\bm{K}^{\infty} contains many negative eigenvalues since signal propagation at finite width tends to reduce the scale of feature kernels . We suspect that this is the cause of the slower dynamics of ensembled predictors for narrower networks in Figure 7 and Figure A.4. Additionally, the term involving ΣKΔ\Sigma^{K\Delta} will generically increase the cross term since the dynamics of Δ\Delta cause its fluctuations to become anti-correlated with the fluctuations in KK. In general, it is challenging to make strong definitive statements about the relative scale of these competing effects on the cross term. However, we can say more about this solution in the lazy limit, where we find that the cross term will generically be positive, leading to larger MSE (Appendix H.2).

G.2 Perturbation Theory in Rates rather than Predictions

In experiments on deep CNNs trained on CIFAR-10 in 7 and A.4, we find that the loss curves for the ensemble averaged predictors are effectively time rescaled by a function of network width. In this section, we argue that a proper way to account for this is to compute a perturbation expansion in the exponent which defines the rate of decay of the training errors. To illustrate the point, we first consider the case of a single training example before describing larger datasets. In this case, we consider the change of variables Δ(t)=e−r(t)y\Delta(t)=e^{-r(t)}y. We now treat rr as an order parameter of the theory with dynamics

Note that this equation is now a linear relation between two order parameters (r(t),K(t)r(t),K(t)), whereas the relation was previously quadratic. In the lazy limit, if K→K−ϵK\to K-\epsilon then r→r−ϵtr\to r-\epsilon t, giving an effective rescaling of training time by 1−ϵK1-\frac{\epsilon}{K}.

The solution to the training prediction errors can be obtained at any time tt by multiplying the initial condition Δ(0)=y\bm{\Delta}(0)=\bm{y} with the transition matrix Δ(t)=T(t)y\bm{\Delta}(t)=\bm{T}(t)\bm{y}, where y\bm{y} are the training targets. In this case, the relevant rate matrix, which would be an alternative order parameter is

where log⁡\log is the matrix logarithm function. Note that in general T(t)\bm{T}(t) admits a Peano-Baker series solution . In the special case where K(t)\bm{K}(t) commutes with Kˉ(t)=1t∫0tdsK(s)\bar{\bm{K}}(t)=\frac{1}{t}\int_{0}^{t}ds\bm{K}(s), we obtain the following simplified formula for the rate matrix R\bm{R}

The benefit of this representation is the elimination of coupled order parameter dynamics which are quadratic in fluctuations (in Δ\bm{\Delta} and K\bm{K}) into a linear dynamical relation between order parameters R\bm{R} and K\bm{K}. An expansion in R\bm{R} will thus give better predictions at long times tt than a direct expansion in Δ\bm{\Delta}. In the lazy γ→0\gamma\to 0 limit, the constancy of K(t)=K\bm{K}(t)=\bm{K} gives the further simplification R=Kt\bm{R}=\bm{K}t. Working with this representation, we have the following finite width expression for the training loss

where <R>∼R∞+1NR1+O(N−2)\left<\bm{R}\right>\sim\bm{R}_{\infty}+\frac{1}{N}\bm{R}^{1}+\mathcal{O}(N^{-2}) is the leading correction to the mean R\bm{R}. In this representation, it is clear that finite width can alter the timescale of the dynamics through a correction to the mean of R\bm{R}, as well as contribute an additive correction from fluctuations. This justifies the study perturbation analysis of rates RNR_{N} as a function of 1/N1/N in Figures 7 and A.4.

Appendix H Variance in the Lazy Limit

We can simplify the propagator equations in the lazy γ→0\gamma\to 0 limit. To demonstrate how to use our formalism, we go through the complete process of inverting the Hessian, however, for this case, this procedure is a bit cumbersome. A simplified derivation for the lazy limit can be found below in section H.1 which relies only on linearizing the dynamics around the infinite width solution. In the γ→0\gamma\to 0 limit, all of the DD tensors vanish and the κ\kappa tensors are constant in time. Thus, it suffices to analyze the kernels restricted to t=0t=0 and study the evolution of the prediction variance Δ(t)\bm{\Delta}(t).

Given these we also have the relevant non-vanishing sensitivity tensors

The propagator of interest is Σq1=U−1[∇q2q22S]U−1⊤\bm{\Sigma}_{\bm{q}_{1}}=\bm{U}^{-1}\left[\nabla^{2}_{\bm{q}_{2}\bm{q}_{2}}S\right]\bm{U}^{-1\top}. We can exploit the block structure of U\bm{U} to find an inverse

where each sub-block can be computed with the Schur-complement formula. Altogether, we multiply through to get the propagator

Two of these blocks corresponding to K,ΔK,\Delta are especially important for characterizing the fluctuations of network predictions. The covariance structure for KK has the form

Next we use the fact that UΔΦ−1=UΔK−1UKΦ−1\bm{U}^{-1}_{\Delta\Phi}=\bm{U}^{-1}_{\Delta K}\bm{U}^{-1}_{K\Phi} and that UΔG−1=UΔK−1UKG−1\bm{U}^{-1}_{\Delta G}=\bm{U}^{-1}_{\Delta K}\bm{U}^{-1}_{KG}, which follows from the block structure of U\bm{U}. Consequently we arrive at the identity

Lastly, we note that, by the Schur-complement formula that UΔK−1=−(I+ΘK)−1DΔK\bm{U}^{-1}_{\Delta K}=-\left(\mathbf{I}+\bm{\Theta}_{K}\right)^{-1}\bm{D}^{\Delta K}. Thus, writing (I+ΘK)ΣΔ(I+ΘK)⊤=DΔKΣK[DΔK]⊤\left(\mathbf{I}+\bm{\Theta}_{K}\right)\bm{\Sigma}_{\Delta}\left(\mathbf{I}+\bm{\Theta}_{K}\right)^{\top}=\bm{D}^{\Delta K}\bm{\Sigma}_{K}[\bm{D}^{\Delta K}]^{\top} as an integral equation, we find

Differentiation with respect to tt and ss gives a simple differential equation

Replacing ΣK=κ\Sigma^{K}=\kappa recovers the equation (7) in the main text.

In this section, we provide a simpler derivation of the lazy limit training error variance dynamics. In this case, we merely perturb the dynamics around its infinite width value Δ(t)=Δ∞(t)+ϵΔ(t)\bm{\Delta}(t)=\bm{\Delta}_{\infty}(t)+\bm{\epsilon}^{\Delta}(t) and K=K∞+ϵK\bm{K}=\bm{K}_{\infty}+\bm{\epsilon}^{K}, and keep terms only linear in these perturbations. The perturbation ϵK\bm{\epsilon}^{K} is fixed in time and the dynamics of ϵΔ(t)\bm{\epsilon}^{\Delta}(t) are

Projecting this equation on the eigenspace of K∞\bm{K}_{\infty} gives

This immediately recovers the final result of the last section

Qualitatively, the process of computing this linear correction (in ϵK\epsilon^{K}) to the dynamics of Δ\bm{\Delta} is identical to the argument utilized in prior work on perturbative feature learning corrections . In that context, the perturbation is caused by small amounts of feature learning, rather than initialization fluctuations.

H.2 Mean Prediction Error Correction in the Lazy Limit

Using a similar heuristic as in the preceeding section, we now consider the correction to the mean predictor <Δμ(t)>\left<\Delta_{\mu}(t)\right> in the lazy limit. Taylor expanding <Δ(t)>\left<\bm{\Delta}(t)\right> in powers of 1/N1/N, we find

Projecting these dynamics onto the eigenspace of the kernel gives

We see that at late sufficiently large tt, that the terms involving ΣK\Sigma^{K} will dominate. We can gain more intuition by considering the special case of a single training data point where the mean error correction has the form

While the term involving ΣK\Sigma^{K} is positive for all tt, K1K^{1} could be positive or negative for a given architecture. If K1K^{1} is positive, then MSE is initially improved at early times but after t>K1ΣKt>\frac{K^{1}}{\Sigma^{K}} the MSE is worse than the infinite width. On the other hand, if K1K^{1} is negative (as we suspect is typically the case), then the MSE will strictly decrease with network width for any time tt.

Appendix I Two Layer Equations and Time/Time Diagonal

For a two layer network trained on a single training point with norm constraint ∣x∣2=D|\bm{x}|^{2}=D, we have the following DMFT action

From these equations, we can compute the entries in the Hessian of the DMFT action SS. Letting q(t)=[Δ(t)K(t)]\bm{q}(t)=\begin{bmatrix}\Delta(t)\\ K(t)\end{bmatrix} and q^(t)=[Δ^(t)K^(t)]\hat{\bm{q}}(t)=\begin{bmatrix}\hat{\Delta}(t)\\ \hat{K}(t)\end{bmatrix}

The covariance matrix of interest (for q(t)\bm{q}(t)) is thus

where [ΘK](t,s)=Θ(t−s)K(s)[\bm{\Theta}_{K}](t,s)=\Theta(t-s)K(s) and [ΘΔ](t,s)=Θ(t−s)Δ(s)[\bm{\Theta}_{\Delta}](t,s)=\Theta(t-s)\Delta(s). The above equations allow one to use the infinite width DMFT dynamics for K(t),Δ(t)K(t),\Delta(t) to compute the finite size fluctuation dynamics of the kernel KK and the error signal Δ\Delta.

In this section, we compute D(t,s)D(t,s) by solving for the sensitivity of order parameters. We start with the DMFT field equations

Now, differentiating both sides with respect to Δ(s′)\Delta(s^{\prime}) gives

We can compute DD Monte carlo by iteratively solving the above equations for each sampled trajectory {h(t),z(t)}\{h(t),z(t)\} . Averaging the necessary fields over the Monte Carlo samples will give us the final expressions for D(t,s)D(t,s).

Similarly, the uncoupled kernel variance κ(t,s)\kappa(t,s) can be evaluated via Monte Carlo sampling for nonlinear networks.

I.2 Test Point Fluctuation Dynamics

We now are in a position to calculate the test/train kernel and test prediction fluctuations. To do this systematically, we augment SS with the test point prediction f⋆f_{\star} and field h⋆h_{\star} and introduce the kernel K⋆(t)=<ϕ(h(t))ϕ(h⋆(t))+g(t)g⋆(t)>K_{\star}(t)=\left<\phi(h(t))\phi(h_{\star}(t))+g(t)g_{\star}(t)\right>. The test prediction f⋆f_{\star} and field h⋆h_{\star} have dynamics

The augmented action for this DMFT has the form

We let q(t)=[Δ(t),f⋆(t),K(t),K⋆(t)]⊤\bm{q}(t)=[\Delta(t),f_{\star}(t),K(t),K_{\star}(t)]^{\top}

Our total covariance matrix / propagator is thus

This is the equation provided in the main text Equation (8).

I.3 Two Layer Linear Network Closed Form

For a linear network on a single data point, we can compute D(t,s)D(t,s) and κ(t,s)\kappa(t,s) analytically. We start from the field equations

We can make a change of variables v+(t)=12(h(t)+z(t))v_{+}(t)=\frac{1}{\sqrt{2}}(h(t)+z(t)) and v−(t)=12(h(t)−z(t))v_{-}(t)=\frac{1}{\sqrt{2}}(h(t)-z(t)). We note that v+(0)=12(u+r)v_{+}(0)=\frac{1}{\sqrt{2}}(u+r) and v−(0)=12(u−r)v_{-}(0)=\frac{1}{\sqrt{2}}(u-r) are independent Gaussians. These functions v+(t),v−(t)v_{+}(t),v_{-}(t) satisfy dynamics

Now, we use the fact that v+(0)=12(u+r)v_{+}(0)=\frac{1}{\sqrt{2}}(u+r) and v−(0)=12(u−r)v_{-}(0)=\frac{1}{\sqrt{2}}(u-r) are independent standard normal random variables to compute K(t)=<h(t)2+z(t)2>=<v+(t)2+v−(t)2>K(t)=\left<h(t)^{2}+z(t)^{2}\right>=\left<v_{+}(t)^{2}+v_{-}(t)^{2}\right>

This operator is causal (D(t,s)=0D(t,s)=0 for s>ts>t) as expected and vanishes as t→0t\to 0. If we take γ→0\gamma\to 0, we have D(t,s)→0D(t,s)\to 0 which agrees with our reasoning that fields h,zh,z only depend on Δ\Delta in the feature learning regime. Since all fields are Gaussian in the linear network case, we can use Wick’s theorem to obtain the exact uncoupled kernel variance in the two layer case.

The v±(t)v_{\pm}(t) functions are those given above. Using the fact that <v+(0)2>=<v−(0)2>=1\left<v_{+}(0)^{2}\right>=\left<v_{-}(0)^{2}\right>=1 allows us to easily compute the single site average above.

Appendix J Multiple Samples with Whitened Data

In this section, we analyze the role that sample number plays in dynamics in a simplified model of a two layer linear network trained on whitened data. Concretely, we assume that xμ⋅xνD=δμν\frac{\bm{x}_{\mu}\cdot\bm{x}_{\nu}}{D}=\delta_{\mu\nu}. The field equations for preactivations hμ(t)h_{\mu}(t) and pregradients z(t)z(t) obey

We will assume the targets have unit norm ∣y∣2=1|\bm{y}|^{2}=1 and we define the projection of Δ\bm{\Delta} onto the target as Δy(t)=y⋅Δ(t)\Delta_{y}(t)=\bm{y}\cdot\bm{\Delta}(t). The other P−1P-1 orthogonal components are denoted Δ⊥(t)\bm{\Delta}_{\perp}(t) so that Δ=Δy(t)y+Δ⊥(t)\bm{\Delta}=\Delta_{y}(t)\bm{y}+\bm{\Delta}_{\perp}(t) with Δ⊥(t)⋅y=0\bm{\Delta}_{\perp}(t)\cdot\bm{y}=0. At infinite width, Δ⊥=0\bm{\Delta}_{\perp}=0 and our field equations become

However, at finite width NN, the off-target predictions Δ⊥\bm{\Delta}_{\perp} fluctuate over random initialization. To model all of the fluctuations simultaneously, we consider the following action

which enforces the constraint that Δμ(t)=yμ−1γ<z(t)hμ(t)>\Delta_{\mu}(t)=y_{\mu}-\frac{1}{\gamma}\left<z(t)h_{\mu}(t)\right> at infinite width. The Hessian over order parameters q=Vec{Δμ(t),Δ^μ(t)}\bm{q}=\text{Vec}\{\Delta_{\mu}(t),\hat{\Delta}_{\mu}(t)\} has the form

We thus get the following covariance for predictions ΣΔ=(γI+D)−1κ[(γI+D)−1]⊤\bm{\Sigma}_{\Delta}=(\gamma\mathbf{I}+\bm{D})^{-1}\bm{\kappa}\left[(\gamma\mathbf{I}+\bm{D})^{-1}\right]^{\top}. We now compute the necessary components of the DD tensor

In the last line, we used the fact that these equations are to be evaluated at the mean field infinite width stochastic process where Δ⊥(t)=0\Delta_{\perp}(t)=0. To compute the sensitivity tensor DD, we find the following equations for our correlators of interest:

We therefore see that the components of DD decouple over indices. In the y\bm{y} direction, we have the following equations

where the correlators must be solved self-consistently. We will provide this solution in one moment, but first, we will look at the orthogonal directions. For the P−1P-1 orthogonal directions, we obtain the explicit formula for DD in each of these directions

Now, we return to DyD_{y}. To solve these equations we utilize the change of variables employed in the single sample case v+(t)=12(hy(t)+z(t)),v−(t)=12(hy(t)−z(t))v_{+}(t)=\frac{1}{\sqrt{2}}(h_{y}(t)+z(t)),v_{-}(t)=\frac{1}{\sqrt{2}}(h_{y}(t)-z(t)) (see Appendix I.3). This orthogonal transformation decouples the dynamics

As a consequence, the field derivatives close

Similarly, we can derive the on-target and off-target uncoupled variances κy(t,s)\kappa_{y}(t,s) and κ⊥(t,s)\kappa_{\perp}(t,s), which satisfy

Using these functions, we arrive at the following variance for each of the PP dimensions

Using the fact that all Δ⊥\Delta_{\perp} variables are independent and identically distributed under the leading order picture, the expected training loss has the form

where <Δy−Δy∞>=1NΔy1(t)+O(N−2)\left<\Delta_{y}-\Delta^{\infty}_{y}\right>=\frac{1}{N}\Delta^{1}_{y}(t)+\mathcal{O}(N^{-2}). We note that the bias correction if O(N−1)\mathcal{O}(N^{-1}) while the variance is O(P/N)\mathcal{O}(P/N). We compare the above leading order theory with and without the bias correction in Appendix Figure A.2.

Appendix K Online Learning

Our technology for computing finite size effects can easily be translated to a setting where the neural network is trained in an online fashion, disregarding the effect of SGD noise. At each step, we compute the gradient over the full data distribution p(x)p(\bm{x}). Focusing on MSE loss, we study the following equation

where K(x,x′;t)K(\bm{x},\bm{x}^{\prime};t) is the dynamic NTK and Δ(x,t)=y(x)−f(x,t)\Delta(\bm{x},t)=y(\bm{x})-f(\bm{x},t) is the prediction error. In general the distribution involves integration over an uncountable set of possible inputs x\bm{x}. To remedy this, we utilize a countable orthonormal basis of functions for the data distribution {ψk(x)}k=1∞\{\psi_{k}(\bm{x})\}_{k=1}^{\infty}. For example, if p(x)p(\bm{x}) were the isotropic Gaussian density for N(0,I)\mathcal{N}(0,\mathbf{I}), then ψk\psi_{k} could be Hermite polynomials. We expand Δ\Delta and KK in this basis ψk\psi_{k}, and arrive at the following differential equation

The Hessian over q={Δμ(t),Δ^μ(t)}\bm{q}=\{\Delta_{\mu}(t),\hat{\Delta}_{\mu}(t)\} is

where DΔ(t,x;s,x′)=<∂∂Δ(s,x′)a(t)ϕ(w(t)⋅x)>D_{\Delta}(t,\bm{x};s,\bm{x}^{\prime})=\left<\frac{\partial}{\partial\Delta(s,\bm{x}^{\prime})}a(t)\phi(\bm{w}(t)\cdot\bm{x})\right> We can use the following implicit rule

The above equations could be solved and then used to compute DΔ(t,x;s,x′)D_{\Delta}(t,\bm{x};s,\bm{x}^{\prime}) which must then be inverted to get the observed prediction variance.

K.2 Linear Activations

At infinite width, we see that the dynamics can be reduced to tracking the projection of the weights w\bm{w} and β\bm{\beta} on the β⋆\bm{\beta}_{\star} direction. The D−1D-1 off-target dimensions vanish β⊥(t)=0\bm{\beta}_{\perp}(t)=0. At infinite width, we arrive at the alignment dynamics studied in prior work

We note that β(t)=β(t)β⋆\bm{\beta}(t)=\beta(t)\bm{\beta}_{\star} and that M\bm{M} has only one special eigenvector β⋆\bm{\beta}_{\star} with eigenvalue m⋆(t)m_{\star}(t). It thus suffices to track evolution in this single direction

We note that this equation is identical to the differential equation for a single training example in Appendix J. Here β⋆−β(t)\beta_{\star}-\beta(t) plays the role of Δy(t)\Delta_{y}(t) and m⋆(t)m_{\star}(t) plays the role of the kernel Ky(t)K_{y}(t). A key observation is the conservation law 4γ2ddtβ(t)2=ddtm⋆(t)24\gamma^{2}\frac{d}{dt}\beta(t)^{2}=\frac{d}{dt}m_{\star}(t)^{2}, from which it follows that m⋆(t)2−4=4γ2β(t)m_{\star}(t)^{2}-4=4\gamma^{2}\beta(t)

This is identical to the differential equations for a single sample (producing prediction f(t)f(t) and kernel K(t)K(t)) if the following substitutions are made

We now proceed to compute finite size corrections starting from the action

Similarly we have to compute the sensitivity tensor

Next, we have to calculate causal derivatives for fields

Following an identical argument as in J, we see that D\bm{D} has block diagonal structure with Dβ⋆(t,s)D_{\beta_{\star}}(t,s) on the β⋆β⋆⊤\bm{\beta}_{\star}\bm{\beta}_{\star}^{\top} direction and D⊥(t,s)D_{\perp}(t,s) in any of the D−1D-1 remaining directions

Similarly, κ(t,s)\bm{\kappa}(t,s) has a similar decomposition

The processes have the following equations at infinite width

As a consequence we note that <w⊥(t)a(s)>=0\left<w_{\perp}(t)a(s)\right>=0 so that κ⊥(t,s)=<a(t)a(s)>\kappa_{\perp}(t,s)=\left<a(t)a(s)\right>. Letting v+(t)=12(wβ⋆(t)+a(t))v_{+}(t)=\frac{1}{\sqrt{2}}(w_{\beta_{\star}}(t)+a(t)) and v−(t)=12(wβ⋆(t)+a(t))v_{-}(t)=\frac{1}{\sqrt{2}}(w_{\beta_{\star}}(t)+a(t)), we find the same decoupled stochastic processes as in Appendix I.3.

We can use these equations to perform the necessary averages for κβ⋆\kappa_{\beta_{\star}} and Dβ⋆D_{\beta_{\star}}. Lastly, we use

to evaluate D⊥(t,s)D_{\perp}(t,s). The observed covariances are just

We note that these expressions are identical to those in Appendix J under the substitution β⋆−β(t)→Δ(t)\beta_{\star}-\beta(t)\to\Delta(t) and D→PD\to P. Thus the expected test risk is

This recovers the variance we obtained in the multiple-sample whitened data case J.

K.3 Connections to Offline Learning in Linear Model

As in the offline case, in Fig. 4 (c) and (d) we see that the variance contribution to test loss ∣β−β⋆∣2|\bm{\beta}-\bm{\beta}_{\star}|^{2} increases with input dimension DD. We note that this perturbative effect to the loss dynamics is reminiscent of the deviations from mean field behavior studied in SGD , though this present work concerns fluctuations driven by initialization variance rather than stochastic sampling of data. In Fig. 4 (e) we show that richer networks have lower variance at fixed NN. Similarly, leading order theory for richer networks more accurately captures their dynamics as D/ND/N increases (Fig. 4 (f)).

Appendix L Deep Linear Networks

We can then automatically differentiate the DMFT action to get the propagator. For example, for a three layer linear network, the full DMFT action has the form

where C1=γΘΔ\bm{C}^{1}=\gamma\bm{\Theta}_{\Delta} and C2=γΘΔ⊙H1+γA\bm{C}^{2}=\gamma\bm{\Theta}_{\Delta}\odot\bm{H}^{1}+\gamma\bm{A} and D1=γΘΔ⊙G2+γB\bm{D}^{1}=\gamma\bm{\Theta}_{\Delta}\odot\bm{G}^{2}+\gamma\bm{B} and D2=γΘΔ\bm{D}^{2}=\gamma\bm{\Theta}_{\Delta}. This above example can be extended to deeper networks. The total size of the block matrices which we compute determinants over is 4PT×4PT4PT\times 4PT for a dataset of size PP trained for TT steps.

Appendix M Discrete Time Dynamics and Edge of Stability Effects

Large step size effects can induce qualitatively different dynamics in neural network training. For instance, if the step size exceeds that required for linear stability with the initial kernel, the kernel can decrease in order to stabilize the dynamics . Alternatively, during training the kernel may exhibit a “progressive sharpening" phase where its top eigenvalue grows before reaching a stability bound set by the learning rate . It is therefore well motivated to study how dynamics in this regime alter finite size effects in neural networks. We will first solve a special model which was considered in prior work : a two layer linear network trained on a single training point. We will then provide the full DMFT equations for the discrete time case and provide an outline for how one could obtain finite size effects in that picture.

In a two layer linear network, the DMFT equations are

The NTK has the form K(t)=<h(t)2+z(t)2>K(t)=\left<h(t)^{2}+z(t)^{2}\right>. We can easily show that the kernel and error have coupled dynamics

These equations define the infinite width evolution of Δ(t)\Delta(t) and K(t)K(t). Already at this level of analysis, we can reason about the evolution of K(t)K(t). In the small η\eta limit, we could disregard terms of order O(η2)\mathcal{O}(\eta^{2}) and arrive at the following gradient flow approximation for K(t)∼21+γ2f(t)2K(t)\sim 2\sqrt{1+\gamma^{2}f(t)^{2}} . This evolution will not reach the edge of stability provided that η<11+γ2y2\eta<\frac{1}{\sqrt{1+\gamma^{2}y^{2}}}. For large γ\gamma and y=1y=1, this leads to the constraint ηγ<1\eta\gamma<1. However, if η\eta exceeds this bound, the gradient flow approximation is no longer reasonable and the system reaches an edge of stability effect as shown in Figure 6.

To calculate the finite size effects, we need to compute κ\kappa and D(t,s)=∂∂Δ(s)<h(t)2+z(t)2>D(t,s)=\frac{\partial}{\partial\Delta(s)}\left<h(t)^{2}+z(t)^{2}\right>. To evaluate these quantities we utilize the same change of variables employed in Appendix I.3. In discrete time, these decoupled equations are

Given Δ(t)\Delta(t), these can be expressed as linear systems of equations. Now, we can easily compute the uncoupled kernel variance

Similarly, we can calculate D(t,s)D(t,s) by using the fact <h(t)2+z(t)2>=<v+(t)2+v−(t)2>\left<h(t)^{2}+z(t)^{2}\right>=\left<v_{+}(t)^{2}+v_{-}(t)^{2}\right>

These can be directly solved as a linear system of equations.

Appendix N Computing Details

Experiments for Figures 3, 6 and 2 were conducted on a Google Colab GPU with JAX. Experiments for Figures 5, A.3, 7 were performed on a NVIDIA SMX4-A100-80GB GPU. The total compute required for all Figures in the paper took around 4 hours. Jupyter Notebooks to reproduce plots can be found at https://github.com/Pehlevan-Group/dmft_fluctuations.