Tractable Function-Space Variational Inference in Bayesian Neural Networks

Tim G. J. Rudner, Zonghao Chen, Yee Whye Teh, Yarin Gal

Introduction

Machine learning models succeed at an increasingly wide range of narrowly defined tasks (Krizhevsky et al., 2012; Mnih et al., 2013; Silver et al., 2016; Jumper et al., 2021) but may fail without warning when used on inputs that are meaningfully different from the data they were trained on (Amodei et al., 2016; Hendrycks et al., 2021; Rudner and Toner, 2021a, b). To deploy machine learning models in safety-critical environments where failures are costly or may endanger human lives, machine learning methods must be reliable and have the ability to ‘fail gracefully.’ A promising tool for incorporating fail-safe mechanisms into machine learning systems, predictive uncertainty quantification allows machine learning models to express their confidence in the correctness of their predictions.

In this paper, we develop a method for obtaining reliable uncertainty estimates in Bayesian neural networks (bnns, Neal (1996)). While bnns have promised to combine the advantages of deep learning and Bayesian inference, existing approaches for approximate inference in bnns fall short of this promise and have been demonstrated to result in approximate posterior predictive distributions that underperform ‘non-Bayesian’ methods both in terms of predictive accuracy and uncertainty quantification—making them of limited use in practice (Ovadia et al., 2019; Foong et al., 2019; Farquhar et al., 2020a; Band et al., 2021). A potential reason for this shortcoming is that commonly used parameter-space inference methods make it difficult to define meaningful priors that effectively incorporate information about the data-generating process into inference.

To avoid this limitation, we follow Sun et al. (2019) and consider a variational objective defined explicitly in terms of distributions over functions induced by distributions over parameters. In contrast to prior works that rely on approximation techniques that prevent such function-space variational objectives to be used with high-dimensional inputs and highly-overparameterized neural networks, we propose a simple estimator of the Kullback-Leibler divergence between distributions over functions that enables us to perform stochastic variational inference. The proposed estimation procedure allows defining priors that explicitly encourage high predictive uncertainty away from the training data as well as priors that reflect relevant information about the task at hand.

We demonstrate that this approach leads to posterior approximations that exhibit significantly improved predictive uncertainty estimates compared to a wide array of state-of-the-art Bayesian and non-Bayesian methods. Figure 1 shows examples of predictive distributions obtained via function-space variational inference on low-dimensional, easy-to-visualize datasets. As can be seen in the figures, the predictive distributions fit the training data well while also exhibiting a high degree of predictive uncertainty in parts of the input space far away from the training data, as desired.

Contributions. We propose a simple estimation procedure for performing function-space variational inference in bnns. The variational method allows for the incorporation of meaningful prior information about the data-generating process into the inference and produces reliable predictive uncertainty estimates. We perform a thorough empirical evaluation in which we compare the proposed approach to a wide array of competing methods and show that it consistently results in high predictive performance and reliable predictive uncertainty estimates, outperforming other methods in terms of predictive accuracy, robustness to distribution shifts, and uncertainty-based detection of distributionally-shifted data samples. We evaluate the proposed method on standard benchmarking datasets as well as on a safety-critical medical diagnosis task in which reliable uncertainty estimation is essential.Our code can be accessed at https://github.com/timrudner/FSVI.

Preliminaries

Instead of seeking to infer an approximate posterior distribution over parameters, we frame variational inference in stochastic neural networks as inferring an approximation to the posterior distribution over functions pf(⋅ ;Θ)∣Dp_{f(\cdot\,;\bm{\Theta})|\mathcal{D}} induced by the posterior distribution over parameters pΘ∣Dp_{\bm{\Theta}|\mathcal{D}}, that is,

where δ(⋅)\delta(\cdot) is the Dirac delta function (Wolpert, 1993). Considering the prior distribution over functions pf(⋅ ;Θ)p_{f(\cdot\,;\bm{\Theta})} induced by a prior distribution over parameters pΘp_{\bm{\Theta}},

and the variational distribution over functions qf(⋅ ;Θ)q_{f(\cdot\,;\bm{\Theta})} induced by a variational distribution over parameters qΘq_{\bm{\Theta}},

we can express the problem of finding a posterior distribution over functions variationally as

which allows us to effectively incorporate meaningful prior information about the underlying data-generating process into training. As discussed by Burt et al. (2021), this variational objective is guaranteed to be well-defined for suitably chosen prior distributions over functions. Specifically, the KL divergence between two distributions over functions generated from different distributions over parameters applied to the same mapping (e.g., the same neural network architecture) is well-defined (i.e., finite) if the KL divergence between the distributions over parameters is finite, since, by the strong data processing inequality (Polyanskiy and Wu, 2017),

Hence, for a likelihood function defined on a finite set of training targets yD\mathbf{y}_{\mathcal{D}} and a suitably defined prior distribution over functions, we can express the variational problem above equivalently as the well-defined maximization problem max⁡qΘ∈QθF(qΘ)\max_{q_{\bm{\Theta}}\in\mathcal{Q}_{{\bm{\theta}}}}\mathcal{F}(q_{\bm{\Theta}}) with

In the next section, we will describe an approximation and estimation procedure that allows scaling function-space variational inference to large neural networks and high-dimensional input data.

Deriving a Tractable Function-Space Variational Objective

The primary obstacle to computing the objective in Equation 6 is the KL divergence from qf(⋅;Θ)q_{f(\cdot;\bm{\Theta})} to pf(⋅;Θ)p_{f(\cdot;\bm{\Theta})}. There are two reasons why the KL divergence in Equation 7 is intractable: First, for bnns or other non-linear models, we do not have access to the probability density functions of the multivariate distributions qf(X;Θ)q_{f(\mathbf{X};\bm{\Theta})} and pf(X;Θ)p_{f(\mathbf{X};\bm{\Theta})}; second, for all but extremely simple input spaces, we are unable to compute the supremum over all possible finite sets of evaluation points. In the remainder of this section, we outline an approach for obtaining an estimator of a locally accurate approximation to the KL divergence that allows for scalable gradient-based optimization of Equation 7.

We first approach the problem of computing the KL divergence between two bnns evaluated at a finite set of points. To do so, we first derive tractable approximations to the distributions over functions qf(X;Θ)q_{f(\mathbf{X};\bm{\Theta})} and pf(X;Θ)p_{f(\mathbf{X};\bm{\Theta})} Next, we show that under these approximations, we are able to obtain a closed-form approximation to the KL divergence and describe a simple Monte Carlo estimator of the supremum in the function-space KL divergence.

To obtain an approximation to the probability distributions of qf(X;Θ)q_{f(\mathbf{X};\bm{\Theta})} and pf(X;Θ)p_{f(\mathbf{X};\bm{\Theta})}, we use a first-order Taylor expansion of the mapping ff about the mean parameters of qΘq_{\bm{\Theta}} and pΘp_{\bm{\Theta}}, respectively, and derive the induced distributions under the linearized mapping.

2 Approximating the Function-Space Kullback-Leibler Divergence

Using this approximation, we obtain an estimator of the variational objective given by

where XCS =˙ {XC(i)}i=1S\mathcal{X}_{\mathcal{C}}^{S}\,\dot{=}\,\{\mathbf{X}_{\mathcal{C}}^{(i)}\}_{i=1}^{S} is a collection of SS sets of context points XC(i) =˙ {x(j)}j=1K\mathbf{X}_{\mathcal{C}}^{(i)}\,\dot{=}\,\{\mathbf{x}^{(j)}\}_{j=1}^{K} jointly sampled from a context distribution pXCp_{\mathcal{X}_{\mathcal{C}}}. Each context set XC(i)\mathbf{X}_{\mathcal{C}}^{(i)} can be viewed as a single Monte Carlo sample from the input space so that the estimator G^(XCS)\hat{G}(\mathcal{X}_{\mathcal{C}}^{S}) provides an SS-sample Monte Carlo estimate of the supremum. While this estimator is crude and only provides a rough approximation to the true supremum, it encourages the variational distribution over functions to match the prior distribution over functions on the sets of context points. The choice of the context distribution pXCp_{\mathcal{X}_{\mathcal{C}}} can be informed by knowledge about the prediction task and should be viewed as a problem-specific modeling choice. Similarly, the numbers of samples SS and KK are hyperparameters to be optimized with a validation set. For details on how pXCp_{\mathcal{X}_{\mathcal{C}}} is chosen for the empirical evaluation in Section 5, see Appendix D.

3 Stochastic Estimation of the Approximate Function-Space Variational Objective

Let qΘq_{\bm{\Theta}} be a Gaussian mean-field variational distribution, let pΘp_{\bm{\Theta}} be an isotropic Gaussian prior, let (XB,yB)(\mathbf{X}_{\mathcal{B}},\mathbf{y}_{\mathcal{B}}) be a mini-batch of the training data, and reparameterize Θ\bm{\Theta} as Θ^(μ,Σ,ϵ(j)) =˙ μ+Σ⊙ϵ(j)\hat{\bm{\Theta}}({\bm{\mu}},\bm{\Sigma},\bm{\epsilon}^{(j)})\,\dot{=}\,{\bm{\mu}}+\bm{\Sigma}\odot\bm{\epsilon}^{(j)}. Using the estimator G^(XCS)\hat{G}(\mathcal{X}_{\mathcal{C}}^{S}) defined above and estimating the expected log-likelihood via Monte Carlo sampling, we obtain a Monte Carlo estimator for the function-space variational objective:

with ϵ(j)∼N(0,IP)\bm{\epsilon}^{(j)}\sim\mathcal{N}(\mathbf{0},\mathbf{I}_{P}) and XCS\mathcal{X}_{\mathcal{C}}^{S} as defined above. This Monte Carlo estimator is biased due to the linearization and context-set approximations but allows for scalable gradient-based stochastic optimization.

Selection of Prior. For all experiments that involve uncertainty quantification, we chose a prior distribution over parameters that induces a prior distribution over functions pf(⋅;Θ)p_{f(\cdot;\bm{\Theta})} and a prior predictive distribution that exhibits a high degree of predictive uncertainty at evaluation points from regions in input space where pXCp_{\mathcal{X}_{\mathcal{C}}} has non-zero support and, under smoothness constraints, on evaluation points in nearby regions. For settings where prior information is encoded in data—for example, in the form of expert demonstrations of robotic manipulation tasks (Rudner et al., 2021) or in the form of pre-trained networks in continual or transfer learning (Rudner et al., 2022)—an empirical prior that reflects this information can be specified. For further details, see Appendix D.

Selection of Context Distribution. The distribution pXCp_{\mathcal{X}_{\mathcal{C}}} allows us to incorporate information about the data-generating process into training and encourage the variational distribution to match the prior over functions in relevant parts of the input space. By taking advantage of the abundance of data available in real-world settings, context distributions can be constructed from large datasets like ImageNet (Krizhevsky et al., 2012), from small but diverse datasets like CIFAR-100, or by using any set of task-related unlabeled data. In our experiments, we choose two types of context distributions. One of the context distributions is constructed from the training data and only contains randomly sampled monochrome images, and one is constructed from a real-world dataset generated from a data distribution related to that of the training data. For example, when training on FashionMNIST, we use KMNIST as the context distribution, and when training on CIFAR-10, we use CIFAR-100 as the context distribution. For further details, see Appendix D.

Posterior Predictive Distribution. After optimizing the variational objective with respect to the parameters of the variational distribution qΘq_{\bm{\Theta}}, we use the fact that we can obtain function draws by sampling from the distribution over parameters to obtain an approximate posterior predictive distribution

where M∗M_{\ast} is the number of Monte Carlo samples used to estimate the predictive distribution.

Related Work

There is a growing body of work on function-space approaches to inference in bnns, deep learning, and applications such as continual learning (Benjamin et al., 2019; Sun et al., 2019; Titsias et al., 2020; Burt et al., 2021; Pan et al., 2020; Ma and Hernández-Lobato, 2021; Rudner et al., 2022).

Previously proposed methods for fsvi in bnns are based on approximate gradient estimators and either replace the supremum in Equation 7 with an expectation (Sun et al., 2019) or do not define an explicit variational objective (Wang et al., 2019). Sun et al. (2019) and Carvalho et al. (2020) use Gaussian process priors over functions for which the function-space variational inference problem is not well-defined (see Section 2.1 and Burt et al. (2021)). More recent work has attempted to circumvent the intractability of the variational objective in Equation 6 by proposing alternative objectives for function-space inference in bnns (Ma et al., 2019; Ober and Aitchison, 2020; Ma and Hernández-Lobato, 2021). Rudner et al. (2022) extend the approach presented in Section 3 to sequential inference problems and apply it to continual learning.

Linear Models.

Immer et al. (2020) and Khan et al. (2019) show that approximate bnn posterior distribution via the Laplace and Generalized-Gauss-Newton approximation corresponds to exact posteriors under linearizations of different models. Unlike in our approach, they use a Laplace approximation and do not perform variational inference and do not optimize the variance parameters. Furthermore, Immer et al. (2020) and Khan et al. (2019) use a neural network model to obtain a parameter maximum a posteriori estimate, but then use a linearization of the neural network model to compute a posterior predictive distribution. In contrast, our work only uses the linearization to obtain an estimator of the variational objective but uses the unlinearized model to construct a posterior predictive distribution.

Pathologies of Variational Inference in Bayesian Neural Networks.

Burt et al. (2021) consider the function-space variational objective in Equation 6 and show that the KL divergence between bnns with different networks architectures are not well-defined. A parallel line of research showed that posterior predictive distributions of shallow bnns with mean-field variational distributions have a limited ability to represent complex covariance structures in function space (Foong et al., 2019, 2020) but that deep bnns do not suffer from this limitation (Farquhar et al., 2020b). Our results are consistent with the findings of Farquhar et al. (2020b) that mean-field variational distributions are able to represent complex covariance structures in function space.

Empirical Evaluation

In this section, we evaluate fsvi on high-dimensional classification tasks that were out of reach for function-space variational inference methods proposed in prior works and compare fsvi to several well-established and state-of-the-art Bayesian deep learning and deterministic uncertainty quantification methods. We show that fsvi (sometimes significantly) outperforms existing Bayesian and non-Bayesian methods in terms of their in-distribution uncertainty calibration and out-of-distribution predictive uncertainty estimation. For a details on models, training and validation procedures, and datasets used, see Appendix D. For a comparison to Sun et al. (2019) on small-scale regression tasks, see Section B.2.

In this set of experiments, we assess the reliability of the uncertainty estimates generated by fsvi. If a bnn trained via fsvi is able to perform reliable uncertainty estimation, its predictive uncertainty will be significantly higher on input points that were generated according to a different data-generating distribution than the training data. For models trained on the FashionMNIST dataset, we use the MNIST and NotMNIST datasets as out-of-distribution evaluation points, while for models trained on the CIFAR-10 dataset, we use the SVHN dataset as out-of-distribution evaluation points.

For models trained on either FashionMNIST or CIFAR-10, we evaluate their in-distribution performance in terms of test accuracy, test log-likelihood, and test calibration. To evaluate the quality of different models’ uncertainty estimates, we compute uncertainty estimates for the pairs FashionMNIST/MNIST, FashionMNIST/NotMNIST, and CIFAR-10/SVHN to and measure for a range of thresholds how well the datasets in each pair can be separated solely based on the uncertainty estimates. This experiment setup follows prior work by van Amersfoort et al. (2020) and Immer et al. (2020). We report the area under the receiver operating characteristic (ROC) curve in Tables 1 and 2.

Predictive Performance and Calibration. To assess in-distribution predictive performance and calibration, we report the test accuracy, negative log-likelihood (NLL), and expected calibration error (ECE) for models trained on FashionMNIST and CIFAR-10 in Tables 1 and 2. On both FashionMNIST and CIFAR-10, fsvi achieves the lowest NLL and either the best or second-best predictive accuracy and ECE, respectively, across all methods. Notably, fsvi significantly outperforms spg (Ma and Hernández-Lobato, 2021), an alternative function-space variational inference method.

Predictive Uncertainty under Distribution Shift. In Tables 1 and 2, we report evaluation metrics that elucidate the reliability of different methods’ predictive uncertainty under distribution shift. fsvi exhibits reliable predictive uncertainty estimates that allow distinguishing between in- and out-of-distribution inputs with high accuracy. As would be expected, we observe that using context distributions that reflect our knowledge about the data-generating process can significantly improve uncertainty quantification under fsvi. For the FashionMNIST experiment, we used the KMNIST dataset, which contains grayscale images of Kuzushiji letters, and for the CIFAR-10 experiment, we used the CIFAR-100 dataset, which contains RGB images of 100 classes. Both KMNIST and CIFAR-100 differ from the OOD datasets (MNIST and NotMNIST and SVHN, respectively) used to compute OOD-AUROC metrics in Tables 1 and 2, but using them as context distributions significantly increased the ability of bnns trained via fsvi to identify distributionally shifted samples. Since the variational objective encourages matching the prior (which we chose to have high variance) on samples from the context distribution can improve uncertainty estimation in regions of the input space far from the training data.

2 Generalization and Reliability of Predictive Uncertainty under Distribution Shift

To assess the reliability of predictive models in deep learning, Ovadia et al. (2019) propose the following desiderata: In order for a model to be considered reliable, it ought to (i) exhibit low predictive uncertainty on training data and high predictive uncertainty on out-of-distribution inputs, (ii) generate predictive uncertainty estimates that allow distinguishing in- from out-of-distribution inputs, and (iii) if possible, maintain high predictive accuracy even under distribution shift. Models that satisfy these desiderata are less likely to make poor, high-confidence predictions and more amenable for use in safety-critical downstream tasks.

To illustrate these desiderata, we follow Ovadia et al. (2019) and consider the rotated MNIST task, where a model is trained on MNIST and evaluated on rotated MNIST digits. The goal is to maintain a high level of predictive accuracy (measured in terms of Brier scores) while exhibiting an increasing level of predictive uncertainty on distribution shifts of increasing magnitude. Figure 2 shows Brier scores (lower is better) and predictive entropy estimates (higher means more uncertain) of four different models. As rotating the MNIST digits gradually shifts the data distributions, we would expect Brier scores to increase (corresponding to worse predictive accuracy) as the rotation angle increases. A model with reliable predictive entropy estimates would only experience a small decrease under distribution shift while exhibiting a large increase in predictive uncertainty. As can be seen in the plot, the Brier scores of fsvi decreases the least, while fsvi’s uncertainty is significantly higher than other models’. To assess the reliability of different uncertainty quantification methods on a more challenging distribution-shift task, we consider corrupted CIFAR-10 inputs under the second-mildest corruption level used in (Ovadia et al., 2019) and report our results in Table 2. Consistent with the rotated MNIST results, fsvi achieves the highest accuracy on the corrupted data.

3 Safety-Critical Uncertainty-Aware Selective Prediction: Diabetic Retinopathy Diagnosis

To evaluate the reliability of the predictive uncertainty of fsvi in a real-world safety-critical setting, we consider the task of diagnosing diabetic retinopathy (DR), a medical condition that can lead to impaired vision, from retina scans (Leibig et al., 2017; Filos et al., 2019; Band et al., 2021). We use two publicly available datasets, EyePACS (2015) and APTOS (2019), each containing RGB images of a human retina graded by a medical expert on the following scale: 0 (no DR), 1 (mild DR), 2 (moderate DR), 3 (severe DR), and 4 (proliferative DR). The Kaggle dataset was collected from patients in the United States, while the APTOS dataset was collected from patients in India using cheaper but more modern scanning devices. We follow Leibig et al. (2017), Filos et al. (2019), and Band et al. (2021) and binarize all examples from both the EyePACS and APTOS datasets by dividing the classes up into sight-threatening diabetic retinopathy—defined as moderate diabetic retinopathy or worse (classes {2,3,4}\{2,3,4\})—and non-sight-threatening diabetic retinopathy—defined as no or mild diabetic retinopathy (classes {0,1}\{0,1\}). This results in a binary prediction task.

To assess the reliability of predictive models when medical training and test data are obtained from different patient populations or collected with the same medical equipment, we follow Band et al. (2021) and use the Kaggle dataset for training and the distributionally shifted APTOS dataset for evaluation. The results are shown in Figure 4, which plot the ROC curves for the binary prediction problems as well as the area under the ROC curve for an uncertainty aware selective prediction task. For further details about the uncertainty-aware selective prediction evaluation protocol, see Section D.4. Figure 4 shows that fsvi performs well on all four tasks and is only outperformed by mc dropout. For full tabular results, see Section B.1.

Conclusion

The paper proposed a scalable and effective approach to function-space variational inference in bnns. We demonstrated that the proposed estimator of the function-space variational objective can be scaled up to high-dimensional data and large neural network architectures and that fsvi exhibits consistently reliable in- and out-of-distribution predictive performance on a wide range of datasets when compared to well-established and state-of-the-art uncertainty quantification methods. We hope that this work will lead to further research into function-space variational inference and the development of more sophisticated data-driven prior distributions over functions.

Acknowledgements

We thank Bryn Elesedy, Bobby He, and Andrew Jesson for feedback on an early draft of this paper. We thank Joost van Amersfoort for helpful discussions about experiment design and implementations. Tim G. J. Rudner is funded by the Rhodes Trust and the Engineering and Physical Sciences Research Council (EPSRC). We gratefully acknowledge donations of computing resources by the Alan Turing Institute.

References

Appendix

Table of Contents

Appendix A Proofs & Derivations

This proof follows steps from Matthews et al. . Consider measures P^\hat{P} and PP both of which define distributions over some function ff, indexed by an infinite index set XX. Let D\mathcal{D} be a dataset and let XD\mathbf{X}_{\mathcal{D}} denote a set of inputs and yD\mathbf{y}_{\mathcal{D}} a set of targets. Consider the measure-theoretic version of Bayes’ Theorem [Schervish, 1995]:

where pX(Y ∣ f)p_{X}(Y\,|\,f) is the likelihood and p(Y)=∫XpX(Y ∣ f)dP(f)p(Y)=\int_{{}^{X}}p_{X}(Y\,|\,f)dP(f) is the marginal likelihood. We assume that the likelihood function is evaluated at a finite subset of the index set XX. Denote by πC:X→C\pi_{C}:{}^{X}\to{}^{C} a projection function that takes a function and returns the same function, evaluated at a finite set of points CC, so we can write

and similarly, the marginal likelihood becomes p(yD)=∫py∣fX(yD ∣ fXD) dPXD(fXD)p(\mathbf{y}_{\mathcal{D}})=\int p_{\mathbf{y}|f_{\mathbf{X}}}(\mathbf{y}_{\mathcal{D}}\,|\,f_{\mathbf{X}_{\mathcal{D}}})\,\textrm{d}P_{\mathbf{X}_{\mathcal{D}}}(f_{\mathbf{X}_{\mathcal{D}}}). Now, considering the measure-theoretic version of the KL divergence between an approximating stochastic process QQ and a posterior stochastic process P^\hat{P}, we can write

where PP is some prior stochastic process. Now, we can apply the measure-theoretic Bayes’ Theorem to obtain

where dQπdPπ(f)\frac{dQ^{\pi}}{dP^{\pi}}(f) is marginally consistent given the projection π\pi. Rearranging, we can get

Finally, this lower bound can equivalently be expressed as

where X\D\mathbf{X}_{\backslash\mathcal{D}} is an infinite index set excluding the finite index set XD\mathbf{X}_{\mathcal{D}}, that is, X\D∩XD=∅\mathbf{X}_{\backslash\mathcal{D}}\cap\mathbf{X}_{\mathcal{D}}=\varnothing, or by Theorem 1 in Sun et al. , we can write

A.2 Distribution under Linearized Function Mapping

where the last line follows from the definition of gΘg_{\bm{\Theta}}. By definition of the covariance, we then obtain

With this result, we obtain the covariance function

For a stochastic function f(⋅ ;Θ)f(\cdot\,;\bm{\Theta}) defined in terms of stochastic parameters Θ\bm{\Theta} distributed according to distribution gΘ=N(m,S)g_{\bm{\Theta}}=\mathcal{N}(\mathbf{m},\mathbf{S}), denote the linearization of the stochastic function f(⋅ ;Θ)f(\cdot\,;\bm{\Theta}) about m\mathbf{m} by

where gΘα=N(mα,Sα)g_{\bm{\Theta}_{\alpha}}=\mathcal{N}(\mathbf{m}_{\alpha},\mathbf{S}_{\alpha}), gΘβ=N(mβ,Sβ)g_{\bm{\Theta}_{\beta}}=\mathcal{N}(\mathbf{m}_{\beta},\mathbf{S}_{\beta}), and

Consider a partition of the set of parameters into sets α\alpha and β\beta and express the linearized mapping as

where Jα(⋅ ;m)\mathcal{J}_{\alpha}(\cdot\,;\mathbf{m}) and Jβ(⋅ ;m)\mathcal{J}_{\beta}(\cdot\,;\mathbf{m}) are the columns of the Jacobian matrix corresponding to the sets of parameters α\alpha and β\beta, respectively, and Θα\bm{\Theta}_{\alpha} and Θβ\bm{\Theta}_{\beta} are the corresponding random parameter vectors.

Appendix B Further Empirical Results

The results below were reproduced from Band et al. using the retina benchmark.

B.2 UCI Regression

Appendix C Illustrative Examples

C.2 Synthetic 1D Regression Datasets

Appendix D Implementation, Training, and Evaluation Details

For fsvi, we used a holdout validation set (10% of the training set) to conduct a hyperparameter search over the prior variance, the number of context points used to evaluate the KL divergence, the context distribution, and the number of Monte Carlo samples used to evaluate the expected log-likelihood. We selected the set of hyperparameters that yielded the highest validation log-likelihood for all experiments. We state the hyperparameters selected for the different datasets below.

For other methods, we used a holdout validation set of the same size and selected the best-performing hyperparameters. We used implementations provided by the authors of mfvi (radial) and swag. All other methods were implemented from scratch unless stated otherwise.

D.2 FashionMNIST vs. MNIST/NotMNIST

We train all model on the FashionMNIST dataset and evaluate the models’ predictive uncertainty performance on out-of-distribution data on the MNIST dataset. Both datasets consist of images of size 28×2828\times 28 pixels. The FashionMNIST dataset is normalized to have zero mean and a standard deviation of one. The MNIST dataset is normalized with the same transformation, that is, using the same mean and standard deviation used for the in-distribution data. We chose FashionMNIST/MNIST instead of MNIST/NotMNIST because the latter is notably easier than the former.

In this experiment, a network architecture with two convolutional layers of 32 and 64 3×33\times 3 filters and a fully-connected final layer of 128 hidden units is used. A max pooling operation is placed after each convolutional layer and ReLU activations are used. We do not use batch normalization. All models are trained for 30 epochs with a mini-batch size of 128 using SGD with a learning rate of 5×10−35\times 10^{-3}, momentum (with momentum parameter 0.9), and a cosine learning rate schedule with parameter 0.050.05.

For fsvi with pXC=p_{\mathbf{X}_{\mathcal{C}}}=random monochrome, we sampled 50% of the context points for each gradient step from the mini-batch and the other 50% according to the method described in Section D.8. For fsvi with pXCp_{\mathbf{X}_{\mathcal{C}}}= KNIST, we used the KMNIST dataset.

D.3 CIFAR-10 vs. SVHN

We train all model on the CIFAR-10 dataset and evaluate the models’ predictive uncertainty performance on out-of-distribution data on the SVHN dataset. Both datasets consist of images of size 32×32×332\times 32\times 3, with RBG channels. The CIFAR-10 dataset is normalized to have zero mean and a standard deviation of one. The SVHN dataset is normalized with the same transformation, that is, using the same mean and standard deviation used for the in-distribution data. The training data is augmented with random horizontal flips (with a probability of 0.5) and random crops (4 zero pixels on all sides).

In this experiment, a standard ResNet-18 network architecture was used. All models are trained for 200 epochs with a mini-batch size of 128 using SGD with a learning rate of 5×10−35\times 10^{-3}, momentum (with momentum parameter 0.9), and a cosine learning rate schedule with parameter 0.050.05.

For fsvi with pXC=p_{\mathbf{X}_{\mathcal{C}}}=random monochrome, we sampled 100% of the context points for each gradient step from the mini-batch and the other 50% according to the method described in Section D.8. For fsvi with pXCp_{\mathbf{X}_{\mathcal{C}}}= CIFAR-100, we used the CIFAR-100 dataset.

D.4 Diabetic Retinopathy Diagnosis

In real-world settings where the evaluation data may be sampled from a shifted distribution, incorrect predictions may become increasingly likely. To account for that possibility, predictive uncertainty estimates can be used to identify datapoints where the likelihood of an incorrect prediction is particularly high and refer them for further review. We consider a corresponding selective prediction task, where the predictive performance of a given model is evaluated for varying expert referral rates. That is, for a given referral rate of γ∈\gamma\in, a model’s predictive uncertainty is used to identify the γ\gamma proportion of images in the evaluation set for which the model’s predictions are most uncertain. Those images are referred to a medical professional for further review, and the model is assessed on its predictions on the remaining (1−γ)(1-\gamma) proportion of images. By repeating this process for all possible referral rates and assessing the model’s predictive performance on the retained images, we estimate how reliable it would be in a safety-critical downstream task, where predictive uncertainty estimates are used in conjunction with human expertise to avoid harmful predictions. Importantly, selective prediction tolerates out-of-distribution examples. For example, even if unfamiliar features appear in certain images, a model with reliable uncertainty estimates will perform better in selective prediction by assigning these images high epistemic (and predictive) uncertainty, therefore referring them to an expert at a lower γ\gamma.

For all methods, experiments are performed using a ResNet-50 network architecture. Training and evaluation scripts as well as model checkpoints can be found at

github.com/google/uncertainty-baselines/.../diabetic_retinopathy_detection.

D.5 Two Moons

In this experiment, we use a multi-layer perceptron (MLP) consisting of two fully-connected layers with 30 hidden units each and tanh activations. We train all models with a learning rate of 10−310^{-3}.

For fsvi, we sampled context points uniformly from ×\times.

D.6 1D Regression

In this experiment, we use a multi-layer perceptron (MLP) consisting of two fully-connected layers with 100 hidden units each and ReLU activations.

For fsvi, we sampled context points uniformly from $$.

D.7 Further Implementation Details

We use the Adam optimizer with default settings of β1=0.9\beta_{1}=0.9, β2=0.99\beta_{2}=0.99 and ϵ=10−8\epsilon=10^{-8} for all experiments. The deterministic neural networks that were used for the ensemble were trained with a weight decay of λ\lambda = 1e-1. mfvi (tempered) was trained with a KL scaling factor of 0.1 to obtain a cold posterior.

D.8 Selection of Context Distribution

We estimate the supremum at every gradient step by sampling a set of context points XC\mathbf{X}_{\mathcal{C}} from a distribution pXCp_{\mathbf{X}_{\mathcal{C}}} at every gradient step. For tasks with image inputs, we construct a distribution pXCp_{\mathbf{X}_{\mathcal{C}}}, defined as a uniform distribution over images with monochromatic channels. To generate a sample from this “monochrome images” distribution, we first take all images in the training data, flatten each channel, and stack the flattened image channels into a single vector each. We then draw a random element (i.e., a pixel) from each channel vector and then use these pixels to generate a monochrome image of a given resolution by setting every channel equal to the value of the pixel that was drawn. For regression tasks with a DD-dimensional input space, pXCp_{\mathbf{X}_{\mathcal{C}}} is defined as a uniform distribution with lower and upper bounds set to the empirical lower and upper bounds of the training data. For further details on the effect of different sampling schemes on the posterior predictive distribution’s performance, see Appendix B.

D.9 Compute Resources

All experiments were carried out on an Nvidia V-100 GPU with 32GB of memory.