Benchmarking Bayesian Deep Learning on Diabetic Retinopathy Detection Tasks

Neil Band, Tim G. J. Rudner, Qixuan Feng, Angelos Filos, Zachary Nado, Michael W. Dusenberry, Ghassen Jerfel, Dustin Tran, Yarin Gal

Introduction

Bayesian deep learning has been applied successfully to a wide range of real-world prediction problems such as medical diagnosis , computer vision , scientific discovery , and autonomous driving .

Despite the demonstrated usefulness of Bayesian deep learning for such practical applications and a growing literature on inference methods , there exists no standardized benchmarking task that reflects the complexities and challenges of safety-critical real-world tasks while adequately accounting for the reliability of models’ predictive uncertainty estimates.

To make meaningful progress in the development and successful deployment of reliable Bayesian deep learning methods, we need easy-to-use benchmarking tasks that reflect the real world and hence serve as a legitimate litmus test for practitioners that aim to deploy their models in safety-critical settings. Further, such tasks ought to be usable without the extensive domain expertise often necessary for appropriate experiment design and data preprocessing. Lastly, any such benchmarking task must include evaluation methods that test for predictive performance and assess different properties of models’ predictive uncertainty estimates, while taking into account application-specific constraints.

In this paper, we propose a set of realistic safety-critical downstream tasks that respect these desiderata and use them to benchmark well-established and state-of-the-art Bayesian deep learning methods. To do so, we consider the problem of using machine learning to detect diabetic retinopathy, a medical condition considered the leading cause of vision impairment and blindness . Unlike in prior works on diabetic retinopathy detection, the benchmarking tasks presented in this paper are specifically designed to assess the reliability of machine learning models and the quality of their predictive uncertainty estimates using both aleatoric and epistemic uncertainty estimates.

Medical diagnosis problems are particularly well-suited to assess reliability due to the severe harm caused by predictive models that make confident but poor predictions (for example, when a disease is not recognized). As a general desideratum, we want a model’s predictive uncertainty to correlate with the correctness of its predictions. Good predictive uncertainty estimates can be a fail-safe against incorrect predictions. If a given data point might result in an incorrect prediction because it is meaningfully different from data in the training set—for example, because it shows signs of the disease not captured there, exhibits visual artifacts, or was obtained using different measurement devices—a good predictive model will express a high level of predictive uncertainty and flag the example for further review by a medical expert.

Contributions. We present the Retina Benchmark: an easy-to-use, expert-guided, open-source suite of diabetic retinopathy detection benchmarking tasks for Bayesian deep learning. In particular, we design safety-critical downstream tasks from publicly available datasets. On these downstream tasks, we evaluate well-established and state-of-the-art Bayesian and non-Bayesian methods on a set of task-specific reliability and performance metrics. Lastly, we provide a modular and extensible implementation of the benchmarking tasks and methods, as well as pre-trained models obtained from an extensive hyperparameter optimization over more than 400 total configurations and evaluation, using over 100 TPU days and 20 GPU days of compute. Code to reproduce our results and benchmark new methods is available at:

Downstream Benchmarking Tasks for Diabetic Retinopathy Detection

In this section, we present two real-world scenarios in diabetic retinopathy detection and describe how we merge two publicly available datasets to design corresponding prediction tasks.

EyePACS Dataset. We construct training datasets for different tasks from the EyePACS dataset, previously used for the Kaggle Diabetic Retinopathy Detection Challenge . It contains high-resolution labeled images of human retinas exhibiting varying degrees of diabetic retinopathy. The dataset consists of 35,126 training, 10,906 validation, and 42,670 test images, each an RGB image of a human retina graded by a medical expert on the following scale: 0 (no diabetic retinopathy), 1 (mild diabetic retinopathy), 2 (moderate diabetic retinopathy), 3 (severe diabetic retinopathy), and 4 (proliferative diabetic retinopathy).

APTOS Dataset. To construct tasks that assess model performance under distribution shift, we use the APTOS 2019 Blindness Detection dataset . The dataset also contains labeled images of human retinas exhibiting varying degrees of diabetic retinopathy, but was collected in India, from a different patient population, using different medical equipment. We use 80% of the images (2,929 images) as a test set and the other 20% (733 images) as a secondary validation set. Moreover, the images are significantly noisier than the images in the EyePACS dataset, with distinct visual artifacts (cf. Figure 7, Section A.8). Each image was graded on the same 0-to-4 scale as the EyePACS dataset.

Prediction Targets. We follow Leibig et al. 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\}). By international guidelines, this is the threshold at which a case should be referred to an ophthalmologist . Example EyePACS retina images from the two classes are shown in Figure 1. Reflecting real-world challenges, the datasets are unbalanced—e.g., for EyePACS, only 19.6%19.6\% of the training set and 19.2%19.2\% of the test set have a positive label—and images have visual artifacts and noisy labels (some labels are incorrect).

Data Preprocessing. Data preprocessing on examples from both the EyePACS and APTOS datasets follows the winning entry of the Kaggle Challenge : Images are rescaled such that retinas have a radius of 300 pixels, are smoothed using local Gaussian blur, and finally, are clipped to 90% size to remove boundary effects. Examples of original and corresponding processed images are provided in Figure 6 (Section A.8). We conduct an empirical study investigating how varying the strength of the Gaussian blur smoothing affects downstream performance and uncertainty quality in Section B.6.

2 Diabetic Retinopathy Detection under Severity Shift

Diabetes and diabetes-related illnesses such as diabetic retinopathy are becoming widespread. Yet cases of sight-threatening diabetic retinopathy are still relatively rare, and scans of retinas exhibiting signs of no or mild diabetic retinopathy are more easily obtainable. As a result, predictive models for detecting diabetic retinopathy may be trained on only a very small number of retina images showing signs of severe or proliferative retinopathy.

We design a prediction task that simulates this setting and allows us to assess the reliability of predictive models when they are evaluated on images that have been assigned a severity higher than any encountered in the training data. Specifically, we train models only on retina images showing signs of at most moderate diabetic retinopathy and evaluate them on retina images showing signs of severe or proliferative diabetic retinopathy. Given that many signs of moderate diabetic retinopathy are similar in appearance to signs of severe or proliferative diabetic retinopathy (just weaker), we would expect a good predictive model to be able to correctly classify the latter, but to exhibit increased predictive uncertainty. There are certain features of diabetic retinopathy progression that are unique to more severe cases, such as vitreous hemorrhage, or bleeding into the vitreous humor . However, we consider uncertainty-aware downstream tasks that tolerate such unfamiliar cases (cf. Section 2.4).

In this Severity Shift task, we partition the EyePACS dataset into a subset containing all retina images labeled as no, mild, or moderate diabetic retinopathy (original classes {0,1,2}\{0,1,2\}) and a subset of retina images labeled as severe or proliferative diabetic retinopathy (original classes {3,4}\{3,4\}). Next, the samples in each subset are binarized (cf. Section 2.1): The subset of retina images showing signs of at most moderate diabetic retinopathy (subset “moderate”) contains images of binarized classes {0,1}\{0,1\}; and the subset of retina images showing signs of severe or proliferative diabetic retinopathy (subset “severe”) only contains the binarized class 1. This results in 33,545 images in the training set, and 40,727 and 3,524 images in the in-domain and distributionally shifted evaluation sets, respectively.

3 Diabetic Retinopathy Detection under Country Shift

Similar to the scarcity of scans of sight-threatening diabetic retinopathy, the availability of retina scans is limited in countries without widespread screening. Hence, a predictive model may be trained on images collected in the United States—where many scans are performed—and used to evaluate scans from another country, where scans are rarer and performed using different medical devices.

We design a prediction task that simulates this setting and allows us to evaluate the reliability of predictive models when the training and test data are not obtained from the same patient population nor collected with the same medical equipment. In this Country Shift task, we train models on retina images from the EyePACS dataset and evaluate them on retina images from the APTOS dataset. We use the entire training and test data provided in the EyePACS dataset and convert the task into binary classification as described in Section 2.1. This results in 35,126 images in the training set, and 42,670 and 2,929 images in the in-domain and distributionally shifted evaluation sets, respectively.

4 Downstream Task: Selective Prediction and Expert Referral

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 as described in Figure 2. 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 τ∈\tau\in, a model’s predictive uncertainty is used to identify the τ\tau 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-\tau) 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 vitreous hemorrhages appear in certain Severity Shift images (cf. Section 2.2), 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 τ\tau. Section A.6 discusses best- and worst-case uncertainty estimates for the selective prediction task.

To assess how well different models’ predictive uncertainty estimates can be used to separate correct from incorrect diagnoses, we perform selective prediction on three different evaluation settings for the prediction problems described in Sections 2.2 and 2.3, to account for the possibility that the evaluation dataset may contain samples from the in-domain distribution, a shifted distribution, or both.

5 Model Diagnostic: Predictive Uncertainty Histograms

We may also investigate how a model’s predictive uncertainty estimates vary with respect to the ground-truth clinical label (0-to-4). For each task (Country or Severity Shift) and each uncertainty quantification method (cf. Section 5), we bin examples by their ground-truth clinical label. Then, for each (task, method, clinical label) tuple, we plot the distribution of predictive uncertainty estimates for correctly and incorrectly predicted examples (in blue and red, respectively). See Section B.1 for further setup details and plots for both tasks. A model that produces reliable uncertainty estimates should assign low predictive uncertainty to examples that it classifies correctly (the blue distribution should have most of its mass near x=0x=0) and high predictive uncertainty to examples that it classifies incorrectly (the red distribution should have its mass concentrated at a higher xx-value).

Related Work

Retina builds on prior works that demonstrated the usefulness of predictive uncertainty estimates in diabetic retinopathy detection and related downstream tasks . We significantly extend the empirical evaluation in Leibig et al. by designing new prediction problems and corresponding safety-critical downstream tasks for diabetic retinopathy detection, benchmarking a wide array of Bayesian deep learning methods, and providing a modular, extensible, and easy-to-use codebase. We also significantly extend Filos et al. (of which this paper is a direct extension; with contributions from some of the authors), which does not consider severity shifts, only compares two variational inference methods, uses an outdated neural network architecture (with only ≈\approx10% of the parameters of the ResNet-50 architecture used in this work), and considers only a small subset of the evaluation procedures included in Retina (cf. Section B.4 for the full set of results).

Previous works have evaluated methods by predictive performance and quality of their predictive uncertainty estimates on curated datasets such as CIFAR-10 and FashionMNIST . Some prior works provide datasets and benchmarks for robustness and uncertainty quantification in real-world settings but have significant shortcomings. Le et al. considers object detection using a real-world dataset but benchmarks only two methods, neither of which can quantify epistemic uncertainty (cf. Section 4), and does not consider distribution shifts. Other works use methods which quantify both epistemic and aleatoric uncertainty, and consider distribution shifts, but use performance metrics which do not assess quality of uncertainty estimates, such as average precision and log-likelihood (cf. Section 6.3). Finally, Koh et al. considers real-world datasets in domain adaptation problems, but restrictively assumes that the training data is composed of multiple training distributions with domain labels, and does not take into account models’ predictive uncertainty.

In contrast, Retina (i) considers real-world safety-critical tasks and accompanying uncertainty-aware metrics in an important application domain, (ii) is composed of large amounts of high-dimensional data (>8080 GB), (iii) compares a larger set of methods than prior works and incorporates both aleatoric and epistemic uncertainty, and is implemented in adherence to the Uncertainty Baselines repositorySee https://github.com/google/uncertainty-baselines. practices for easy future use and extension, making it easier to benchmark other Bayesian deep learning methods not only on the tasks presented but also on a range of other datasets.

Uncertainty Estimation

Predictive models’ total uncertainty can be decomposed into aleatoric and epistemic uncertainty. A model’s aleatoric uncertainty is an estimate of the uncertainty inherent in the data (e.g., due to noisy inputs or targets), whereas a model’s epistemic uncertainty is an estimate of the uncertainty due to constraints on the model (e.g., due to model misspecification) or the training process (e.g., due to convergence to bad local optima) . Optimal uncertainty estimates would be perfectly correlated with the model error. Hence, because both aleatoric and epistemic uncertainty may contribute to an incorrect prediction, total uncertainty is our uncertainty measure of choice. For a model with stochastic parameters Θ\bm{\Theta}, pre-likelihood outputs f(X;Θ)f(\mathbf{X};\bm{\Theta}), and a likelihood function p(y∗ ∣ x∗; θ)p(\mathbf{y}_{\ast}\,|\,\mathbf{x}_{\ast};\,{\bm{\theta}}), the model’s predictive uncertainty can be decomposed as

where the expectation is taken with respect to the distribution over the model parameters, H(⋅)\mathcal{H}(\cdot) is the entropy functional, and I(y∗; Θ)\mathcal{I}(\mathbf{y}_{\ast};\,\bm{\Theta}) is the mutual information between the model parameters and its predictions .

In binary classification settings with classes {0,1}\{0,1\}, the total predictive uncertainty is given by

Methods

Estimating a model’s predictive uncertainty in terms of both aleatoric and epistemic uncertainty requires a distribution over predictive functions. Such a distribution over predictive functions can be obtained by treating the parameters of a neural network as random variables and inferring a posterior distribution p(θ ∣ D)p({\bm{\theta}}\,|\,\mathcal{D})—a distribution over the network parameters conditioned on a set of training data D=(XD,yD)\mathcal{D}=(\mathbf{X}_{\mathcal{D}},\mathbf{y}_{\mathcal{D}})—according to the rules of Bayesian inference. Neural networks with such distributions over the network parameters—referred to as Bayesian neural networks (bnn)—induce distributions over functions that are able to capture both aleatoric and epistemic uncertainty . Unfortunately, computing a posterior distribution over the parameters of a neural network according to the rule of Bayesian inference is analytically intractable and requires the use of approximate inference methods . Below, we describe baseline and state-of-the-art methods for which we implemented standardized and optimized runscripts that are readily extensible for experimentation and deployment in application settings.

2 Variational Inference in Bayesian Neural Networks

Variational inference is an approximate inference method that seeks to sidestep the intractability of exact posterior inference over the network parameters by framing posterior inference as a variational optimization problem. In particular, variational inference in neural networks seeks to find an approximation to the posterior distribution over parameters by solving the optimization problem

where Q\mathcal{Q} is a variational family of distributions and pp is a prior distribution.

Gaussian Mean-Field Variational Inference. If p =˙ pΘp\,\dot{=}\,p_{\bm{\Theta}} and q =˙ qΘq\,\dot{=}\,q_{\bm{\Theta}} are distributions over parameters, Q\mathcal{Q} is the family of mean-field (i.e., fully-factorized) Gaussian distributions, and the prior distribution over parameters pΘp_{\bm{\Theta}} is also a diagonal Gaussian, the resulting variational objective is amenable to stochastic variational inference and can be optimized using stochastic gradient methods . Henceforth, we refer to bnn inference methods that make these variational assumptions as mean-field variational inference. To optimize this objective, the expectation is estimated using Monte Carlo sampling and the network parameters are reparameterized as Θ =˙ μ+σ⊙ϵ\bm{\Theta}\,\dot{=}\,{\bm{\mu}}+\bm{\sigma}\odot\bm{\epsilon} with ϵ∼N(0,I)\bm{\epsilon}\sim\mathcal{N}(\mathbf{0},\mathbf{I}). Throughout, we use the flipout estimator to reduce the variance of the gradient estimates, and temper the Kullback-Leibler divergence term in the variational objective .

Radial-Gaussian Mean-Field Variational Inference. Radial-Gaussian mean-field variational inference uses the same variational objective, prior, and variational distribution as standard Gaussian mean-field variational inference, but uses an alternative gradient estimator to obtain an improved signal-to-noise ratio in the gradient estimates. Specifically, the network parameters are reparameterized as Θ =˙ μ+σ⊙ϵ∣∣ϵ∣∣2⋅∣r∣\bm{\Theta}\,\dot{=}\,{\bm{\mu}}+\bm{\sigma}\odot\frac{\bm{\epsilon}}{||\bm{\epsilon}||_{2}}\cdot|r| with ϵ∼N(0,I)\bm{\epsilon}\sim\mathcal{N}(\mathbf{0},\mathbf{I}) and r∼N(0,1)r\sim\mathcal{N}(0,1).

Function-Space Variational Inference. Rudner et al. proposed a tractable function-space variational objective for Bayesian neural networks. If p =˙ pf([XD,XI];Θ)p\,\dot{=}\,p_{f([\mathbf{X}_{\mathcal{D}},\mathbf{X}_{\mathcal{I}}];\bm{\Theta})} and q =˙ qf([XD,XI];Θ)q\,\dot{=}\,q_{f([\mathbf{X}_{\mathcal{D}},\mathbf{X}_{\mathcal{I}}];\bm{\Theta})} are distributions over functions evaluated at the training inputs XD\mathbf{X}_{\mathcal{D}} and at a set of inducing inputs XI\mathbf{X}_{\mathcal{I}}, Q\mathcal{Q} is the family of distributions over functions induced by some distribution over network parameters, and the Kullback-Leibler divergence between distributions over functions evaluated at [XD,XI][\mathbf{X}_{\mathcal{D}},\mathbf{X}_{\mathcal{I}}] is approximated by a linearization of the neural network mapping, then the resulting variational objective is amenable to stochastic variational inference . In Retina, we define a Gaussian mean-field distribution over the final layer of the neural network and reparameterize the parameters as Θ =˙ μ+σ⊙ϵ\bm{\Theta}\,\dot{=}\,{\bm{\mu}}+\bm{\sigma}\odot\bm{\epsilon} with ϵ∼N(0,I)\bm{\epsilon}\sim\mathcal{N}(\mathbf{0},\mathbf{I}).

Rank-1 Parameterization. Dusenberry et al. propose a rank-1 parameterization of Bayesian neural networks, where each weight matrix involves only a distribution on a rank-1 subspace, that is, each stochastic weight matrix is defined as Wk′=Wk⊙rksk⊤\mathbf{W}^{\prime}_{k}=\mathbf{W}_{k}\odot\mathbf{r}_{k}\mathbf{s}_{k}^{\top}, where Wk\mathbf{W}_{k} is a deterministic set of weights, and rk\mathbf{r}_{k} and sk\mathbf{s}_{k} are random vectors of parameters. Variational distributions over rk\mathbf{r}_{k} and sk\mathbf{s}_{k} and a Dirac delta distribution over Wk\mathbf{W}_{k} for all layers kk are obtained by optimizing a variational objective.

3 Model Ensembling

Deep Ensembles. A deep ensemble is a mixture of multiple independently-trained deterministic neural networks. As such, unlike bnns, deep ensembles do not explicitly infer a distribution over the parameters of a single neural network. Instead, they marginalize over multiple deterministic models to obtain a predictive distribution that captures both aleatoric and epistemic uncertainty. We construct deep ensembles from multiple map neural networks trained with different random seeds.

Ensembles of Bayesian Neural Networks. Ensembles of Bayesian neural networks are mixtures of multiple independently-trained Bayesian neural networks. They can account for the possibility that any individual approximate posterior distribution obtained via variational inference may be a poor approximation to the exact posterior distribution and may hence yield a poor predictive distribution. A common issue in the Bayesian deep learning literature is that ensembles are frequently compared to single models, often due to computational constraints. In Retina, we provide a unified comparison and construct ensembles for all predictive models, including bnns.

Retina Benchmark

Network Architecture. We use a ResNet-50 architecture for all experiments . A sigmoid transformation is applied to the final linear layer of all networks to obtain class probabilities corresponding to the outcomes of the binary classification problems described in Sections 2.2 and 2.3.

Validation Data, Hyperparameter Tuning, and Monte Carlo Estimation. Reliable uncertainty estimation on data points from shifted distributions is the central challenge for Bayesian deep learning methods. In training and evaluating such methods, practitioners must decide how they should choose validation data: specifically, in which settings they would benefit from using “out-of-distribution” data points for hyperparameter tuning. We consider two real-world settings: (i) No distributionally shifted data is available during hyperparameter tuning. This setting reflects scenarios in which practitioners do not know what data or distributional shift they might encounter during deployment and hence cannot make assumptions about it at training time. (ii) Shifted validation data is available for hyperparameter tuning. This setting reflects scenarios in which practitioners may intend to train a model on data collected from one subpopulation and deploy it on data collected from another subpopulation, but are able to acquire a small number of examples from the deployment subpopulation for use in tuning to improve generalization. Prior works on out-of-distribution detection and uncertainty quantification have considered setting (ii), but have not provided a comparative analysis, which would inform practitioners on when they ought to collect shifted validation data for tuning. We rigorously investigate the two settings across downstream tasks in Section B.4. In the main paper, we report results for models tuned under setting (i). Lastly, for all evaluations, we use five Monte Carlo samples per model to estimate predictive means (e.g., the mc dropout ensemble with K=3K=3 ensemble members uses a total of S=15S=15 Monte Carlo samples).

The aim of the Retina Benchmark is to adequately represent the challenges of real-world distributional shift, and rigorously assess the reliability of (Bayesian) uncertainty quantification in deep learning. Our selective prediction downstream tasks demonstrate two real-world use cases:

Tuning Referral Thresholds. On the Severity Shift task, models demonstrate reasonable uncertainty estimates: predictive performance increases monotonically with an increasing referral rate τ\tau. Therefore, practitioners can infer which referral rate will lead to a desired predictive performance, or infer the performance for a predetermined referral rate (i.e., respecting a budget of expert time).

Detecting Low-Quality Predictive Uncertainty. On the Country Shift task, most methods fail: predictive performance on the shifted dataset declines as τ\tau increases, indicating that the quality of uncertainty estimates is no better than random referral. Importantly, this failure is not reflected in the standard performance measure for retinopathy diagnosis, the receiver operating characteristic (ROC) curve —the area under the ROC curve (AUC) is higher on the shifted evaluation dataset than the in-domain dataset—meaning that a practitioner using only AUC might wrongly conclude that these models would perform well as part of an automated diagnosis pipeline (cf. Figure 2) on distributionally shifted data.

For each method, we assess both AUC and accuracy as a function of the referral rate τ\tau, evaluating the models’ predictions for the (1−τ)(1-\tau) proportion of cases on which they are most certain, as indicated by their predictive uncertainty estimates. We additionally examine predictive uncertainty histograms for each task, method, and ground-truth clinical label (cf. Section 2.5, Section B.1) to determine if methods have particularly good or bad uncertainty estimates at particular severity levels. We also investigate other metrics to assess the reliability of models’ uncertainty estimates, including expected calibration error and out-of-distribution detection AUC, in Section B.4.

2 Severity Shift

On the Severity Shift task (Figure 4, Table 2), models are trained on EyePACS images that show signs of at most moderate diabetic retinopathy. We assess their ability to generalize to images showing signs of severe or proliferative retinopathy. Surprisingly, we find that models generalize well from cases with no worse than moderate diabetic retinopathy (in-domain) (Figure 4(a)) to severe cases (Figure 4(b)), improving their AUC under the distribution shift.

Methods Generalize Reasonably Well Under Severity Shift. Reliable predictive uncertainty estimates correlate with predictive error, and therefore we would expect a model’s performance (e.g., measured in terms of accuracy or AUC) to increase as more examples on which the model exhibits high uncertainty are referred to an expert. On both the in-domain and Severity Shift evaluation sets (Figures 4(c) and (d)), models demonstrate reasonable uncertainty in that accuracy monotonically increases as τ\tau increases. This highlights two ways that practitioners may use selective prediction to prepare models for a real-world deployment in the presence of potential distribution shifts. First, given a performance target (e.g., ≥95%\geq 95\% accuracy) the referral curve can be used to determine the minimum τ\tau achieving this target, estimating a medical experts’ workload. Second, for a maximum acceptable referral rate (e.g., a clinic has medical experts to handle referral of τ≤20%\tau\leq 20\% of patients) the referral curve can be used to determine the optimal τ\tau value and the corresponding performance. For monotonically increasing referral curves, the optimal τ\tau is uniquely the maximum acceptable referral rate.

Taking into Account Epistemic Uncertainty Can Improve Reliability. On the Severity Shift task (Figure 4(d)) many models achieve near-perfect accuracy well before all examples have been referred. For example, mc dropout, which incorporates both epistemic and aleatoric uncertainty (cf. Section 4), achieves 100% predictive accuracy near the 50% referral rate—nearly 20% lower than the referral rate at which a deterministic neural network (map), which only represents aleatoric uncertainty, achieves this level of accuracy. Other variational inference methods underperform map, underscoring the importance of continued work on approximate inference in bnns.

Predictive Uncertainty Histograms Identify Harmful Uncertainty Quantification. In Figure 8 (Section B.1), we find that map, rank-1, and mfvi generate worse uncertainty estimates than other methods on the shifted data (labels 3 and 4); many of their incorrect predictions are assigned low predictive uncertainty (i.e., the red distribution is concentrated near 0). These include false negatives with low uncertainty which are particularly dangerous in automated diagnosis settings (cf. Figure 2), as a medical expert would not be able to catch the model’s failure to recognize the condition.

3 Country Shift

In the Country Shift task (Figure 5, Table 1), we consider the performance of models trained on the US EyePACS dataset and evaluated under distributional shift, on the Indian APTOS dataset . The left two plots of Figure 5 present the ROC curves of methods evaluated on the in-domain (a) and Country Shift (b) evaluation datasets. The black dot in Figures 5(a) and (b) denotes the minimum sensitivity–specificity threshold for the deployment of automated diabetic retinopathy diagnosis systems set by the British National Health Service (NHS) . On the in-domain test dataset, only the mc dropout variants meet the NHS standard; on the APTOS dataset, essentially all methods surpass the standard.We investigate this in Section B.5 and find that class proportions do not account for the improved predictive performance on APTOS, implying other contributing factors such as demographics or camera type. Hence, practitioners using only the ROC curve and its AUC (cf. Table 1) might conclude that their model generalizes under the distribution shift although the ROC curve provides no information on the application of uncertainty estimates to real-world scenarios (cf. Figure 2).

Selective Prediction Can Indicate Failures in Uncertainty Estimation. Unlike the ROC curve, the selective prediction metric conveys how a model would perform in an automated diagnosis pipeline in which the reliability of models’ uncertainty estimates directly impacts performance (cf. Figure 2). Recall that if a model generates reliable predictive uncertainty estimates, the AUC should increase as more patients with uncertain predictions are referred for expert review. This mechanism is illustrated well by the application of mfvi to the Country Shift task (Figure 5(d) and Table 1), since the AUC improves from an initial 91.4%91.4\% up to 93.8%93.8\% when referring 50% of the patients, but then deteriorates as the model is forced to refer patients on which it is both certain and correct. In contrast, other models’ AUCs trend downwards; using uncertainty to refer patients actively hurts model performance on this shifted dataset.

Different Prediction Tasks Yield Different Method Rankings. In Figure 5(c), variational inference methods, including mc dropout, fsvi, and deep ensemble, outperform map inference. This highlights that rankings are task-dependent, and underscores the importance of generic evaluation frameworks to enable rapid benchmarking on many tasks.

Conclusions

The deployment of modern machine learning models in safety-critical real-world settings necessitates trust in the reliability of the models’ predictions.

To encourage the development of Bayesian deep learning methods that are capable of generating reliable uncertainty estimates about their predictions, we introduced the Retina Benchmark, a set of safety-critical real-world clinical prediction tasks which highlight various shortcomings of existing uncertainty quantification methods. We demonstrate that by taking into account the quality of predictive uncertainty estimates, selective prediction can help identify whether methods might fail when deployed as part of an automated diagnosis pipeline (cf. Figure 2), whereas standard metrics such as ROC curves cannot.

While no single set of benchmarking tasks is a panacea, we hope that the tasks and evaluation methods presented in Retina will significantly lower the barrier for assessing the reliability of Bayesian deep learning methods on safety-critical real-world prediction tasks.

Acknowledgments and Disclosure of Funding

References

Supplementary Material

Table of Contents

Appendix A Implementation, Training, and Evaluation Details

Reproducibility in machine learning is often hampered by the wide variety of experimental artifacts made available in papers. Perhaps the most common approach is a GitHub dump of experimental code lacking documentation and testing. This common practice fails to enforce a rigorous standard across works: for example, experiment protocol on cross-validation, access to distributionally shifted validation data, and various tweaks in optimization such as learning rate annealing.

The Retina Benchmark is implemented in the open-sourced Uncertainty Baselines repository. All models implemented in this repository conform to explicit design principles intended to facilitate easy extension and reproduction of dataset loading utilities, metrics, and evaluation.

Each model baseline (e.g., map, mc dropout, fsvi) is implemented in its own self-contained experiment pipeline. This minimizes external dependencies, and therefore provides researchers and practitioners an immediate starting point for experimenting with a particular model. For example, Tran et al. extended the RETINA codebase to include experiment pipelines for state-of-the-art Vision Transformers pretrained on the ImageNet-21K dataset. Datasets are implemented as lightweight wrappers around TensorFlow Datasets . Users that wish to extend our benchmark with new datasets (e.g., clinical practitioners that wish to apply our methods on their own diabetic retinopathy tasks) can follow our custom implementation of the APTOS data loader, which constructs the dataset from raw images and a CSV containing metadata, and applies the preprocessing used by the winner of the EyePACS Kaggle competition . Dataset implementation can be found here.https://github.com/google/uncertainty-baselines/tree/main/uncertainty_baselines/datasets

Framework Agnosticity.

Retina is framework-agnostic. For example, fsvi is implemented in JAX, a variant of MC Dropout is in PyTorch (though we use in this work a TensorFlow variant to simplify TPU tuning), and other models in raw TensorFlow . This interoperability means that users can easily incorporate our datasets and evaluation utilities, including an arrangement of robustness and uncertainty metrics such as selective prediction, out-of-distribution detection, and expected calibration error.

Reproducibility.

All models include testing, and all results are reported over multiple seeds. For each method (e.g., mc dropout or mfvi), downstream task (Country and Severity Shift), and tuning assumption (whether or not distributionally shifted validation data is available for tuning), we sweep over at least 32 hyperparameter configurations. Instead of using a domain-specific and limiting tuning framework for this, we simply provide hyperparameters through Python flags, and implement for convenience of the user the ability to specify automatic logging to TensorBoard and Weights & Biases, an increasingly popular deep learning experiment management service .

A.2 Class Imbalance Adjustment

We compensate for the class imbalance discussed in Section 2 by reweighing the cross-entropy portion of each objective function, placing more weight on the minority class based on the relative class frequencies in each mini-batch of MM samples, p(k)mini-batchp(k)_{\text{mini-batch}} :

where kk is the class of sample ii. We also tried using constant class weights, but found that this resulted in lower overall performance.

A.3 Mean-Field Variational Inference Implementation

We employ a set of standard optimizations to improve training stability for the mfvi and radial-mfvi methods. We fix the mean of the prior to that of the variational posterior, which causes the KL term to only penalize the standard deviation of the weight posterior, and not its mean. We use flipout for lower-variance gradients in convolutional layers and the final dense layer , and KL annealing using a cyclical schedule, following . Finally, for radial-mfvi, the prior’s standard deviation is by default set to the He initializer standard deviation 2/fan_in\sqrt{2/\text{fan\_in}} .

A.4 Uncertainty Estimation and Related Work

Some other works consider uncertainty estimation in medical imaging. Wang et al. uses test-time augmentation for uncertainty estimation, but captures only aleatoric uncertainty. considers uncertainty estimation with a Monte Carlo dropout model but does not isolate how their various measures of uncertainty correspond to epistemic or aleatoric uncertainty. None of the above works contribute and open-source tasks designed to emulate real-world distribution shifts, nor do they implement and benchmark a significant number of baseline uncertainty quantification models considering both aleatoric and epistemic uncertainty.

A.5 Receiver Operating Characteristic Curves

The ROC curve (e.g., see Figure 5(a) and (b)) illustrates the diagnostic ability of a binary classification system as a function of the discrimination threshold. The curve is created by plotting the true positive rate (that is, the sensitivity) against the false positive rate (that is, 1−specificity1-\text{specificity}). The quality of the ROC curve can be summarized by the area under the curve, which ranges from 0.50.5 (chance level) to 1.01.0 (perfect classification).

A.6 Selective Prediction

For the purposes of selective prediction, a model with optimal uncertainty estimates on a given dataset would have uncertainty perfectly correlate rank-wise with the model error. For example, the image on which the model has the highest error should be assigned the highest uncertainty, the image with the second highest error should be assigned the second highest uncertainty, and so on. On the other hand, the worst possible uncertainty estimates are random, which would be uninformative to referral.

Finally, we explain in more detail the dip observed at the right side of the referral curves using AUC as the base metric (e.g., Figure 5(c) and (d)). At relatively high referral rates τ\tau, models begin to refer examples on which they are both confident and correct. This results in the referral curve decreasing. At the highest τ\tau values (the last few examples), for many models, nearly all remaining predictions are correct with high certainty, and the AUC increases.

A.7 Hyperparameter Tuning

We provide full tuning details so that users of Retina will be able to reproduce our results.

All tuning scripts across all methods, tasks (Country and Severity Shift), and tuning procedures (on in-domain validation AUC and area under the selective prediction accuracy curve using the joint validation dataset, described in Section B.3) are documented in the Uncertainty Baselines repository.https://github.com/google/uncertainty-baselines/tree/main/baselines/diabetic_retinopathy_detection

We considered model selection for each of the models on each of the two tasks (Country and Severity Shift) using two different validation metrics: in-domain validation AUC, and area under the accuracy referral curve constructed using both in-domain and distributionally shifted validation data. We describe the reasoning behind the latter metric in Section B.3. We used this validation performance to select the best hyperparameter setting and retrained a configuration for each combination of model, task, and validation tuning metric for 66 random seeds. We evaluated single models by averaging performance over those seeds, and evaluated ensembles by randomly sampling ensembles of size 33 without replacement from the 66 available models, and averaging over 66 such ensemble constructions. As described in Section 6.1, for evaluation, we use five Monte Carlo samples per model to estimate predictive means (e.g., the mc dropout ensemble with K=3K=3 ensemble members uses a total of S=15S=15 Monte Carlo samples).

The majority of methods were tuned on TPU v2-8 nodes. mfvi had particularly high memory requirements which required the use of TPU v3-8 nodes to achieve a reasonable batch size and stable training. Evaluation was performed on NVIDIA A100 GPUs with 40 GB memory, though GPUs with standard sizes (e.g., >6 GB) will be sufficient to run evaluation and inference with the models in the benchmark, e.g., using the model checkpoints. Approximately 100 TPU days and 20 GPU days were used collectively across the initial hyperparameter tuning, fine-tuning with selected configurations, and evaluation across the various tasks. Though a significant cost, we hope that our open-sourcing of all code along with hyperparameter sweep details and checkpoints will significantly decrease future consumption of researchers interested in designing deep models for diabetic retinopathy, along with Bayesian deep learning researchers using our configurations to inform their hyperparameter tuning, or our generally applicable evaluation utilities.

A.8 EyePACS and APTOS Input Data Examples

Appendix B Further Empirical Results

In the figures below, predictive uncertainty (cf. Section 4) is displayed as a normalized density for correct (blue) and incorrect (red) predictions. All histograms are normalized and are displayed with the same range on the xx- and yy-axis. Some bars of the histograms are cut off because the plots are zoomed-in along the yy-axis to improve legibility. See Section 2.5 for a description of predictive uncertainty histograms as a model diagnostic tool, including a discussion of the expected behavior of reliable models. See Section 6 for a discussion of the results for single models on the shifted datasets.

B.2 Tuning without Distributionally Shifted Data: Country Shift Accuracy.

We provide referral curves on accuracy for Country Shift with in-domain validation tuning in Figure 14.

B.3 Tuning in the Presence of Distributionally Shifted Data

In prior work in Bayesian deep learning, little emphasis has been placed on the standardization of a training and evaluation protocol; in particular, the assumption of whether a model has access to distributionally shifted validation data for hyperparameter tuning is often changed on an ad-hoc basis across studies.

This is a significant assumption, and researchers in Bayesian deep learning should be expected to outwardly declare their tuning procedure—in particular access to distributionally shifted data—as is done in works such as Prior Networks . This will permit researchers and practitioners to more fairly compare the performance of methods based on results reported in their respective papers.

We investigate what impact this assumption—access to distributionally shifted validation data—has on downstream performance across all our tasks, and on held-out in-domain, distributionally shifted, and joint (in-domain combined with distributionally shifted) evaluation datasets. We find that it has a significant impact on metrics commonly used to assess robustness and uncertainty quantification, including area under referral curves (Figure 15) and expected calibration error.

To consider the performance of our baseline models under this assumption, we construct a metric that conveys both in-domain and distributionally shifted performance. In particular, we construct an accuracy referral curve on a combined set of in-domain and distributionally shifted validation examples. Because the in-domain validation dataset is significantly larger than the distributionally shifted dataset for both of the tasks, we upsample the shifted dataset to avoid the signal from the in-domain examples overwhelming that from the shifted examples. We construct an upsampled shifted dataset by first duplicating the shifted validation dataset as many times as possible without exceeding the size of the in-domain validation dataset, and then randomly sampling examples from the shifted validation dataset without replacement until the upsampled shifted dataset contains the same number of examples as the in-domain validation dataset. We construct the “balanced” joint validation dataset as the union of the in-domain validation and upsampled shifted datasets. We construct a “balanced” accuracy referral curve using this balanced joint validation dataset, sweeping over τ\tau to obtain all possible partitions of the dataset into “referral” and “non-referral”. We then tune on the area under this curve.

B.4 Complete Tabular Results

We report additional tabular results for standard predictive performance and robustness (expected calibration error), referral metrics, and out-of-distribution detection across the Severity and Country Shift tasks, considering hyperparameter tuning on either in-domain validation AUC or the joint validation metric (cf. Section B.3), in Tables 2-11.

B.5 Effect of Class Balancing the APTOS Dataset (Figure 16 and 17).

We additionally investigated to what extent the change in class distribution—in terms of the ground-truth clinical labels ranging from 0 (No DR) to 4 (Proliferative DR)—contributed to the higher performance of models in AUC, and weaker performance of models in selective prediction on the APTOS dataset (the distributionally shifted dataset in the Country Shift task) than the in-domain test dataset.

In order to normalize for the change in class distribution, we constructed a variant of the APTOS dataset with the same clinical class proportions as the in-domain EyePACS dataset. This was done by randomly sampling APTOS examples from each class, weighted by the empirical class probability of the EyePACS dataset, until reaching 10,000 samples.

In Figure 16, we see that the ROC curves of models on the rebalanced APTOS dataset is shifted further towards the upper left as compared to the original APTOS dataset. This suggests that the class proportions of the original APTOS dataset were not the reason why models obtained stronger ROC performance on APTOS than the in-domain test set—on the contrary, introducing the in-domain class proportions in the class-balanced dataset improves model performance.

In Figure 17, we observe that the selective prediction performance of models on this rebalanced APTOS dataset is slightly better than on the original APTOS dataset, but the ordering of models does not notably change, and performance is still significantly worse at high referral thresholds than on the in-domain data.

This supports the claim that factors other than simply a changed class distribution, such as meaningful shifts in equipment or patient demographics, result in both stronger predictive performance at 0%0\% of data referred and poor quality of uncertainty estimates in the shifted setting.

B.6 Effect of Preprocessing on Downstream Tasks

Preprocessing played an important role in the EyePACS Kaggle challenge . Here, we investigate how changes in preprocessing affect downstream predictive performance and uncertainty quantification.

In the above experiments, we used the preprocessing procedure of the Kaggle competition winner which consisted of the following steps:

Rescaling the images such that the retinas have a radius of 300 pixels,

Subtracting the local average color, computed using Gaussian blur, and finally,

Clipping the images to 90% size to remove “boundary effects”.

While (1) and (3) are (somewhat) standard techniques used to make the data more amenable for use in non-convex optimization, the standard deviation hyperparameter of the Gaussian blur kernel in (2) presupposes some amount of expert knowledge as the size of the standard deviation governs how visible certain visual artifacts are. As such, varying it has a dramatic visual effect on the preprocessed image, and likely required significant tuning.

In the preprocessing procedure, the standard deviation of the kernel is computed as σ=(target_radius/blur_constant)\sigma=(\texttt{target\_radius}/\texttt{blur\_constant}), where by default, target_radius=300\texttt{target\_radius}=300 and blur_constant=30\texttt{blur\_constant}=30.

Decreasing the blur_constant results in a larger kernel standard deviation, and hence the local average color at each pixel location is computed using a larger window. This ultimately results in the preservation of more signal as well as more noise in the input image (because lower-frequency patterns are subtracted). See Figure 18 for examples of unprocessed retina images along with processed images with various blur constants.

We test the downstream performance of MAP estimation (a deterministic model), a Deep Ensemble, MC Dropout, and an MC Dropout Ensemble on the Country and Severity Shift prediction tasks, varying the blur_constant ∈{5,10,20,30}\in\{5,10,20,30\}.

Severity Shift: Varying Blur Constant (Figure 19, Section B.6). On the in-domain evaluation dataset, higher blur_constant (corresponding to stronger smoothing) tends to perform better across map and mc dropout, single and ensembled models, and the various referral thresholds. However, on the Severity Shift (distributionally shifted evaluation dataset), the mc dropout variants perform better with lower blur_constant. This highlights the importance for practitioners to test changes in experimental settings, including preprocessing, across a variety of uncertainty quantification methods.

Country Shift: Varying Blur Constant (Figure 20, Section B.6). Similarly to the Severity Shift results, higher blur_constant tends to perform better on the in-domain evaluation data across methods and referral rates. Notably, on the distributionally shifted APTOS data, Deep Ensemble outperforms MC Dropout Ensemble, and blur_constant=20\texttt{blur\_constant}=20 significantly improves performance from the default blur_constant=30\texttt{blur\_constant}=30 for Deep Ensemble between referral rates 0.4 and 0.7. For example, for Deep Ensemble at τ=0.7\tau=0.7, we observe 82.2±2.582.2\pm 2.5 AUC with blur_constant=20\texttt{blur\_constant}=20 versus 67.4±5.667.4\pm 5.6 AUC with blur_constant=30\texttt{blur\_constant}=30.