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 -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 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 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 fluctuate in a neighborhood of . 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 average of observable across initializations, which we denote by , admits an expansion of the form whose leading terms are
where denotes an average over the Gaussian distribution and the function contains cubic and higher terms in the Taylor expansion of around . The terms shown include all the leading and sub-leading terms in the series in powers of . The terms in ellipses are at least 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): . This object can be shown to capture leading order fluctuation statistics (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 . To calculate such corrections, in App. E, we provide a complete description of Hessian and its inverse (the propagator) for a depth- 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 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 and 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 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 .
For MSE loss, the prediction error covariance satisfies a differential equation (App. H)
where are the errors at infinite width for eigenmode .
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 . To gain insight, we first study a simplified setting where the data distribution is a single training example and single test point in a two layer network. We will track and the test prediction . To identify the dynamics of these predictions we need the NTK on the train point, as well as the train-test NTK . 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 , are Heaviside step functions and and quantify sensitivity of the kernel to perturbations in the error signal . Lastly and are the uncoupled variances of and and is the uncoupled covariance of .
In Fig. 3, we plot the resulting theory (diagonal blocks of from Equation 8) for two layer neural networks. As predicted by theory, all average squared deviations from the infinite width DMFT scale as . Similarly, the average kernels and test predictions change by a larger amount for larger (equation (I.1)). The experimental variances also match the theory quite accurately. The variance of the train error 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 . As the curve approaches , where 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 , which increases in . 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 and the orthogonal variables inherit initialization variance, which we represent as and . In Fig. 4 (a)-(b) we show this approximate solution across varying and varying (see Appendix J for and formulas). We see that variance of train point predictions increases with the total number of points despite the signal of the target vector being fixed. In this model, the bias correction is always but the variance correction is . The fluctuations along the orthogonal directions begin to dominate the variance at large . Fig. 4 (b) shows that as 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 (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 and . We use learning rate and feature learning strength . The infinite width mean field equations for the prediction and the kernel are (App. M)
For small , the equations are well approximated by the gradient flow limit and for small corresponds to a discrete time linear model. For large , the kernel progressively sharpens (increases in scale) until it reaches 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- ReLU CNN trained on CIFAR-10 (experimental details in App. B, G.2) In Fig. 7, we train an ensemble of independently initialized CNNs of each width . 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 G.2. To test this hypothesis, we fit the ensemble training loss curve to exponential function where is a constant. We plot the fit as a function of result in Fig. 7(e). For large , we see the leading behavior is linear in , but begins to deviate at small as a quadratic function of , suggesting that second order effects become relevant around .
In App. Fig. A.4, we train a smaller subset of CIFAR-10 where we find that is well approximated by a correction, consistent with the idea that higher sample size drives the dynamics out of the leading order picture. We also analyze the effect of on variance in this task. In App. Fig. A.5, we train models with varying . Increased 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 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 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 by examining cubic and higher derivatives of 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/P, the resulting expressions are not realistically computable in deep networks trained on large dataset size for long times since the number of Hessian entries scales as 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 . Each element in the ensemble of networks is trained on identical batches presented in identical order. For the Figure 7 experiments, the raw learning rate is scaled as with (note that mean field theory requires scaling the raw learning rate linearly with since the raw NTK is ). For Figure A.5, the learning rate is . We find that choosing gives approximately conserved training times across (though distinct representation dynamics). The Figure A.4 shows the dynamics of fitting training points with full batch gradient descent and .
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 networks. From the action 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 near zero-source ()
To average over the initial weights, we introduce a Fourier representation of the Dirac-Delta function . 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 we again introduce Dirac-delta functions
Analogous constraints for and are enforced with conjugate variables . After introducing these variables, we find that the moment generating functional has the form
where is the 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 ) of any observable average that depends on DMFT order parameters . At any width the observable average takes the form
As discussed in the main text, the limit gives where by a steepest descent argument . We assume that ’s Hessian is negative semidefinite so that and Taylor expand around the saddle point giving . We note that the remainder function contains only cubic and higher powers of . The variable will be order . This will allow us to verify that additional terms are suppressed in powers of . Expanding both the numerator and denominator’s integrands in powers of , we find
where represents an average over the Gaussian fluctuation . We see that the series in the denominator contains terms of the form while the numerator depends on terms of the form . In either of these power series, the -th term can contribute at most
since contributes only cubic and higher terms. Thus each term in the numerator and denominator’s series contains increasing powers of . Concretely, each of the two series have terms of order . Thus any quantity of the form admits a ratio of power series in powers of . One could truncate each of the series in the numerator and denominator to a desired order in . 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 which can always be obtained with the truncation after the terms linear in for any observable .
We will now analyze the fluctuation statistics of our order parameters around the saddle point 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 can be no larger than comes from vanishing of odd moments for in the unperturbed distribution. Thus the leading expression for only depends on and not on .
D.2 Mean Deviation from DMFT
Although the square displacement from DMFT only depended on and not on , we note that the average order parameter displacement does receive a correction that depends on the perturbed potential
where in the last line we used Stein’s lemma (Gaussian integration by parts) for the Gaussian distribution over . Note that since the derivative of the cubic term in gives a quadratic function of , whose average must be . 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: , since
due to the arguments above which showed that and that . Therefore, in the leading order picture, it is safe to associate 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 where
This choice will become apparent shortly.
To calculate the full propagator , we will assume invertibility of the upper block and use this in the Schur complement
We seek a physically sensible inverse where the variance of is vanishing . This leads to the following sub-propagator
Thus given , we can solve for and ultimately for the full propagator . The relevant entries in and are given by those second derivatives calculated above. We note that each of the field derivatives needed for can be computed implicitly from the field dynamics. For example, for the derivatives we have
Appendix F Solving for the Propagator
In this section we sketch out the required steps to obtain the propagator .
Step 3: After populating the entries of the block matrix for the Hesssian , we then calculate the propagator 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 ). We will now give a detailed set of instructions about how the infinite width limit for is solved (step 1 above). This corresponds to the algorithm of Bordelon & Pehlevan 2022 to solve the saddle point equations .
Step 3: For each sample, solve integral equations for and .
These will be samples from the single site distribution for
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 at width . Letting the indices enumerate all entries of the order parameters in (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 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 to evaluate the four point averages in terms of the propagator , which was provided in the preceding section E. In practice computing even the full set of second derivatives for the DMFT action to get 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 , using the formalism in the preceeding section, we can compute the correction to the average network prediction error due to finite size. We let represent the average of errors over an ensemble of width networks.
where is the leading covariance (propagator element) between the kernel and prediction error . We see that the average kernel (which depends on the finite width ) plays an important role in characterizing the timescales of the average prediction dynamics. Once this equation is solved for , the square loss at width and time has the form
We will now comment on the structure of the cross term in this above solution. First, if and is negligible then the average errors at finite width will decay more rapidly than the infinite width model. However, we suspect that in general, 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 will generically increase the cross term since the dynamics of cause its fluctuations to become anti-correlated with the fluctuations in . 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 . We now treat as an order parameter of the theory with dynamics
Note that this equation is now a linear relation between two order parameters (), whereas the relation was previously quadratic. In the lazy limit, if then , giving an effective rescaling of training time by .
The solution to the training prediction errors can be obtained at any time by multiplying the initial condition with the transition matrix , where are the training targets. In this case, the relevant rate matrix, which would be an alternative order parameter is
where is the matrix logarithm function. Note that in general admits a Peano-Baker series solution . In the special case where commutes with , we obtain the following simplified formula for the rate matrix
The benefit of this representation is the elimination of coupled order parameter dynamics which are quadratic in fluctuations (in and ) into a linear dynamical relation between order parameters and . An expansion in will thus give better predictions at long times than a direct expansion in . In the lazy limit, the constancy of gives the further simplification . Working with this representation, we have the following finite width expression for the training loss
where is the leading correction to the mean . In this representation, it is clear that finite width can alter the timescale of the dynamics through a correction to the mean of , as well as contribute an additive correction from fluctuations. This justifies the study perturbation analysis of rates as a function of in Figures 7 and A.4.
Appendix H Variance in the Lazy Limit
We can simplify the propagator equations in the lazy 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 limit, all of the tensors vanish and the tensors are constant in time. Thus, it suffices to analyze the kernels restricted to and study the evolution of the prediction variance .
Given these we also have the relevant non-vanishing sensitivity tensors
The propagator of interest is . We can exploit the block structure of 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 are especially important for characterizing the fluctuations of network predictions. The covariance structure for has the form
Next we use the fact that and that , which follows from the block structure of . Consequently we arrive at the identity
Lastly, we note that, by the Schur-complement formula that . Thus, writing as an integral equation, we find
Differentiation with respect to and gives a simple differential equation
Replacing 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 and , and keep terms only linear in these perturbations. The perturbation is fixed in time and the dynamics of are
Projecting this equation on the eigenspace of gives
This immediately recovers the final result of the last section
Qualitatively, the process of computing this linear correction (in ) to the dynamics of 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 in the lazy limit. Taylor expanding in powers of , we find
Projecting these dynamics onto the eigenspace of the kernel gives
We see that at late sufficiently large , that the terms involving 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 is positive for all , could be positive or negative for a given architecture. If is positive, then MSE is initially improved at early times but after the MSE is worse than the infinite width. On the other hand, if is negative (as we suspect is typically the case), then the MSE will strictly decrease with network width for any time .
Appendix I Two Layer Equations and Time/Time Diagonal
For a two layer network trained on a single training point with norm constraint , we have the following DMFT action
From these equations, we can compute the entries in the Hessian of the DMFT action . Letting and
The covariance matrix of interest (for ) is thus
where and . The above equations allow one to use the infinite width DMFT dynamics for to compute the finite size fluctuation dynamics of the kernel and the error signal .
In this section, we compute by solving for the sensitivity of order parameters. We start with the DMFT field equations
Now, differentiating both sides with respect to gives
We can compute Monte carlo by iteratively solving the above equations for each sampled trajectory . Averaging the necessary fields over the Monte Carlo samples will give us the final expressions for .
Similarly, the uncoupled kernel variance 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 with the test point prediction and field and introduce the kernel . The test prediction and field have dynamics
The augmented action for this DMFT has the form
We let
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 and analytically. We start from the field equations
We can make a change of variables and . We note that and are independent Gaussians. These functions satisfy dynamics
Now, we use the fact that and are independent standard normal random variables to compute
This operator is causal ( for ) as expected and vanishes as . If we take , we have which agrees with our reasoning that fields only depend on 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 functions are those given above. Using the fact that 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 . The field equations for preactivations and pregradients obey
We will assume the targets have unit norm and we define the projection of onto the target as . The other orthogonal components are denoted so that with . At infinite width, and our field equations become
However, at finite width , the off-target predictions fluctuate over random initialization. To model all of the fluctuations simultaneously, we consider the following action
which enforces the constraint that at infinite width. The Hessian over order parameters has the form
We thus get the following covariance for predictions . We now compute the necessary components of the 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 . To compute the sensitivity tensor , we find the following equations for our correlators of interest:
We therefore see that the components of decouple over indices. In the 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 orthogonal directions, we obtain the explicit formula for in each of these directions
Now, we return to . To solve these equations we utilize the change of variables employed in the single sample case (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 and , which satisfy
Using these functions, we arrive at the following variance for each of the dimensions
Using the fact that all variables are independent and identically distributed under the leading order picture, the expected training loss has the form
where . We note that the bias correction if while the variance is . 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 . Focusing on MSE loss, we study the following equation
where is the dynamic NTK and is the prediction error. In general the distribution involves integration over an uncountable set of possible inputs . To remedy this, we utilize a countable orthonormal basis of functions for the data distribution . For example, if were the isotropic Gaussian density for , then could be Hermite polynomials. We expand and in this basis , and arrive at the following differential equation
The Hessian over is
where We can use the following implicit rule
The above equations could be solved and then used to compute 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 and on the direction. The off-target dimensions vanish . At infinite width, we arrive at the alignment dynamics studied in prior work
We note that and that has only one special eigenvector with eigenvalue . 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 plays the role of and plays the role of the kernel . A key observation is the conservation law , from which it follows that
This is identical to the differential equations for a single sample (producing prediction and kernel ) 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 has block diagonal structure with on the direction and in any of the remaining directions
Similarly, has a similar decomposition
The processes have the following equations at infinite width
As a consequence we note that so that . Letting and , we find the same decoupled stochastic processes as in Appendix I.3.
We can use these equations to perform the necessary averages for and . Lastly, we use
to evaluate . The observed covariances are just
We note that these expressions are identical to those in Appendix J under the substitution and . 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 increases with input dimension . 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 . Similarly, leading order theory for richer networks more accurately captures their dynamics as 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 and and and . This above example can be extended to deeper networks. The total size of the block matrices which we compute determinants over is for a dataset of size trained for 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 . We can easily show that the kernel and error have coupled dynamics
These equations define the infinite width evolution of and . Already at this level of analysis, we can reason about the evolution of . In the small limit, we could disregard terms of order and arrive at the following gradient flow approximation for . This evolution will not reach the edge of stability provided that . For large and , this leads to the constraint . However, if 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 and . To evaluate these quantities we utilize the same change of variables employed in Appendix I.3. In discrete time, these decoupled equations are
Given , these can be expressed as linear systems of equations. Now, we can easily compute the uncoupled kernel variance
Similarly, we can calculate by using the fact
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.