Bayesian Deep Ensembles via the Neural Tangent Kernel
Bobby He, Balaji Lakshminarayanan, Yee Whye Teh
Introduction
Given a prior distribution over the parameters, we can define the posterior over , , using Bayes’ rule and subsequently the posterior predictive distribution at a test point :
The posterior predictive is appealing as it represents a marginalisation over weighted by posterior probabilities, and has been shown to be optimal for minimising predictive risk under a well-specified model . However, one issue with the posterior predictive for NNs is that it is computationally intensive to calculate the posterior exactly. Several approximations to have been introduced for Bayesian neural networks (BNNs) including: Laplace approximation ; Markov chain Monte Carlo ; variational inference ; and Monte-Carlo dropout .
In this work, we will relate deep ensembles to Bayesian inference, using recent developments connecting GPs and wide NNs, both before and after training. Using these insights, we devise a modification to standard NN training that yields an exact posterior sample for in the infinite width limit. As a result, when ensembled together our modified baselearners give a posterior predictive approximation, and can thus be viewed as a Bayesian deep ensemble.
One concept that is related to our methods concerns ensembles trained with Randomised Priors to give an approximate posterior interpretation, which we will use when modelling observation noise in regression tasks. The idea behind randomised priors is that, under certain conditions, regularising baselearner NNs towards independently drawn “priors” during training produces exact posterior samples for . Randomised priors recently appeared in machine learning applied to reinforcement learning and uncertainty quantification , like this work. To the best of our knowledge, related ideas first appeared in astrophysics where they were applied to Gaussian random fields . However, one such condition for posterior exactness with randomised priors is that the model is linear in . This is not true in general for NNs, but has been shown to hold for wide NNs local to their parameter initialisation, in a recent line of work. In order to introduce our methods, we will first review this line of work, known as the Neural Tangent Kernel (NTK) .
NTK Background
Wide NNs, and their relation to GPs, have been a fruitful area recently for the theoretical study of NNs: we review only the most salient developments to this work, due to limited space.
First introduced by Jacot et al. , the empirical NTK of is, for inputs , the kernel:
and describes the functional gradient of a NN in terms of the current loss incurred on the training set. Note that depends on a random initialisation , thus the empirical NTK is random for all .
Jacot et al. showed that for an MLP under a so-called NTK parameterisation, detailed in Appendix A, the empirical NTK converges in probability to a deterministic limit , that stays constant during gradient training, as the hidden layer widths of the NN go to infinity sequentially. Later, Yang extended the NTK convergence result to convergence almost surely, which is proven rigorously for a variety of architectures and for widths (or channels in Convolutional NNs) of hidden layers going to infinity in unison. This limiting positive-definite (p.d.) kernel , known as the NTK, depends only on certain NN architecture choices, including: activation, depth and variances for weight and bias parameters. Note that the NTK parameterisation can be thought of as akin to training under standard parameterisation with a learning rate that is inversely proportional to the width of the NN, which has been shown to be the largest scale for stable learning rates in wide NNs .
Lee et al. built on the results of Jacot et al. , and studied the linearised regime of an NN. Specifically, if we denote as the network function at time , we can define the first order Taylor expansion of the network function around randomly initialised parameters to be:
where and is the randomly initialised NN function.
The results of Lee et al. showed that in the infinite width limit, with NTK parameterisation and gradient flow under squared error loss, and are equal for any , for a shared random initialisation . In particular, for the linearised network it can be shown, that as :
and thus as the hidden layer widths converge to infinity we have that:
We can replace with the generalised inverse when invertibility is an issue. However, this will not be a main concern of this work, as our methods will add regularisation that corresponds to modelling observation/output noise, which both ensures invertibility and alleviates any potential convergence issues due to fast decay of the NTK eigenspectrum .
For a generic kernel , Lee et al. observed that this limiting distribution for does not have a posterior GP interpretation unless and are multiples of each other, .
As mentioned in Section 1, previous work has shown that there is a distinct but closely related kernel , known as the Neural Network Gaussian Process (NNGP) kernel, such that at initialisation in the infinite width limit and . Thus Eq. (6) with tells us that, for wide NNs under squared error loss, there is no Bayesian posterior interpretation to a trained NN, nor is there an interpretation to a trained deep ensemble as a Bayesian posterior predictive approximation.
Proposed modification to obtain posterior samples in infinite width limit
Lee et al. noted that one way to obtain a posterior interpretation to is by randomly initialising but only training the parameters in the final linear readout layer, as the contribution to the NTK from the parameters in final hidden layer is exactly the NNGP kernel .Up to a multiple of last layer width in standard parameterisation. is then a sample from the GP posterior with prior kernel NNGP, , and noiseless observations in the infinite width limit i.e. . This is an example of the “sample-then-optimise” procedure of Matthews et al. , but, by only training the final layer this procedure limits the earlier layers of an NN solely to be random feature extractors.
2 Comparison of predictive distributions in infinite width
Having introduced the different ensemble training methods considered in this paper: NNGP; deep ensembles; randomised prior; and NTKGP, we will now compare their predictive distributions in the infinite width limit with squared error loss. Table 1 displays these limiting distributions, , and should be viewed as an extension to Equation (16) of Lee et al. .
For , . Similarly, for , .
Here, when we write for p.d. kernels , we mean that is also a p.d. kernel. One consequence of Proposition 2 is that the predictive distribution of an ensemble of NNs trained via our NTKGP methods is always more conservative than a standard deep ensemble, in the linearised NN regime, when the ensemble size . It is not possible to say in general when this will be beneficial, because in practice our models will always be misspecified. However, Proposition 2 suggests that in situations where we suspect standard deep ensembles might be overconfident, such as in situations where we expect some dataset shift at test time, our methods should hold an advantage.
3 Modelling heteroscedasticity
Following Lakshminarayanan et al. , if we wish to model heteroscedasticity in a univariate regression setting such that each training point, , has an individual observation noise then we use the heteroscedastic Gaussian NLL loss (up to additive constant):
4 NTKGP Ensemble Algorithms
5 Classification methodology
For classification, we follow recent works which treat classification as a regression task with one-hot regression targets. In order to obtain probabilistic predictions, we temperature scale our trained ensemble predictions with cross-entropy loss on a held-out validation set, noting that Fong and Holmes established a connection between marginal likelihood maximisation and cross-validation.
Experiments
Due to limited space, Appendix I will contain all experimental details not discussed in this section.
We begin with a toy 1D example , using homoscedastic . We use a training set of points partitioned into two clusters, in order to detail uncertainty on out-of-distribution test data. For each ensemble method, we use MLP baselearners with two hidden layers of width 512, and erf activation. The choice of erf activation means that both the NTK and NNGP kernel are analytically available . We compare ensemble methods to the analytic GP posterior using either or as prior covariance function using the Neural Tangents library .
Figure 1 compares the analytic NTKGP posterior predictive with the analytic NNGP posterior predictive, as well as three different ensemble methods: deep ensembles, RP-param and NTKGP-param. We plot 95% predictive confidence intervals, treating ensembles as one Gaussian predictive distribution with matched moments like Lakshminarayanan et al. . As expected, both NTKGP-param and RP-param ensembles have similar predictive means to the analytic NTKGP posterior. Likewise, we see that only our NTKGP-param ensemble predictive variances match the analytic NTKGP posterior. As foreseen in Proposition 2, the analytic NNGP posterior and other ensemble methods make more confident predictions than the NTKGP posterior, which in this example results in overconfidence on out-of-distribution data.Code for this experiment is available at: https://github.com/bobby-he/bayesian-ntk.
Flight Delays
We now compare different ensemble methods on a large scale regression problem using the Flight Delays dataset , which is known to contain dataset shift. We train heteroscedastic baselearners on the first 700k data points and test on the next 100k test points at 5 different starting points: 700k, 2m (million), 3m, 4m and 5m. The dataset is ordered chronologically in date through the year 2008, so we expect the NTKGP methods to outperform standard deep ensembles for the later starting points. Figure 2 (Left) confirms our hypothesis. Interestingly, there seems to be a seasonal effect between the 3m and 4m test set that results in stronger performance in the 4m test set than the 3m test set, for ensembles trained on the first 700k data points. We see that our Bayesian deep ensembles perform slightly worse than standard deep ensembles when there is little or no test data shift, but fail more gracefully as the level of dataset shift increases.
Figure 2 (Right) plots confidence versus error for different ensemble methods on the combined test set of 5100k points. For each precision threshold , we plot root-mean-squared error (RMSE) on examples where predictive precision is larger than , indicating confidence. As we can see, our NTKGP methods incur lower error over all precision thresholds, and this contrast in performance is magnified for more confident predictions.
MNIST vs NotMNIST
We next move onto classification experiments, comparing ensembles trained on MNIST and tested on both MNIST and NotMNIST.Available at http://yaroslavvb.blogspot.com/2011/09/notmnist-dataset.html Our baselearners are MLPs with 2-hidden layers, 200 hidden units per layer and ReLU activation. The weight parameter initialisation variance is tuned using the validation accuracy on a small set of values around the He initialisation, , for all classification experiments. Figure 3 shows both in-distribution and out-of-distribution performance across different ensemble methods. In Figure 3 (left), we see that our NTKGP methods suffer from slightly worse in-distribution test performance, with around 0.2% increased error for ensemble size . However, in Figure 3 (right), we plot error versus confidence on the combined MNIST and NotMNIST test sets: for each test point , we calculate the ensemble prediction and define the predicted label as , with confidence . Like Lakshminarayanan et al. , for each confidence threshold , we plot the average error for all test points that are more confident than . We count all predictions on the NotMNIST test set to be incorrect. We see in Figure 3 (right) that the NTKGP methods vastly outperform both deep ensembles and RP methods, obtaining over 15% lower error on test points that have confidence , compared to all baselines. This is because our methods correctly make much more conservative predictions on the out-of-distribution NotMNIST test set, as can be seen by Figure 4, which plots histograms of predictive entropies. Due to the simple MLP architecture and ReLU activation, we can compare ensemble methods to analytic NTKGP results in Figures 3 & 4, where we see a close match between the NTKGP ensemble methods (at larger ensemble sizes) and the analytic predictions, both on in-distribution and out-of-distribution performance.
CIFAR-10 vs SVHN
Finally, we present results on a larger-scale image classification task: ensembles are trained on CIFAR-10 and tested on both CIFAR-10 and SVHN. We conduct the same setup as for the MNIST vs NotMNIST experiment, with baselearners taking the Myrtle-10 CNN architecture of channel-width 100. Figure 5 compares in distribution and out-of-distribution performance: we see that our NTKGP methods and RP-fn perform best on in-distribution test error. Unlike on the simpler MNIST task, there is no clear difference on the corresponding error versus confidence plot, and this is also reflected in the entropy histograms, which can be found in Figure 8 of Appendix I.
Discussion
We built on existing work regarding the Neural Tangent Kernel (NTK), which showed that there is no posterior predictive interpretation to a standard deep ensemble in the infinite width limit. We introduced a simple modification to training that enables a GP posterior predictive interpretation for a wide ensemble, and showed empirically that our Bayesian deep ensembles emulate the analytic posterior predictive when it is available. In addition, we demonstrated that our Bayesian deep ensembles often outperform standard deep ensembles in out-of-distribution settings for both regression and classification tasks.
In terms of limitations, our methods may perform worse than standard deep ensembles when confident predictions are not detrimental, though this can be alleviated via NTK hyperparameter tuning. Moreover, our analyses are planted in the “lazy learning” regime , and we have not considered finite-width corrections to the NTK during training . In spite of these limitations, the search for a Bayesian interpretation to deep ensembles is of particular relevance to the Bayesian deep learning community, and we believe our contributions provide useful new insights to resolving this problem by examining the limit of infinite-width.
A natural question that emerges from our work is how to tune hyperparameters of the NTK to best capture inductive biases or prior beliefs about the data. Possible lines of enquiry include: the large-depth limit , the choice of architecture , and the choice of activation . Finally, we would like to assess our Bayesian deep ensembles in non-supervised learning settings, such as active learning or reinforcement learning.
Acknowledgments and Disclosure of Funding
We thank Arnaud Doucet, Edwin Fong, Michael Hutchinson, Lewis Smith, Jasper Snoek, Jascha Sohl-Dickstein and Sheheryar Zaidi, as well as the anonymous reviewers, for helpful discussions and feedback. We also thank the JAX and Neural Tangents teams for their open-source software. BH is supported by the EPSRC and MRC through the OxWaSP CDT programme (EP/L016710/1).
References
Appendix A Recap of standard and NTK parameterisations
For completeness, we recap the difference between standard and NTK parameterisations & initialisations for an MLP in this section.
On the other hand, under standard parameterisation, the recurrence relation of the NN is:
with and at initialisation. Commonly used initialisation schemes like LeCun or He fall into this category.
We see that the different parameterisations yield the same distribution for the functional output at initialisation, but give different scalings to the parameter gradients in the backward pass. Sohl-Dickstein et al. have recently explored further variants of these parameterisations.
Appendix B Proofs
For our purposes, it will be sufficient to prove convergence of finite-dimensional marginals, jointly, for arbitrary sets of inputs . Note that previous work has already shown that .
The proof that relies on Lévy’s Convergence theorem and the Cramér-Wold device (Theorem 29.4 of ). Using these results it is sufficient to show, denoting as the characteristic function of a random variable , that:
where , defined as the difference between Eqs. (23) & (22), can be shown to be using the Bounded Convergence theorem and the empirical NTK convergence results, and by noting that proofs of NTK convergence are all done on a layer-by-layer basis.
B.2 Proof of Proposition 2
We will prove the case for as the case for is similar, and one can replace inversions of and with generalised inverses if need be.
Let be an arbitrary test set. We will first show . It will suffice to show that is a p.s.d. matrix. But it is not hard to check that:
Likewise, to show we can check that:
and is the contributions to the NTK from parameters before the final layer as before. Finally, we need to define as:
The notation denotes the generalised inverse. follows from standard properties of generalised Schur complements, as does the fact that , which is required for Eq. (26) to hold.
Appendix C Alternative constructions of NTKGP baselearners
A possible alternative construction would be if one could (approximately) sample a fixed , and set:
Appendix D Regularisation in the NTKGP and RP training procedures
As stated in Lemma 3 of Osband et al. , suppose we are in a Bayesian linear regression setting with linear map , model for i.i.d., and parameter prior . Then, having observed training data , solving the following optimisation problem returns a posterior sample :
Appendix E Additional ensemble algorithms
Here, we present our ensemble algorithms for NTKGP-Lin (Algorithm 2) and NTKGP-fn (Algorithm 3), to complement the NTKGP-param algorithm that was presented in Section 3.4.
Note also that for the NTKGP-fn it is unreasonable to assume that the linearised NN dynamics will hold true for the duration of training because, unlike in NTKGP-param (Algorithm 1) we regularise towards the origin not the initialised parameters.
Appendix F Aggregating predictions from ensemble members
For completeness, we now describe how to aggregate predictions from ensemble members. Given a test point , for each baselearner NN , we suppose we have a probabilistic prediction obtained from the NN output. We then treat the ensemble as a uniformly-weighted mixture model over baselearners and combine predictions as . For our Bayesian deep ensembles, we can view this aggregation as a Monte Carlo approximation of the GP posterior predictive with NTK prior.
For classification tasks, this aggregation is exactly an average of predicted probabilities. For regression tasks, the prediction is a mixture of normal distributions, and we follow Lakshminarayanan et al. by approximating the ensembled prediction as a single Gaussian with matched moments. That is to say, if , then we approximate by for and .
Appendix G Comparison of memory and computation costs for ensemble methods
There is only a negligible training-time computational overhead for our NTKGP methods compared to other ensemble methods , for a training set of fixed size (e.g. MNIST, CIFAR-10). This is because one can obtain and store our fixed additive JVPs in a single pass over the training data. For test-time constrained applications, one can employ ensemble distillation for our NTKGP ensembles as one would for standard deep ensembles.
In terms of memory, both NTKGP and RP methods require storage of extra sets of parameters in order to compute the untrainable additive functions and regularise in parameter space, displayed in Table 2 (right). However, the activations of the extra forward pass in the Randomised prior function method need not be stored. And moreover, forward mode JVPs are composed alongside the primitive operations that comprise the forward pass, so the memory requirements incurred by the extra JVP are independent of the NN depth for our NTKGP methods. Note that the memory bottleneck for large NNs is most often from the need to store activations for the backward pass and not from storing parameter sets, hence our NTKGP ensembles are not affected by the main memory bottleneck for large NNs, relative to standard deep ensembles.
It is worth noting that our Bayesian deep ensembles still retain the distributability of standard deep ensembles. Moreover, our computational and memory costs still scale linearly in dataset size and parameter space dimension, enabling us to work with large scale datasets like Flight Delays .
Finally, in this section we only compare the costs associated to different ensemble methods. Ensembles methods are known to be computationally expensive and there has been recent interest in the community to derive new methods that reduce such costs. However, at the time of writing, deep ensembles are state-of-the-art for uncertainty quantification tasks , and hence we believe a comparison of costs between ensemble methods is most appropriate for this work.
Appendix H Scaling for one-hot targets in classification
In Figure 6 we see a different results to Figure 5, as here our NTKGP methods suffer slightly on in-distribution performance but also outperform the baselines methods on out-of-distribution detection. This highlights the importance of the regression target scale when considering classification tasks, and moreover reflects a general theme in our experiments of the trade-off between more aggressive predictions (that tend to perform better on in-distribution) and more conservative predictions (that tend to perform better on out-of-distribution). In our classification methodology, larger values tend to lead to more confident predictions. We point out that this is an issue that affects all ensemble methods and is not limited to our Bayesian ensembles.
Appendix I Experimental Details & additional plots
We set ensemble size , and train on full batch GD with learning rate for 50,000 iterations under standard parameterisation in Neural Tangents , with & , for defined as in Appendix A. In Figure 7 we evaluate the impact of the ensemble size on this toy problem for different ensemble methods. We find that, of the two methods that approximate the analytic NTKGP mean predictor (c.f. Table 1), the approximation of the analytic mean predictor for NTKGP-param degrades compared to RP-param at small ensemble sizes, although the predictive uncertainties are well matched even at small ensemble sizes. The degradation in mean predictor is unsurprising as there is more (untrainable) noise in the initialised NTKGP baselearners. One simple possible solution to this problem, which we leave for future work, is to use separate baselearners for the mean and uncertainty predictions, like in Ciosek et al. .
I.2 Flight Delays
Our baselearners are MLPs with 4 hidden layers, 100 hidden units per layer and ReLU activations, and we use standard parameterisation with & , and choose ensemble size . We train for 10 epochs with learning rate 0.001, batch size 100 and Adam . For all experiments, all ensemble methods apart from standard deep ensembles are regularised according to Appendix D, with weight decay strength set to for standard deep ensembles.
We use a validation set of size 50k that is sampled uniformly from the training set of size 700k, and early stop baselearner NNs based on validation set loss. Inputs and targets are standardised so that the training data is zero mean and unit variance.
I.3 MNIST vs. NotMNIST
For all image classification experiments, we use a split for the train-validation sets needed for temperature scaling.
Baselearners are MLPs with 2-hidden layers, 200 hidden units per layer and ReLU activations. We standardise data to have mean and standard deviation across flattened pixels.
For all ensemble methods, we use standard parameterisation with fixed bias standard deviation , observation noise and tune weight variance on a small linear scale around . We set observation noise for . We train for 20 epochs with batch size 100, learning rate 0.001 and Adam . We do not early stop for any classification experiment, and use the final trained baselearners throughout.
For the analytic NTKGP results, we use the NTK in NTK parameterisation, and use the same observation noise and bias variance as for ensemble methods. However, we fix and also do not tune target scale (set to the base value described in Appendix A) due to computational resources. We also use only half the test sets both for MNIST and NotMNIST due to resource requirements, keeping the ratios of test sizes consistent in order for the error versus confidence plot Figure 3 (right) to be comparable. To compute test and out-of-distribution predictions, having obtained the optimal temperature scale and analytic NTKGP predictions in logit space, , we approximate the softmax class probability predictions: , by a Monte Carlo ensemble approximation with 100 samples.
For all classification ensemble methods, we temperature scale on validation cross entropy for 5 epochs with batch size 100 and learning rate 0.1, whereas for analytic NTKGP we temperature scale for 1000 epochs on full batch size 6000. Like above, we approximate the analytic NTKGP validation predictions (for temperature scaling) by a Monte Carlo ensemble, this time of size 10. We found the various temperature scaling training hyperparameter considerations here to be unimportant to achieve convergence, due to the fact that the temperature scale is a scalar value.
I.4 CIFAR-10 vs SVHN
Baselearners are Myrtle-10 CNNs with 100 channel width and ReLU activations. We use and set observation noise . Like for MNIST we tune on a small linear scale around . We train using SGD, with momentum parameter 0.9, for 100 epochs and learning rate , which is decayed to after 80 epochs. In the first 5 epochs we raise the learning rate in linear increments from to . We use batch size 125. During training we apply random crops and horizontal flips before standardisation. We do not compare to the analytic NTKGP for the Myrtle-10 CNN due to resource requirements.
Figure 8 displays entropy histograms for ensembles trained on CIFAR-10 and tested on in distribution CIFAR-10 test data and out-of-distribution SVHN test data, corresponding to the same experiments as in Figure 5. As we can see, there is a much less noticeable difference between ensemble methods compared to the simpler MNIST vs NotMNIST case.