Separation of Scales and a Thermodynamic Description of Feature Learning in Some CNNs

Inbar Seroussi, Gadi Naveh, Zohar Ringel

Introduction

Identifying slow or relevant variables is an essential step in analyzing large-scale non-linear systems. In the context of deep neural networks (DNNs), these should be some combinations of the individual weights that are weakly fluctuating and obey a closed set of equations. One potential set of such variables is the DNNs’ outputs themselves. Indeed, in the limit of infinitely over-parameterized DNNs these provide an elegant picture of deep learning based on a mapping to Gaussian Processes (GPs). However, these GP limits miss out on several qualitative aspects, such as feature learning and the fact that real-world DNNs are not nearly as over-parameterized as required for the GP description to hold . Obtaining a useful set of slow variables for describing deep learning at finite over-parameterization is thus an important open problem in the field.

Several works provide guidelines for this search. Noting that GP limits can have surprisingly good performance and that over-parameterization is natural to deep learning we are inclined to keep some elements of the GP picture. One such element is to work in function space and study pre-activation and outputs instead of weights whose posterior distribution becomes complicated even in the GP limit . Another element is the layer-wise composition of hidden layer kernels , using which one generates the output kernel of the GP . Such a layer-wise picture is also harmonious with the idea that DNN layers should not correlate strongly, to prevent co-adaptation . Recently, it was shown that in some limited settings, making the GP kernel "dynamical" or flexible, so that it adapts to the dataset, can account for differences between infinite and finite DNNs . Still, the task of finding an explicit and general set of equations describing this flexibility in DNNs remains unsolved.

In this work, we identify such slow variables and use these to derive an effective theory for deep learning capable of capturing various finite channel/width (C/NC/N) effects in convolutional neural networks (CNNs) and fully connected neural networks (FCNs) such as feature learning. We argue that:

For, C,N≫1C,N\gg 1 the erratic behavior of specific channels/neurons averages out and hidden layers coupled to each other only through two “slow" variables per layer: The second moment of the pre-activations (pre-kernel), K(l)K^{(l)}, and the second moment of activations (post-kernel), Q(l)Q^{(l)}, of the llth layer. Furthermore, for mean square error (MSE) loss, FCNs in the so-called mean-field (MF) scaling (where the last layer weights are scaled down) or CNNs with a large read-out layer fan-in behave effectively as a GP with a data-aware kernel determined by the second moment of pre-activations in the penultimate layer.

In settings where the kernels have a large density of dominant eigenvalues, the posterior (or trained) pre-activations fluctuate in a nearly Gaussian manner. Following this we use a multivariate Gaussian variational approximation for the posterior pre-activations (pre-kernels) and derive explicit equations (equations of state), for the covariance matrices governing these pre-activations.

We identify an emergent feature learning scale (FLS) denoted by χ\chi, proportional to the train MSE times n2n^{2} over CC (or NN). This scale controls the difference between the finite C,NC,N output kernel (QfQ_{f}) and its C,N→∞C,N\rightarrow\infty limit and in this sense reflects feature learning. Due to the n2n^{2} factor, χ\chi can be O(1)O(1) or larger even for C≫1C\gg 1, e.g. for CNN architectures (see Fig. 1 panel c). The same holds, with CC replaced by NN, for FCNs in the MF scaling . Unlike perturbation theory , our theory tracks all orders of χ\chi and treats only 1/C,1/N1/C,1/N perturbatively. The separation of scales between χ\chi and 1/C,1/N1/C,1/N is thus central to our analysis. Its manifestation is the fact that feature learning shifts and stretches the dynamical variables in the theory (the pre-activations) in a considerable manner yet barely spoils their Gaussianity.

The predictions of our approach are tested on several toy and real-world examples using direct analytical approaches and numerical solutions to the equations of state. Our analysis takes a physics viewpoint on this complex non-linear problem. Rigorous mathematical proofs are left as an open problem for future research.

We note that there are several works showing evidence that the spectrum of the empirical weight correlation matrix show various tail effects and spikes . While in deeper layers we focus on pre-activations, the spectrum of input layer weights we obtained, is Gaussian but not independent as in Ref. . Hence it can produce a variety of spectral distributions for the covariance matrix, similar to the aforementioned ones. We note a recent interesting work arguing that the test-loss depends only on the mean and variance of hidden activations. There, however, the setting is of a fixed trained DNN and the statistics are over the input measure rather than over the DNN parameters as in our case. While quantitatively different, our approach is similar in spirit to the phenomenological layer-wise Gaussian Processes put forward in a recent work . Additional approaches for finite-width include perturbative correction around the infinite width limit to leading or higher orders . There is however mounting evidence from bounds on GP limits , numerical experiments , as well as the current work, that such perturbative expansions have slow convergence in practical regimes. In contrast our EoS are useful both numerically (see Sec. 4 and 2.2) and analytically (see Sec. 2.3) and in addition allow us to model pre-activation distributions in the wild via our pre-kernels (see Sec. 2.4).

Our theory can be applied to any finite number of convolutional, dense, or pooling layers. To illustrate its main aspects, let us focus on an LL-layer fully connected model with width NlN_{l} (l∈[1..L−1]l\in[1..L-1]),

The main object we analyze is the equilibrium distribution of the GD+Noise algorithm in function space. Adopting physics notation, this distribution can be written as p(f∣Dn)=e−S/Z(Dn)p(\bm{f}|\mathcal{D}_{n})=e^{-\mathcal{S}}/\mathcal{Z}(\mathcal{D}_{n}), where Z(Dn)=∫e−S\mathcal{Z}(\mathcal{D}_{n})=\int e^{-\mathcal{S}} is the partition function, and S\mathcal{S} is the action or negative log-posterior (see also Methods). Taking a Bayesian perspective, this probability distribution can also be viewed as a posterior distribution given measurements of yμy_{\mu} having Gaussian noise with variance σ2\sigma^{2} and a prior given by a finite-width random DNN. As shown in Supp. Mat. (1) our partition function is governed by the following action

Here, however, our focus is at large but finite width (Nl≫1N_{l}\gg 1). In this more complex regime, several corrections may appear: (i) The pre-activations’ average and covariance may deviate from those of a random DNN. (ii) Q(l)Q^{(l)}, the covariance of activations in the l−1l-1 layer, would not solely determine the covariance of pre-activations on the downstream layer ll, as upstream effects between h(l+1)\bm{h}^{(l+1)} and h(l)\bm{h}^{(l)} come into play. (iii) Inter-channel (or inter-neuron in the fully-connected case) and inter-layer correlations may appear. (iv) The fluctuations of pre-activations may deviate from that of a Gaussian. A priori, all these corrections may play similarly dominant roles, thereby making analysis cumbersome.

Results

This leaves us with corrections of types (i) and (ii). Interestingly, following these corrections to all orders leads to a tractable mean-field picture of learning. The latter is an augmentation of the standard correspondence between GPs and DNNs at infinite width (NNGP): Pre-activations in different layers or channels/neurons remain uncorrelated and Gaussian. Correlations only appear between different data-points (and latent pixels for CNNs) within the same layer and channel/neuron. We henceforth denote the covariance of pre-activations and activations at layer ll (up to normalization by the variance of the weights) by K(l)K^{(l)} and Q(l)Q^{(l)} and refer to these as pre-kernel and post-kernel, respectively. However in the NNGP viewpoint, Q(l)Q^{(l)} is simply proportional to K(l)K^{(l)} and fully determined by the upstream kernel (Q(l−1)Q^{(l-1)}) whereas here K(l)K^{(l)} and Q(l)Q^{(l)} differ and moreover depend both on the upstream and downstream kernels.

As their lack of dependence on width suggests, the first equation together with the definitions of the post-kernels Q(l),QfQ^{(l)},Q_{f} are already present in the strict GP limit (Nl→∞N_{l}\rightarrow\infty). They are, respectively, the GP inference formula and standard kernel recursive equations of random DNNs with \erf\erf activation. The remaining equations are, to the best of our knowledge, novel and follow the changes to the pre-kernels and post-kernels at finite NlN_{l}. These could be solved analytically in some simple cases (see subsection 2.3 for the case of two-layer CNN). We note that for non-anti-symmetric activation, one will also need to track the mean of each layer’s pre-activation (see Supp. Mat. (5)).

To get a qualitative impression of their role, one can consider the case where the penultimate layer (l=L−1l=L-1) is linear, in which case Qf=σL2K(L−1)Q_{f}=\sigma_{L}^{2}K^{(L-1)}. Consequently, ∂[Qf]μ′,ν′∂[K(L−1)]μ,ν=σL2δμμ′δνν′\frac{\partial[Q_{f}]_{\mu^{\prime},\nu^{\prime}}}{\partial[K^{(L-1)}]_{\mu,\nu}}=\sigma_{L}^{2}\delta_{\mu\mu^{\prime}}\delta_{\nu\nu^{\prime}}, (where δμν\delta_{\mu\nu}, with double index refers here to the Kronecker delta) and thus the second equation simplifies to

where δ=(y−fˉ)/σ2\bm{\delta}=(\bm{y}-\bar{\bm{f}})/\sigma^{2}. We note in passing that even for a non-linear penultimate layer, a similar term will arise from the expansion of QfQ_{f} in K(L−1)K^{(L-1)} to linear order. From the above form, several insights can be drawn.

First, we argue that the above equation implies that the trained DNN is more susceptible to changes along δ\bm{\delta} than the DNN at NL−1→∞N_{L-1}\rightarrow\infty. Noting how Qf−1Q_{f}^{-1} enters the action (Eq. 2), it controls the stiffness associated with fluctuations in f{\bm{f}}. Hence Qf−1Q_{f}^{-1} makes fluctuations in the direction of δ\bm{\delta} more likely than they are according to (Q(L−1))−1(Q^{(L-1)})^{-1}. Since δ\bm{\delta} measures the discrepancy in train predictions, this effect reduces the discrepancy by making the DNN more responsive in these directions than it is at NL−1→∞N_{L-1}\rightarrow\infty. The second term, proportional to [Qf+σ2In]−1\left[Q_{f}+\sigma^{2}I_{n}\right]^{-1} amounts to a negligible reduction in fluctuations along eigenvectors of QfQ_{f} corresponding to eigenvalues which are larger than σ2\sigma^{2}.

Using Eq. (5) one can also identify the aforementioned emergent feature learning scale (or FLS) namely, χ=NL−1−1δTQ(L−1)δ\chi=N_{L-1}^{-1}\bm{\delta}^{\mathsf{T}}Q^{(L-1)}\bm{\delta}. This scale represents the magnitude of the leading term when one Taylor expands QfQ_{f} in 1/NL−11/N_{L-1}. When χ=O(1)\chi=O(1) or larger there is a significant change in the eigenvalues of QfQ_{f} compared to Q(L−1)Q^{(L-1)} which indicates feature learning. On the other hand, when this quantity is small, we are closer to the GP regime (see Supp. Mat. (1.6)). To asses the scaling χ\chi, one can consider the common situation where δ\bm{\delta} has some non-negligible overlap with dominant eigenvectors of Q(L−1)Q^{(L-1)} whose eigenvalues are on the scale λ\lambda. Here we find χ≈λ⋅MSE/σ2⋅n/NL−1\chi\approx\lambda\cdot\text{MSE}/\sigma^{2}\cdot n/N_{L-1}, where MSE denotes the mean train MSE which enters here via ∣∣δ∣∣2/n||\bm{\delta}||^{2}/n. Due to its explicit nn dependency, and for λ=O(n)\lambda=O(n) at large nn — χ\chi maybe O(1)O(1) even at very large NL−1N_{L-1} and/or when the average MSE is rather small.

Figure 1 panel (c) shows the value of NL−1N_{L-1} (or CL−1C_{L-1}) at which χ=1\chi=1 (i.e. NL−1N_{L-1} or CL−1C_{L-1} at which feature learning becomes a dominant effect) as a function of nn for several DNNs we study. The scale separation, demonstrated there by the fact that χ\chi can be O(1)O(1) in regions where 1/Nl1/N_{l}’s is negligible, is central to our analytical approach.

This scale χ\chi is also the reason that naive perturbation theory in 1/Nl1/N_{l} fails at large nn , as it treats χ\chi and O(1/NL−1)O(1/N_{L-1}) on the same footing, since they both have a single negative power of NL−1N_{L-1}. In contrast, our EoS treat the FLS non-perturbatively.

Last we stress that the EoS provide us with a concrete effective GP description for the entire DNN as well as its hidden layers. A priori one would expect that the normality of pre-activations, a large C,NC,N trait, will be lost at finite C,NC,N. Yet we find that pre-activations remain Gaussian and accommodate strong feature learning effects while maintaining accurate predictions. This unexpectedly simple behavior opens various reverse engineering possibilities wherein one infers the effective kernels from experiments and uses their spectrum and eigenvectors to rationalize about the DNN (see also Fig. 5).

2 Numerical Demonstration: 3 Layer FCN

Next, we test the agreement between the above results and statements and actual trained DNNs, starting from the 3-layer FCN defined in Eq. (1) with L=3L=3. We focus here on a student-teacher setting with n=512n=512 or 10241024 training data points drawn from iid Gaussian distributions with unit variance along each input dimension. The target was generated by a randomly drawn teacher FCN of the same type only with N1=N2=1N_{1}=N_{2}=1. The student was trained using an analog scaling to the MF scaling , wherein the output layer weights are scaled down by a factor of 1/Nl1/\sqrt{N_{l}}. Whereas for the CNNs discussed below, this choice of scaling was not required, for FCNs we found it necessary for getting any appreciable feature learning at N1=N2=N≫1N_{1}=N_{2}=N\gg 1 (see Fig. 1 panel (c)).

As described in the methodology section 4, we trained 2020 FCNs using the GD+Noise algorithm until they reached equilibrium. We use these trained FCNs to calculate various average quantities under our partition function (Eq. (2)). Specifically, we focused on: (i) The normalized train-loss on the scale of σ2\sigma^{2}, namely MSE/(σ2∑μyμ2)\text{MSE}/(\sigma^{2}\sum_{\mu}y^{2}_{\mu}) (ii) The eigenvalues (λi,i∈[1..d]\lambda_{i},i\in[1..d]) of the average Σ\Sigma, where we average over neurons, training seeds, and training time (the latter within the equilibrium region). (iii) The normalized overlap (α\alpha) between the discrepancy in prediction on the training set times the target, namely α=σ−2∑μ=1n(fˉμ−yμ)yμ/(∑μyμ2)\alpha=\sigma^{-2}\sum_{\mu=1}^{n}(\bar{f}_{\mu}-y_{\mu})y_{\mu}/(\sum_{\mu}y^{2}_{\mu}). We then used a JAX-based numerical solver for the EoS and compared it with the experiment.

As the results of Fig. 1 panel (b) show, the predictions of our EoS for all these three quantities converged well as we increased NN. Furthermore, they do so in a region where they differ considerably from their associated GP limit. Indeed, as shown in the Supp. Mat. (3) the top Σ\Sigma eigenvalue came out 2-3 times larger than it is in the GP limit. The associated eigenvector corresponded to the first layer weights of the teacher (w∗\bm{w}^{*}). The rest of the eigenvalues remained at their GP limit values. Put together, this is a clear sign of strong feature learning.

Notably, however, this notion of feature learning does not involve compression. Indeed, since Σ\Sigma has the same variance as in the GP limit for directions perpendicular to w∗\bm{w}^{*}, it does not compress the input by projecting it solely on the label relevant direction (w∗\bm{w}^{*}). Instead, it exaggerates the fluctuation of student weights along, w∗\bm{w}^{*} thereby making it statistically more likely that hcμ(1)h^{(1)}_{c\mu} and hcν(1)h^{(1)}_{c\nu} with opposite sign of w∗⋅xμ\bm{w}^{*}\cdot\bm{x}_{\mu} and w∗⋅xν\bm{w}^{*}\cdot\bm{x}_{\nu} will be further apart in the space of pre-activations.

Next, we study how χ\chi behaves as a function of nn and CC (or NN) for different architectures. Fig. 1 shows the value of NN (or CC) at which χ=1\chi=1. As χ\chi contains a single inverse power of CC at ten times this value, χ\chi would be 0.10.1 and thus indicate only minor feature learning effects in our EoS. As N,CN,C diminish from this latter value, our EoS yield increasingly stronger feature learning effects. We find that both for CNNs in the standard scaling and for FCNs with MF scaling, the crossover to feature learning happens well within the validity region of our mean-field decoupling (i.e. large NN or CC). In contrast, FCN with standard scaling shows this crossover when N=O(1)N=O(1), which is outside the scope of our theory. In this aspect, we comment that there is evidence that FCNs with standard scale are inferior to those with mean-field scaling and perform similarly to GPs .

3 Analytical Solution of the EoS - Two Layer CNN

Having tested our EoS numerically, we turn to show they lend themselves, in simple settings, to a fully analytical calculation. Amongst other things, this will flesh out the non-perturbative nature of our results. To this end, we consider a simple non-linear CNN with 2 layers. Though bounds have been derived , we are not aware of any analytical predictions for the performance of finite non-linear 2-layer DNNs, let alone CNNs. It is therefore a natural first application of our approach. Specifically, we consider

The equations of state are given by (See Supp. Mat. (3))

Here we denote the discrepancy from the target by δ=(y−fˉ)/σ2\bm{\delta}=(\bm{y}-\bar{\bm{f}})/\sigma^{2}. The above equations for Σss′\Sigma_{ss^{\prime}} and δμ{\delta}_{\mu} could be solved numerically and compared with DNN training experiments. The results are shown in Fig. (2) in solid lines and match empirical values well.

To obtain fully analytical results, we proceed with several approximations for large nn. First, we approximate the spectrum of the matrix [Qf]μν[Q_{f}]_{\mu\nu} based on its continuum kernel version Qf(x,x′)Q_{f}(\bm{x},\bm{x}^{\prime}). This is closely related to the equivalent kernel approximation, which we adopt here along with its leading order correction . Similarly, we use large nn to replace the double summation ∑νμδμδν[Qf]μν\sum_{\nu\mu}{\delta}_{\mu}{\delta}_{\nu}[Q_{f}]_{\mu\nu} by two integrals over the measure from which xμ\bm{x}_{\mu} are drawn (dμd\mu). See Supp. Mat. (3.1) for further details and a discussion of the fully connected case (N=1N=1).

The latter approximation also fleshes out the importance of the FLS (χ\chi). Technically, the two summations provide for the n2n^{2} scaling and δμ\delta_{\mu} is related to the MSE via tμ=(yμ−fˉμ)/σ2t_{\mu}=(y_{\mu}-\bar{f}_{\mu})/\sigma^{2} (see also Supp. Mat. (3)). The FLS controls the deviation between the pre-kernel and post-kernel of the penultimate layer with our EoS. Hence, in particular, it implies deviations from GP where these are not equal.

Following our approximations, the equations acquire the full rotation symmetry of the data-set measure, which amount to an independent orthogonal transformation of each xi\bm{x}_{i}. Furthermore, as shown in Supp. Mat. (3), at large SS, fˉ(x)\bar{f}(\bm{x}) (the continuum function representing fμf_{\mu}) is linear given a linear target, y(x)y(\bm{x}), regardless of Σ\Sigma and hence so is δ(x){\delta}(\bm{x}). The above symmetry then implies that δ(x){\delta}(\bm{x}) is only a function of w∗⋅xi\bm{w}^{*}\cdot\bm{x}_{i} and furthermore takes the simple form δ(x)=αy(x){\delta}(\bm{x})=\alpha y(\bm{x}). The quantity α\alpha thus measures the overlap between the discrepancy in predictions (t{\bm{t}}) and the target. Following this, the EoS are reduced to a non-linear equation in a single variable α\alpha

Solving the above equation for α\alpha, one obtains l∗l_{*} and hence Σ\Sigma and also QfQ_{f} (via Eq. (7)). Using the obtained QfQ_{f} one can calculate the DNN’s predictions on the test-set. The effect of the FLS is evident in the second equation where it controls the deviations from the GP limit. Here we also recall that α2\alpha^{2} is the train MSE over σ4\sigma^{4}, thus χ2\chi_{2} as defined above contains the MSE factor mentioned in the introduction.

To test the theoretical predictions, we trained two such CNNs, with n={800,1600};S=64;N=20n=\{800,1600\};S=64;N=20 and varying channel number. Fig. 2, left panel, shows the empirical test-set values for α\alpha (dots) compared with a numerical solution of the equations of state (solid lines) and their analytical solution (dashed lines). For the latter, we obtained Σ\Sigma analytically and performed the resulting GP inference with QfQ_{f} numerically. The inset tracks the train-set results which, in this case, are fully analytical and involve no numerical GP inference. Both predictions match empirical values quite well, even in the regime where test root MSE is roughly half that of a Gaussian Process (C→∞C\rightarrow\infty). The right panel shows the input layer weights, dotted with w∗/∣w∗∣\bm{w}^{*}/|\bm{w}^{*}| and with a normalized random vector. These remain Gaussian up to minor statistical noise. Further details can be found in the methods section.

To emphasize the non-perturbative nature of our Eq. (111), let us assume for the sake of negation that they agree with first order perturbation theory in 1/C1/C (as in Refs. ). If so, we may replace α\alpha in the above expression for χ2\chi_{2} by its GP value, as it already contains one negative power of CC and hence receives no further corrections at that order. Numerics show this value αGP=0.558\alpha_{GP}=0.558 for n=1600n=1600. Plugging this in, one obtains l∗=2S[1−633.2/C]−1l_{*}=\frac{2}{S}[1-633.2/C]^{-1}. Clearly, this logic leads to a contradiction unless C≫633.2C\gg 633.2. In contrast, our theory provides highly accurate predictions for n=1600,C=320n=1600,C=320 and C=640C=640 well away from where 2S[1−633.2/C]−1\frac{2}{S}[1-633.2/C]^{-1} admits a perturbation theory in 1/C1/C. In Supp. Mat. (6.2) we report additional results on l∗l_{*} over its GP value.

4 Extensions to Deeper CNNs and Subsets of Real-World Data-sets

For truly deep CNNs and real-world datasets, obtaining fully analytical predictions for DNN performance is a challenging task even in the C→∞C\rightarrow\infty limit. Still, the EoS could be solved numerically and compared with experimental values. Furthermore, the quantities which underlie them could be examined and reasoned upon. We do so here in two richer settings, a 3-layer CNN trained with a teacher CNN and the Myrtle-5 CNN trained on a subset of CIFAR-10.

Our first setting extends that of the previous subsection by having an extra activated layer and a non-linear target function. Specifically, we consider a student CNN defined by

As the first test of our theory, we examine the fluctuations of pre-activations in the input and middle layers of the trained student CNN and check their normality. Specifically, for the input weights wc\bm{w}_{c} we obtain the histogram (over channels, equilibrium samples, and seeds) of wc⋅w∗/∣w∗∣\bm{w}_{c}\cdot\bm{w}^{*}/|\bm{w}^{*}|, where w∗\bm{w}^{*} is the teacher input weight and the histogram of wc⋅wr/∣wr∣\bm{w}_{c}\cdot\bm{w}_{r}/|\bm{w}_{r}| where wr\bm{w}_{r} is a random vector. Teacher overlap has a variance of 0.2540.254 here, whereas random overlap variance, averaged over choices of wr\bm{w}_{r}’s, was 0.0390.039 with an std of 0.00430.0043. For the hidden layer, we obtain the histogram of hc′(2)⋅h(2),∗/∣h(1),∗∣\bm{h}^{(2)}_{c^{\prime}}\cdot\bm{h}^{(2),*}/|\bm{h}^{(1),*}| where h(2),∗\bm{h}^{(2),*} are the teacher’s pre-activations as well as hc′(2)⋅hr(2)/∣hr(2)∣\bm{h}^{(2)}_{c^{\prime}}\cdot\bm{h}^{(2)}_{r}/|\bm{h}^{(2)}_{r}| where hr(2)\bm{h}^{(2)}_{r} is the pre-activation of a different randomly chosen teacher. Teacher overlap variance here was, 64.464.4 whereas average student variance was, 2.32.3 with an std of 0.120.12. Fig. 4 shows the associated histograms along with their fit to a Gaussian. The large and consistent differences in the variance of the fluctuations between teacher directions and random directions show that we are deep in the feature learning regime. Remarkably, the fluctuations remain almost perfectly Gaussian. The larger variance along teacher directions implies that by drawing DNNs from the trained DNN ensemble and diagonalizing their empirical covariance matrices, one is more likely to find dominant eigenvalues along these teacher directions.

We turn to verify the EoS and rationalize on the behavior pre-kernels. To this end, we average the empirical pre-activations, over channels and training seeds, to obtain an estimator for the pre-kernel and post-kernel of h(2)h^{(2)} (i.e. Q(2)Q^{(2)} and K(2)K^{(2)}) and that of the input weights (Σ\Sigma). We then obtain ∂DKL(K(2)∣∣Q(2))/∂Σ\partial D_{\text{KL}}(K^{(2)}||Q^{(2)})/\partial\Sigma analytically using the 3rd equation from Eqs. 4, plugging in the empirical Σ\Sigma. Finally, we compare the empirical Σ\Sigma with that obtained from the last equation from Eqs. (4). Fig. 5 left panel plots the eigenvalues of Σ\Sigma, our predictions (Σpred\Sigma_{\text{pred}}), and the post-kernel of the input layer which is simply σw−2S0IS0\sigma_{\text{w}}^{-2}S_{0}I_{S_{0}}, showing a good match between the first two. Figure (5) right panel plots the eigenvalues of Q(2)Q^{(2)} as predicted from K(2)K^{(2)}, compared with its empirical value.

Next, we trained the myrtle-5 CNN, capable of good performance and containing both pooling layers and ReLU activations, with C=256C=256 on a subset of CIFAR-10 (n=2048n=2048). As shown in Fig. 6, pre-activations show a strong deviation of trained DNNs from non-trained DNNs or DNNs at infinite channel/width, and at the same time show quite a good fit to Gaussian in most cases. This opens the possibility of reverse engineering the pre-kernels governing this trained network and using them to rationalize about the DNN, for instance by identifying their dominant eigenvectors.

The 2nd layer (as well as the input layer, (see Supp. Mat. (6.3))) show deviations from Gaussianity in the leading eigenvalue. This is expected since the kernels of these layers show quite a dilute dominant spectrum whereas VGA requires a contribution from many adjacent modes (see Supp. Mat. (1.3)). Interestingly, despite this non-Gaussianity in the leading eigenvalue of layers 1 and 2, Gaussainity is restored in the downstream layers 3 and 4. Correlations across layers and across channels within the same layer are very weak (largely on the order of 10−310^{-3}) and fully consistent with the mean-field decoupling underlying this work. Further technical details are found in Supp. Mat. (6.3).

Discussion

In this work, we presented what is, to the best of our knowledge, a novel mean-field framework for analyzing finite deep non-linear neural networks in the feature learning regime. Central to our analysis was a series of mean-field approximations, revealing that pre-activations are weakly correlated between layers and follow a Gaussian distribution within each layer with a pre-kernel K(l)K^{(l)}. Using the latter together with the post-kernel Q(l)Q^{(l)} induced by the upstream layer, explicit equations-of-state (EoS) governing the statistics of the hidden layers were given. These enabled us to derive, for the first time, analytical predictions for the performance of non-linear CNNs and deep non-linear FCNs in the feature learning regime. We further note that our EoS generalize straightforwardly to combined CNN-FCN architectures, pooling layers, and models with multiple outputs. Our theory can also be viewed from a Bayesian perspective, the GP process represented by the equation we find is a good approximation to the true posterior distribution generated by a large but finite width Bayesian neural network.

Various aspects of this work invite further study. Empirically, it would be interesting to better characterize the scope of models for which Langevin dynamics (or potentially ensemble-averaged NTK dynamics) leads to Gaussian pre-activations and overall GP-like behavior. Probing the “feature-learning-load“ of each layer, by experimentally measuring the differences between the kernels Q(l)Q^{(l)} and K(l)K^{(l)}, may also provide insights on generalization, transfer learning, and pruning, thus complementing other diagnostic tools suggested recently . For instance, transferring a layer with a small feature learning load may provide little benefit, and pruning a channel having a large overlap with a leading eigenvalue of K(l)−Q(l)K^{(l)}-Q^{(l)} may be harmful.

From the theory side, it is desirable to develop analytical techniques for solving the EoS as well as guarantees regarding the existence and uniqueness of solutions. In particular, exploring the possibility of spontaneous symmetry breaking of internal symmetries such as weight inversion. Providing a mathematical underpinning for the approximations involved here may lend itself to developing performance bounds on the Langevin algorithm and Bayesian neural networks. Similarly, one can consider using the empirical effective kernel (QfQ_{f}) as a starting point to develop GP-based bounds on performance. Last, it is interesting to explore the approach to equilibrium of the training dynamics and adapt the approximations carried here to the NTK setting .

Methods

Here we present the main ingredients of our theory, leading to the EoS we find. Further details can be found in the Supp. Mat.

First, we provide decoupling of Eq. (2) into layer-wise neuron-wise terms, wherein each of the terms depends on the upstream and downstream layers only through channel-averaged second moments of activations and pre-activations (pre-kernel and post-kernel). Further details are found in Supp. Mat. (1).

Consider the non-linear terms in the action 2 which couple the different layers. This coupling is mediated through the channel/width-averaged quantities: indeed h(1)\bm{h}^{(1)} depends on h(2)\bm{h}^{(2)} through the channel/width averaged square term in h(2)\bm{h}^{(2)}, h(2)\bm{h}^{(2)} depends on h(1)\bm{h}^{(1)} through the average of ϕ(h(1))ϕ(h(1))\phi(\bm{h}^{(1)})\phi(\bm{h}^{(1)}), and h(3)\bm{h}^{(3)} depends on h(2)\bm{h}^{(2)} through the average of ϕ(h(2))ϕ(h(2))\phi(\bm{h}^{(2)})\phi(\bm{h}^{(2)}) and so forth. For Nl≫1N_{l}\gg 1 we expect these to be weakly fluctuating and well approximated by their mean-field values. This behavior propagates till the output layer, and in particular implies that the outputs f\bm{f} fluctuate in a Gaussian manner, as previously conjectured . As for the dependency of h(L−1)\bm{h}^{(L-1)} on the f\bm{f} variables, it is not through a channel/width averaged quantity. However, we find that in various scenarios, such as FCNs with MF scaling or CNNs with large NN, the fluctuations of f\bm{f} are suppressed enabling us to replace f\bm{f} by its average (see Supp. Mat. (1.7) and (2.2)). Following this, we obtain our mean-field action,

Notably, any coupling between the different layers is only through static mean-field quantities, namely the pre-kernels and-post kernels. In addition, all neuron-neuron couplings (and similarly, channel-channel couplings for CNNs) have been removed.

1.2 Intra-Layer Decoupling

Despite the simplified inter-layer coupling and intra-layer neuron coupling, the mean-field actions are still non-quadratic for all layers but the output layer. This non-linearity couples all the h(l)\bm{h}^{(l)} variables for the same neuron (channel in the CNN case) in a way that is roughly all-to-all in the data-point index. In atomic and nuclear physics, similar circumstances are well described by self-consistent Hartree-Fock approximations . In our setting, this approximation is directly analogous to a variational Gaussian approximation (VGA). In Supp. mat. (4) we argue that in the typical case where the diagonal of K(l)K^{(l)} is much larger than the off-diagonal elements, the VGA is well controlled. Technically, we do so by showing, order by order in perturbation theory, that the diagrams accounted for by the VGA approximation dominate all other perturbation theory diagrams. In Supp. Mat. (3) we also establish this using different means for S0≫1S_{0}\gg 1 for the specific case of two-layer CNN with a single activated layer. We further comment that the VGA is exact for deep linear DNNs.

Accordingly, we now look for the Gaussian distribution, governed by a kernel K(l)K^{(l)} which is the closest to the above non-quadratic action. In models with many hidden layers, this leads to the following "inverse kernel shift" behavior for 1<l<L−21<l<L-2

For non-anti-symmetric ones, see Supp. Mat. (5).

2 Experimental Details

Hyperparameters. For the 2-layer CNN experiments, we used S=64S=64, N=20N=20, and varying channel number. The training parameters (noise and weight-decay) were tuned such that σ2=0.1\sigma^{2}=0.1 and weight variance of 2.02.0 over fan-in, for both layers at n=0n=0. The target was drawn once for all experiments using i.i.d. Gaussian centered random ai∗a^{*}_{i} and ws∗w^{*}_{s} with variances 1/N1/N and 1/S1/S respectively.

For the 3-layer CNN experiments, we took S1=50,S0=30,N=2S_{1}=50,S_{0}=30,N=2. The training parameters (noise and weight-decay) were scaled such that σ2=0.005\sigma^{2}=0.005 and weight variance of 2.02.0 over fan-in for the inputs and hidden layer with no training data (at initialization). The weight variance of the read-out layer was 1515 over the fan-in. The target was drawn again once for all experiments from a teacher CNN with C=1C=1.

For all the myrtle-5 experiments, we used n=2048,C=256n=2048,C=256 and ReLU activation. The training parameters (noise and weight-decay) were scaled such that σ2=0.005\sigma^{2}=0.005 and weight variance of 2.02.0 over fan-in for all layers with no training data (at initialization).

For all the FCN experiments, we used equal width (N1=N2N_{1}=N_{2}) and weight decay corresponding to variance, σw2=σa2=2\sigma_{\text{w}}^{2}=\sigma_{\text{a}}^{2}=2 (with no training data) in the regular scaling. For the MF scaling, we took σa2=2/N2\sigma_{\text{a}}^{2}=2/N_{2}. The target was drawn again once for all experiments from a teacher CNN with N1=N2=1N_{1}=N_{2}=1. Specifically when calculating the emergent scale, we used σa2=2/256\sigma_{\text{a}}^{2}=2/256 independent of N2N_{2}.

Equilibrium sampling. To obtain weakly correlated samples from the equilibrium distribution of the trained CNNs we used the following procedure. For the 2 and 3-layer CNNs, we used an adaptive learning rate scheduler: For the first 100100 epochs we used a learning rate lr0/10lr_{0}/10, then we crank up the learning rate to lr0lr_{0}. As of epoch 50005000, every 10001000 epoch we estimate the fluctuations of the train-loss and check for spikes - events in which the train-loss was 5-times larger than the standard deviation in the past 500500 epochs. If a spike is observed, the learning rate is reduced by a factor of 0.70.7. This continues until 5000050000 epochs pass without any events. Then the learning rate is reduced again by a factor of two and remains fixed. Samples from these final stages were treated as equilibrium samples. We further checked that (i) different initialization seeds trained with this protocol reached the same train-loss statistics. (ii) No further reduction in train-loss occurred after the final learning rate reduction. For several runs, we also verified that increasing the last reduction of learning rate by an additional factor of 22 did not have any appreciable effect on the loss. The initial lr0lr_{0} was ∼1e−4\sim 1e-4 (w.r.t. a standard mean reduction MSE loss) and the final learning rate was typically ∼1e−5\sim 1e-5. The runs terminated at epoch 300000300000.

For the myrtle-5 CNN trained on CIFAR-10, we first ran several runs for 300000300000 epochs using the above procedure and examined those that reached the lowest train-loss. We then generated a fixed scheduler based on those more successful instances, running up to 4e54e5 epochs. We again verified that further lowering the final learning rate has no appreciable effect on the training loss and that different seeds reach similar final train-loss. This ensures that we are indeed sampling from a valid equilibrium distribution.

For the 3-layer CNN and Myrtle-5, we found that auto-correlation times of pre-activations change considerably between the layers. While the read-out layer typically had an auto-correlation time of the order of 10310^{3} epochs (at the lowest learning rates) the auto-correlation times for the input layers could reach ∼106\sim 10^{6} or larger values. To overcome this issue, when analyzing pre-activations of these deeper DNNs we took an ensemble containing 9898 and 234234 different initialization seeds for the 3-layer CNN and Myrtle-5 respectively.

For the 3-layer FCN we used a fixed scheduler which starts at 1/21/2 the maximal stable learning rate and reduces the learning rate by factors of 22 at 100,1e5,1e6,3e6100,1e5,1e6,3e6 epochs and by a factor of 44 at 4e6,5e64e6,5e6 epochs (factor of 128128 in total). Equilibrium sampling was done between 6e6−7e66e6-7e6 epochs.

Numerical solution of the equations of state. For the 2-layer CNN, the equations of state were solved using Newton-Krylov method, which does not require explicit gradients. To facilitate convergence, we adopted an annealing procedure: For C∼1000C\sim 1000, we obtain the solution using a GP initial value (x0x_{0}) for Σ\Sigma. The optimization outcome was then used as x0x_{0} for the next lower value of CC. Using 12 CPU cores, this optimization took several hours. After obtaining Σss′\Sigma_{ss^{\prime}} as a function of CC, the resulting kernel [Qf]μν[Q_{f}]_{\mu\nu} was used in standard GP inference to obtain ff on the test-set. For the 3-layer FCN we used a more efficient JAX-based code to generate the kernels and kernel derivatives involved in the EoS, but otherwise followed the same procedure. Optimization took between several minutes to a few hours on one Titan-X GPU, depending on parameters.

Supplementary Information - Separation of Scales and a Thermodynamic Description of Feature Learning in Some CNNs

Derivation of the Mean-field Equations for a Fully Connected Network

In this section, we derive the equations of state for deep fully connected NNs with a finite number of layers, LL. In Sec. 6 we provide a sketch of the derivation for CNN architecture. The analysis can also be extended to other architectures, including pooling layers and skip connections.

The model is composed of a LL layer NN having L−1L-1 activated hidden layers, and one linear readout layer. Specifically, we consider

Our main object of interest is the following equilibrium distribution of the Langevin dynamics algorithm with noise strength σ2\sigma^{2} and weight decay written in function space (i.e. in terms of the DNNs outputs)

In practice, we sample from this distribution using Gradient descent (GD), at small learning rates, together with weight decay and noise on each weight derivative. The weight decay parameters are σl2/Nl−1\sigma_{l}^{2}/N_{l-1} for layer ll. The first term on the r.h.s. is given by

where f=(f1,...,fn)\bm{f}=(f_{1},...,f_{n}) is viewed now as a random variable following the NN outputs on all different training points. The average ⟨...⟩w\langle...\rangle_{\bm{w}} is over the weights w\bm{w} of the network at equilibrium. The weights’ distribution at equilibrium can be obtained explicitly. At Nl→∞N_{l}\rightarrow\infty, for l∈[1,L]l\in[1,L] these distributions (Eq. (21) and Eq. (20)) tend to a GP, however our interest here is at finite NlN_{l}.

Eq. (20) and Eq. (21) can also be understood from a Bayesian perspective. Eq. (20) can be viewed as the posterior distribution assuming each sample is generated by a neural network model as in Eq. (19) and is corrupted by an additive i.i.d. Gaussian noise with variance σ2\sigma^{2}. The prior distribution over the weights of the network is taken to be Gaussian with variance σl2/Nl−1\sigma_{l}^{2}/N_{l-1}.

To obtain a more explicitly "layer-wise" representation of Eq. 21, we next condition over the pre-activations (h(l)\bm{h}^{(l)}) of each layer using Bayes’ formula, we obtain the following Markov representation of Eq. (21)

The hidden layers probabilities above are defined as follows:

where we write for short hiμ(l)=hi(l)(xμ)h_{i\mu}^{(l)}=h_{i}^{(l)}(\bm{x}_{\mu}) for l∈{1,...,L−1}l\in\{1,...,L-1\}, the latter being the random variables describing argument of the activation function (pre-activation) of the llth layer at neuron ii, on the xμ\bm{x}_{\mu} data-point. Later we will use this Markov structure of the distribution to decouple the different layers.

We continue our analysis by using the Fourier identity, which replaces the above delta functions by auxiliary fields, t,{m(l)}l=1L−1\bm{t},\{\bm{m}^{(l)}\}^{L-1}_{l=1}. To this end, the probability distribution over the network output given the input data can be written as follows.

where we adopt here physics notation, and define, S(f,t,{m(l),h(l)}l=1L−1)\mathcal{S}\left(\bm{f},\bm{t},\{\bm{m}^{(l)},\bm{h}^{(l)}\}^{L-1}_{l=1}\right), as the action associated with this distribution. Collecting all terms, the action is defined as follows

To obtain the above expression, we performed Gaussian integration over the weights of all layers. In addition, we define the following matrices:

In the following section, we derive the inter-layer mean-field decoupling. In this context, the form of the action in the presence of the auxiliary fields (Eq. (27)) turns out to be useful.

2 Mean-Field Decoupling

Our first step is to understand the dependence of h(l)\bm{h}^{(l)} on h(l−1)\bm{h}^{(l-1)} for all l∈[2,L−1]l\in[2,L-1]. Following the mean-field idea introduced in Eq. (32) we first subtract average quantities and rewrite the action of the jjth neuron of the llth layer such that S(l)=∑j=0Nl−1Sj(l)\mathcal{S}^{(l)}=\sum^{N_{l}-1}_{j=0}\mathcal{S}_{j}^{(l)}:

Following Eq. (34) for ll and l−1l-1, we gather all the h(l−1)h^{(l-1)} dependent terms and obtain the mean-field action of the l−1l-1th layer:

2.2 Mean-Field Decoupling - Output Layer

We next consider the coupling between the final/output layer and the penultimate layer. As done previously in the analysis of the hidden layers, we rewrite the action of the top layer as

Our mean-field decoupling is therefore valid at NL−1≫1N_{L-1}\gg 1 but finite. We stress that even in this regime, feature learning effects can still be of order, 11 as these are controlled by an emergent scale (χ\chi) containing positive powers of nn. See Fig. 1(c) in the main text.

For the L−1L-1 layer, we obtain from Eq. (35) the following action:

where the mean field average of t\bm{t} is then,

We can also now identify the last layer of the mean-field action:

Combining all the layers Eq. (35), Eq. (39), and Eq. (42), we can write the mean-field action:

3 Variational Gaussian approximation

For the first layer, one is free to choose either pre-activations or the weights themselves as the variables, as these two are linear functions of one another. Here we will use the weights, Wi:(1)W^{(1)}_{i:}, (the iith row of the matrix W(1)W^{(1)}) and denote the covariance matrix of these weights as Σ\Sigma. The pre-kernel matrix of the input layer pre-activations is then given by K(1)=XnΣXnTK^{(1)}=X_{n}\Sigma X^{\mathsf{T}}_{n}.

The KL divergence between the above distribution and a Gaussian distribution for all j∈[0,Nl−1]j\in[0,N_{l}-1] with covariance K(l)K^{(l)} is then

Note that, this is indeed the optimal distribution since taking the second derivative yields −[K(l)]−2/2-[K^{(l)}]^{-2}/2 and since K(l)K^{(l)} is positive definite as a covariance matrix −[K(l)]−2/2-[K^{(l)}]^{-2}/2 is negative definite, therefore it is a global minimum.

3.2 l=1𝑙1l=1

We turn to find the pre-kernel of the input layer. Here, we find it more convenient to work with the covariance matrix of the weights rather than the pre-activation, as it is a more compact object having fewer indices. We follow the same procedure and minimize the KL divergence for all i∈[0,N1)i\in[0,N_{1}) to find the closest Gaussian distribution with covariance Σ\Sigma,

taking the derivative with respect to Σ\Sigma

Plugging in the above expression for A(2)A^{(2)} and combining with Eq. (46) we have that:

where K(1)=XnΣXnTK^{(1)}=X_{n}\Sigma X^{\mathsf{T}}_{n} and Q(L)=QfQ^{(L)}=Q_{f}.

4 Equations of State (EoS)

Collecting all the results above, we obtain the following closed set of equations determining all pre-kernel and post-kernel as well as the average output of the Langevin algorithm, fˉ\bar{\bm{f}}: namely

where NL=1N_{L}=1, Q(L)=QfQ^{(L)}=Q_{f}, and K(1)=XnΣXnTK^{(1)}=X_{n}\Sigma X^{\mathsf{T}}_{n}. While not explicitly apparent, the above expression does converge to the GP limit as N=Nl→∞N=N_{l}\rightarrow\infty for all l∈[2,L−1]l\in[2,L-1]. Indeed, for very large NN, K(L−1)=Q(L−1)+O(1/N)K^{(L-1)}=Q^{(L-1)}+O(1/N). Consequently, the term [[Q(l)]−1(K(l)−Q(l))[Q(l)]−1][[Q^{(l)}]^{-1}(K^{(l)}-Q^{(l)})[Q^{(l)}]^{-1}] is O(1/N)O(1/N). The term on the r.h.s. thus vanishes as 1/N1/N.

We note that the second term in Eq. (53) and Eq. (54) has a more profound meaning in terms of the information transfer between the pre-kernel and post-kernel. Looking at the KL divergence between two centered multidimensional Gaussian with kernel K(l)K^{(l)} and Q(l)Q^{(l)} of the same dimension m=Nnm=Nn

Taking the derivative with respect to K(l−1)K^{(l-1)} for l∈[3,L−1]l\in[3,L-1] and with respect to Σ\Sigma for l=2l=2, we then obtain that:

where in the second transition we rearranged terms and used the cyclical property of the trace. Substituting this relation leads to Eq. (4) provided in the main text.

4.2 Auxiliary field correlation

The second moment correlation then can be found by taking twice the derivative of the following free entropy and taking the source field to zero:

Plugging this expression back in Eq. (60), and using the fact that all matrices are symmetric, we have that

5 An Emergent Scale

Here we identify a quantity (χ\chi) whose scale characterizes the amount of feature learning in the DNN and its deviations from the GP limit. More specifically, when χ\chi becomes comparable to 11, strong feature learning effects appear and perturbation theory in 1/Nl1/N_{l} becomes impractical. Since χ\chi would consist of a non-trivial combination of factors, we refer to it as an emergent scale. To define this scale, we work within our EoS. and ask when QfQ_{f} changes in a noticeable manner as we lower NlN_{l} from the Nl→∞N_{l}\rightarrow\infty at fixed σl2\sigma_{l}^{2} (the GP limit). More technically, we next solve the EoS using perturbation theory in Nl−1N^{-1}_{l} and estimate the magnitude of the leading O(1/Nl)O(1/N_{l}) terms we obtain.

For simplicity, we focus on a setting where the penultimate layer is linear, due to this linearity we obtain Qf=σL2K(L−1)Q_{f}=\sigma_{L}^{2}K^{(L-1)} and so:

Following the aforementioned perturbation theory in 1/Nl1/N_{l}, we perform a first-order Taylor expansion of QfQ_{f} yielding

To evaluate the magnitude of the first term we multiply by δ\bm{\delta} from both sides to obtain

We define χ=1NL−1δTQ(L−1)δ\chi=\frac{1}{N_{L-1}}\bm{\delta}^{\mathsf{T}}Q^{(L-1)}\bm{\delta}, and find:

As we argue below, the last term on the r.h.s is small compared to the first two. Putting it aside and recalling that the first term on the right-hand side is the zeroth order term in 1/Nl1/N_{l}, we find that χ\chi controls the ratio between the zeroth order contribute and the first order perturbative correction. Hence, once χ\chi becomes order 11, first-order perturbation in 1/Nl1/N_{l} becomes inaccurate. However, it can be further checked that a second-order perturbation will contain a χ3\chi^{3} contribution, and hence this is not just a problem in first-order perturbation theory. Rather, it is that low order perturbation theory becomes inaccurate.

Next, we argue that having non-negligible χ\chi also implies feature learning, in the sense that QfQ_{f} changes from its Nl→∞N_{l}\rightarrow\infty value. Indeed, the quantity we are examining (δTQfδ\bm{\delta}^{\mathsf{T}}Q_{f}\bm{\delta}) involves both the discrepancy (δ\bm{\delta}) and the kernel QfQ_{f}. Thus, a-priory may change just because δ\bm{\delta} changes. However, δ\bm{\delta} is a function of QfQ_{f} via δ=[Qf+σ2I]−1y\bm{\delta}=[Q_{f}+\sigma^{2}I]^{-1}\bm{y}. Thus, it cannot undergo any change if QfQ_{f} remains inert. Thus, we conclude that a change to δTQfδ\bm{\delta}^{\mathsf{T}}Q_{f}\bm{\delta} must come from a change in QfQ_{f} and hence, by our definition, from feature learning. This combined with the previous paragraph also shows that strong feature learning effects are beyond the practical reach of straightforward perturbation theory.

We turn to estimate the magnitude of the term we neglected namely

Noting Qf=σL2K(L−1)Q_{f}=\sigma_{L}^{2}K^{(L-1)} and that, following our EoS., K(L−1)=Q(L−1)+O(1/Nl)K^{(L-1)}=Q^{(L-1)}+O(1/N_{l}), we can rewrite the above term up to 1/Nl21/N_{l}^{2} corrections as

which is indeed negligible compared to χ\chi at large NL−1N_{L-1}.

6 Estimating Corrections to Mean-field Results

In the mean-field derivation, when focusing on the two last layers, we neglected the term

In this section, we study the effect of this term in perturbation theory and when it can be neglected.

Specifically, we treat the above term as a perturbation over the mean-field limit and calculate its leading order effect on the mean-field average of, (f−y)/σ2(\bm{f}-\bm{y})/\sigma^{2} which coincides with the mean-field average of t\bm{t} which is directly related to the MSE. To expose this matter in its simplest form, we shall assume that the penultimate layer is linear. Consequently Qf=σL2K(L−1)Q_{f}=\sigma^{2}_{L}K^{(L-1)} and in addition, VGA becomes exact (see Eq. 33 and definition of Qf(h(L−1))Q_{f}(\bm{h}^{(L-1)})).

Turning to an action formulation, we focus on the following mean-field action of the h(L−1),f\bm{h}^{(L-1)},\bm{f} and t\bm{t} variables together with the perturbation namely,

Perturbation in the ΔSf\Delta\mathcal{S}_{f} term yields the following zeroth and first-order contributions

A key point is that making ϵ\epsilon smaller makes χ\chi bigger. Hence, our mean-field treatment works within the feature learning regime. Indeed, the emergent scale is defined as δTQ(L−1)δ/NL−1\bm{\delta}^{\mathsf{T}}Q^{(L-1)}{\bm{\delta}}/N_{L-1} and hence does not involve ϵ\epsilon directly. Decreasing ϵ\epsilon can only make δ{\bm{\delta}} larger since δ=[Qf+σ2In]−1y{\bm{\delta}}=[Q_{f}+\sigma^{2}I_{n}]^{-1}{\bm{y}}. Thus, decreasing ϵ\epsilon increases feature learning while making mean-field corrections more negligible. On this note, we comment that in our FCN experiments we found, similarly to Ref. , that mean-field scaled FCNs had better test performance compared to those which used standard scaling similar to

Mean-Field Equations for CNNs

In this section, we provide a sketch of the derivation of the equations of state for deep CNNs highlighting the differences between CNN and FCN. For simplicity, we follow here the derivation of a three-layer CNN. The generalization to any number of layers is straightforward. The model we consider is a three-layer CNN having two activated convolutional layers and one linear readout layer. Specifically, we consider

where the matrix XnX_{n} represents all the samples, {xμ}μ=1n\{\bm{x}_{\mu}\}^{n}_{\mu=1}, and the first term on the r.h.s. is given by

where f=(f1,...,fn)\bm{f}=(f_{1},...,f_{n}) is viewed now as a random variable following the CNN outputs on all different training points, and ⟨...⟩a,v,w\langle...\rangle_{\bm{a,v,w}} denote average over the weights a,v,w\bm{a,v,w}. The above probability can be viewed as the prior induced on f\bm{f} by a finite random DNN with i.i.d. Gaussian weights and variance σw2/S0,σv2/(S1C1)\sigma_{\text{w}}^{2}/S_{0},\sigma_{\text{v}}^{2}/(S_{1}C_{1}) and σa2/(NC2)\sigma_{\text{a}}^{2}/(NC_{2}) respectively for each layer. At C1,C2→∞C_{1},C_{2}\rightarrow\infty such priors tend to a GP, however our interest here is at finite C1,C2C_{1},C_{2}.

By conditioning on the pre-activation output (h(2)\bm{h}^{(2)}), Eq. (76) can be re-written as

where the hidden layers probabilities used above are

the tensor h(2)\bm{h}^{(2)} consists of all hicμ(2)h^{(2)}_{ic\mu} the latter being the random variables describing outputs of the second layer at hidden pixel ii, channel cc on the xμ\bm{x}_{\mu} data-point.

Similar to the FCN, we continue our analysis by using the Fourier identity, which replaces the delta function with auxiliary fields t,m\bm{t},\bm{m}. The resulting action (Z=∫e−SZ=\int e^{-\mathcal{S}}) is then,

Where the “channel” post-kernels for CNN contain also summation overstrides and are defined as:

We comment that by averaging over the auxiliary fields, t,m\bm{t},\bm{m}, one obtains the following equivalent form, containing only the pre-activations and the outputs

We continue with the auxiliary variables and derive the inter-layer mean-field decoupling. Where for CNN the number of channels C1,C2C_{1},C_{2} plays the role of width in FCN. As for FCN, we note that in Eq. (79), the hidden layer and the output layer depend on their respective upstream layers only through the "channel" post-kernels. Performing our mean-field decoupling as in Sec. 5.2 we obtain the resulting action

where the post-kernel can now be defined as the mean-field average of the “channel” post-kernel via the above mean-field action distribution:

The correlation functions similar to the FCN are then,

The last layer of auxiliary field correlation is

this is easily derive from Eq. (82). The average of t\bm{t} using the mean-field action is then,

Despite reducing the full system into decoupled systems per channel and layer, the resulting action for all but the top layer is still non-Gaussian. Following the justifications discussed in the main text and for the FCN, we approximate the latter using the variational Gaussian approximation (VGA). Assuming for simplicity an antisymmetric activation function such as \erf\erf, the CNN has an internal symmetry, making each pre-activation positive output as likely as its negative. Also, at large enough C1,C2C_{1},C_{2} we do not expect spontaneous symmetry breaking, thus our VGA will involve only a centered Gaussian. Specifically, we denote the optimal covariance of h(2)\bm{h}^{(2)}, Kμjνj′(2)K^{(2)}_{\mu j\nu j^{\prime}} as the pre-kernel of the second layer and the optimal covariance of w\bm{w}, Σss′\Sigma_{ss^{\prime}}, is a connected to the first layer pre-kernel, K(1)=XnΣXnTK^{(1)}=X_{n}\Sigma X_{n}^{\mathsf{T}}.

Following the analysis in subsection. 5.3 with the above modification for CNN, we obtain the following closed set of equations determining all pre-kernels as well as the outputs namely

Following the FCN case, we again examine the EoS for the penultimate layer assuming that it is linear, solve them using perturbation theory in 1/Cl1/C_{l}, and use the magnitude of the correction to estimate the scale at which feature learning becomes important. For our CNNs we have that [Qf]μν=σL2N−1∑jKμj,νj(L−1)[Q_{f}]_{\mu\nu}=\sigma_{L}^{2}N^{-1}\sum_{j}K^{(L-1)}_{\mu j,\nu j} which yields the following equation of state for K(L−1)K^{(L-1)}

Next, we perform a leading order perturbation theory in 1/Cl1/C_{l} (or equivalently in the second term on the r.h.s) yielding

As justified in the FCN case, we focus on the second term on the r.h.s. and look again at δTQfδ\bm{\delta}^{\mathsf{T}}Q_{f}\bm{\delta} yielding,

We define the ratio of the second to the first term as the emergent scale. Notably for N=1N=1 it coincides with the definition for FCNs. In addition to considering the 2-layer CNN studied in the main text and estimating it using the same approximations, it results in the same scale.

2 Estimating Mean-field corrections - CNNs

In the FCN case, we found that an additional ingredient, on top of large NlN_{l}, is needed for our mean-field decoupling to hold - either mean-field scaling or a target with support only on weak QfQ_{f} eigenvalues. For the CNNs we have studied, and quite possibly for a much larger family of CNNs, this additional ingredient comes naturally from the read-out layer averages over NN latent pixels in the penultimate layers. Since these are expected to be somewhat independent, one can hope that summing over these terms is similar to increasing the number of channels by a factor of NN (the number of pixels in the penultimate layer). Here we establish this more concretely. Similar to the FCN we calculate the average of the discrepancy tμt_{\mu} up to the second order:

For CNN second order term can be further simplified by using Wick theorem for pre-activations of the last layer when again we consider the case of linear activation function

A similar derivation to 5.6 yields the following correction to tˉμ\bar{t}_{\mu} ((fμ−yμ)/σ2(f_{\mu}-y_{\mu})/\sigma^{2})

To simplify this expression we next note that for data-sets in which for each xμ\bm{x}_{\mu} there exists a "symmetry-partner" point wherein all coordinates in the fan-in of the ii’th latent pixels are flipped - the second, trace-like, term must vanish for i≠ji\neq j. This is due to the fact that KfK_{f} is invariant under the action of the associated symmetry, whereas K∗i,∗j(L−1)K^{(L-1)}_{*i,*j} receives a minor sign whenever i≠ji\neq j. As n→∞n\rightarrow\infty we expect this symmetry to be approximately realized as it is a symmetry of the underlying measure from which xμ\bm{x}_{\mu} are drawn. Following this, we remove i≠ji\neq j terms from the summation.

Next we notice that due to the approximate translation symmetry of the dataset, at large nn

becomes independent of ii. We thus replace it by its average and perform the remaining summation over, ii which now involves only the first term to obtain

Recall that [Qf]μν=σL2N−1∑jKμj,νj(L−1)[Q_{f}]_{\mu\nu}=\sigma_{L}^{2}N^{-1}\sum_{j}K_{\mu j,\nu j}^{(L-1)}. This resulting expression is very similar to its FCN version (with standard scaling) with one crucial difference, which is the appearance of the aforementioned 1/N1/N factor. More specifically, the first summation is smaller than the zeroth term (∑μtμ2\sum_{\mu}t_{\mu}^{2}). The scale controlling the mean-field decoupling is therefore 1CL−1N\Tr[Kf−1Qf]\frac{1}{C_{L-1}N}\Tr[K_{f}^{-1}Q_{f}]. Thus, at large NN, we can have a reliable mean-field decoupling even when n=CL−1n=C_{L-1}.

Toy Example - One Hidden Layer

requiring enough data-points to resolve the target sets n>N+Sn>N+S (number of target parameters) while staying within the over-parameterized regime implies n<CNSn<CNS. We further consider the large-scale “thermodynamic” limit, where C,S,N,n≫1C,S,N,n\gg 1.

Similarly to the 3 layer case, the equations of states here are given by,

These can be viewed as non-linear equations in the S(S−1)/2S(S-1)/2 variables making up the symmetric matrix Σ\Sigma. Having these variables determines the tμt_{\mu} variables directly.

We begin with approximating the GP inference appearing in the last equation. This can be represented as δ=[y−fˉ]/σ2{\bm{\delta}}=[\bm{y}-\bar{\bm{f}}]/\sigma^{2} where, fˉ=Qf[Qf+σ2In]−1y\bar{\bm{f}}=Q_{f}[Q_{f}+\sigma^{2}I_{n}]^{-1}\bm{y} following the standard GP prediction formula (for the training set). At large nn, fˉ\bar{\bm{f}} can be approximated using the equivalence kernel (EK) approximation together with its perturbative corrections (non-perturbative approaches in 1/n1/n could also be considered ). Within this approximation scheme, one considers [Qf]μν=Qf(xμ,xν)[Q_{f}]_{\mu\nu}=Q_{f}(\bm{x}_{\mu},\bm{x}_{\nu}) as the continuum operator Qf(x,y)Q_{f}(\bm{x},\bm{y}), diagonalized on the data-set measure (dμd\mu), leading to the following formula for fˉ(x)\bar{f}(\bm{x}) in the strict EK limit

where λ,ϕλ(x)\lambda,\phi_{\lambda}(\bm{x}) are the eigenvalues and eigenfunctions of QfQ_{f} and yλ=∫dμxy(x)ϕλ(x)y_{\lambda}=\int d\mu_{x}y(\bm{x})\phi_{\lambda}(\bm{x}). To obtain an explicit formula for, tˉ(x)=[y(x)−fˉ(x)]/σ2\bar{t}(\bm{x})=[y(\bm{x})-\bar{f}(\bm{x})]/\sigma^{2}, we proceed by solving the eigenvalue problem. To this end, we first consider the kernel action on a general linear function (w′⋅zj\bm{w^{\prime}}\cdot\bm{z}_{j}), where zj,w′\bm{z}_{j},\bm{w^{\prime}} are vectors of size SS for all jj, and zj\bm{z}_{j} is drawn from the dataset measure,

where p(w)p(\bm{w}) is a centered Gaussian with covariance matrix Σ\Sigma. In the second transition, we undo the kernel integral, where dμzd\mu_{z} is a Gaussian measure. We then exchange the order of integration and sum over the weights and the data. We now do the integration over the data,

where on the r.h.s. we noted that as S≫1S\gg 1, \normw2\norm{\bm{w}}^{2} is weakly fluctuating and close to its mean. Next, we perform the ∫dwp(w)\int d\bm{w}p(\bm{w}) integral which, following this mean-field replacement, is now of the same type as the previous one. Overall this yields

Using again S≫1S\gg 1 we replace 2xjTΣxj2\bm{x}_{j}^{\mathsf{T}}\Sigma\bm{x}_{j} by its mean under ∫dμx\int d\mu_{x}. Following this, we obtain that the action of QfQ_{f} preserves the space of linear function. Furthermore, we see that diagonalizing QfQ_{f} in this subspace reduces to diagonalizing Σ\Sigma. We thus reach the conclusion that since y(x)y(\bm{x}) is a linear function, so must be fˉ(x)\bar{f}(\bm{x}) and hence tˉ(x)\bar{t}(\bm{x}).

We turn to the quantity ∑μνAμν(2)∂[Qf]μν∂Σss′\sum_{\mu\nu}A_{\mu\nu}^{(2)}\frac{\partial[Q_{f}]_{\mu\nu}}{\partial\Sigma_{ss^{\prime}}} appearing in the equation of state for Σ\Sigma. For the moment, we omit fluctuation piece of Aμν(2)A_{\mu\nu}^{(2)} and replace it by −δμδν-{\delta}_{\mu}{\delta}_{\nu}. Below, we will show that the contribution of this fluctuation term is negligible at large, nn which is our focus here. Following this, we approximate the two summations with the following expression

Next, we argue that δ(x)=∑ibi(w∗⋅xi){\delta}(\bm{x})=\sum_{i}b_{i}(\bm{w}^{*}\cdot\bm{x}_{i}) at large nn, where bib_{i}’s are some real numbers. Indeed, at large nn when replacing summations by integrals and all matrices by their continuum kernels, the full symmetry of the measure from which xμ\bm{x}_{\mu}’s are drawn becomes manifest in the equations. The latter amounts to an independent orthogonal rotation (O(S)O(S)) of each, xi\bm{x}_{i} which leaves w∗\bm{w}^{*} invariant. Recalling the previous result, that fˉ(x)\bar{f}(\bm{x}) is linear, together with this symmetry, implies that fˉ(x)=∑ici(w∗⋅xi)\bar{f}(\bm{x})=\sum_{i}c_{i}(\bm{w}^{*}\cdot\bm{x}_{i}) where cic_{i}’s are some real numbers. Consequently, δ(x)=σ−2[y(x)−fˉ(x)]\delta(\bm{x})=\sigma^{-2}[y(\bm{x})-\bar{f}(\bm{x})] is of the same form.

Using the above ansatz for, tˉ(x)\bar{t}(\bm{x}) we can solve for the r.h.s. of Eq. (102). To this end, we again rewrite the r.h.s. by undoing the kernel integral and exchanging the order of the integration and the sum,

Next we use again the fact that \normw2\norm{\bm{w}}^{2} is weakly fluctuating at large SS to perform the remaining integration over w\bm{w} and obtain

Here one can also see a different justification for the VGA underlying our equations of state. Indeed, Eq. (102) is exactly the non-linear term in the mean-field-decoupled-action for the input layer. At large, nn it leads to Eq. 103 where replacing \normw2\norm{\bm{w}}^{2} by its expectation value makes the term quadratic. As the first term in that action is quadratic in, w\bm{w} the overall action becomes Gaussian, as the VGA assumes.

Following the above simplification, the equation for Σ\Sigma becomes

revealing that only the eigenvector along w∗\bm{w}^{*} is affected by training. At large, SS we may thus replace \Tr[Σ]\Tr[\Sigma] by σw2\sigma_{\text{w}}^{2} rendering the above an explicit formula for Σ\Sigma.

Next, we return to the eigenvalue equation for QfQ_{f} (Eq. 101) and use the fact that w∗\bm{w}^{*} is an eigenvalue (l∗l_{*}) of Σ\Sigma

Taking next the EK predictions along with its leading correction yields

where Cˉn\bar{C}_{n} is the posterior covariance in the EK limit given by ∑λ1λ−1+n/σ2\sum_{\lambda}\frac{1}{\lambda^{-1}+n/\sigma^{2}}, where λ\lambda are Qf(x,y)Q_{f}(\bm{x},\bm{y})’s eigenvalues in the C→∞C\rightarrow\infty limit. In practice, we estimated Cˉn\bar{C}_{n} numerically by diagonalizing large kernels. For N,S=20,64N,S=20,64 found Cˉ800=0.13\bar{C}_{800}=0.13 and Cˉ1600=0.078\bar{C}_{1600}=0.078.

Notably since δ(x)\delta(\bm{x}) came out proportional to y(x)y(\bm{x}) we obtained a simplified form for δ(x)\delta(\bm{x}) containing only a single free parameter (α\alpha)

In addition we may now replace all the \normb2\norm{\bm{b}}^{2} factor appearing above with α2\norma∗2\alpha^{2}\norm{\bm{a}^{*}}^{2}.

Altogether, this yields the following non-linear equation for the scalar quantity α\alpha

Solving the above equation for α\alpha, one obtains Σ\Sigma and QfQ_{f} using the equations of state. Using, QfQ_{f} we can calculate the DNNs predictions on the test set using standard GP inference. We note by passing that one can also estimate the result of this GP inference using EK. However, we found that just keeping a leading perturbative correction to the EK result, as done for the training set, resulted in 20−30%20-30\% discrepancies when compared to exact the GP inference formula. As estimating GP inference was not a main focus of the current work, we instead opted to perform this last GP inference on the test set numerically. In principle, other analytical methods for estimating GP inference could be used here . Finally, we note that when estimating αtrain\alpha_{\text{train}} and αtest\alpha_{\text{test}} on a specific dataset they are defined as

where Train{\it{Train}} and Test{\it{Test}} refer to samples taken from the train and test datasets, respectively.

Last we turn to discuss the fluctuation term we omitted given by

We wish to compare its contribution to that of C−1\Tr[δδT∂[Qf]μν∂Σss′]C^{-1}\Tr[{\bm{\delta}}{\bm{\delta}}^{\mathsf{T}}\frac{\partial[Q_{f}]_{\mu\nu}}{\partial\Sigma_{ss^{\prime}}}]. To this end, we write QfQ_{f} in terms of its eigenvectors (vk\bm{v}_{k}) and eigenvalues (λk\lambda_{k})

aiming for an order of magnitude estimation, we perform the following two approximations: First, we approximate the eigenvalue by the leading nn eigenvalues of the continuum kernel time nn. Second, we take this continuum kernel to be GP kernel. The latter is justified by the fact that the feature-learning effects we found are large, but still do not correspond to an order of magnitude change. Following this, we obtain NSNS degenerate eigenvalues equal to nλ∞n\lambda_{\infty} and the corresponding vk\bm{v}_{k}’s span all linear functions on input space (sampled on the training set).

Next, we estimate how adding such terms to the previous computation affects the equation for Σ−1\Sigma^{-1}. First, we note that nλ∞∼n/(NS)∼1n\lambda_{\infty}\sim n/(NS)\sim 1, in our two experiments. Next, we imagine repeating the computation of the previous section, with these extra terms corresponding to the various different vkvkT\bm{v}_{k}\bm{v}_{k}^{\mathsf{T}}. Notably each such term would enter the computation in the same exact manner to tˉ\bar{\bm{t}} (see Eq. 102) namely

the only two differences are that (i) ∣∣t∣∣2=α2n||\bm{t}||^{2}=\alpha^{2}n whereas ∣∣[vk]∣∣2=1||[\bm{v}_{k}]||^{2}=1 (hence a factor nn on the r.h.s. was lost compared to Eq. 102) and (ii) vk\bm{v}_{k} can have any dependence on xi\bm{x}_{i} and not only through w∗⋅xi\bm{w}^{*}\cdot\bm{x}_{i}. For concreteness, let us span the continuum version of these vk\bm{v}_{k} by [xi]s[\bm{x}_{i}]_{s}. Summed together, all these NSNS eigenvalues will end up augmenting the r.h.s. Eq. 106 into

comparing the first and last term on the r.h.s we find it is negligible for CS(nλ∞+σ2)≫nCS(n\lambda_{\infty}+\sigma^{2})\gg n. Notably, even for our n=800,S=64n=800,S=64 experiment at C=80C=80, we find this last term is negligible.

Here, we argue that the replacement involved in Eq. (102) is valid for n≫NSn\gg\sqrt{N}S. We further comment on some implications this has for the fully-connected case (N=1N=1). To show this, we consider a summation of the form

undoing the kernel integral as we have done before (see for example Eq. (99)) yields

For simplicity, we present the analysis for N=1N=1, and report the results for general NN by symmetry. Re-focusing on the relevant summation,

we consider the average (denoted below by a bar) and variance of II over the dataset, {xμ}μ=1n\{\bm{x}_{\mu}\}_{\mu=1}^{n} where each sample is drawn from the measure dμxd\mu_{x}, conditioning on the values of w\bm{w} and u\bm{u}. This yield,

To estimate the scale of these quantities, we focus simplicity on the regime where the error function is linear. This can be generalized by taking into account perturbative corrections, but these do not change the scale. Following this approximation and taking into account centered Gaussian i.i.d. measure with variance 11 for, dμxd\mu_{x} one finds,

Since u\bm{u} and w\bm{w} are high- dimensional vectors where w∼N(0,ISσ2/S)\bm{w}\sim\mathcal{N}(0,I_{S}\sigma^{2}/S) and u\bm{u} can be taken to be fixed with O(1)O(1) norm in our setting. Therefore, the norm of w\bm{w} concentrates on its average value, σw2\sigma_{\text{w}}^{2}, which is of order one, with fluctuation of order S−1/2S^{-1/2}. Hence, the variance fluctuation are of order 1/n21/n^{2} and mean fluctuation are of order max⁡(1/S,1/n)\max(1/{S},1/n). Thus, n≫Sn\gg{S} is required for replacing the summation by an integral for N=1N=1 (the fully connected case). Similar analysis can be done for general NN leading to Var(I)=O(1/(Nn2))Var(I)=O(1/(Nn^{2})), and Iˉ=O(1/(SN))\bar{I}=O(1/(SN)) which then requires n≫NSn\gg{\sqrt{N}S}, where we assumed, without loss of generality, that uj=aj∗w∗\bm{u}_{j}=a^{*}_{j}\bm{w}^{*} and \norma∗2=1\norm{\bm{a}^{*}}^{2}=1, as in the previous section. For our CNN model N=20,S=64N=20,S=64 and n=800,1600n=800,1600 hence this approximation is reasonable.

Let us consider the implications this has on the fully connected case. In the regime where n≫Sn\gg S, the behavior of the MSE will change drastically compared to N≫1N\gg 1. Indeed, as our previous results show, for our CNN experiments, the train MSE (over σ4\sigma^{4}) reaches 0.420.4^{2} of the corresponding at n=1600≫Sn=1600\gg S, and for the GP this quantity is order 11. Thus, while being small, this MSE is far from negligible. Specifically, in our CNN experiments, the emergent scale (α2n2/(CNS)\alpha^{2}n^{2}/(CNS)) is order 11.

In contrast, at N=1N=1, much like in experiments , the self-consistent equation predicts a negligible GP-DNN performance gap down to CC or order 11. Moreover, for n≫Sn\gg S the GP (which has a uniform prior over all the SS possible linear functions) actually performs very well. Specifically, the EK approximation yields a train MSE (over σ4\sigma^{4}) of order 10−510^{-5}. The emergent scale (α2n2/(CS)\alpha^{2}n^{2}/(CS)) is of the same order at 1/C1/C.

Several conclusions could be drawn here: (i) Taking N=1N=1 in our experiments, there is no separation of scales between the emergent scale and 1/C1/C. (ii) In a related manner, no appreciable label/target-aware feature learning will take place down to the scale where our inter-layer mean-field breaks down (C∼1C\sim 1).

Validity of the Variational Gaussian Approximation

Here, we provide some analytical support to the validity of the Gaussian variational approximation, used to obtain the equation of state. We present the variational treatment from a perturbation theory approach, as a partial summation of a subset of all perturbative corrections. We identify a qualitative difference between this subset and other perturbative corrections that we neglect. We apply our analysis to one of the typical hidden layers in the mean-field limit. This is easily generalized to all layers due to the recursive structure of the problem.

In Eq. (43), we introduce the mean-field probability distribution:

where ⟨...⟩K\langle...\rangle_{K} is averaging with respect to the Gaussian measures induced by the kernel KK. For the layer below, This yields the following self-consistent equation for the post-kernel, QQ and the pre-kernel, KK:

We now show in what sense this approximation is valid, i.e. when can we approximate the mean-field distribution by a Gaussian distribution. We start by calculating the interacting Green function of the process (second moment of the process).

where ⟨…⟩\langle\ldots\rangle is the connected expectation with respect to SMF\mathcal{S}_{\text{MF}} and ⟨…⟩0\langle\ldots\rangle_{0} is the connected expectation with respect to S0,MF.\mathcal{S}_{0,\text{MF}}. Here, connected mean cumulant moments w.r.t the variables hhhh and ΔSMF\Delta\mathcal{S}_{\text{MF}}. We perform our perturbation analysis, for simplicity, for the monomial activation function ϕ(x)=xk\phi(x)=x^{k}, with kk finite, as a characteristic example. This can be generalized to other smooth activation functions. The perturbative correction term, to the free moment, ⟨hμhν(ΔSMF)j⟩0\langle h_{\mu}h_{\nu}\left(\Delta S_{\text{MF}}\right)^{j}\rangle_{0} is as follows:

where k2jk^{2j} is due to the choice of hαh_{\alpha} from each monomial activation function, the j!j! is due to the arrangement in pairs of the remaining fields. We denote by V(k)α1β1=⟨hα1k−1hβ1k−1⟩0Aα1β1V(k)_{\alpha_{1}\beta_{1}}=\langle h_{\alpha_{1}}^{k-1}h_{\beta_{1}}^{k-1}\rangle_{0}A_{\alpha_{1}\beta_{1}}. Plugging back in Eq. (126) leads to the following self-consistent equation:

Plugging the definition of the matrix V, we have that,

The second transition is using Gaussian integration by parts. The resulting equation is very similar to the mean-field equation we find. The difference is that here instead of derivative by KK, the pre-kernel, we have a derivative of the post-kernel. In addition, the expectation is also with respect to the free theory. Indeed, the full variational treatment is self-consistent or, equivalently stated, it takes into account a larger set of diagrams (terms in perturbation theory) which amount to renormalizing the 2-point function from QQ to KK. We argue however that doing so only improves the overall accuracy. Indeed, the expansion of Eq. (128), is the same as one would get from a Gaussian action consisting of S0,MF\mathcal{S}_{0,\text{MF}} plus a quadratic term of the form hαjhβjk⟨hαjk−1hβjk−1⟩Aαjβjh_{\alpha_{j}}h_{\beta_{j}}k\langle h_{\alpha_{j}}^{k-1}h_{\beta_{j}}^{k-1}\rangle A_{\alpha_{j}\beta_{j}}. The variational Gaussian approximation essentially finds the closest Gaussian distribution. Therefore, it can only improve upon this simpler approximation we took here.

Variational Gaussian Approximation for ReLU Activation

Here, we extend the previous VGA treatment, which assumed centered distributions, to non-centered ones. Indeed, for antisymmetric activation functions, the pre-activations appear schematically as h(l+1)=vϕ(h(l))\bm{h}^{(l+1)}=\bm{v}\phi(\bm{h}^{(l)}), thus vϕ(h(l))=−vϕ(−h(l))\bm{v}\phi(\bm{h}^{(l)})=-\bm{v}\phi(-\bm{h}^{(l)}). Since v\bm{v} appears only in quadratic order in the action, we find that h(l)\bm{h}^{(l)} is as likely as −h(l)-\bm{h}^{(l)} and its ensemble average is strictly zero. However, for ϕ=ReLU\phi=\text{ReLU}, it is not the case. This requires us to extend the previous treatment by including extra variational parameters for the mean.

Concretely, let us focus on the VGA for the input layer of 3 layers, CNN. Our VGA for the probability is now defined by the variance of the Gaussian in weight space, Σss′\Sigma_{ss^{\prime}} as well as the mean χs\chi_{s}. Repeating the previous analysis one obtains

where wc∈RS0\bm{w}_{c}\in R^{S_{0}} are the weights of channel, cc and Qμj1νj2(2)Q^{(2)}_{\mu j_{1}\nu j_{2}} is now defined by

which one can reduce to a one-dimensional integral following Ref. . Obtaining an explicit expression, or potentially a perturbation expansion in χ⋅xμ,i\bm{\chi}\cdot\bm{x}_{\mu,i}, is left for future work.

Taking the derivative of the KL-divergence with respect to Σ\Sigma and equating it to zero, one obtains

Similarly as χ\bm{\chi} one obtains the additional equation

where χμj(2)\chi^{(2)}_{\mu j} is the average of hμjc′(2)h_{\mu jc^{\prime}}^{(2)} (for any c′c^{\prime}).

Further details on the numerical experiments

Here we report on several additional numerical results for FCNs. In particular, we provide details about the full spectrum of Σ\Sigma, the equilibration process, and the Gaussianity measures. We also report on some numerical experiments with FCN in the standard scaling.

We conducted further experiments with the above FCNs (d=64,n=1024,σ2=0.001d=64,n=1024,\sigma^{2}=0.001, teacher-student with Nl=1N_{l}=1 for the teacher) however with standard scaling (σa2=2\sigma_{\text{a}}^{2}=2, regardless of NlN_{l}) rather than "MF" scaling (σa2=2/N2\sigma_{\text{a}}^{2}=2/N_{2}). Here, we expect feature learning to diminish . Considering the numerical solution of our EoS, we found that they essentially remain close to the GP limit for N1=N2=64N_{1}=N_{2}=64. Specifically, at N1=N2=64N_{1}=N_{2}=64 leading Σ\Sigma eigenvalue came out 0.032570.03257, whereas in the GP limit we obtain 2/d=0.031252/d=0.03125. This is consistent with the fact that the emergent scale here (Figure 1. panel (b) main text) is of the order of 1/N21/N_{2}. Figure 7 and Figure 8 presents the results for different width NlN_{l}.

2 2-layer CNN experiment

Figure 9 shows the top Σ\Sigma eigenvalues normalized by σw2/d\sigma_{\text{w}}^{2}/d, this is a complementary figure to Figure 1(b) introduced in the main text. Clearly, as the width increase, the eigenvalues of Σ\Sigma are getting closer to the GP kennel eigenvalues.

3 Myrtle-5 CNN on subsets of CIFAR-10 experiment

Here, we report the statistics of the pre-activations in Fourier space similar to Figure 6 in the main text, but now for all layers and projected on the 1st, 3rd and 10th eigenvectors of the Fourier space covariance matrix. Let us begin by giving some more details on the procedure for deriving these quantities. The pre-activations at some layer is of the form hμ,c,x,yh_{\mu,c,x,y} where μ\mu is a data point index, cc is a channel index, and x,yx,y are pixel coordinates. We choose some wavenumber kk and transform these to Fourier space to yield hμ,ckh^{k}_{\mu,c} which summarizes contributions from all pixels. We then compute the covariance matrix of these hμ,ckh^{k}_{\mu,c} (an n×nn\times n matrix), averaging across channels and seeds. Finally, we project hμ,ckh^{k}_{\mu,c} on some eigenvector of the covariance matrix, and these are the quantities whose statistics we report in figures 10, 11.

There are several empirical observations to be made here:

Gaussianity generally increases as we go deeper into the network from the input to the output.

Gaussianity generally increases as we project on higher index eigenvectors (going from left to right across the columns).

Deviations from Gaussianity can appear in several ways: e.g. as multi-modality (e.g. top left panel), or as excessive kurtosis (e.g. 2nd-row left column).

References