Plex: Towards Reliability using Pretrained Large Model Extensions

Dustin Tran, Jeremiah Liu, Michael W. Dusenberry, Du Phan, Mark Collier, Jie Ren, Kehang Han, Zi Wang, Zelda Mariet, Huiyi Hu, Neil Band, Tim G. J. Rudner, Karan Singhal, Zachary Nado, Joost van Amersfoort, Andreas Kirsch, Rodolphe Jenatton, Nithum Thain, Honglin Yuan, Kelly Buchanan, Kevin Murphy, D. Sculley, Yarin Gal, Zoubin Ghahramani, Jasper Snoek, Balaji Lakshminarayanan

Reliability as a Goal for Artificial Intelligence

Over the past few years, the deep learning approach to artificial intelligence (AI) has made significant progress on benchmark tasks across domains such as computer vision (Dosovitskiy et al., 2021) and natural language processing (Raffel et al., 2020; Brown et al., 2020). With this progress, there is unfettered excitement about the potential of AI to have a transformative impact on applications ranging from medical diagnoses with a human-in-the loop, to using AI to detect online misinformation and toxicity, and to perhaps a pathway for artificial general intelligence. While hypothesizing about this potential is important, we highlight that the typical tasks where deep learning has been most successful have been carefully devised to fit within narrow boundaries—for example, a focus on predictive performance with test inputs close to the data on which the model was trained.

To go beyond these limitations, we argue that the ability of models to make reliable decisions is critical to the deeper integration of AI in the real world. Here, we define reliability as the ability for a model to work consistently across real-world settings. We borrow the term from reliability engineering (Barlow and Proschan, 1975; O’Connor and Kleyner, 2012), a discipline of engineering involving risk assessment, testability, and fault tolerance. Related nomenclature include robustness (Russell et al., 2015), safety (Amodei et al., 2016; Everitt et al., 2018; Hendrycks et al., 2021b), calibration (Dawid, 1982), credibility (D’Amour et al., 2020), and trustworthiness (Avin et al., 2021), each with their own broad and intersecting scopes.

It is common practice in machine learning research to focus on measures of performance based on the accuracy on a test set drawn from the same distribution as the training set, the so-called independent and identically distributed (i.i.d.) assumption. However, this does not capture the real-world deployment of AI systems, where often the testing environment is very different from the training environment, and where the tasks only indirectly involve accuracy measures. The emphasis in our paper is on how reliable an AI system is in broad array of scenarios. We posit three general categories of desiderata for reliable AI systems: they should represent their own uncertainty, they should generalize robustly to new scenarios, and they should be able to efficiently adapt to new data.

Importantly, the aim for a reliable model is to do well in all of these areas simultaneously out-of-the-box without requiring any customization for individual tasks (Figure 1):

Uncertainty involves imperfect or unknown information where it is impossible to exactly describe an existing state (Ghahramani, 2015). Predictive uncertainty quantification enables practitioners to know when to trust the model’s predictions, thereby enabling graceful failures when the model is likely to be wrong. A variety of measurements can be used to quantify the quality of uncertainty such as expected calibration error (Naeini et al., 2015; Nixon et al., 2019), which measures how well the model’s confidence aligns with its accuracy. Providing a quantification of uncertainty also enables better decision making (Parmigiani and Inoue, 2009); one popular setting is selective prediction where a model may defer its prediction to human experts when it is not confident. Another popular task is open set recognitionEarlier work sometimes referred to this setting as out-of-distribution detection. However, in recent work, the term “out-of-distribution” is used as a unifying term for “non I.I.D” test data which includes any type of shift, e.g., covariate or subpopulation shift. We use the term open set recognition to denote detection for shifts where test inputs belong to semantic classes not encountered in the training data. , where the model encounters inputs from new classes at test time that were not seen during training, and the goal is to reliably detect that such inputs do not belong to any of the training classes.

Robust Generalization involves an estimate or forecast about an unseen event (Abraham and Ledolter, 1983; Tran et al., 2020). Prediction quality is typically measured using accuracy (e.g., top-1 error for classification problems and mean squared error for regression problems) and proper scoring rules such as log likelihood and Brier score (Gneiting and Raftery, 2007). In the real world, we care not only about metrics on new data obtained from the same distribution the model was trained on (i.i.d.), but also about robustness, as measured by metrics on data under out-of-distribution shifts such as covariate or subpopulation shift.

Adaptation involves probing the model’s abilities over the course of its learning process. Benchmarks typically evaluate on static datasets with pre-defined train-test splits. However, in many applications, we are interested in models that can quickly adapt to new data and efficiently learn with as few labeled examples as possible. Examples include few-shot learning (Li et al., 2006), where the model learns from a small set of examples; active learning (Settles, 2009), where the model not only learns but also participates in acquiring the data to learn from; and lifelong learning (Thrun, 1998), where the model learns over a sequence of tasks and must not forget about relevant information for previous tasks.

While there has been initial progress on specific tasks within reliability, there remain critical limitations. First, prior work has typically focused on narrow settings. For example, much focus has been on accuracy and calibration error on ImageNet and its corrupted versions as a benchmark task (Hendrycks and Dietterich, 2019; Ovadia et al., 2019). Similarly, the literature on open set recognition has led to the development of methods that focus solely on improving open set recognition, often at the expense of other downstream tasks. This has led to a fragmentation of research techniques, and it remains unknown how well advances on these individual tasks connect to a broader set of tasks and datasets. The goal of this work is to move towards a notion of general reliability by seeking models that work well across many decision-making tasks. rather than needing to design, train, and tune a model for each individual task. The approach is philosophically inspired by the trend in natural language processing where evaluating on multiple tasks (cf. GLUE (Wang et al., 2018a)) has encouraged the community to focus on developing general-purpose techniques.

2 Approaches to Reliability

Prior work has investigated a variety of approaches to improve narrower definitions of reliability. From the literature, several overarching dimensions arise (Tran et al., 2020)—such as the importance of model size; model inductive biases (e.g., architecture); and the combination of multiple models (e.g., ensembles and Bayesian neural networks). There is not yet an understanding of how these dimensions interact (and within current literature, it can be a surprise that combining certain methods is detrimental, e.g., Wen et al. (2020a)). We investigate how to combine these dimensions to form an overall model for reliability.

Modern AI is trending towards training a single large model on a large and diverse data set, known as pretraining, and then applying the model to a wide variety of related downstream tasks (Brown et al., 2020; Kolesnikov et al., 2020; Radford et al., 2021; Thoppilan et al., 2022). This often improves over task-specific state-of-the-art in predictive performance, with many considering such large scale models to represent a “paradigm shift” in ML (Bommasani et al., 2021). Large-scale pretrained models have also significantly improved state-of-the-art on narrower tasks such as accuracy and calibration under covariate shift (Minderer et al., 2021) and open set recognition (Fort et al., 2021; Ren et al., 2021). Given these initial promising results, we use large pretrained models as a building block for reliability. However, large models are compute-intensive and often operate at a different scale than has been studied in many uncertainty and robustness papers. This warrants revisiting existing recipes in this new context, where we focus on methods that scale efficiently to large models.

3 Contributions

First, we define and evaluate reliability in a comprehensive fashion. We use 10 types of tasks in order to capture the three reliability areas—uncertainty, robust generalization, and adaptation—and so that the tasks measure a diverse set of desirable properties in each area. Together the tasks comprise 40 downstream datasets across vision and natural language modalities: 14 datasets for finetuning (including few-shot and active learning-based adaptation) and 26 datasets for out-of-distribution evaluation. As part of the tasks, we build two new datasets in order to assess areas that we found missing in the literature: ImageNet Real-H for large-scale label uncertainty evaluation; and NaLUE for uncertainty in conversational language understanding.

To improve reliability, we develop ViT-Plex and T5-Plex, building on large pretrained models for vision (ViT; Dosovitskiy et al. (2021)) and language (T5; Raffel et al. (2020)), respectively. We train Plex over model sizes up to 1 billion parameters and pretraining dataset sizes of up to 4 billion examples. Figure 3 illustrates Plex’s performance on selected tasks comparing to existing state-of-the-art, which typically is a model specialized for that task. Plex achieves new state-of-the-art on many of the 40 datasets. Importantly, Plex achieves strong performance across all tasks using the out-of-the-box model output without requiring any custom designing or tuning for each task.

Evaluating Reliability

To experiment with a variety of models and tasks, we set up a pipeline for training and evaluation; see Figure 4. It involves three steps: pretrain on a large diverse dataset, finetune on a given dataset’s training split, and evaluate over both in-distribution and out-of-distribution datasets using the appropriate task metrics. We experiment with a variety of methods to improve reliability during both pretraining and finetuning.

We evaluate a model’s reliability using 10 types of tasks, which we categorize as follows. More detailed descriptions of metrics in each task and datasets are provided in Section C.

Calibration assesses the quality of a model’s predicted confidence over a population (Dawid, 1982). It quantifies how well model confidence (the predictive probability of correctness) aligns with model accuracy (the observed probability of correctness). We compute expected calibration error (Naeini et al., 2015) and calibration AUROC (Kivlichan et al., 2021) on 14 image and 10 text datasets.

Selective prediction jointly assesses a model’s predictive performance and quality of uncertainty estimates, by abstaining from making predictions when the model is uncertain (El-Yaniv and Wiener, 2010). There is no popular metric that’s the de facto standard in the literature, so we experiment with two: Oracle Collaborative Accuracy measures accuracy where any deferred predictions are sent to an oracle as a proxy for human++AI collaboration (Kivlichan et al., 2021); and selective prediction by rejection rate traces a curve of predictive accuracies over the percentage of examples that can be rejected (Band et al., 2021). We examine selective prediction on 4 image and 10 text datasets.

Open set recognition assesses how well a model can detect examples belonging to none of the training classes (Geng et al., 2020). This happens in the case of semantic (class) shift, where the test input and label distributions both change, particularly in a structured manner where the output classes change from train to test. For an overview of distribution shifts, see Figure 5, where we train on a distribution p\textscind(x,y)p_{\textsc{ind}}(\mathbf{x},y) and we evaluate on another distribution. For example, the training set may consist only of dog images and the new input is a coyote image. We compute AUROC and use maximum softmax probability as a simple and general detection score.

Label uncertainty is a type of uncertainty inherent in the data labels. This is a form of data uncertainty—irreducible output noise—which is distinct from uncertainty arising from the choice of model (Dusenberry et al., 2020b). Label uncertainty can arise when, for example, human raters disagree about the label for an ambiguous input. If human disagreement is encoded as a label distribution, we can directly compare our model’s predictive distribution to it. We analyze label uncertainty on two image datasets.We focus on label uncertainty in an OOD evaluation-only context, as collecting label distributions can be expensive and therefore the datasets are small relative to single-label training sets. It is not exclusive to the distribution shift setting as training sets may also carry a label distribution per input.

1.2 Robust Generalization

In-distribution generalization assesses how well a model can make predictions after finetuning on a downstream dataset. In particular, we examine accuracy, negative log-likelihood, and Brier score on the in-distribution test splits of 5 image and 3 text datasets. For binary classification tasks, we also look at AUROC and AUPRC.

Covariate shift refers to scenarios where the distribution of inputs changes while the conditional distribution of outputs is unchanged (Sugiyama and Kawanabe, 2012). For example, the training set may include natural dog images and the new input is a drawing of a dog. We use the same metrics as those used for assessing in-distribution generalization.

Subpopulation shift refers to scenarios where the distribution of interest is only part of the full distribution seen during training. A common assumption is that data is sampled from individual subpopulation distributions, which are themselves sampled from a meta population distribution (Santurkar et al., 2020; Yuan et al., 2022). For example, the training set may include natural dog images and the test set only includes terriers. In this setting, we aim to improve predictive performance on unseen or long-tail subpopulations.

1.3 Adaptation

Active learning assesses a model’s ability to not only learn over a fixed set of data points, but also participate in knowing which data points to learn from in the first place. This procedure assesses a model’s label efficiency, where label annotations may be scarce, and so we would like to maximize performance while minimizing the number of labeled data points used (Settles, 2009). We assess accuracy over the total number of acquired examples and apply margin sampling for multi-class uncertainty sampling.

Few-shot learning assesses how well a model can make predictions downstream with only a few training examples (Li et al., 2006). We use 9 datasets and evaluate the settings of 1-shot, 5-shot, 10-shot, and 25-shot (xx-shot means xx examples per class).

With few-shot uncertainty, we examine calibration and open set recognition in the few-shot regime. We use all 9 datasets for few-shot learning in order to evaluate calibration, and we use those with OOD datasets for open set recognition. We also perform zero-shot open set recognition by using Mahalanobis distance scoring to detect whether an input is out-of-distribution based on the model’s representation layer (Lee et al., 2018).

2 Datasets

We selected a broad suite of 40 downstream datasets under the tasks, each ranging from several thousand to a million examples. We outline the datasets for each modality.

We’re motivated to capture datasets spanning natural web images, specialized domains that are likely rare or unseen in large pretrained models, and with both small and large sizes. To do so, we use 11 datasets which we describe below.

CIFAR-10 is a dataset of web images with a training set of 50,000 examples and a test set of 10,000 examples (Krizhevsky et al., 2009). Following Dosovitskiy et al. (2021), we use 99% of the training set for training and 1% for validation.

CIFAR-100 is a dataset of web images with a training set of 50,000 examples and a test set of 10,000 examples. Following Dosovitskiy et al. (2021), we use 99% of the training set for training and 1% for validation.

ImageNet1K is an image dataset organized according to the WordNet hierarchy, with a training set of roughly 1.2 million examples and a test set of 50,000 examples (Deng et al., 2009). Following Dosovitskiy et al. (2021), we use 98% of the training set for training and 2% for validation.

RETINA is a set of benchmarking datasets containing retina scans exhibiting varying degrees of diabetic retinopathy, a medical condition that can result in a loss of eyesight (Band et al., 2021). We chose RETINA as an example of difficult transfer, since retina images are quite different from natural web images used for pretraining. RETINA includes two types of splits: a “Country Shift” split with an in-distribution test set and a test set exhibiting covariate shift; and a “Severity Shift” split with an in-distribution test set and a test set containing labels not included in the training data, representing more severe types of diabetic retinopathy.

We use 7 datasets with a range from 1,880 to 8,144 training examples: Describable Textures Dataset (Cimpoi et al., 2014), UC Merced (Yang and Newsam, 2010), Caltech 101 (Fei-Fei et al., 2004), Oxford-IIIT Pets (Parkhi et al., 2012), Colorectal Histology (Kather et al., 2016), Caltech-UCSD Birds 200 (Welinder et al., 2010), and Cars196 (Krause et al., 2013).

Distribution shift is a common challenge for image problems, and so we cover multiple types for a total of 19 datasets. Table 1 provides an outline. Most notable, there is little work for evaluating label uncertainty, so we propose a large-scale dataset which we call ImageNet ReaL-H. ImageNet ReaL recollects human ratings for the original ImageNet test set (Beyer et al., 2020), and we use its raw data of individual ratings to construct a label distribution representing rater uncertainty for each image. For ImageNet1K, we use 7 datasets for different covariate shifts: ImageNet-A, ImageNet-C, ImageNet-R, ImageNetV2, ImageNet-Vid-Robust, ObjectNet, and YTBB Robust.

2.2 Text

For text, we consider real-world decision making tasks that are known to deploy machine learning models: natural language inference, toxic comments detection, and conversational language understanding. Natural language inference and toxic comments are binary classification tasks that map a (pair of) natural language sentences to a binary category: entailment or no entailment, and toxic or non-toxic, respectively. Conversational language understanding is a task common in chatbot design, where the model maps a natural language query to a multi-token prediction of user intents: for example, “I want to order dinner using Uber Eats” →\to 3-token prediction of (FoodDelivery, Uber, Order).

For natural language inference, we use the Multi-Genre Natural Language Inference (MNLI) corpus which consists of 433k sentence pairs from a diverse collection of genres (fiction, government report, news magazine articles, etc.) (Williams et al., 2017).

For toxic comments detection, we use the WikipediaTalk corpus (Wulczyn et al., 2017) which is composed of roughly 200k English Wikipedia talk page comments between Wikipedia editors across the world.

For conversational language understanding, a large-scale corpus for evaluating uncertainty quantification is lacking. We propose a new dataset Natural Language understanding Uncertainty Evaluation (NaLUE) that is a relabelled and aggregated version of three large NLU corpuses: CLINC150 (Larson et al., 2019), Banks77 (Zhang et al., 2021) and HWU64 (Liu et al., 2021). NaLUE contains 50k+ utterances spanning 18 verticals, 77 domains, and roughly 260 intents. For this task, the model needs to map each utterance to a 3-token sequence of (vertical name, domain name, intent name).

In terms of data distribution, MNLI has a balanced distribution both across the genre and across the label class. NaLUE exhibits a slight skewness toward some popular domains for chatbot development (e.g, banking customer service requests). On the other hand, the toxic comments datasets often exhibit extreme label imbalance. For example, ∼\sim10% of the examples in Wikipedia Talk Corpus examples have positive labels, since most online content is not toxic (Kivlichan et al., 2021).

Natural language is diverse, fast evolving, and rich in long-tail linguistic phenomena. Therefore out-of-distribution examples, particularly long-tail subpopulations, are pervasive in the real-world deployment environment. In Table 2, we outline a total of 7 out-of-distribution challenge sets. Most notably, we construct three new out-of-distribution shifts for NaLUE. NaLUE-tail contains utterances from 28 low-frequency intents categories in NaLUE. NaLUE Standard-OOS and NaLUE Near-OOS contain utterances that describe out of the scope services, differing in their closeness in distribution to NaLUE.

Plex: Pretrained Large model Extensions

Plex is the result of an extensive study of the reliability of large pretrained models and their complementarity with existing reliability methods. ViT-Plex and T5-Plex use several key ingredients:

Base Transformer architecture. We adopt the Transformer standard of an alternating sequence of attention and feedforward layers. We build on T5 1.1 (Raffel et al., 2020) for text as a Transformer in an encoder-decoder setup where the raw text is tokenized with SentencePiece, and on Vision Transformer (Dosovitskiy et al., 2021) for images in an encoder-only setup where the raw images are effectively tokenized into patches.

Model size. We investigate 3 scales of the model size in ViT-Plex: Small (ViT-Plex S; ∼\sim22 million parameters) has patch size 32, 384 embedding size, 1536 feedforward size, 12 residual blocks, and 6-headed attention; Base (ViT-Plex B; ∼\sim87 million parameters) has patch size 32, 768 embedding size, 3072 feedforward size, 12 residual blocks, and 12-headed attention; and Large (ViT-Plex L; ∼\sim325 million parameters) has patch size 32, 1024 embedding size, 4096 feedforward size, 24 residual blocks, and 16-headed attention.

For T5-Plex, we consider three model sizes: Small (T5-Plex S; ∼\sim77 million parameters) has 512 embedding size, 8 encoder / decoder blocks, and 6-headed attention; Base (T5-Plex B; ∼\sim250 million parameters) has 768 embedding size, 12 encoder / decoder blocks, and 12-headed attention; and Large (T5-Plex L; ∼\sim880 million parameters) has 1024 embedding size, 24 encoder / decoder blocks, and 16-headed attention.

Pretraining dataset size. For vision, we scale pretraining from ImageNet-21K to the JFT web dataset on up to 4B images. This mirrors recent work on scaling vision models (Zhai et al., 2021; Pham et al., 2021). For language, we use the C4 dataset which consists of hundreds of gigabytes of English text scraped from the web (Raffel et al., 2020).

Efficient ensembling. Ensembles and Bayesian neural nets have shown to be very effective for uncertainty and robustness (Ovadia et al., 2019; Dusenberry et al., 2020a; Band et al., 2021). To apply ensembling scalably, we use BatchEnsemble (BE) (Wen et al., 2020b) and experiment with its use on both the attention and feedforward layers or on only the feedforward layer. For faster training, we only apply BatchEnsemble at a select number of later layers, similar to mixture of experts models (Riquelme et al., 2021). In both ViT-Plex and T5-Plex, we apply no dropout.

Last layer changes. We experiment with two approaches that modify the model’s final layer to improve reliability, given a fixed representation (a.k.a. deterministic uncertainty quantification setting (Van Amersfoort et al., 2020)). First, we use a Gaussian process (GP) last-layer, which improves distance-awareness of the decision surface by increasing uncertainty far away from the training representations. We use the GP layer implementation proposed by Liu et al. (2020). In addition, to model input-dependent label noise in datasets with many output classes we apply the Heteroscedastic (Het) method of Collier et al. (2021). For detailed background, see Section D.

What to apply in pretraining versus finetuning. We apply efficient ensembling during both pretraining and finetuning. For last-layer methods, we find pretraining benefits can be obtained primarily from only applying the method during finetuning, so we restrict them to that setting (detailed in Section 4.2). In addition, due to compute constraints, we exclusively focus on the finetuning-only setting for T5-Plex. That is, T5-Plex models are initialized from the official pretrained T5 checkpoints, and we apply efficient ensembling and last layer changes during finetuning.

Few-shot protocol. As an alternative to logistic regression on the final layer of frozen representations, we experiment with gradient descent over all parameters; on 5-shot and higher, training the full model with gradient descent can lead to significant gains (for example, 2-3% accuracy boost on ImageNet). We also experiment with a GP or Heteroscedastic last layer as an alternative to a linear last layer.

Summary of Results and Scaling Trends

Figure 3 displays the largest variants of Plex’s performance compared to existing specialized state-of-the-art on a diverse collection of reliability tasks. Plex not only sets new state-of-the-art on many tasks but Plex also unifies reliability performance under one general model for vision and language respectively. Here, we validate several takeaways as we ablate to understand Plex’s ingredients.

Scaling model size improves reliability. Figure 6 displays ViT-Plex and T5-Plex over varying model sizes. We compute a reliability score which is a normalized average over all task metrics: 139 for vision and 54 for language (see Section B for details). We also display reliability scores for individual reliability areas, which are the averages separately over the uncertainty, robust generalization, and adaptation tasks. Classical machine learning theory would suggest that a larger model translates to more overfitting and might therefore be less reliable as it may be overconfident and less robust. However, we find that scale improves overall performance across tasks.

Scaling pretraining dataset size improves reliability. ViT-Plex L with JFT performs better than ViT-Plex L with ImageNet-21K (Figure 6). In Table 3, we also perform an ablation by comparing pretraining on JFT with 300M examples to JFT with 4B examples. We pretrain on up to 8M steps with batch size 4096, which is up to 8X more steps than we typically use for pretraining; each result is a separate run using a tuned learning rate schedule. ImageNet 10-shot accuracy is always better on JFT 4B under the same number of training steps. The models also converge faster with the smaller JFT 300M, reaching a performance limit, whereas JFT 4B keeps improving.

BatchEnsemble improves pretraining. For vision, we run ablations at the fixed setting of L pretrained with JFT, and we use both B and L sizes for text, which are highly competitive settings. Figure 7 displays the ranking across tasks for each model. Methods are applied either during both pretraining and finetuning, or only during finetuning given a pretrained model checkpoint (“None→\toHet” denotes pretraining with no changes and finetuning with Het on top). All the methods displayed improve over a baseline without ensembling or last layer changes (“None”). BatchEnsemble is consistently the best for pretraining.

Last-layer methods improve finetuning. The best ranked models for the vision and language tasks use all of Plex’s ingredients: Het on top of a pretrained BE for vision and GP on top of a BE for language. In particular, for T5-Plex ablations, BE++GP and BE tend to have the strongest performance. From more detailed per-dataset analysis in Section I, BE++GP and BE perform well on MNLI and NaLUE with BE++GP performing slightly better; notably, they outperform even an expensive deep ensemble baseline which also performs well on MNLI and NaLUE. BE++GP outperforms None on Toxic Comments while a Monte Carlo Dropout baseline performs best on that task. T5-Plex L also outperforms T5-Plex B, which indicates the benefit of scale not only in Figure 6’s normalized average score but in their average ranking.

Downstream dataset size has no obvious pattern with reliability. Figure 8 decomposes the performance of ViT-Plex L by analyzing it as a function of the downstream dataset’s size of training set. That is, we aggregate reliability performance for each training dataset separately, ranging from a size of 1,880 examples with Describable Textures Dataset (dtd) to 1.2 million examples with ImageNet. There is no clear pattern with respect to size. On the other hand, the datasets with lower reliability tend to be more different from the distribution of natural images in JFT: UC Merced is a remote sensing dataset of map areas, and Colectoral Histology is a histology dataset of human colorectal cancer. Pet images are not uncommon in JFT, but most breeds in Oxford IIT Pets aren’t in JFT. This suggests reliability performance is possibly more connected to the distribution shift between pretraining and downstream dataset than the downstream dataset’s size.

Given that there are many tasks that make up reliability, here we analyze relationships between the individual tasks. We’re motivated by the following question: Is there an inherent underlying evaluation metric that is indicative of reliability? Specifically, can reliability be predicted from pretraining performance, without any downstream finetuning or adaptation?

In Figure 9, we analyze the correlation of validation loss on the pretraining dataset with each downstream metric. Most of the tasks correlate highly, meaning that the pretraining performance is an important predictor of downstream performance. The one exception are two uncertainty metrics—calibration error and calibration AUROC—which may make sense intuitively given that calibration as a property about models is not tied to predictive performance. Interestingly, AUROC, which refers to the metric for open set detection (an uncertainty task), is strongly coupled with predictive performance (Pearson correlation of 0.99!), more than, say, selective prediction performance (OC-AUC); this suggests that for open set recognition, prediction may be more important than uncertainty. Few-shot accuracy also correlates strongly with pretraining performance, and more examples leads to higher correlation (25>10≈5>125>10\approx 5>1-shot). Accuracy and NLL correlate less strongly than few-shot: this is likely because they are measured on both in- and out-of-distribution datasets whereas few-shot is only evaluated on in-distribution test splits.

Surprisingly, we can find an even stronger answer to the question: training loss on pretraining data is predictive of reliability. Table 4 shows Pearson correlation of upwards of 0.97 for the reliability score, prediction, and adaptation areas. Uncertainty is the least correlated but still quite strong at 0.76. This suggests that to perform well for reliability, simply fitting the data—that is, having high model capacity and the ability to efficiently train to utilize that capacity—is one of the most essential ingredients.

2 The Effect of Pretraining vs Finetuning

How do the effects we find above differ as model ingredients are applied during the pretraining phase vs finetuning phase? Here, we analyze the performance of models depending on choices made separately during the two phases. For example, can the benefits of BatchEnsemble be obtained when only applied during finetuning given a model pretrained without any changes?

To broadcast the weights of a pretrained model UU into those of a downstream model DD, we follow the initialization scheme for DD and replace any weights common to UU and DD by those in UU. For example, a BatchEnsemble obtained by broadcasting a pretrained vanilla model will inherit its slow weights from the model, but fast weights will be initialized randomly. Consistent with all finetuning experiments, the head layer is also reinitialized rather than inherited from the pretrained model.

Figure 10 displays results on ImageNet. BatchEnsemble when applied during both pretraining and finetuning strictly outperforms the finetuning-only variant across the 3 metrics. We also ran a naive deep ensemble baseline and find that this is similarly the case. On the other hand, for both GP and Heteroscedastic, applying the method only during finetuning roughly matches the performance of the method when applied during both pretraining and finetuning. Therefore we apply BatchEnsemble for both pretraining and finetuning, and we restrict using last-layer methods for finetuning. See Section H for more analysis of the heteroscedastic last layer for vision tasks.

3 Ensemble Scaling

Ensembles of neural networks (Hansen and Salamon, 1990; Lakshminarayanan et al., 2017), which aggregate the predictions of several instances of the same model class, provide a remarkably simple, yet effective, way to improve the performance of the base model. This holds true not only for predictive performance but also, crucially, for robustness and uncertainty quantification (Ovadia et al., 2019; Gustafsson et al., 2020; Wen et al., 2020b). However, ensembles require an increasing amount of compute as one increases the size of the ensemble. We’re motivated to understand this axis of model scaling, ensemble size, as it compares to an alternative, popular axis of model scaling: increasing width and depth.

Figure 11 shows that naive ensembling consistently provides better performance as the number of ensemble members increases. However, this comes at significant computational cost. Scaling up the model size from S to B or from B to L leads to an ∼\sim4x increase in compute, and larger single models tend to outperform the ensembles of 4 smaller models. This motivates the importance of efficient ensembles in Plex: BatchEnsemble adds minimal extra compute and outperforms the baseline without ensembles consistently in the ranking across tasks (Figure 7).

Reliability Task Results

In this section, we examine performance on individual tasks across the areas of uncertainty, robust generalization, and adaptation.

To make an accurate assessment of the quality of a model’s predictive uncertainty, we consider a set of tasks, each highlighting a different property of models’ predictive uncertainty estimates.

TL;DR Plex improves calibration over a baseline without any changes, especially on out-of-distribution. Scaling model size and pretraining dataset size can improve calibration.

Intuitively, calibration is about the “accuracy” of a model’s uncertainty estimates. It reflects how well the model’s confidence, which quantifies the predictive probability of correctness, aligns with the model’s accuracy, which is the observed probability of correctness (Dawid, 1982). We investigate two different measures of calibration.

Expected calibration error (ECE) (Naeini et al., 2015) is a binning metric that computes the average difference in confidence and accuracy within different confidence bins. We evaluate ECE of vision models on 3 in-distribution datasets—CIFAR-10, CIFAR-100, ImageNet—and 6 out-of-distribution—CIFAR-10H, ImageNet ReaL-H, ImageNet-A, ImageNet-C, ImageNet-R, and ImageNet-V2. We also evaluate ECE on T5-Plex over 2 in-distribution datasets—MNLI-matched and NaLUE—and 3 out-of-distribution datasets—MNLI-mismatched, HANS, and NaLUE-tailWe excluded toxicity detection datasets which exhibits severe label imbalance. For these datasets, ECE as a population-average metric is not a suitable measure. For example, a naive model that always predict the majority label will trivially achieve high accuracy and low ECE..

Figure 12 displays the average calibration error of vision models on both in- and out-of-distribution evaluation datasets. Calibration errors on in-distribution are relatively low in general (less than 2.5% ECE). Plex improves calibration error over using no changes (None) by roughly a reduction in half on average; on out-of-distribution, it is roughly a 3% improvement whether the models both use JFT or both use ImageNet21K. Regarding model size, we find calibration error improves over S, B, and L. Another significant difference is in pretraining dataset size from ImageNet-21K to JFT, where we find JFT models consistently outperform ImageNet-21K pretrained models.

A similar trend can be observed in the language domain. As shown in Figure 13, compared to the baseline without changes (None), Plex improves the model’s calibration performance by roughly 3% on in- and out-of-distribution. On the other hand, a model’s uncertainty performance is impacted more by the type of uncertainty methods than the model size. For example, the Plex B model is on average stronger than the None L model.

Calibration AUROC considers the binary classification problem of predicting from a model’s uncertainty on a given input whether its prediction for the same input (i.e., the class associated with the highest class probability) is correct. It is the area under the ROC curve for this binary classification task, and it aggregates the model’s predictive performance over all possible confidence thresholds. Comparing to ECE, Calibration AUROC only evaluates the uncertainty score’s ranking performance, that is, whether a model consistently assigns higher uncertainty to incorrect predictions rather than correct predictions, and such an uncertainty score can be used as a signal for prediction correctness with good precision-recall and ROC performance (Krishnan and Tickoo, 2020; Kivlichan et al., 2021). Table 5 presents Calibration AUROC on vision and text datasets. We see a similar pattern as ECE where larger models generally perform better, but the pattern is slightly weaker. For example, Plex L performs best on only 4/6 datasets.

1.2 Selective Prediction

TL;DR The model extensions in Plex provide a significant improvement for selective prediction, enabling the ability to predict at much lower error rates when deferring just a small fraction of predictions.

Deep learning models are typically evaluated under the lens of average predictive accuracy. However, this doesn’t account for the real-world cost of mistakes in deployed models today. Often the risk associated with a mistake can outweigh the benefit of being correct and thus, in expectation, it can be better to not predict at all under some confidence level or defer to a more expensive procedure (e.g., a human expert). This motivates selective prediction, which includes predictive performance as part of a larger decision-making scenario in which the model may abstain from making certain predictions (El-Yaniv and Wiener, 2010).

Selective prediction performance can be evaluated using several approaches, and there has not been much standardization in the literature. We investigate two that use model uncertainty, through joint human-model collaboration and through rejection rates.

Many real world use cases of AI allow for a model to defer a subset of predictions to a human expert. Oracle Collaborative Accuracy and AUROC measure the performance of an oracle–model collaboration system, where the oracle acts a proxy for human experts in order to automate evaluation. The model sends predictions with high predictive uncertainty to an oracle subject to a fixed referral budget (e.g., only 1% of all queries can be referred to the oracle). We compute Oracle Collaborative Accuracy and AUROC over a range of budgets and evaluate on both vision and language modalities.

For vision, Table 6 displays Oracle Collaborative AUROC with a review budget of 0.5% on three vision datasets over different variants of Plex. Surprisingly, the AUCs are quite high: on ImageNet for example, Plex L attains 0.98, and the models all achieve greater than 0.9 across the datasets. This implies that the ability to defer can enable lower error, which can be useful for higher-risk applications. Plex consistently performs best, even when taking into account different model and dataset sizes. We also examined the metric over review budgets of 1%, 2%, and 5%, and found the results to be unchanged. This suggests that the results on 0.5% are the model’s limiting performance with an oracle in the loop.

For language, Figure 15 reports results across three different model sizes (S, B, L) and for eight methods: a baseline without changes (None), MC Dropout (MCD), Deep Ensemble (DE), Gaussian process (GP), Batch Ensemble (BE), Heterostochastic (Het) and two method combinations Gaussian process ensemble (DE+GP) and BE-GP (Plex) (see Tables 19 and 20 in the Appendix for detailed results). We show the rankings of these methods under different types of test data distribution (i.e. in-domain, OOD and tail-population). We first see that across different methods, DE++GP, BE++GP (Plex), BE, and MCD tend to have the strongest performance. In particular, DE++GP almost always dominates the other methods on MNLI and NaLUE, and remains competitive in the case of label imbalance (i.e., Toxic Comments). However, DE++GP is an expensive method that costs 10x more in memory and compute and therefore is not competitive in scale (a more thorough analysis is in Section 4.3). On the other hand, among the more efficient, single-model methods (i.e., Plex, BE and MCD), BE and Plex perform well on MNLI and NaLUE (notably, outperform the most expensive DE), while MCD stands out in the Toxic Comments. The above observations suggest that, when the training examples are drawn from a relatively simple distribution, quantifying output-layer uncertainty alone is sufficient to attain strong performance. However, when there are pathologies in the data distribution (e.g., extreme label imbalance and high label noise), quantifying the uncertainty within the model’s intermediate representations (e.g., via some form of perturbation like BE) becomes important.

Next we investigate how a model’s performance is impacted by the model size. For model size scaling, we evaluate BE++GP (Plex), MCD, and None. We evaluate the performance of each method under progressively larger architectures, S, B, and L, and observe how the behavior changes across the method and with respect to the architecture size. Figure 15(d)-15(f) summarizes the rankings of uncertainty methods organized by the sizes of the architecture. As shown, comparing across architecture sizes, we see a larger architecture almost always leads to stronger performance in collaboration. This trend remains largely consistent even under distributional shifts and in tail groups. On the other hand, a model’s uncertainty performance is impacted more by the type of uncertainty methods. That is, the MCD and BE++GP models are on average stronger than None models, regardless of architecture size. Finally, within larger architectures (i.e., T5 B and T5 L), BE++GP generally outperforms MCD, and MCD outperforms None. This trend is broken in two situations, (1) in Toxic Comments, MCD strongly outperforms all other architectures, (2) in small architecture, BE-GP is often the poorer performing model, and None achieves the best performance in multi-token prediction problems (NaLUE). Finally, between the larger architectures (i.e., T5 B v.s. T5 L), Plex (BE++GP) generally outperforms MCD, and MCD outperforms None. This trend is broken in two situations, (1) in Toxic Comments, MCD strongly outperforms all other architectures, (2) among the small architectures, the None baseline achieves the best performance in multi-token prediction problems (i.e., NaLUE).

Rejection AUC measures a model’s performance in the scenario where it is permitted to not predict on some fraction of the data for which the model is most uncertain, e.g., up to 10% of all queries. Unlike the collaboration metric, it does not assume the query is passed to an oracle as rejected inputs may be subject to further review for broader decision making such as to improve data collection. Rejection AUC is computed from predictive uncertainty (e.g., predictive entropy, predictive variance, confidence) and a chosen predictive performance evaluation metric (e.g., accuracy, AUROC, AUPRC) at different rejection rates. For a rejection rate of τ∈\tau\in, the τ×100\tau\times 100% of data points in the evaluation set for which the model is most uncertain are identified. Those inputs are then rejected, and we assess the predictions on the remaining (1−τ)×100(1-\tau)\times 100% of data points. Accuracy–Rejection AUC is given by computing the area under the Accuracy–Rejection Rate curve. Additional results for AUROC and AUPRC—Rejection AUCs can be found in Section F.

Table 7 displays overall performance for a selection of finetuned models that were pretrained on ImageNet-21K, and Figure 17 displays the full accuracy–rejection curves for the Country and Severity Shift out-of-distribution tasks. We find that that Plex generally improves performance over using no changes on both datasets on three out of four evaluation tasks and that Het L performs as well or better than Plex L on all four tasks.

Examining the quality of different models’ predictive uncertainty estimates for out-of-distribution evaluation more carefully, we find that on the Country Shift out-of-distribution (covariate shift) task, all models exhibit non-monotonically increasing accuracy–rejection rate curves. Since the rejected examples are selected based on the models’ predictive uncertainty estimates, this behavior indicates that beyond certain rejection-rate thresholds, the models tend to be underconfident about correct predictions and overconfident about mistakes. We do not observe this pattern on the Severity Shift out-of-distribution (semantic shift) task or on the Country and Severity Shift in-distribution tasks (see Section F).

1.3 Open-set recognition (OSR)

TL;DR A simple scoring technique (maximum softmax probability) works well across open set recognition problems. Over image and text datasets, Plex provides a consistent improvement for new state-of-the-art, and without custom changes per dataset.

Open-set recognition, also called out-of-distribution detection or anomaly detection, aims to detect samples from new classes that are not included in training. In contrast to OOD generalization where the test example belongs to the same in-distribution training classes (x,y)(x,y), y∈Y\textscindy\in Y_{\textsc{ind}}, OSR aims to detect test example (x,y)(x,y), where y∈Y\textscoody\in Y_{\textsc{ood}}.

To evaluate OSR performance, we design an uncertainty score Uθ(x)U_{\theta}(x) that indicates the likelihood of an input xx being OOD, and θ\theta denotes the classification model’s parameters. As a default for all OSR problems, we apply the most commonly used uncertainty score: 1−MSP1-\text{MSP}, where MSP is the maximum over predicted softmax probabilities (Hendrycks and Gimpel, 2017). A mixture of in-distribution examples and OOD examples are used as the test set. For each test example, the model generates an uncertainty score. Comparing the uncertainty score with its ground truth OOD label (0 indicates the example is from in-distribution data, and 1 indicates the example is from OOD data), we use AUROC to measure how well the uncertainty score separates the two groups. Note that the model is not finetuned using any OOD examples; test OOD examples are only used for evaluation. We discuss results on alternative uncertainty scores in Section E.

Comparing models’ OSR performance across tasks. We evaluate the performance of Plex and multiple baseline models. Based on the similarity between the in- and out-of-distribution data, we can classify OSR tasks under two categories (Table 8): (a) far-OOD detection and (b) near-OOD detection. In Table 9, we observe that first, larger models have better performance. Second, among all the models, Plex performs the best for most takes. Third, for ImageNet2012 vs Places365, models pretrained with ImageNet-21K (I21K) outperform the same models but pretrained on JFT.

Open-set intent detection based on large language pretrained models. To study OSR performance of different models in the language domain, we design an intent detection task for detecting natural utterances that are out of the scope (OOS) services. For the in-distribution dataset NaLUE, the model needs to map each utterance input xx to a 3-token sequence of y=(y1,y2,y3)=(vertical name,domain name,intent name)y=(y_{1},y_{2},y_{3})=(\text{vertical name},\text{domain name},\text{intent name}). To study both far-OOD performance and near-OOD performance, we construct two out-of-the-scope datasets for NaLUE:

NaLUE Standard-OOS: completely out-of-domain queries, based on CLINC150-OOS. The standard-OOS queries do not share any of the (vertical name, domain name, intent name) with the in-domain query.

NaLUE Near-OOS: in-domain, out-of-scope queries. It is created by sampling 20% intent from each known domain as OOD data. The near-OOS examples can share the same vertical name / domain name (but definitely different intent name) with the in-domain query.

The model outputs a sequence, and we compute the uncertainty score based on the conditional softmax probability p(yl∣y<l,x,θ)p(y_{l}|y_{<l},x,\theta), using the conditional entropy. The conditional entropy is a standard uncertainty score for sequences (Malinin and Gales, 2021),

In Table 10, models including None, GP, Het, MCD, BE, BE-GP, Deep Ensemble (DE), DE-GP are evaluated for their detection performance on Standard-OOS and Near-OOS. BE and DE-GP outperform the other models, followed by MCD.

1.4 Label Uncertainty

TL;DR We propose label uncertainty as an important challenge, with a new large-scale dataset and propose an evaluation metric. The heteroscedastic last-layer particularly helps for label uncertainty.

It is typical to encounter label noise in non-academic datasets. However, benchmark datasets are often carefully curated to reduce label noise due to unreliable or disagreeing annotators. We use the existing CIFAR-10H dataset, which is a relabelled version of the CIFAR-10 test set with on average >>50 crowd-sourced noisy labels per input (Peterson et al., 2019).

There is a notable lack of label uncertainty datasets in the literature, and we found limiting performance on CIFAR-10 as our models achieve >>99% accuracy on the original CIFAR-10 test set. Therefore we build a larger dataset called ImageNet ReaL-H, which leverages the raw annotations of ImageNet ReaL (Beyer et al., 2020). When raw labels were available for an image we followed the same averaging procedure as for CIFAR-10H to produce a soft label. There are some cases where no raw annotations were available. These cases correspond to an image where a set of ML models all agreed with the original ImageNet label for the image and thus the image was not sent to human annotators for re-labelling. In these cases we took the one-hot ImageNet label as the soft label (equivalent to all human annotators agreeing the original ImageNet label was correct).

To evaluate a model’s ability to capture label uncertainty, consider a divergence measure D\mathcal{D} between probability distributions,

where p\textscdata(⋅)p_{\textsc{data}}(\cdot) is the true label distribution for the data point and p\textscmodel(⋅)p_{\textsc{model}}(\cdot) is the model’s distribution. This evaluation metric only reaches 0 when the model captures the true label distribution for every example. We use KL divergence and approximate the label distribution using an empirical distribution over the multiple label samples per input.

We observe in Figure 18 that the heteroscedastic models (Het in finetuning-only and Plex) outperform the None model. Plex does best on ImageNet ReaL-H excluding DE which uses more compute, but Plex performs worse on CIFAR-10H.

2 Robust Generalization

An important aspect of model reliability is it’s ability to make accurate predictions when the test data distribution changes. In this section, we examine model robustness to different forms of data distribution shift. We look at performance both in-distribution and under two out-of-distribution shifts: covariate shift and subpopulation shift.There are 2 distribution shift types which we do not cover here: semantic (class) shift and label uncertainty. Semantic (class) shift is not possible to measure predictive performance under without an open vocabulary model such as recent image-text models (Radford et al., 2021). Therefore we restrict studies of class shift to the open set recognition task (Section 5.1.3). Label uncertainty is covered in Section 5.1.4.

TL;DR Ablation experiments show that pretraining does not help for diabetic retinopathy diagnosis on images, but the Plex model extensions do give a benefit. Across language tasks, pretraining, model scale, and the Plex extensions all contribute to better in-domain generalization.

We look at accuracy and NLL across the in-distribution test splits of each dataset after finetuning on their respective training splits.

For RETINA in Table 11, we compare Plex to the best-performing ResNet-50 baseline results based on a wide array of uncertainty quantification methods (Gal and Ghahramani, 2016; Blundell et al., 2015; Rudner et al., 2021, 2022; Dusenberry et al., 2020a; Farquhar et al., 2020) reported in Band et al. (2021) and find that pretraining on neither I21K or JFT results in an improvement in predictive performance across all evaluation metrics compared to the state-of-the-art ResNet-50 trained from scratch. This failure may be due to the fact that the retina scans used for the diagnosis tasks differ significantly in appearance from the images in the pretraining datasets and as such may be too dissimilar to convey meaningful inductive biases into the finetuned neural network. Nevertheless, the Plex extensions do still provide a benefit over the pretrained model. See Figure 19 for CIFAR and ImageNet.

In the Appendix, Table 19 summarizes the performance of different uncertainty methods, and Table 20 evaluates the models’ performance across three different scales (i.e., T5small{}_{\texttt{small}}, T5base{}_{\texttt{base}}, T5large{}_{\texttt{large}}). Finally, Figures 35 and 36 in the Appendix visually summarize the relative performance ranking across methods and scales. Comparing across methods, uncertainty methods (specifically, BE, DE, GP, and their combinations) provides improved generalization performance when compared to the None baseline. Specifically, DE+GP provides the strongest performance across almost all tasks. Among the more efficient methods, Plex (BE+GP) consistently ranked among the top performing methods, while the performance of GP, MCD and BE varies across the datasets (Table 19). Comparing across model scale, the methods’ generalization performance generally improves as the model scale increases, with Plex (BE+GP) the method benefiting the most strongly from scaling, delivering on average the strongest performance at T5large{}_{\texttt{large}}.

2.2 Covariate shift

TL;DR In general, Plex significantly improves performance across metrics under covariate shift, even when it doesn’t outperform on in-distribution test data. Model size (larger models perform better under covariate shift) and pretraining data are consistently major factors.

In order to be robust under covariate shift, a model should be able to reliably make correct predictions on noisy, corrupted, and otherwise distribution-shifted inputs. In this section, we evaluate robustness using a variety of datasets that exhibit different types of covariate shift.

RETINA’s Country Shift task is concerned with diagnosing diabetic retinopathy from retina scans obtained using different medical equipment and a different patient population. In Table 12, we compare Plex to the best-performing ResNet-50 baseline results based on a wide array of methods to improve uncertainty quantification (Gal and Ghahramani, 2016; Blundell et al., 2015; Rudner et al., 2021, 2022; Dusenberry et al., 2020a; Farquhar et al., 2020) reported in Band et al. (2021) and find that pretraining on either I21K or JFT results in an improvement in predictive performance across all evaluation metrics compared to the state-of-the-art ResNet-50 trained from scratch. This result is in contrast with corresponding results on the Country Shift in-distribution task, for which the best-performing models trained from scratch outperform all pretrained models. In this case, pretraining does improve generalization under covariate shift even when training from scratch yields better predictive performance on in-distribution evaluation tasks. The Plex entensions also yield a significant improvement both in accuracy and log-likelihood.

Figure 20 displays performance over multiple types of covariate shift as a function of model scale (number of parameters): image corruptions (ImageNet-C (Hendrycks and Dietterich, 2019)), natural adversarial examples (ImageNet-A (Hendrycks et al., 2021c)), artistic renditions (ImageNet-R (Hendrycks et al., 2021a)), and “natural” shift obtained by following ImageNet’s data collection procedure to obtain a new test set (ImageNet-V2 (Recht et al., 2019; Taori et al., 2020)). We find that BE ViT L outperforms None methods, that pretraining on JFT outperforms models trained on other pretraining datasets, and that increasing model scale consistently leads to improved performance under shifts across metrics. Unlike in some sections, for performance under shifts, we consistently find that pretraining on JFT outperforms that on ImageNet-21K.

ImageNet-Vid-Robust and YTBB-Robust (Shankar et al., 2021) are benchmark datasets for measuring robustness to natural perturbations in images by using subsequent frames from video sources with human curation. On the robust accuracy (pm-k accuracy) metric, in Figure 21, None and BE L models perform equally well, and performance improves with model size. Models pretrained on JFT outperform those pretrained on I21K, with BE B JFT even outperforming None L I21K and BE L I21K.

In natural language applications, it is common for a trained NLP model to be deployed to an environment that exhibits significant style and topical drift when compared to training data. For example, a NLI model trained on written literature is used to analyze in-person dialogues, or a toxic comment detection model trained on U.S. web forums is deployed to non-U.S. news websites (Williams et al., 2017; Kivlichan et al., 2021). To understand the robustness of large uncertainty models with respect to these common language drifts, we evaluate the models’ prediction performance on out-of-distribution splits of the NLI and Toxic Comment tasks (MNLI-mismatched and Civil Comments, respectively). Specifically, MNLI-mismatched represents a low-degree shift scenario where models trained on written and spoken genres (e.g., government report, fiction, phone conversations) are tested on different subtypes of similar genre (e.g., fundraising letters, in-person conversation, 9/11 calls). On the other hand, Civil Comments represents a high-degree shift scenario where models trained on web forum conversations before 2015 were tested on news websites comments from a different period (2015-2017). Figure 22 compares the performance of nine models based on T5base{}_{\texttt{base}}, and also that of three representative efficient methods (None, MCD and Plex) across model scales (i.e., T5small{}_{\texttt{small}}, T5base{}_{\texttt{base}}, T5large{}_{\texttt{large}}) (see Appendix Table 35-36 for detailed results). As shown, the models’ performance in out-of-distribution generalization correlates well with their in-domain performance, with DE-GP being the most competitive ensemble model, and BE / Plex the most competitive efficient method for MNLI, and MCD the most efficient method for Toxic Comments. Comparing between model scales, larger models in general lead to stronger out-of-distribution performance, with Plex being the most competitive method for moderate-to-large size models (T5base{}_{\texttt{base}}, T5large{}_{\texttt{large}}).

2.3 Subpopulation shift

TL;DR Plex improves generalization under subpopulation shift relative to models omitting pretraining, efficient ensembling, or last layer changes, for both language and vision datasets.

We next study performance on shifts where data is composed of subpopulations which may be rare or unseen during training. A common assumption is that subpopulation shift data is drawn from a distribution of distributions: subpopulations are drawn from a population distribution, and each subpopulation has its own data-generating distribution (Santurkar et al., 2020; Yuan et al., 2022). Since data in practice is often generated by individuals or groups who may have different data distributions, and models ideally generalize to unseen individuals and groups, this setting can be used to measure a natural notion of reliability. We study Plex under multiple ways of partitioning data into subpopulations.

One method for producing data with subpopulation shift is partitioning a standard dataset into subpopulations such that each subpopulation contains semantically similar examples. We leverage the Semantically Partitioned CIFAR10 and CIFAR100 datasets from Yuan et al. (2022) to study how Plex performs on new subpopulations. In particular, we measure classification accuracy with 30 and 100 subpopulations unseen during training, studying how accuracy varies among subpopulations.

From Figure 3, Plex greatly outperforms the baseline presented in Yuan et al. (2022) for tail subpopulation shift accuracy (we report 25th percentile among subpopulations). The baseline does not leverage large-scale pretraining or other changes Plex introduces. To better understand how Plex’s ingredients lead to improvements in accuracy under subpopulation shift, we perform a series of ablations shown in Figure 23. Comparing Plex to training with BatchEnsemble for both pretraining and finetuning (‘BE→BE’), the latter produces a slight degradation in performance for both median and tail subpopulations, especially for CIFAR100, where performance is less saturated. This indicates that last layer changes offer some improvement but do not fully explain Plex’s improvement over baseline in this setting. Looking at ‘None→BE’, standard pretraining slightly degrades performance relative to BatchEnsemble pretraining, but using BatchEnsemble for finetuning still offers a significant boost over completely standard training (‘None’); ensembling appears to be more important in this setting than last layer changes. Comparing Plex and ‘None’, we see a significant difference in performance, but the difference between Plex’s performance and that of the baseline presented in Figure 3 remains much larger. This indicates that the difference in pretraining is an important driver of Plex’s improvements on unseen and long-tail subpopulations, since the baseline does not pretrain on large-scale data.

A well-observed issue in machine learning models is their tendency to learn decision rules that rely on spurious, non-causal patterns that exhibit strong statistical correlation with the outcome in the training data, i.e., shortcut learning (Geirhos et al., 2020). In particular, previous theoretical studies show that a model’s tendency to learn spurious patterns is rooted in the overparameterized model’s tendency to flexibly adapt to the (biased) empirical distribution of the data, and cannot be fully addressed by increasing model size (Bommasani et al., 2021). However, recent empirical studies also show that, when the data distribution exhibits suitable diversity such that it contains a small amount of counterexamples where the spurious pattern doesn’t hold, large pretrained models can still lead to improved robustness (Tu et al., 2020). Therefore, we investigate if the large pretrained models’ robustness indeed improves with model scale, and if the uncertainty methods bring additional gain on top of the None baselines.

For the language modality, we evaluate the performance of T5 models on two spurious-correlation subpopulations: HANS for natural language inference, and CivilComments-Identities for toxic comment detection (McCoy et al., 2019; Borkan et al., 2019). Specifically, HANS evaluates the NLI model’s robustness against non-causal heuristic patterns that are empirically associated with sentence entailment (e.g., lexical overlap) (McCoy et al., 2019). On the other hand, CivilComments-Identities contains comments that mention certain gender, sexual orientation, ethnic or religious identities, which empirically exhibits differential distribution of toxicity label with respect to their surface-level textual features. A toxic detection model that relies on these surface-level identity mentions can lead to unintended consequences, unjustly reinforcing existing social stereotypes to disadvantaged identity groups (Borkan et al., 2019).

Figure 24 summarizes the performance of different uncertainty methods, and across three different scales (i.e., T5small{}_{\texttt{small}}, T5base{}_{\texttt{base}}, T5large{}_{\texttt{large}}) (See Appendix Figures 36 and 35 for detailed results). Compared to the None baseline, the uncertainty methods (BE, DE, GP and their combinations) provide significant improvements across all aspects of model performance (i.e., generalization, selective prediction, and uncertainty calibration). For generalization and selective prediction performance, the ensemble of SNGP models (DE-GP) and Plex (BE-GP) perform the best among ensemble and non-ensemble methods, respectively. Furthermore, within each method class, the subpopulation generalization improves as the model scale increases, with Plex (BE+GP) being the most competitive method for moderate- and large-size models. Finally, the trend in uncertainty calibration is less consistent, we see that among all methods, GP and BE are on average the most calibrated ensemble and non-ensemble method, while the model scale does not seem to have a consistent impact to calibration performance.

3 Adaptation

In this section, we examine the models’ ability to adapt to new data. We focus on data-efficient learning (small data generalization), where the goal is to attain high reliability performance with only a small set of examples. Our study focuses on ViT vision models, since adaptation has not yet been studied for T5 text models (Raffel et al., 2020).

TL;DR Large pretrained models are good active learners. Plex not only provides a significant boost in initial performance, but it also finds examples in order to adapt and improve at a faster rate than active learning models without pretraining.

The goal in active learning (AL) (Cohn et al., 1996; Settles, 2009) is to maximize label efficiency for machine learning models where label annotations are scarce. The success of AL has been demonstrated on a range of real-world problems where labels are expensive to acquire, e.g. computer vision (Gal et al., 2017; Citovsky et al., 2021), natural language (Thompson et al., 1999; Siddhant and Lipton, 2018), speech (Hakkani-Tür et al., 2002; Riccardi and Hakkani-Tur, 2005) and robotics (Martinez-Cantin et al., 2007; Wang et al., 2018b). Here, we investigate Plex for AL on the image domain and focus on four datasets: CIFAR-10 and CIFAR-100 which are standard in active learning research (Tran et al., 2019; Hu et al., 2019; Song et al., 2019); and we scale active learning to a larger dataset, ImageNet, which is less commonly evaluated (Emam et al., 2021; Beluch et al., 2018). Finally, we evaluate active learning on Places365 (roughly 1.8 million examples in the training pool) which is also less explored.

We adopt a standard setup of AL for multi-class single-label image datasets (Section 2.2), where we assume an initial model, a training pool of unlabeled images, and a budget of total number of labels to acquire. AL operates in a loop where in each round, labels are acquired for examples with the highest acquisition score from the training pool, the model is subsequently finetuned on the image-label pairs observed so far, and the training pool is updated to remove the newly labeled images. Once we exhaust the budget on label acquisition, the cycle of AL stops and the final model is obtained.

For the label acquisition strategy, we use margin sampling (Margin) as a representative AL approach (Scheffer et al., 2001; Roth and Small, 2006), and it has been found competitive on AL tasks with deep learning models (Citovsky et al., 2021). Margin uses the difference between the highest and second highest predicted probabilities (Scheffer et al., 2001) to score informativeness: the smaller the difference, the more uncertain the model is about an example and the more informative it is expected to be. As a baseline, we compare to uniformly randomly sampling from the training pool (Uniform).

To better simulate realistic label acquisition settings with parallel annotators, we adopt the batch active learning setup (Settles, 2009) for both Margin and Uniform. For Margin, the examples with the top-KK margin scores are chosen for each acquisition round. For Uniform, we randomly select KK examples without replacement.

Figure 25 displays the AL results on CIFAR-10, CIFAR-100, ImageNet and Places365. We compare Margin and Uniform with Plex models pretrained on ImageNet21K or JFT, or without any pretraining, i.e. randomly initialized Plex models. We set initial data size to be 2×2\times number of classes, max training set size to be 20×20\times number of classes, and acquisition batch size to be 0.5×0.5\times number of classes.

Figure 25 shows that AL with pretrained models significantly outperform models without pretraining across all tasks. We observed two notable effects. First, pretraining has a significant initial boost in performance, where for example Plex starts at 20% for CIFAR-10 without pretraining and 50-60% with pretraining. Second, pretraining results in faster adaptation in terms of the accuracy gain for each labelled example: for example, it takes roughly 200 examples on CIFAR-10 to go from 20-30% whereas it takes roughly 10 examples for pretrained Plex to go from 50-60%.

Pretrained models also enables better performance from better AL acquisition methods. Without pretraining, we observe no gain in performance using Margin comparing to Uniform. However, with pretrained models, Margin almost always outperforms Uniform. Notably, for Plex models pretrained on JFT, Margin requires 37.5%, 31.6%, 35.9% and 10.3% less labeled data than Uniform respectively on CIFAR-10, CIFAR-100, ImageNet and Places365 in order to achieve the best accuracy obtained by Uniform.

For CIFAR-100 and ImageNet2012, AL with models pretrained on I21K outperforms that on JFT. We hypothesize that CIFAR-100 and ImageNet are more “in-distribution” in I21K than JFT (Section 5.3.2 also shows this pattern). However, Places365 is an OOD dataset for both I21K and JFT, where JFT is much larger than I21K. Likely due to better representations pretrained on a larger and more diverse dataset, AL with Plex models pretrained on JFT performs better than I21K. On a related note, Evci et al. (2022) observed how domain shift can impact finetuning and hypothesized that the ability to leverage pretrained representations contributes to the effectiveness of finetuning; and Tamkin et al. (2022) observed related findings for pretrained models focused on task ambiguity.

3.2 Few-shot Learning

TL;DR Plex’s BatchEnsemble representation improves few-shot and over increasing model and pretraining dataset scales. Full-model training can work better than linear evaluation if the goal is to maximize performance from limited examples.

In few-shot learning, we examine how well representations learned during pretraining enable fast downstream adaptation. We use a linear evaluation protocol where we extract features from the models’ pre-logits layer and train a multinomial logistic regression model using L-BFGS. Unlike linear regression which is also popular (e.g., Dosovitskiy et al. (2021)), logistic regression produces a categorical distribution, enabling one to measure the quality of the few-shot model’s uncertainty estimates; we take advantage of this in Section 5.3.3.

Figure 26 displays test set accuracy and negative log likelihood (NLL) averaged over 9 datasets: birds, caltech, cars, cifar100, col hist, dtd, imagenet, pets, and uc merced. The choice of model translates to a consistent 1-2% improvement in accuracy depending on the few-shot setting. The BatchEnsemble representation is among the top performing that gives high predictive performance.

The size of the model and pretraining dataset also has a noticeable effect on few-shot performance. While ImageNet-21K is 10 times smaller than JFT, models pretrained on ImageNet-21K consistently outperform those pretrained on JFT across prediction and uncertainty tasks, with the exception of calibration. We hypothesize this is because ImageNet-21K is closer in distribution to the datasets we evaluate on average. It is also surprising that when finetuning over the full ImageNet dataset (Section 5.2), the conclusion is reversed with JFT models outperforming ImageNet-21K. This suggests a transition point when the pretrained models are trained on a sufficiently high number of examples.

Intuitively, linear evaluation is about “representation learning”, assessing how well a pretrained model’s fixed representation layer generalizes to many different scenarios. This setup may not be optimal for real-world transfer if the goal is simply to maximize performance from limited examples, without constraints on the approach to doing so.

In Figure 27, we adopt the usual finetuning protocol and experiment with training the model’s full set of parameters given the pretrained checkpoint. We set up a few-shot version of ImageNet1k for this study.https://www.tensorflow.org/datasets/catalog/imagenet2012_fewshot On 1-shot, linear evaluation outperforms full-model training which may not be surprising given the high possibility of overfitting. Starting from 5-shot (and which is likely to hold for even lower shot settings), full-model training surpasses linear evaluation with a consistent 2-3% improvement in accuracy.

3.3 Few-shot Uncertainty

TL;DR Plex performs well on both few-shot and zero-shot open set recognition, particularly using larger model sizes and on ImageNet21K. Few-shot calibration remains a challenging problem.

Few-shot learning is an important and practical setting, for which the literature almost exclusively considers accuracy to measure performance. Few-shot uncertainty quantification is also a key problem: compared to finetuning, few-shot models are more sensitive to the examples they’ve seen and by extension, much less accurate. Thus few-shot models should be more likely to be uncertain, and the quality of their uncertainty estimates must be good in order to know when to refrain from making a prediction.

In Figure 28, we measure uncertainty performance on three metrics: ECE and Calibration AUROC for calibration, averaged over all 9 datasets used for few-shot learning, and OOD AUROC for open set recognition. Open set recognition requires an OOD dataset in addition to a test split, so we only evaluate it for two datasets: CIFAR-100 with an OOD dataset of CIFAR-10; and ImageNet1K with an OOD dataset of Places365. Calibration performance is roughly comparable across methods. BatchEnsemble has a noticeable improvement on open set recognition. The scale of the pretraining dataset has no significant effect on calibration, where the trend is weaker than what’s seen on calibration of finetuned models (Section 5.1.1). The choice of pretraining dataset does have a noticeable effect on open set recognition, where like in Section 5.3.2, Plex on ImageNet21K surprisingly does better than JFT models.

Given that the pretrained model itself provides rich and robust representations, we can use that as a feature extractor and conduct zero-shot open set recognition without any training examples. To do so, we can take advantage of Mahalanobis distance (Maha) (Lee et al., 2018) as a detection score: Maha measures the distance between the test input and the fitted training distribution in the embedding space. It operates on a fixed representation layer and does not require operating on softmax outputs with a newly trained last layer. The training distribution is fitted using a class conditional Gaussian N(μk,Σ),k=1,2,…,K\mathcal{N}(\mathbf{{\bm{\mu}}}_{k},{\bm{\Sigma}}),k=1,2,\dots,K to each of the KK in-distribution classes based on the embedding z{\bm{z}}. We estimate the mean vectors and covariance matrix as: μk=1Nk∑i:yi=kzi\mathbf{{\bm{\mu}}}_{k}=\frac{1}{N_{k}}\sum_{i:y_{i}=k}{\bm{z}}_{i}, for k=1,…,K,k=1,\dots,K, and Σ=1N∑k=1K∑i:yi=k(zi−μk)(zi−μk)T{\bm{\Sigma}}=\frac{1}{N}\sum_{k=1}^{K}\sum_{i:y_{i}=k}\left({\bm{z}}_{i}-\mathbf{{\bm{\mu}}}_{k}\right)({\bm{z}}_{i}-\mathbf{{\bm{\mu}}}_{k})^{T}. Note that class-conditional means μk\mathbf{{\bm{\mu}}}_{k} are independent for each classes, while the covariance matrix Σ\Sigma is shared by all classes to avoid under-fitting errors. For a test input, Mahalanobis distances from a test input to each of the fitted KK in-distribution Gaussian distributions N(μk,Σ)\mathcal{N}(\mathbf{{\bm{\mu}}}_{k},{\bm{\Sigma}}) is computed, and the minimum of the distances over all classes is used as the uncertainty score, MD(z)=min⁡k{(z−μk)TΣ−1(z−μk)}.\text{MD}({\bm{z}})=\min_{k}\{\left({\bm{z}}-\mathbf{{\bm{\mu}}}_{k}\right)^{T}{\bm{\Sigma}}^{-1}\left({\bm{z}}-\mathbf{{\bm{\mu}}}_{k}\right)\}.

We also experiment with a Relative Mahalanobis distance variant (RMaha) (Ren et al., 2021), a modified version of Mahalanobis distance which corrects for the background confounding effect using another Gaussian distribution fitted using entire training data ignoring class labels. The uncertainty score is defined as RMDk(z)=MD(z)−MD0(z)\text{RMD}_{k}({\bm{z}})=\text{MD}({\bm{z}})-\text{MD}_{0}({\bm{z}}), where MD0(z)\text{MD}_{0}({\bm{z}}) indicates the Mahalanobis distance to a distribution fitted to the entire training data not considering the class labels: N(μ0,Σ0)\mathcal{N}(\mathbf{{\bm{\mu}}}_{0},{\bm{\Sigma}}_{0}), where μ0=1N∑i=1Nzi{\bm{\mu}}_{0}=\frac{1}{N}\sum_{i=1}^{N}{\bm{z}}_{i} and Σ0=1N∑i=1N(zi−μ0)(zi−μ0)T{\bm{\Sigma}}_{0}=\frac{1}{N}\sum_{i=1}^{N}\left({\bm{z}}_{i}-{\bm{\mu}}_{0}\right)({\bm{z}}_{i}-{\bm{\mu}}_{0})^{T}. This is a good proxy for the background distribution.

Zero-shot results are displayed in Table 13. Interestingly, the zero-shot setting achieves performance close to that of the finetuned setting. For example, in the challenging near-OOD task CIFAR-100 vs CIFAR-10, the best zero-shot AUROC achieved by Plex L (I21K) is 0.915, while its performance with finetuning is 0.934, and the best model which uses JFT and finetuning achieves 0.954 (Table 9). For the far-OOD task CIFAR-100 vs SVHN, the best AUROC achieved by zero-shot is 0.933, and the best model which uses JFT and finetuning achievs 0.938.

The Relative Mahalanobis distance also significantly outperforms the raw Mahalanobis distance, suggesting that there are background features that confound the raw Mahalanobis distance. In fact, the raw Mahalanobis distance does not work well for most of the pretrained models, mostly having AUROC lower than 0.9 and varies depending on the model types.

Conclusion

We presented a framework for thinking about reliable deep learning, and provided a number of tasks and datasets for stress-testing the reliability of models through multiple tasks such as the ability to quantify its confidence in predictions, be robust to distribution shifts, and adapt quickly to new distributions. To improve reliability, we also developed Plex, pretrained large model extensions for vision (ViT-Plex) and language (T5-Plex) that significantly improve the reliability in deep learning. Our techniques are broadly applicable for other large models (cf. LaMDA (Thoppilan et al., 2022), PaLM (Chowdhery et al., 2022)) and wider suite of tasks, cf. BIG-bench (Srivastava et al., 2022).

Acknowledgements

We thank Ben Adlam, Dilip Krishnan, Ed Chi, Neil Houlsby, Rif A. Saurous, and Sharat Chikkerur for helpful discussions and feedback on earlier drafts of the paper. We also thank Tom Small and Ajay Nainani for assisting with visualizations.

References

A Author Contributions

Andreas Kirsch: Active Learning only – joint with Joost van Amersfoort. Implemented Active Learning loop, various acquisition functions and datasets, and ran experiments on CIFAR-10/CIFAR-100. Advised on experimental setup for larger datasets, and wrote active learning section.

Balaji Lakshminarayanan: Helped with direction, narrative, getting resources, advising folks on experiment setup (open set recognition, few-shot adaptation, last layer), helped with writing (abstract, intro, contributions) and revising other sections.

Clara Huiyi Hu: Implemented BE + GP in language domain, executed associated experiments and helped building post-analysis colabs. On the vision side, contributed to sweeping and debugging BE/GP variants.

Dustin Tran: Came up with work’s vision, lead overall writing and experiment design, coordinated individual leads and contributors. Designed and tuned Plex, various figures and plotting code, and core infra setup.

Du Phan: Lead and designed the infrastructure to evaluate uncertainty methods on the language domain. Implemented methods None, DE, MCD, GP, DE-GP, Het for language and performed most corresponding experiments across MNLI, ToxicComments, NaLUE tasks.

D. Sculley: Advisor on project, including its direction and the paper narrative.

Honglin Yuan: Initial experimental setup for CIFAR subpopulation shift experiments, including pipeline for converting data into format to be consumed by uncertainty baselines codebase.

Jasper Snoek: Helped with direction, narrative, getting resources, helped with writing, editing, aggregating results, various plots.

Jeremiah Liu: Lead for language domain efforts. Developed all tasks, datasets, evaluation metrics for language. Implemented initial t5x infra and baseline (None) experiments. Work with eng co-lead Du and teammates Clara and Jie in finishing executing experiments and generating visualizations. Wrote the sections on language. On the vision side, implemented the ViT-GP model, conducted initial pretraining experiments. Assisted editing the section on SNGP method.

Jie Ren: Lead work on open set recognition (vision and text), GPs, and assisted with few-shot. Helped with engineering work including building vision finetuning models, running vision and text experiments. Made the demo figures for vision and text. Helped with Mahalanobis distance-based active learning method.

Joost van Amersfoort: See Andreas Kirsch (joint contribution). Co-wrote compute workflow tutorial for collaborators.

Karan Singhal: Vision subpopulation shift setup and experiments: integrated code for CIFAR-10/100 subpopulation evaluation into codebase, set up, ran, and tuned experiments with different ingredients (Det, BE, BE+HET, BE+GP). Ran ablations: omitting last-layer changes, omitting ensembling during pretraining, ImageNet-only pretraining, no pretraining. Created ablations figure and added discussion. Made detailed edits to paper.

Kehang Han: Responsible for finetuning experiments involving GP and few-shot experiments. Created Imagenet2012Fewshot TFDS (https://www.tensorflow.org/datasets/catalog/imagenet2012_fewshot). Primary author of Sections 5.3.2 and 5.3.3. Helped with GP model ablations.

Kelly Buchanan: Contributed to early formation of project and codebase, particularly around tasks involving structured output evaluation and uncertainty.

Kevin Murphy: Advisor on project, including its direction and the paper narrative.

Mark Collier: Responsible for most experimental results involving heteroscedastic model e.g. Het upstream, Het→\toHet, None→\to Het, BE→\toBE++Het. Primary author of section Section H. Implemented CIFAR-10H and ImageNet ReaL-H label uncertainty tasks. Primary author of Section 5.1.4. Ran sensitivity analysis (together with Mike) to the number of JFT pretraining epochs and volume of pretraining data (JFT4B).

Mike Dusenberry: Engineering lead. Designed, built, and owned infrastructure for Plex, including reproducible training scripts, utilities, experiment scaling components, and open-sourcing efforts. Trained and tuned None ViT baselines (e.g., L) to match existing results. Scaled up experiments to JFT-4B. Ablated JFT dataset size and number of epochs with Mark Collier. Ran experiments and wrote Section 5.2.2 for generalization to covariate distribution shifts. Advised and worked on infra and experiments for the active learning section. Co-wrote a compute workflow tutorial for collaborators. Created a demo notebook for ViT-Plex. Open-sourced code and model checkpoints.

Neil Band: Co-led design of experiments and narrative for integration of RETINA experiments, including selective prediction and distribution generalization. Led RETINA implementation and experiments. Co-wrote section on selective prediction and compute workflow tutorial for collaborators. Led open-sourcing of RETINA distribution shift tasks, baseline model checkpoints, evaluation, and plotting utilities.

Nithum Thain: Contributed to experiments analyzing Deep Ensembles and their tradeoff with compute.

Rodolphe Jenatton: Developed the code to evaluate deep ensembles in https://github.com/google/uncertainty-baselines and ran associated experiments (together with Nithum Thain). Developed the code to evaluate models and ensembles in https://github.com/google-research/robustness_metrics. Developed the code to wrap the pjit-based https://github.com/google-research/vmoe in https://github.com/google/uncertainty-baselines, enabling the evaluation of sparse MoE models, namely E3, V-MoE and ensembles thereof—ran corresponding experiments. Contributed to the shaping of the narrative of the ensembling section (together with Zelda Mariet). Helped with the computation of the FLOPs. Improved finetuning numbers by suggesting use of longer finetuning schedule.

Tim G. J. Rudner: Co-led design of experiments and narrative for integration of RETINA experiments, including selective prediction and distribution generalization. Contributed to all RETINA experiments. Led writing of section on selective prediction and contributed to writing of section on calibration AUROC.

Yarin Gal: Advisor on project, including its direction and the paper narrative.

Zachary Nado: Designed and contributed to core infrastructure that was the major platform to iterate on our models, tasks, and other results.

Zelda Mariet: Ran BatchEnsemble experiments, ran ensemble diversity analysis, drew up vision for the ensembling analysis Section 4.3, proposed and ran the upstream vs downstream comparison (Section 4.2). Wrote the code for experimental analysis, hyperparameter tuning, visualization and summarization. Code refactoring for the overall project.

Zi Wang: Led the active learning effort including initial exploration as an engagement with internal teams as a close collaboration between Google Reliable Deep Learning and the Oxford OATML group (with Mike, Andreas, Joost, and Dustin). Contributed to refactoring BE and None to facilitate better integration with AL. Contributed to active learning code (debugging, adding new features etc). Extensively experimented with AL on RDL models, conducted analyses and obtained/wrote final results. Added discussions on the relations between prior learning and pretraining.

Zoubin Ghahramani: Advisor on project, including its direction and the paper narrative.

B Details behind Key Figures

For the task radar plots of Figure 3, we list references used for task-specific state-of-the-art in Table 14.

For the model scaling plot of Figure 6, we aggregate all task metrics under a single scalar between 0 and 100. In order to do this, we normalize all metrics to be between 0 and 100; we then compute an unweighted average. Most metrics are already bounded between 0 and 100: for example, accuracy, expected calibration error (we do 100−ECE100-\text{ECE} so higher is better), Calibration AUROC, and OOD AUROC. The one exception are scoring rules such as log-loss and Brier score. Because the output distributions are discrete, log-loss has a lower bound of 0 and an upper bound given by the highest entropy distribution (uniform). Therefore we rescale scoring rule values based on their lower and upper bounds so that they’re now between 0 and 100 and so that higher is better.

C Additional Details on Tasks & Datasets

We summarize the list of metrics in Table 15.

For out-of-distribution evaluation in vision, we use the following datasets.

CIFAR-10: CIFAR-10-C (Hendrycks and Dietterich, 2019).

CIFAR-100: CIFAR-100-C (Hendrycks and Dietterich, 2019).

ImageNet1K: ImageNet-A (Hendrycks et al., 2021c), ImageNet-C (Hendrycks and Dietterich, 2019), ImageNetV2 (Recht et al., 2019), ImageNet-Vid-Robust, YTBB Robust (Shankar et al., 2021), ObjectNet (Barbu et al., 2019), and ImageNet-R (Hendrycks et al., 2021a).

RETINA: RETINA’s Country Shift dataset (Band et al., 2021). We train models on images of retinas obtained from patients in the United States (EyePACS) and evaluate trained models on images of retinas obtained from patients in India using different collection equipment (APTOS).

RETINA: RETINA’s Severity Shift dataset (Band et al., 2021). We train models on images of retinas exhibiting no worse than mild diabetic retinopathy, and consider a shifted evaluation dataset with images of moderate diabetic retinopathy or worse. The evaluation data contains features not contained in the training images, such as vitreous hemorrhages. The motivation for this shift is that images of retinas with more severe retinopathy are relatively scarce and that it is likely for a model to be trained only on more widely-available images of retinas exhibiting mild cases of diabetic retinopathy.

Label uncertainty. We use CIFAR-10H (Peterson et al., 2019) which captures human uncertainty over labels for CIFAR-10 dataset. We also construct a larger-scale variant, which we call ImageNet ReaL-H. Individual human ratings were recollected for the original ImageNet test set, available as raw data from ImageNet ReaL (Beyer et al., 2020), and we use them to newly construct soft label targets.

Subpopulation shift. We use Semantically Partitioned CIFAR-10/100 (Yuan et al., 2022) for vision subpopulation shift. CIFAR-10/100 test data is partitioned into semantically similar subpopulations, where each subpopulation has its own data-generating distribution sampled from a meta subpopulation distribution. We aim to improve predictive performance on tail subpopulations.

For covariate and semantic shift in language, we use:

the MNLI-mismatched (Williams et al., 2017) data as the OOD set for NLI, which contains sentence pairs from 5 genres that are distinct from those in MNLI training data.

the CivilComments corpus (Borkan et al., 2019) as the OOD set for toxic comment prediction, which consists of one million public comments appearing on approximately 50 English-language news sites across the world.

We also analyze language subpopulation shift where long-tail groups from the in-distribution set are a desired area for generalization.

HANS (McCoy et al., 2019) eval datasets for NLI, which contains template-generated examples attacking the surface-level heuristics that the neural models are found to rely on when predicting entailment relationships.

CivilCommentsIdentity (Borkan et al., 2019) for toxic comments, which is a subset of CivilComments that has explicit mention of social identities (e.g., muslim, LGBTQ, etc) that the model are often found to generate mispredictions.

NaLUE-tail dataset for CLU, which is a subset of NaLUE corresponding to utterances from 28 low-frequency intents categories.

D Details of Plex ingredients

BatchEnsembles (Wen et al., 2020b) approximate deep ensembles (Lakshminarayanan et al., 2017), but reduce their computational and memory costs by sharing weights across the ensemble members. The weight matrix Wi{\bm{W}}_{i} of any given ensemble member ii is written as the Hadamard product of a shared weight matrix W0{\bm{W}}_{0} and a local rank-1 matrix risi⊤r_{i}s_{i}^{\top}:

The vectors r and s are commonly referred to as fast weights.

Unless otherwise stated, Plex applies BE to all layers in the last 2 residual blocks of the network. This idea follows work for mixture of experts (Riquelme et al., 2021).

Unlike ensemble approaches, SNGP proposed by Liu et al. (2020) focuses on improving the uncertainty quality of a neural network given a fixed representation (a.k.a. deterministic uncertainty quantification setting (Van Amersfoort et al., 2020). When applied to a DNN without pretraining, SNGP enhances the DNN uncertainty property by applying spectral normalization to the hidden weights, and replaces the output layer from a dense layer to a random-feature Gaussian process (GP) layer. That is, given hidden representations h(x)h({\bm{x}}), the GP layer enables scalable computation of a GP posterior by applying a random feature approximation ϕ\phi to the predictive function and then a Laplace approximation to the predictive variance:

where (W,b)({\bm{W}},{\bm{b}}) are frozen random weights of the random feature embedding ϕ(x)=cos⁡(Wh(x)+b)\phi({\bm{x}})=\cos({\bm{W}}h({\bm{x}})+{\bm{b}}), and Φ⊤Φ=∑iϕ(xi)ϕ(xi)⊤\Phi^{\top}\Phi=\sum_{i}\phi({\bm{x}}_{i})\phi({\bm{x}}_{i})^{\top} is the covariance of the random feature embedding estimated using the training data.

Liu et al. (2020, 2022) show that this combined technique improves the model’s awareness of the semantic distance between the test and train examples on the data manifold, leading to improved performance in calibration and out-of-domain detection. When applied to a large pretrained DNN, we find it sufficient to only use the last-layer Gaussian process (i.e., omit the spectral normalization regularization), as the pretrained embedding has already provided a semantic-distance-aware representation of the data.

Heteroscedastic last layers are designed to model input-dependent label noise/label uncertainty (a.k.a. aleatoric uncertainty (Kendall and Gal, 2017)) that is present in the data. We use the Heteroscedastic (Het) last layer introduced by Collier et al. (2020, 2021) who place a multivariate Gaussian distribution over the logits in a standard DNN classifier. A low-rank approximation to the K×KK\times K covariance matrix (KK = number of classes/outputs) is made when KK is large and (Collier et al., 2021) further develop a parameter efficient version of the method with parameterization inspired by BE to enable scaling to tens of thousands of classes.

Among the above methods, None, Het, GP, and BE are both memory and compute efficient, since they only require a single forward pass from a single model to compute the output distribution. We also experiment with deep ensembles (Section 4.3), which is the most expensive in both memory and compute, since they require forward passes from multiple trained models, as well as Monte Carlo Dropout (Gal and Ghahramani, 2016).

E Additional Results for Open Set Recognition

While we use MSP as the basic uncertainty score in the main text, we also would like to discuss other alternative uncertainty scores, and how is their performance improved by the Plex models. We study the following 4 more uncertainty scores,

Mahalanobis distance (Maha) (Lee et al., 2018) which measures the distance between the test input and the fitted training distribution in the embedding space. The training distribution is fitted using a class conditional Gaussian. The uncertainty score is the distance.

Relative Mahalanobis distance (Ren et al., 2021) is a modified version of Mahalanobis distance which corrects for the background confounding effect using another Gaussian distribution fitted using entire training data ignoring class labels.

Entropy of the softmax probability. High entropy suggests high uncertainty. So the uncertainty score is entropy.

Maximum over the logits (MaxLogit) (Hendrycks et al., 2019) which uses the maximum of the un-normalized logits as the confidence score. Then the uncertainty score is then the negative of the MaxLogit.

We study the performance of five different OOD detection methods, MSP, Maha, RMaha, Entropy and MaxLogit, across the various tasks and model types. As shown in Table 16, Mahalanobis distance outperforms all other scores, and Relative Mahalanobis distance performs the second best.

Despite the good performance, one drawback of Mahalanobis distance based method is that its computational time is linear to the number of classes. For the ImageNet2012 vs Places365 task, the in-distribution data ImageNet2012 has 1,000 classes where by definition of Mahalanobis distance method, we need to fit 1K Gaussians for each of the classes and compute the distance between the test input to each of the fitted Gaussian. It becomes very time consuming and not scalable. Therefore in that task, we study the lightweight OOD scores, including MSP, Entropy, and MaxLogit. Interestingly, we noticed that MaxLogit is the best among the three methods for single models (None, GP, None→\toGP) except for Het and None→\toHet, while Entropy is the best for ensemble models (BE, DE based models).

Overall we suggest using Mahalanobis distance based methods to achieve the best performance, but in case there is a computation budget, MaxLogit is a good choice for single models and Entropy is a good choice for ensemble models.

F Additional Results for Selective Prediction

We performed model selection on pretrained models finetuned with different hyperparameters using the area under the in-distribution selective-prediction accuracy curve. We evaluated the predictive performance of the finetuned models on in-distribution and distributionally shifted data. In particular, we computed models’ predictive accuracy, negative log-likelihood, expected calibration error, and the area under the selective prediction curve obtained by computing the area under the ROC curve for selective prediction referral rates from 0% to 99% using different evaluation metrics (accuracy and, where applicable, AUROC and AUPRC).

G Extensions to Sparse Mixtures of Experts

In this section, we discuss extensions of some approaches of this paper in the light of recent advances in sparse mixtures of experts models (sparse MoEs).

Sparse MoEs constitute a family of models that rely on conditional computation (Bengio et al., 2013, 2015) to combine multiple submodels—referred to as experts—in an input-dependent fashion. Sparse MoEs have successfully been applied both in NLP (Shazeer et al., 2017; Lepikhin et al., 2021; Fedus et al., 2022) and more recently in computer vision (Riquelme et al., 2021; Lou et al., 2021). The goal of conditional computation is to grow the number of parameters of the model while maintaining its training and inference costs constant, by only activating a particular subset of experts given an input. This is in contrast with classical neural networks that use all their parameters for each input.

In the context of this paper, we focus on Allingham et al. (2021) that studied the robustness and uncertainty estimates of several extensions of ViT endowed with conditional computation, namely:

V-MoE: The authors of Riquelme et al. (2021) extended ViT to sparse MoEs by placing experts at the level of the MLPs of the transformer architecture. Since ViT operates on image patches, the conditional computation of V-MoEs also operates on image patches. More specifically, the image patches are sparsely routed to KK out of EE available experts (in practice, E=32E=32). The experts are selected by computing some gating weights in combination with a top-KK strategy, allowing for an end-to-end training procedure. When only a single expert is selected (K=1K=1), the training and inference costs of V-MoEs almost match those of a standard ViT model (Riquelme et al., 2021), while leading to improved predictive performance at both pretraining and finetuning time.

In the next section, we will use the acronym MoE to refer to V-MoE.

Deep ensembles of V-MoEs: Allingham et al. (2021) explored the combination of the static ensembling of deep ensembles and the conditional computation of V-MoEs. They found that both effects are complementary. In particular, over tasks where either V-MoEs or deep ensembles are known to perform well (e.g., respectively, few-shot classification and OOD detection), the combined approach inherits from the best of both worlds, while matching the cost of standard deep ensembles.

In what follows, we will use the acronym [\textscMoE]4[\textsc{MoE}]_{4} to refer to a deep ensemble formed by four MoEs where both the upstream and downstream models were varied. In the situation where only the downstream models were varied for a single fixed pretrained model, we use the notation \textscMoE→[\textscMoE]4\textsc{MoE}\rightarrow[\textsc{MoE}]_{4}. We will use the same convention for the model without any changes, None. While the rest of the paper uses ensembles of size 3, we consider in this section ensembles of size 4 to be faithful to the setting of Allingham et al. (2021).

E3: The authors of Allingham et al. (2021) further designed an efficient ensemble approach, referred to as E3, to overcome the computational burden of the naive ensembling of V-MoEs described above. In a nutshell, E3 jointly learns an ensemble of smaller sparse MoEs where all layers not equipped with experts (e.g., attention layers) are shared across the ensemble members. Interestingly, E3 features some of the extensions of Plex due to some of its conceptual similarities with BE. Allingham et al. (2021) observed that E3 tends to lie either on, or close to, the Pareto frontiers of several metrics—e.g., NLL, Brier score and few-shot classification error—versus the computational cost of the models, as measured by FLOPs.

G.2 Results

We present the results in Table 18. They extend the evaluation from Allingham et al. (2021) to the additional metrics of this paper related to the computer vision tasks. We summarize the results with the prediction, uncertainty and adaptation scores.

In agreement with the conclusions reported in Allingham et al. (2021), we can see that [\textscMoE]4[\textsc{MoE}]_{4} outperformsA closer inspection at the results shows that the prediction score of [\textscMoE]4[\textsc{MoE}]_{4} is slightly worse than that of [\textscNone]4[\textsc{None}]_{4} because of isolated, slightly worse performance on CIFAR-10. the standard deep ensemble [\textscNone]4[\textsc{None}]_{4} while having the same computational cost. This confirms the complementary of static ensembling with the adaptivity of sparse MoEs within the broader evaluation of this paper. More generally, combining ensembling and sparse MoEs seems to be a promising direction to improve the reliability of models.

As observed in Riquelme et al. (2021), MoE has strong performance in fewshot learning tasks, which is highlighted here by the adaptation scores of MoE and [\textscMoE]4[\textsc{MoE}]_{4}. Moreover, Allingham et al. (2021) further observed that MoE did not perform well in terms of calibration, as illustrated by its lower uncertainty score compared with None. The efficient ensemble E3 has performance comparable to deep ensembles formed from a single pretrained model (i.e., \textscNone→[\textscNone]4\textsc{None}\rightarrow[\textsc{None}]_{4} and \textscMoE→[\textscMoE]4\textsc{MoE}\rightarrow[\textsc{MoE}]_{4}), while being substantially cheaper. This conclusion echoes the take-home messages from Allingham et al. (2021).

H Analysis of Heteroscedastic Last Layer

We add a heteroscedastic output layer to the base deterministic model and the BE base model. We assess whether enabling the network to model input-dependent (heteroscedastic) label noise results in better performance on datasets known to have noisy labels.

In Figure 30 and Figure 31 the heteroscedastic model, when applied on top of the base deterministic model, provides performance gains on the JFT dataset while performance is neutral on ImageNet-21K.

H.2 Impact of heteroscedastic last layer downstream

In Figure 32, Figure 33 and Figure 34 we look at the core in-distribution performance metrics for downstream vision datasets. Het→\toHet refers to a model trained with a heteroscedastic output layer upstream on JFT and finetuned on the target dataset with a heteroscedastic head. None→\toHet is similar but where the cheaper and simpler baseline deterministic model without a heteroscedastic head is used upstream. BE→\toBE + Het refers to pretraining using a Batch Ensemble model upstream on JFT and finetuning with a BE model with a heteroscedastic head downstream.

The outperformance of a deterministically pre-trained model finetuned with a heteroscedastic last layer, indicates that the gains from a heteroscedastic model downstream can be realized without heteroscedastic upstream pretraining. This is a surprising result given that the heteroscedastic model improves JFT performance upstream. However this result suggests that when the representations are transferred from this upstream model the Deterministic model is sufficient. Given that heteroscedastic model is more expensive than the deterministic model and requires tuning the temperature and the rank for the low-rank approximation hyperparameters this is an encouraging result.

None→\toHet outperforms None on all metrics for ImageNet and CIFAR-100. Performance on CIFAR-10 is saturated and similar. This demonstrates that applying the heteroscedastic head on the downstream model improves downstream performance. Further, see Figure 7 which shows that on a wide variety of reliability metrics beyond in-distribution accuracy, NLL and ECE None→\toHet outperforms None. We will also see later that the heteroscedastic method models label uncertainty better than baseline models. Heteroscedastic can be combined with BE upstream pretraining for excellent results. The results in the above table and the more comprehensive summary results in Figure 7 show that the heteroscedastic head can be combined with a BE model to give the best performance single model on vision tasks (this is Plex). BE→\toBE++Het (Plex) combines the epistemic uncertainty model capability of Batch Ensembles with the label uncertainty modeling of the heteroscedastic method. Plex approaches Deep Ensembles in terms of overall reliability, despite being substantially cheaper. Therefore we can recommend this model as our best single model for use by practitioners.

I Summarization of Language Results

We first compare the performance across types of uncertainty methods, fixing the architecture size to T5-base. We compare performances in prediction, uncertainty calibration, and human-model collaboration, across all datasets (MNLI, NaLUE and Toxic Comments) and all splits (In-domain, OOD, and tail-population). Table 20 reports the full results, and Figure 35 summarizes the rankings of uncertainty methods under each type of population shift (in-domain v.s. OOD v.s. tail-population). Among all methods, DE++GP, Plex (i.e., BE++GP), BE, and MC Dropout tend to have the strongest performance. In particular, DE++GP almost always dominates the other methods on MNLI and NaLUE, and remains competitive in the case of label imbalance (i.e., Toxic Comments). However, DE++GP is an expensive method that costs x10 more in memory and compute and therefore is not competitive in scale (a more thorough analysis is in Section 4.3). On the other hand, among the more efficient, single-model methods, BE and Plex perform well on MNLI and NaLUE (notably, outperform the most expensive DE), while MCD stands out in the Toxic Comments. The above observations suggest that, when the training example has a simple distribution, quantifying output-layer uncertainty alone is sufficient to attain strong performance. However, when there are pathologies in the data distribution (e.g., extreme label imbalance), quantifying the uncertainty within the model’s intermediate representations (e.g., via some form of perturbation like BE) becomes important.

We then investigate how a model’s uncertainty performance is impacted by the architecture size. For model size scaling, we evaluate Plex, None, and MC Dropout, the three best-performing and efficient methods in the previous study. We evaluate the performance of each method under three progressively larger architectures: T5 S, T5 B, and T5 L, and observe how the behavior changes across the method and with respect to the architecture size. Table 20 reports the full results, and Figure 36 summarizes the rankings of uncertainty methods organized by the sizes of the architecture. As shown, comparing across architecture sizes, we see a larger architecture almost always leads to stronger performance in collaborative performance. This trend remains largely consistent even when out-of-distribution.