Asymptotics of representation learning in finite Bayesian neural networks
Jacob A. Zavatone-Veth, Abdulkadir Canatar, Benjamin S. Ruben, Cengiz Pehlevan
Introduction
The expressive power of deep neural networks critically depends on their ability to learn to represent the features of data . However, the structure of their hidden layer representations is only theoretically well-understood in certain infinite-width limits, in which these representations cannot flexibly adapt to learn data-dependent features . In the Bayesian setting, these representations are described by fixed, deterministic kernels . As a result of this inflexibility, recent works have suggested that finite Bayesian neural networks (henceforth BNNs) may generalize better than their infinite counterparts because of their ability to learn representations .
Theoretical exploration of how finite and infinite BNNs differ has largely focused on the properties of the prior and posterior distributions over network outputs . In particular, several works have studied the leading perturbative finite-width corrections to these distributions . Yet, the corresponding asymptotic corrections to the feature kernels, which measure how representations evolve from layer to layer, have only been studied in a few special cases . Therefore, the structure of these corrections, as well as their dependence on network architecture, remain poorly understood. In this paper, we make the following contributions towards the goal of a complete understanding of feature learning at asymptotically large but finite widths:
We argue that the leading finite-width corrections to the posterior statistics of the hidden layer kernels of any BNN with a linear readout layer and Gaussian likelihood have a largely prescribed form (Conjecture 1). In particular, we argue that the posterior cumulants of the kernels have well-defined asymptotic series in terms of their prior cumulants, with coefficients that have fixed dependence on the target outputs.
We explicitly compute the leading finite-width corrections for deep linear fully-connected networks (§4.1), deep linear convolutional networks (§4.2), and networks with a single nonlinear hidden layer (§4.3). We show that our theory yields quantitatively accurate predictions for the result of numerical experiment for tractable linear network architectures, and qualitatively accurate predictions for deep nonlinear networks, where quantitative analytical predictions are intractable.
Our results begin to elucidate the structure of learned representations in wide BNNs. The assumptions of our general argument are satisfied in many regression settings, hence our qualitative conclusions should be broadly applicable.
Preliminaries
In our analysis, we fix an arbitrary training dataset of examples. We define the input and output Gram matrices of this dataset as and , respectively. For analytical tractability, we consider a Gaussian likelihood for
where is an inverse temperature parameter that sets the variance of the likelihood and . We then introduce the Bayes posterior over parameters given these data:
we denote averages with respect to this distribution by . By tuning , one can then adjust whether the posterior is dominated by the prior () or the likelihood (). We will mostly focus on the case in which the input dimension is large and the training dataset can be linearly interpolated; the low-temperature limit then enforces the interpolation constraint.
2 The Gaussian process limit
In this limit, for built out of compositions of most standard neural network architectures, the prior over function values tends to a Gaussian process (GP) . Moreover, with our choice of a Gaussian likelihood, the posterior over function values also tends weakly to the posterior induced by the limiting GP prior . The kernel of the limiting GP prior is given by the deterministic limit of the inner product kernel of the postactivations of the final hidden layer,
multiplied by the prior variance . For a broad range of network architectures, can be computed recursively . For brevity, we define the kernel matrix evaluated on the training data: .
Elementary perturbation theory for finite Bayesian neural networks
We first present our main result, which shows that the form of the leading perturbative correction to the average hidden layer kernels of a BNN is tightly constrained by the assumptions that the readout is linear, that the cost is quadratic, and that the GP limit is well-defined.
Consider a BNN of the form (1), with posterior (3). Assume that this network admits a well-defined GP limit as discussed in §2.2. Let be a hidden layer observable, that is, a function of the hidden layer activations that is not a function of the readout weights . Assume that tends in probability to a finite, deterministic limit under the posterior in the GP limit.
Then, the posterior cumulants of this observable admit well-behaved asymptotic series at large widths in terms of its joint prior cumulants with the postactiviation kernel . In particular, the asymptotic expansion of the posterior mean has leading terms
where . Here, the cumulants of the kernels are computed with respect to the prior, and are themselves given by asymptotic series at large widths. The ellipsis denotes terms that are of subleading order in the inverse hidden layer widths.
In Appendix B, we derive this result perturbatively by expanding the posterior cumulant generating function of in powers of the deviations of and from their deterministic infinite-width values. There, we also give an asymptotic formula for the posterior covariance of two observables. However, the resulting perturbation series may not rigorously be an asymptotic series, and this method does not yield quantitative bounds for the width-dependence of the terms. We therefore frame it as a conjecture. We note that similar methods can be applied to compute asymptotic corrections to the posterior predictive statistics; we comment on this possibility in Appendix G.
The leading output-dependent correction has several interesting features. First, it includes a factor of , reflecting the fact that inference in wide Bayesian networks with many outputs is qualitatively different from that in networks with few outputs relative to their hidden layer width . If does not tend to zero with increasing , the infinite-width behavior is not described by a standard GP . Moreover, we note that the matrix is invertible at any finite temperature, even when is singular. Therefore, provided that one can extend the GP kernel by continuity to non-invertible , Conjecture 1 can be applied in the data-dense regime as well as the data-sparse regime . Furthermore, we observe that the correction depends on the outputs only through their Gram matrix . This result is intuitively sensible, since with our choice of likelihood and prior the function-space posterior is invariant under simultaneous rotation of the output activations and targets. Finally, is transformed by factors of the matrix , hence the correction depends on certain interactions between the output similarities and the GP kernel .
2 High- and low-temperature limits of the leading correction
To gain some intuition for the properties of the leading finite-width corrections, we consider their high- and low-temperature limits. These limits correspond to tuning the posterior (3) to be dominated by the prior or the likelihood, respectively. At high temperatures (), expanding as a Neumann series (see Appendix A and ) yields
At low temperatures (), the behavior of differs depending on whether or not is of full rank. Assuming for simplicity that it is invertible, we have
in the non-invertible case there are additional contributions involving projectors onto the null space of . Therefore, the leading-order low temperature correction depends on the difference between the target and GP kernels, while the leading non-trivial high temperature correction depends on their sum.
Learned representations in tractable network architectures
Having derived the general form of the leading perturbative finite-width correction to the average feature kernels, we now consider several example network architectures. For these tractable examples, we provide explicit formulas for the feature-learning corrections to the hidden layer kernels, and test the accuracy of our theory with numerical experiments.
This result simplifies further at low temperatures, where, by the result of §3.2, we have
in the regime in which is invertible. We thus obtain the simple qualitative picture that the low-temperature average kernels linearly interpolate between the input and output Gram matrices. In Appendix F, we show that this limiting result can be recovered from the recurrence relation derived through other methods by Aitchison , who did not use it to compute finite-width corrections. We note that the low-temperature limit is peculiar in that the mean predictor reduces to the least-norm pseudoinverse solution to the underlying underdetermined linear system ; we comment on this property in Appendix G.
We can gain some additional understanding of the structure of the correction by using the eigendecomposition of . As is by definition a real positive semidefinite matrix, it admits a unitary eigendecomposition with non-negative eigenvalues . In this basis, the average kernel is
We now seek to numerically probe how accurately these asymptotic corrections predict learned representations in deep fully-connected linear BNNs. Using Langevin sampling , we trained deep linear networks of varying widths, and compared the difference between the empirical and GP kernels with theory predictions. We provide a detailed discussion of our numerical methods in Appendix I. In Figure 1, we present an experiment with a 2-layer linear neural network trained on the MNIST dataset of handwritten digit images using the Neural Tangents library . We find an excellent agreement with our theory, confirming the inverse scaling with width and linear scaling with depth for the deviations from GP kernel.
2 Deep linear convolutional networks
With the given readout strategy, the two-index feature map kernel appearing in Conjecture 1 is related to the four-index kernel of the last hidden layer by . We discuss other readout strategies in Appendix C, but use this vectorization strategy in our numerical experiments.
As shown by Xiao et al. , the infinite-width four-index kernel obeys the recurrence
with base case . This gives convolutional linear networks a sense of spatial hierarchy that is not present in the fully-connected case: even at infinite width, the kernels include iterative spatial averaging.
In Appendix C, we derive the kernel covariances appearing in Conjecture 1. As in the fully-connected case, this computation is easy to perform with the aid of Isserlis’ theorem. The general result is somewhat complicated, but things simplify under the assumption that readout is performed using vectorization. Then, one finds that
where we have defined for brevity. Thus, the correction to the convolutional kernel is quite similar to that obtained in the fully-connected case. To this order, the difference between these network architectures manifests itself largely through the difference in the infinite-width kernels. In Appendix C, we show that a similar simplification holds if readout is performed using global average pooling over space.
As we did for fully-connected networks, we test whether our theory accurately predicts the results of numerical experiment, using the MNIST digit images illustrated in 2(a-d). We consider a network with one-dimensional (Figure 2e and f) and two-dimensional (Figure 3) convolutional hidden layers, trained to classify MNIST images (see Appendix I for details of our numerical methods). As shown in Figure 2(e, f) (Figure 3(a,b) for 2D convolutions), we again obtain good quantitative agreement between the predictions of our asymptotic theory and the results of numerical experiment. In Figure 3c, we directly visualize the learned feature kernels for 2D convolutional layers, illustrating the good agreement between theory and experiment. Therefore, our asymptotic theory can be applied to accurately predict learned representations in deep convolutional linear networks.
3 Networks with a single nonlinear hidden layer
where expectations are taken over the -dimensional Gaussian random vector , which has mean zero and covariance . Unlike for deeper nonlinear networks, here there are no finite-width corrections to the prior expectations .
Though these expressions are easy to define, it is not possible to evaluate the four-point expectation in closed form for general Gram matrices and activation functions , including ReLU and erf. This obstacle has been noted in previous studies , and makes it challenging to extend approaches similar to those used here to deeper nonlinear networks. For polynomial activation functions, the required expectations can be evaluated using Isserlis’ theorem (see Appendix A). However, even for a quadratic activation function , the resulting formula for the kernel will involve many elementwise matrix products, and cannot be simplified into an intuitively comprehensible form.
Learned representations in deep nonlinear networks
In the preceding section, we noted that analytical study of learned representations in deep nonlinear BNNs is generally quite challenging. Here, we use numerical experiments to explore whether any of the intuitions gained in the linear setting carry over to nonlinear networks. Concretely, we study how narrow bottlenecks affect representation learning in a more realistic nonlinear network. We train a network with three hidden layers and ReLU activations on a subset of the MNIST dataset . Despite its analytical simplicity, ReLU is among the activation functions for which the covariance term in Conjecture 1 cannot be evaluated in closed form (see §4.3). However, it is straightforward to simulate numerically. Consistent with the predictions of our theory for linear networks, we find that introducing a narrow bottleneck leads to more representation learning in subsequent hidden layers, even if those layers are quite wide (Figure 4). Quantitatively, if one increases the width of the hidden layers between which the fixed-width bottleneck is sandwiched, the deviation of the first layer’s kernel from its GP value decays roughly as with increasing width, while the deviations for the bottleneck and subsequent layers remain roughly constant. In contrast, the kernel deviations throughout a network with equal-width hidden layers decay roughly as (Figure 4). These observations are qualitatively consistent with the width-dependence of the linear network kernel (8), as well as with previous studies of networks with infinitely-wide layers separated by a finite bottleneck . Keeping in mind the obstacles noted in §4.3, precise characterization of nonlinear networks will be an interesting objective for future work.
Related work
Our work is closely related to several recent analytical studies of finite-width BNNs. First, Aitchison argued that the flexibility afforded by finite-width BNNs can be advantageous. He derived a recurrence relation for the learned feature kernels in deep linear networks, which he solved in the limits of infinite width and few outputs, narrow width and many outputs, and infinite width and many outputs. As discussed in §4.1 and in Appendix F, our results on deep linear networks extend those of his work. Furthermore, our numerical results support his suggestion that networks with narrow bottlenecks may learn interesting features.
Moreover, our analytical approach and the asymptotic regime we consider mirror recent perturbative studies of finite-width BNNs. As noted in §3 and Appendix B, we make use of the results of Yaida , who derived recurrence relations for the perturbative corrections to the cumulants of the finite-width prior for an MLP. However, Yaida did not attempt to study the statistics of learned features; the goal of his work was to establish a general framework for the study of finite-width corrections. Bounds on the prior cumulants of a broader class of observables have been studied by Gur-Ari and colleagues ; these results could allow for the identification of observables to which Conjecture 1 should apply. Finally, perturbative corrections to the network prior and posterior have been studied by Halverson et al. and Naveh et al. , respectively. Our work builds upon these studies by perturbatively characterizing the internal representations that are learned upon inference.
Following the appearance of our work in preprint form, Roberts et al. announced an alternative derivation of the zero-temperature limit of Conjecture 1 for MLPs; we have adopted their terminology of hidden layer observables. As in Yaida ’s earlier work, they rely on sequential perturbative approximation of the prior over preactivations as the hidden layers are marginalized out in order from the first to the last. While our elementary perturbative argument for Conjecture 1 does not require assuming a particular network architecture for the hidden layers, it takes as input information regarding the prior cumulants that would have to be approximated using such methods. Moreover, the approach of layer-by-layer approximation to the prior could enable a fully rigorous version of Conjecture 1 to be proved on an architecture-by-architecture basis .
Our work, like most studies of wide BNNs , focuses on the regime in which the sample size is held fixed while the hidden layer width scale tends to infinity, i.e., . One can instead consider regimes in which is not negligible relative to , in which the posterior would be expected to concentrate. The behavior of deep linear BNNs in this regime was recently studied by Li and Sompolinsky , who computed asymptotic approximations for the predictor statistics and hidden layer kernels. In Appendix F, we show that our result (9) for the zero-temperature kernel can be recovered as the limit of their result. As the dataset size appears only implicitly in our approach, we leave the incorporation of large- corrections as an interesting objective for future work. We note, however, that alternative methods developed to study the large- regime cannot overcome the obstacles to analytical study of deep nonlinear networks encountered here.
Conclusions
In this paper, we have shown that the leading perturbative feature learning corrections to the infinite-width kernels of wide BNNs with linear readout and least-squares cost should be of a tightly constrained form. We demonstrate analytically and with numerical experiments that these results hold for certain tractable network architectures, and conjecture that they should extend to more general network architectures that admit a well-defined GP limit.
Limitations. We emphasize that our perturbative argument for Conjecture 1 is not rigorous, and that we have not obtained quantitative bounds on the remainder for general network architectures. It is possible that there are non-perturbative contributions to the posterior statistics that are not captured by Conjecture 1; non-perturbative investigation of feature learning in finite BNNs will be an interesting objective for future work . More broadly, we leave rigorous proofs of the applicability of our results to more general architectures and of the smallness of the remainder as objective for future work. As mentioned above, one could attempt such a proof on an architecture-by-architecture basis . Alternatively, one could attempt to treat all sufficiently sensible architectures uniformly . Furthermore, we have considered only one possible asymptotic regime: that in which the width is taken to infinity with a finite training dataset and small output dimensionality. As discussed above in reference to the work of Aitchison and Li and Sompolinsky , investigation of alternative limits in which output dimension, dataset size, depth, and hidden layer width are all taken to infinity with fixed ratios may be an interesting subject for future work.
Acknowledgments and Disclosure of Funding
We thank B. Bordelon for helpful comments on our manuscript. JAZ-V acknowledges partial support from the NSF-Simons Center for Mathematical and Statistical Analysis of Biology at Harvard and the Harvard Quantitative Biology Initiative. This work was further supported by the Harvard Data Science Initiative Competitive Research Fund, the Harvard Dean’s Competitive Fund for Promising Scholarship, and a Google Faculty Research Award. The authors declare no conflict of interest.
References
Appendix A Preliminary technical results
In this appendix, we review useful technical results upon which our calculations rely.
Let be a zero-mean Gaussian random vector. Then, Isserlis’ theorem states that
where the sum is over all pairings of and the product is over all pairs contained in . In particular, for , we have
In physics, Isserlis’ theorem is often known as Wick’s probability theorem .
A.2 Neumann series for matrix inverses near the identity
The Neumann series is the generalization of the geometric series to bounded linear operators, including square matrices. In particular, let be a square matrix. Then, we have
provided that the series converges in the operator norm . We will use this result without concern for rigorous convergence conditions, as we are interested only in asymptotic expansions.
A.3 Series expansion of the log-determinant near the identity
Let be a square matrix, and let be a small parameter. Then, we have
assuming that the series converges. We will not concern ourselves with rigorous convergence conditions, as we will use this expansion formally.
The base case is given by Jacobi’s formula :
and the fact that commutes with , we find that the claim holds by induction. As , this implies the desired Maclaurin series.
Appendix B Perturbation theory for wide Bayesian neural networks with linear readout
We fix an arbitrary training dataset of examples, and use a Gaussian likelihood , where
is a quadratic cost. We then introduce the Bayes posterior
averages with respect to this distribution will be denoted by .
We define the postactivation feature map kernel
and write for the kernel evaluated on the training set. For brevity, we will frequently abbreviate throughout this appendix.
Our starting point is the partition function of the Bayes posterior (3) for the network (1), including a source term for the (generically matrix-valued) observable :
where denotes all of the parameters except for the readout weight matrix and expectation is taken with respect to the Gaussian prior. The logarithm of the partition function is the posterior cumulant generating function of the observable , with
We first show that the readout layer can be integrated out exactly. As the source term is independent of , Fubini’s theorem yields
The expectation over is a Gaussian integral, hence it is easy to evaluate exactly:
where we abbreviate and introduce the matrices and . Here, we have used the fact that the matrix is invertible at any finite temperature. By the Weinstein–Aronszajn identity ,
where we introduce the (non-constant) kernel matrix
as mentioned above, we abbreviate for brevity. By the push-through identity ,
hence, using the cyclic property of the trace,
where we have defined the normalized Gram matrix of the outputs
B.2 Perturbative expansion
We now consider how this expression behaves in the large-width limit. We assume that this limit is well-defined in the sense that the readout kernel tends in probability to the constant GP kernel , and that the observable similarly tends to a deterministic limit . Then, we formally write and as their infinite-width limits plus corrections which are small at large hidden layer widths:
where the parameter is used to track powers of the small deviations.
We first expand the term resulting from integrating out the readout layer into its infinite-width limit and a finite-width correction. We define the constant matrix
which is invertible at any finite temperature. Then, by the Woodbury identity , we have,
Noting that that both and are , we expand the logarithm of the partition function as
We can then see that the -th cumulant is , hence the -th posterior cumulant of will be . Specifically, we can read off the posterior mean
To make further progress, we expand in powers of . Using the Neumann series for the matrix inverse (see Appendix A), we have
and, using the series expansion of the log-determinant near the identity (see Appendix A), we have
The leading term is simple because it is linear in . Then, keeping only the leading non-trivial corrections and recognizing that
Appendix C Explicit covariance computations in deep linear networks
In this appendix, we detail how to compute the prior covariances appearing in (5) for the hidden layer kernels of deep linear fully-connected and convolutional networks.
for any . By Isserlis’ theorem (see Appendix A), we have
for the second moments of the kernels at each layer. This recurrence relation is in principle exactly solvable for any finite width, but we are interested only in its leading-order behavior at large widths. In particular, we can read off that
Moreover, one can see by Isserlis’ theorem that the third and higher cumulants will be . Substituting this result into (5) with the hidden layer kernel as the observable of interest, we obtain the expression (8) given in the main text.
C.2 Convolutional linear networks
In this subsection, we derive the prior cumulants required to compute corrections to the average feature kernels of deep convolutional linear networks. As described in the main text, following the setup of Novak et al. and Xiao et al. , we consider a network consisting of linear convolutional layers followed by a fully-connected linear readout layer. For simplicity, we assume circular padding and no internal pooling. As discussed in Novak et al. , this setup could be easily extended to other padding strategies, strided convolutions, and average pooling in intermediate layers.
The hidden layer activations are then defined through the recurrence
with base case . We fix the prior distribution of the filter elements to be
where is a weighting factor that sets the fraction of receptive field variance at location (and is thus subject to the constraint ). For inputs and , we introduce the hidden layer kernels
We will first compute the prior mean and covariance of these four-indexed kernels, and then address how to handle readout across space.
As shown by Xiao et al. , the prior mean obeys the recurrence
Moreover, as in the fully-connected case considered in the preceding section, we have
for the second prior moments of the kernels. As in the fully-connected case, these recurrence relations could in principle be solved exactly, but we are only interested in their large-width behavior. Using the forward recurrence for the GP kernels, we can easily read off that
which can then be substituted into the desired cross-layer covariance:
We now address the question of how to read out the convolutional layer activities across space. Following Novak et al. , we consider two strategies: vectorization and projection. With vectorization, the output of the final convolutional layer is flattened into a -dimensional vector before readout, i.e., or . The two-index feature map kernel appearing in Conjecture 1 is then related to the four-index convolutional hidden layer kernel analyzed above via
With projection, the feature map is formed by contracting the final convolutional layer with a fixed vector , i.e.,
Examples of common projection readout strategies include global average pooling () and single-pixel subsampling ( for some desired location ). These readout approaches endow the network with differing properties under spatial transformations; global average pooling has the particular property of making the output translation-invariant.
where we have defined for notational convenience. As elsewhere, for the two-index kernel determined by the chosen readout strategy. Depending on the chosen readout strategy, this general expression can be simplified dramatically. In particular, for vectorization or global average pooling, the correction does not depend on the particular form of .
To show this for vectorization (the strategy used in our experiments), we substitute the definition of from (C.31) and the expression for the cross-layer kernel covariance from (C.2) into the general expression for the correction to obtain
thanks to the normalization constraint . We now notice that is a symmetric matrix, and that the kernel remains invariant under the simultaneous exchange of indices and . Then, substituting in the expression for the same-layer kernel covariance (C.2), it is easy to show that the correction reduces to
This yields the expression given in the main text.
For projection, an analogous simplification is possible in the case of global average pooling (). Substituting the definition of from (C.33) and expression for the cross-layer kernel covariance (C.2) into the correction, we have
Substituting in the expression for the same-layer kernel covariance (C.2), it is again easy to show that the correction reduces to
For projection strategies other than global average pooling (more precisely, for strategies for which is not constant), the sum over indices in the cross-layer covariance is not independent of the shift, hence we cannot simplify the correction in a similar fashion. This can be seen explicitly when treating the case of single-pixel subsampling ( for some desired location ). In this case, the correction reduces to
Unlike for vectorization or for projection using global average pooling, this expression is manifestly dependent on the form of .
Appendix D Direct computation of the average hidden layer kernels of a deep linear MLP
In this appendix, we provide a self-contained derivation of the average hidden layer kernels of a deep linear fully-connected network (MLP). This derivation relies upon neither the results of Appendices B and C nor those of Yaida .
where the “effective action” for the preactivations and Lagrange multipliers is
As described in Appendix B, source terms can be added to the effective action to allow computation of various averages. For deep linear networks, it is convenient to scale the source terms by an overall factor of , for which we must correct when computing the averages:
For an MLP, our task is therefore to integrate out the preactivations and corresponding Lagrange multipliers. We will do so sequentially from the first layer to the last, keeping terms up to the desired order at each step, akin to the approach of Yaida . So long as and are fixed and small relative to the width of the hidden layers, this is a consistent perturbative approach, as noted by Yaida .
D.2 General form of the perturbative layer integrals for a deep linear network
In this section, we evaluate the general form of the integrals required to perturbatively marginalize out a given layer of a deep linear network to . These integrals are generically of the form
We will proceed by evaluating the integrals for invertible, and then infer the general case by a continuity argument. We treat the quartic term perturbatively, and all other terms directly. Writing
the leading term in the integral over is
Then, the quartic correction to the integral over is proportional to
where we write .
We now must integrate over . The leading term is simply
and, by analogy to the corresponding four-point average for ,
Then, the correction to the integral over is proportional to
by analogy with the corresponding quartic expectation for .
We must now expand our results in . The inverses of the matrices and have Neumann series
and we write . Then, using the series expansion of the log-determinant, we find that the logarithm of the leading term expands as
while the quartic correction simplifies to
Combining these results, we find that the result of integrating out the layer to is
As this result is a continuous function of , as the set of full-rank positive definite matrices is dense in the space of positive semidefinite matrices, this result holds for all positive-semidefinite .
We now further expand this result in . This yields
hence we find that the logarithm of the leading term yields
After some straightforward but tedious algebra, the quartic term reduces to
Combining these results, we find that the result of integrating out the layer is
Again, this result is continuous in , hence it holds even if is rank-deficient.
D.3 Perturbative computation of the partition function of a deep linear network
We now apply the results of Appendix D.2 to compute the partition function for a deep linear network to the desired order. Our starting point is the effective action before any of the layers have been integrated out, including a source term:
Applying the results of Appendix D.2 with
we find that the effective action after integrating out the first layer is
Assuming that the network has more than one hidden layer, if we now again apply the results of Appendix D.2 with
we find that the effective action after integrating out the first two layers is
Then, by induction, we can see that we can iterate this procedure to integrate out all of the hidden layers, yielding
Applying the results of Appendix D.2 one final time with
D.4 Computing the average hidden layer kernels of a deep linear network
With the relevant partition function in hand, we can finally compute the average hidden layer kernels. In particular, we can immediately read off that
To obtain the expression listed in the main text, we note that
mirroring the width dependence found by Yaida in his study of the prior of deep linear networks.
Appendix E Average kernels in a deep feedforward linear network with skip connections
In this appendix, we show that Conjecture 1 holds perturbatively for a linear feedforward network with arbitrary skip connections, following the method of Appendix D. Concretely, we consider a network defined as
Upon integrating out the weights, we obtain an effective action for the preactivations and the corresponding Lagrange multipliers of
we find that the effective action after integrating out the first layer is
we find that the effective action after integrating out the first two layers of the network is
where the coupling constants and effective source obey the recurrences
we find the source-dependent terms in the logarithm of the partition function are
E.2 Computing the average hidden layer kernels
With the source-dependent terms of the relevant partition function in hand, we can compute the average hidden layer kernels for a feedforward linear network with arbitrary skip connections. We can immediately read off that
From the form of these recurrences, we can see that
Appendix F Comparison to the results of Aitchison [10] and Li and Sompolinsky [16]
In this appendix, we compare our results for the average kernels of deep linear networks to those of Aitchison and Li and Sompolinsky .
We first show that our result (9) for the low-temperature limit of the average kernels of a deep linear network can be recovered from the results of Aitchison . Working in what corresponds to the zero-temperature limit of our setup, Aitchison derives the following implicit recurrence
and solve the recurrence relations order-by-order using the resulting Neumann series
We now consider the leading finite-width correction. For the last hidden layer, we obtain
after dropping all terms that are of and multiplying on the left and right by . For the first hidden layer, we have
Based on the form of these recurrences, we make the ansatz that the solution is of the form
Substituting the expression for into the condition resulting from the recurrence relation centered on , we find that we must have
hence we can iterate this process backwards to the second hidden layer, yielding
F.2 Comparison to the results of Li and Sompolinsky [16]
We now show that our result (9) for the low-temperature limit of the average kernels of a deep linear network can be recovered as a limiting case of the result of Li and Sompolinsky . Their result for the zero-temperature kernel in the limit with , , , and is, in our notation,
Here, the orthogonal matrix is the matrix of eigenvectors of
for the pseudoinverse of , and the scalars are in turn defined in terms of the eigenvalues as
we note that Li and Sompolinsky use variables .
As we are interested in the limit , it is useful to write the implicit equation for as
hence we expect as . Thus, we have
in the limit in which and . Therefore, combining this result with that of the previous subsection, our result (9) agrees with those of Aitchison and of Li and Sompolinsky in the appropriate limit. Whether the full result of Li and Sompolinsky agrees with that of Aitchison is an interesting question, but is well beyond the scope of the present work.
Appendix G Predictor statistics and generalization in deep linear networks
Though the main focus of our work is on the asymptotics of representation learning, we have also computed the leading finite-width corrections to the predictor statistics. Though one can derive the analogy of Conjecture 1 for the predictor statistics of a general BNN with linear readout, the resulting formula is not particularly illuminating. We will therefore present results only for linear networks. As was true of the hidden layer kernels of deep linear networks, this calculation can be performed either using methods similar to those described in Appendix B or Appendix D. As the steps are largely identical to those calculations, we only briefly summarize the results.
In short, we fix a test dataset of examples, and define the Gram matrices
Introducing appropriate source terms to allow us to compute predictor statistics, we then proceed perturbatively as before, assuming that the combined input Gram matrix
is invertible. Again, the final result can be extended to the case in which this matrix is not invertible by a continuity argument.
Our notation in this appendix will follow that of Appendix B rather than Appendix D in that we will introduce matrices
to denote the blocks of the infinite-width kernel of the last hidden layer, rather than introducing scalar parameters to represent the products of variances. This will make our expressions somewhat more compact than they would be under the conventions of Appendix D.
Defining the matrix , we find that the mean predictor can be written compactly as
The mean and covariance of the training set predictor can be obtained by setting and to in the above expressions.
G.2 Bias-variance decompositions and the low-temperature limit
These results allow us to define thermal bias-variance decompositions of the form
for the mean training and test errors. However, the resulting expressions are not particularly illuminating except in the low-temperature limit . We will focus on the regime in which (and thus ) is invertible, in which the underlying linear system is underdetermined and the training set can be interpolated. In this regime, , and the mean predictor reduces to the least-norm pseudoinverse solution to the linear system, with mean training and test predictions of
respectively. The training and test set covariances have low-temperature limits of
respectively. Then, it is easy to see that both and are , while
Thus, at least to leading order, width affects the low-temperature test error only through the variance term. Substituting in the definition of , we find that to leading order the test error decreases with increasing width if
and increases with increasing width otherwise. This small-initialization condition is the generalization of that found by Li and Sompolinsky to our asymptotic regime.
G.3 Effects of alternative regularization temperature-dependence
In this appendix, we comment on the possibility of alternative temperature-dependent posteriors. This possibility arises from the interpretation of the Bayes posterior (3) as the equilibrium distribution of the Langevin dynamics
Then, if we assume a low-temperature power-law dependence for simplicity, we find that the zero-temperature limits of the training set predictor mean and covariance are
respectively, while those of the test set mean and covariance are
respectively. Therefore, taking yields sensible zero-temperature infinite-width behavior for a linear network of any depth in the underdetermined regime.
Appendix H Derivation of the average kernels for a depth-two network
In this appendix, we derive the average feature kernel for a network with a single (possibly nonlinear) hidden layer and a linear readout. This derivation is a simple extension of the perturbative derivation of Conjecture 1 in Appendix B, using the fact that the size of the terms in the expansion for two-layer networks can be directly controlled in terms of the inverse hidden layer width.
Concretely, we consider a network defined as
Our task is to control the prior cumulants of the hidden layer feature kernel
We can use the fact that the rows of are independent and identically distributed under the prior to obtain
at any hidden layer width . Similarly, we can easily see that
where , and that higher cumulants are . Then, we can directly apply the result of Appendix B to conclude that
for . Depending on the nonlinearity, this result may be continuous in , and therefore extensible to the non-invertible case via a continuity argument. In particular, as noted in Appendix D, this holds for a linear network.
To gain some intuition for how different choices of nonlinear activation function affect the learned representations, we consider the case in which is diagonal. In this special case, the four-point term simplifies dramatically. In particular, we have
Moreover, applying the Sherman-Morrison formula , we have
Appendix I Numerical methods
In this appendix, we describe the numerical methods used in our experiments. We perform our simulations by sampling network parameters at each time step of the Langevin update G.21 after some large burn-in period when the loss function stabilizes around a fixed number. We used Euler-Maruyama method to obtain the discretized Langevin equation:
where is a standard Gaussian random variable sampled i.i.d. at each time step and is the time step. The first, second and last terms represent the weight decay, the gradient descent update and the stochastic Wiener process, respectively.
We used the Neural Tangents framework and PyTorch deep learning library to generate the neural networks and trained them according to the discretized full-batch Langevin update rule. A typical burn-in time was iterations and after that the parameters were sampled over iterations where we chose a learning rate of . Simulations have been performed on a cluster with NVIDIA Tesla V100 GPU’s with 32 GB RAM and a typical simulation run took depending on the architecture and the network width. All code used throughout this work can be reached at https://github.com/Pehlevan-Group/finite-width-bayesian/.
All figures shown here are results of a single instance of a trained neural network on a fixed dataset. Since we performed all our experiments with , we observed that the different initializations of a network did not influence the final posterior mean due to the weight decay and long burn-in periods.
Throughout all experiments, the MNIST digits were downsized from pixels to pixels without distorting the original digits. This was done to accelerate the training process since large input dimensions would take an order of magnitude more time to obtain well estimated posterior means. We considered -dimensional outputs corresponding to one-hot encoded digits. Both inputs and labels were ordered according to their class. Figure 2 shows an example of MNIST digits and the input and output Gram matrices.