Just Train Twice: Improving Group Robustness without Training Group Information

Evan Zheran Liu, Behzad Haghgoo, Annie S. Chen, Aditi Raghunathan, Pang Wei Koh, Shiori Sagawa, Percy Liang, Chelsea Finn

Introduction

The standard approach of empirical risk minimization (ERM)—training machine learning models to minimize average training loss—can produce models that achieve low test error on average but still incur high error on certain groups of examples (Hovy & Søgaard, 2015; Blodgett et al., 2016; Tatman, 2017; Hashimoto et al., 2018; Duchi et al., 2019). These performance disparities across groups can be especially pronounced in the presence of spurious correlations. For example, in the task of classifying whether an online comment is toxic, the training data is often biased so that mentions of particular demographics (e.g., certain races or religions) are correlated with toxicity. Models trained via ERM then associate these demographics with toxicity and thus perform poorly on groups of examples in which the correlation does not hold, such as non-toxic comments mentioning a particular demographic (Borkan et al., 2019). Similar performance disparities due to spurious correlations occur in many other applications, including other language tasks, facial recognition, and medical imaging (Gururangan et al., 2018; McCoy et al., 2019; Badgeley et al., 2019; Sagawa et al., 2020a; Oakden-Rayner et al., 2020).

Following prior work, we formalize this setting by considering a set of pre-defined groups (e.g., corresponding to different demographics) and seeking models that have low worst-group error (Sagawa et al., 2020a). Previous approaches typically require annotations of the group membership of each training example (Sagawa et al., 2020a; Goel et al., 2020; Zhang et al., 2020). While these approaches have been successful at improving worst-group performance, the required training group annotations are often expensive to obtain; for example, in the toxicity classification task mentioned above, each comment has to be annotated with all the demographic identities that are mentioned.

In this paper, we propose a simple algorithm, Jtt (Just Train Twice), for improving the worst-group error without training group annotations, instead only requiring group annotations on a much smaller validation set to tune hyperparameters. Jtt is composed of two stages: we first identify training examples that are misclassified by a standard ERM model, and then we train the final model by upweighting the examples identified in the first stage. Intuitively, this procedure exploits the observation that sufficiently-regularized ERM models tend to incur high worst-group training error (and subsequently high worst-group test error). This makes selecting misclassified examples an effective heuristic for identifying examples from groups that ERM models fail on, such as minority groups. Since the final classifier upweights such examples, it performs better on such groups and achieves better minority group performance.

We evaluate Jtt on two image classification datasets with spurious correlations, Waterbirds (Wah et al., 2011; Sagawa et al., 2020a) and CelebA (Liu et al., 2015) and two natural language processing datasets, MultiNLI (Williams et al., 2018) and CivilComments-WILDS (Borkan et al., 2019; Koh et al., 2021). We use the versions of Waterbirds, CelebA, and MultiNLI from Sagawa et al. (2020a), where in Waterbirds, the label waterbird or landbird spuriously correlates with water in the background; in CelebA, the label blond or non-blond spuriously correlates with binary gender; and in MultiNLI, the label spuriously correlates with the presence of negation words. In CivilComments-WILDS, where the input is online comments, the label toxic, non-toxic spuriously correlates with the mention of particular demographics, as discussed above. Our method outperforms ERM on all four datasets, with an average worst-group accuracy improvement of 16.2%, while maintaining competitive average accuracy (only 4.2% worse on average). Furthermore, despite having no group annotations during training, Jtt closes 75% of the gap between ERM and group DRO, which uses complete group information on the training data.

We then empirically analyze Jtt. First, we analyze the examples identified by Jtt and show that Jtt upweights groups on which standard ERM models perform poorly, e.g., minority groups that do not have the spurious correlation (such as waterbirds on land in the Waterbirds dataset). Second, we show that having validation group annotations is essential for hyperparameter tuning for Jtt and other related algorithms.

Finally, we compare Jtt with the distributionally robust optimization (DRO) algorithm that minimizes the conditional value at risk (CVaR). CVaR DRO aims to train models that are robust to a wide range of potential distribution shifts by minimizing the worst-case loss over all subsets of the training set of a certain size (Duchi et al., 2019). This objective does not require training group annotations, and it can be optimized by dynamically upweighting training examples with the highest losses in each minibatch (Levy et al., 2020). Though CVaR DRO and Jtt share conceptual similarities—they both upweight high loss training points and do not require training group information—the difference is that Jtt upweights a static set of examples, while CVaR DRO dynamically re-computes which examples to update. Empirically, we find that Jtt empirically substantially outperforms CVaR DRO on worst-group accuracy in the datasets we tested.

Related Work

In this paper, we focus on group robustness (i.e., training models that obtain good performance on each of a set of predefined groups in the dataset), though other notions of robustness are also studied, such as adversarial examples (Biggio et al., 2013; Szegedy et al., 2014) or domain generalization (Blanchard et al., 2011; Muandet et al., 2013). Approaches for group robustness fall into the two main categories we discuss below.

Several approaches leverage group information during training, either to combat spurious correlations or handle shifts in group proportions between train and test distributions. For example, Mohri et al. (2019); Sagawa et al. (2020a); Zhang et al. (2020) minimize the worst-group loss during training; Goel et al. (2020) synthetically expand the minority groups via generative modeling; Shimodaira (2000); Byrd & Lipton (2019); Sagawa et al. (2020b) reweight or subsample to artificially balance the majority and minority groups; Cao et al. (2019, 2020) impose heavy Lipschitz regularization around minority points. These approaches substantially reduce worst-group error, but obtaining group annotations for the entire training set can be extremely expensive.

Another line of work studies worst-group performance in the context of fairness (Hardt et al., 2016; Woodworth et al., 2017; Pleiss et al., 2017; Agarwal et al., 2018; Khani et al., 2019). While these works also aim to improve the worst-group loss, they explicitly focus on equalizing the loss across all groups.

We focus on the setting where group annotations are expensive and unavailable on the training data, and potentially only available on a much smaller validation set. Many approaches for this setting fall under the general DRO framework, where models are trained to minimize the worst-case loss across all distributions in a ball around the empirical distribution (Ben-Tal et al., 2013; Lam & Zhou, 2015; Duchi et al., 2016; Namkoong & Duchi, 2017; Oren et al., 2019). Pezeshki et al. (2020) modify the dynamics of stochastic gradient descent to avoid learning spurious correlations. Sohoni et al. (2020) automatically identify groups based by clustering the data points. Kim et al. (2019) propose an auditing scheme that searches for high-loss groups defined by a function within a pre-specified complexity class and postprocess the model to minimize discrepancies identified by the auditor. Khani et al. (2019) minimize the variance in the loss across all data points to encourage lower discrepancy in the losses across all possible groups. Another approach is to directly learn how to reweight the training examples either using small amounts of metadata (Shu et al., 2019) or automatically via meta-learning (Ren et al., 2018).

Most closely related to Jtt are several approaches that also train a pair of models, where the performance of the first model is used to help train the second model (Yaghoobzadeh et al., 2019; Utama et al., 2020; Nam et al., 2020). We compare Jtt with one such approach, called Learning from Failure (LfF) (Nam et al., 2020). In LfF, the first model is intentionally biased and tries to identify minority examples where the spurious correlation does not hold. The identified examples are then upweighted while training the second model. This approach interleaves the updates of both models and requires an intentional biasing with the the first model. In contrast, our approach of Jtt is simpler, though conceptually similar: we only identify points to upweight once (i.e. no interleaved updates which generally destabilize training), and we just perfom standard ERM with regularization to identify points without any artificial biasing. Empirically, despite its simplicity, Jtt performs better than LfF.

Concurrently, Creager et al. (2021) also proposed a method that leverages similar intuition to Jtt. This work first uses the errors of a standard ERM model to infer group labels, similar to Jtt. Then, they learn a model that is invariant to the predicted labels.

Preliminaries

We consider the setting of classifying an input x∈Xx\in\mathcal{X} as a label y∈Yy\in\mathcal{Y}. We are given nn training points {(x1,y1),…,(xn,yn)}\{(x_{1},y_{1}),\ldots,(x_{n},y_{n})\}. Our goal is to learn a model fθ:X→Yf_{\theta}:\mathcal{X}\rightarrow\mathcal{Y}, parameterized by θ∈Θ\theta\in\Theta. We measure performance across a set of pre-defined groups G\mathcal{G}. Each point (x,y)(x,y) belongs to some group g∈Gg\in\mathcal{G} and we evaluate classifiers on their worst-group error defined as follows:

where l0−1(x,y;θ)=1[fθ(x)≠y]l_{0-1}(x,y;\theta)=\mathbf{1}[f_{\theta}(x)\neq y] is the 0-1 loss.

We are interested in the setting where we do not have group annotations on training points because they are expensive to obtain. Our goal is to achieve good worst-group error at test time without training group annotations. However, we are given a small validation set of mm points with group annotations {(x1,y1,g1),…,(xm,ym,gm)}\{(x_{1},y_{1},g_{1}),\ldots,(x_{m},y_{m},g_{m})\}. These group annotations allow us to compute the worst-group validation error, which we use to tune hyperparameters.

In our experiments, we primarily consider the setting where each group g=(a,y)∈Gg=(a,y)\in\mathcal{G} is defined by the label yy and a spurious attribute a∈Aa\in\mathcal{A} that spuriously correlates with the label (i.e., G=A×Y)\mathcal{G}=\mathcal{A}\times\mathcal{Y}). Figure 1 illustrates the four groups on the Waterbirds dataset, where the background spuriously correlates with the label.

2 Comparisons

Here, we describe four other algorithms that we use as comparisons in this paper: (i) Empirical risk minimization (ERM), which is the standard approach for training machine learning models by minimizing the average training loss (Section 3.2.1). (ii) A distributionally robust optimization (DRO) method for minimizing the conditional value at risk (CVaR), which seeks to minimize error over all groups above a certain size (Duchi et al., 2019), and is a natural approach to training models with low worst-group error without group annotations (Section 3.2.2). (iii) Learning from Failure (LfF) (Nam et al., 2020), a recent approach that is conceptually similar to Jtt (Section 3.2.3). (iv) Group DRO (Sagawa et al., 2020a), which—unlike all of the preceding methods—uses training group annotations, and can therefore be considered as an oracle method that upper bounds the performance we might expect from methods that do not use training group annotations (Section 3.2.4).

2.2 Distributionally robust optimization of the conditional value at risk (CVaR DRO)

Instead of minimizing the expected loss over the empirical training distribution, distributionally robust learning algorithms define an uncertainty set over distributions that are within some distance of the empirical training distribution, and then minimize the expected loss over the worst-case distribution in this uncertainty set (Duchi et al., 2019).

In this paper, we study a classic instance of this type of worst-case loss known as the conditional value at risk (CVaR) at level α∈(0,1]\alpha\in(0,1], which corresponds to an uncertainty set that contains all α\alpha-sized subpopulations of the training distribution (Rockafellar & Uryasev, 2000). The idea is that the worst loss over α\alpha-sized subpopulations upper bounds the worst-group loss over the (unknown) groups in G\mathcal{G} when α\alpha is close to the size of the smallest group in G\mathcal{G}.

Note that the CVaR objective is equivalent to the average loss incurred by the α\alpha-fraction of training points that have the highest loss.

2.3 Learning from Failure (LfF)

LfF attempts to automatically upweight examples from challenging groups, such as those where the spurious correlation does not hold. It does this by learning two models fB(y∣x;θB)f_{B}(y\mid x;\theta_{B}) and fD(y∣x;θD)f_{D}(y\mid x;\theta_{D}), parameterized by θB\theta_{B} and θD\theta_{D}.

The first model fBf_{B} is trained with ERM using generalized cross-entropy (GCE) loss (Zhang & Sabuncu, 2018):

where q∈[0,1)q\in[0,1) is a hyperparameter. Compared to standard cross-entropy loss, the gradient of GCE loss upweights examples where fB(yi∣xi;θB)f_{B}(y_{i}\mid x_{i};\theta_{B}) is large, which intentionally biases fBf_{B} to perform better on easier examples and poorly on examples group challenging groups.

The second model is also trained with ERM, using cross-entropy loss, where each example (xi,yi)(x_{i},y_{i}) is reweighted by a factor of:

The hope is that early in training, log⁡fB(yi∣xi)\log{f_{B}(y_{i}\mid x_{i})} will be smaller than log⁡fD(yi∣xi)\log{f_{D}(y_{i}\mid x_{i})} on the easier examples, which leads to smaller weights on the easier examples and larger weights on the challenging examples.

2.4 Group distributionally robust optimization (Group DRO)

where ngn_{g} is the number of training points with group gi=gg_{i}=g.

Jtt: Just Train Twice

We now present Jtt, a simple two-stage approach that does not require group annotations at training time. In the first stage, we train an identification model and select examples with high training loss. Then, in the second stage, we train a final model while upweighting the selected examples.

The key empirical observation that Jtt builds on is that sufficiently low complexity ERM models tend to fit groups with easy-to-learn spurious correlations (e.g., landbirds on land and waterbirds on water in the Waterbirds dataset), but not groups that do not exhibit the same correlation (e.g., waterbirds on land) (Sagawa et al., 2020a). We therefore use the simple heuristic of first training an identification model f^id\hat{f}_{\text{id}} via ERM and then identifying an error set EE of training examples that f^id\hat{f}_{\text{id}} misclassifies:

Next, we train a final model f^final\hat{f}_{\text{final}} by upweighting the points in the error set EE identified in step one:

Overall, training Jtt is summarized in Algorithm 1. In practice, to restrict the capacity of the identification model, we only train it for TT steps, where TT is a hyperparameter (line 1). This prevents it from potentially overfitting the training data and yielding an empty error set. To implement the upweighted objective (8), we simply upsample the examples from the error set by λup\lambda_{\text{up}} (line 3) and train the final model on the upsampled data (line 4). Specifically, in each epoch of training, we sample each example from the error set λup\lambda_{\text{up}} times and all other examples only once.

Experiments

In our experiments, we first demonstrate that Jtt substantially improves worst-group performance compared to standard ERM models (Section 5.2). We also show that it recovers a significant fraction of the performance gains yielded by group DRO, which, as discussed in Section 3, is an oracle that relies on group annotations on training examples. We then present empirical analysis of Jtt, including the analysis of the error set (Section 5.3), exploration on the role of the validation set (Section 5.4), and comparison with CVaR DRO (Section 5.5).

We study four datasets in which prior work has observed poor worst-group performance due to spurious correlations (Figure 2). Full details about these datasets are in Appendix B.

Waterbirds (Wah et al., 2011; Sagawa et al., 2020a): The task is to classify images of birds as “waterbird” or “landbird”, and the label is spuriously correlated with the image background, which is either “land” or “water.”

CelebA (Liu et al., 2015): We consider the task from Sagawa et al. (2020a) of classifying the hair color of celebrities as “blond” or “not blond.” The label is spuriously correlated with gender, which is either “male” or “female.”

MultiNLI (Williams et al., 2018): Given a pair of sentences, the task is to classify whether the second sentence is entailed by, neutral with, or contradicts the first sentence. We use the spurious attribute from Sagawa et al. (2020a), which is the presence of negation words in the second sentence; due to the artifacts from the data collection process, contradiction examples often include negation words.

CivilComments-WILDS (Borkan et al., 2019; Koh et al., 2021): The task is to classify whether an online comment is toxic or non-toxic, and the label is spuriously correlated with mentions of certain demographic identities (male, female, White, Black, LGBTQ, Muslim, Christian, and other religion). We use the evaluation metric from Koh et al. (2021), which defines 16 overlapping groups (a,toxic)(a,\emph{toxic}) and (a,non-toxic)(a,\emph{non-toxic}) for each of the above 8 demographic identities aa, and report the worst-group performance over these groups.

We aim to answer two main questions: (1) How does Jtt compare with other approaches that also do not use training group information? (2) How does Jtt compare with approaches that do use training group information?

To answer the first question, we compare Jtt with ERM, CVaR DRO, and a recently proposed approach called Learning from Failure (LfF) (Nam et al., 2020). To answer the second question, we compare Jtt with group DRO (Sagawa et al., 2020a), an oracle that uses training group annotations. For details about these approaches, see Section 3. Note that on CivilComments, group DRO cannot be directly applied on the 16 defined groups, since it is not designed for overlapping groups. Instead, our group DRO minimizes worst-group loss over 4 groups (y,a)(y,a), where the spurious attribute aa is a binary indicator of whether any demographic identity is mentioned and the label yy is toxic or non-toxic. We tune the hyperparameters of all approaches based on worst-group performance on a small validation set with group annotations.

2 Main Results

Table 1 reports the average and worst-group accuracies of all approaches. Compared to other approaches that do not use training group information, Jtt consistently achieves higher worst-group accuracy on all 4 datasets. Additionally, Jtt performs well even relative to approaches that use training group information. In particular, Jtt recovers a significant portion of the gap in worst-group accuracy between ERM and group DRO, closing 75% of the gap on average. As a caveat, we note that simple label balancing also achieves comparable worst-group accuracy to group DRO on CivilComments.

Jtt’s worst-group accuracy improvements come at only a modest drop in average accuracy, averaging only 4.2% worse than the highest average accuracy on each dataset. This drop is consistent with Sagawa et al. (2020a), which observes a tradeoff between average and worst-group accuracies.

3 Error set analysis

We find it surprising that just a small amount of group information on the validation set can allow Jtt to achieve high worst-group accuracy with no knowledge of the groups on the training set. We now probe into how Jtt achieves such high worst-group accuracy. In order to perform this analysis, we use the group annotations on the training data to closely examine what examples are upweighted in the error set identified in the first step of Jtt, though we don’t use such training group annotations for training Jtt.

To start, we define the worst group as the group on which the standard ERM model achieves the lowest test accuracy, when tuned for worst-group validation accuracy. We analyze how well the error set captures this worst group. To do this, we measure precision, the fraction of examples in the error set that belong to the worst group, recall, the fraction of the worst group examples that are included in the error set, and the empirical rate, the rate at which the worst group examples appear in the training data.

As reported in Table 2, we observe that the error set contains worst-group examples at a much higher rate (precision) than they appear in the training dataset (empirical rate). Worst-group examples appear in the error set 2.2x to 15.9x more frequently in the error set than in the training data, across the 4 datasets. In other words, the worst group is significantly enriched in the error set compared to the training dataset, which may explain why Jtt has much better worst-group performance over ERM. Additionally, the error set has high worst-group recall, ranging from 67.1% to 96.9% and averaging to 86.4% across the 4 datasets. Together, these results indicate that the worst group is included in the error set at relatively high both precision and recall.

Empirically, ERM performs poorly on several groups, not just on a single worst group. We therefore next examine what other groups the examples in the error set belong to, beyond the worst group. For each group, we compute two metrics: (i) enrichment, defined as how much more frequently examples from a group appear in the error set than in the training data (i.e., the precision of the group divided by the empirical rate of the group); (ii) the test accuracy that ERM achieves on this group, when tuned for worst-group validation accuracy.

Tables 3 to 5 and Table 12 in Appendix B.4 report these results for Waterbirds, CelebA, MultiNLI, and CivilComments respectively. We observe that the enrichment roughly inversely correlates with ERM’s test accuracy on that group: examples from low performance groups are included at high rates in the error set relative to the empirical rate. This may help Jtt perform better across all groups that ERM performs poorly on, which in turn improves worst-group accuracy.

Finally, we note that while the groups with high enrichment often correspond to groups where the spurious correlation does not hold, this is not always the case. In particular, the waterbird on water background group in Waterbirds and the blonde female group in CelebA have high enrichments, even though the spurious correlation holds in these groups. We hypothesize that this occurs due to label imbalance, since the waterbird label and blonde label are relatively rare and appear only in 23% and 15% of the training examples, respectively. Empirically, upweighting examples from these groups is indeed important for Jtt’s worst-group test accuracy. When we remove all waterbird on water background examples from the error set, Jtt’s worst-group test accuracy drops by 6%. However, while the group composition of the error set (i.e., the fraction of the error set in each group) is important, the exact examples inside the error set do not seem to matter. When we replace each error set example with another example from the same group on Waterbirds, worst-group test accuracy drops by only 0.7%. These two results are shown in Table 6.

4 Hyperparameter Tuning and the Role of the Validation Set

In all of our experiments, we tune the algorithm and model hyperparameters based on the worst-group accuracy on the validation set. In general, across all methods, we found hyperparameter tuning in this fashion to be critical. Table 7 shows that even for CVaR DRO, LfF, and Jtt, which all try to improve worst-group accuracy without relying on training group annotations, the worst-group test accuracies on Waterbirds and CelebA plummet when the hyperparameters are tuned for average accuracy on the validation set, instead of worst-group accuracy on the validation set. Importantly, this means that even though these methods do not require training group annotations, they still require validation group annotations in order to have high worst-group test accuracy. Existing methods for improving worst-group accuracy generally rely on some form of this assumption (either by explicitly tuning for validation group accuracy, or by assuming access to a validation set that is balanced by groups). Removing this reliance is an important direction for future work; we discuss this further in Section 6.

Compared to ERM, Jtt has two additional hyperparameters: the number of epochs to train the identification model TT and the upweight factor λup\lambda_{\text{up}}. As an illustration of the sensitivity to hyperparameters, Figure 3 shows how the worst-group accuracy of Jtt’s final model changes as we vary TT between 20 and 100 epochs on Waterbirds. Worst-group accuracy is high when TT is between 40 and 60, but drops when TT is too small or too large.

In the main experiments in Section 5.2, we use the default validation sets provided with each of the datasets. Using group annotations on these validation sets is already cheaper than training group annotations, as these default validation sets are 2–10x smaller than their corresponding training sets. However, we additionally test if Jtt can continue to achieve high worst-group performance using even smaller validation sets to further reduce the cost of obtaining group annotations on those sets. On Waterbirds and CelebA, we reduce the validation set size by a factor of 1x (no reduction), 15\frac{1}{5}x, 110\frac{1}{10}x, and 120\frac{1}{20}x and tune Jtt’s hyperparameters based on worst-group accuracy on the reduced validation set. We find that Jtt continues to achieve high worst-group accuracy, even when reducing the validation set size by 110\frac{1}{10}x and 120\frac{1}{20}x, amounting to only 119 and 993 total examples, on Waterbirds and CelebA respectively.

5 Comparison with CVaR DRO

In this section, we explore the relation between Jtt and CVaR DRO. Recall that the CVaR objective in Equation 3 is the average loss incurred by the α\alpha-fraction of training examples with the highest loss. We can view minimizing this objective as upweighting this α\alpha-fraction of examples while ignoring the remaining examples. In this way, Jtt is conceptually similar to CVaR DRO: both upweight training points with high loss, without requiring group annotations of training points. However, their empirical performance is widely different: CVaR DRO offers only small gains in worst-group accuracy over ERM, while Jtt offers substantial gains. One key difference is that in Jtt, the set of points that get upweighted EE is computed once during stage 1, and then held fixed. In contrast, minimizing the CVaR objective involves dynamically computing the α\alpha-subset of points with the highest loss at each step, upweighting them and updating the model, and then repeating to update the α\alpha-subset. As we show next, ablating this key difference from Jtt substantially degrades worst-group performance.

We start by observing that the performance of Jtt drops when we dynamically recompute the error set EE, instead of only computing EE once using the identification model. Concretely, we study a variant of Jtt on the Waterbirds dataset: as usual, we first train an identification model for T=50T=50 epochs, but then while training the final model, every KK epochs, we dynamically update the error set EE as the errors of the final model over the training set. Setting KK to be ∞\infty—which means that we only compute the error set EE once after training the identification model for TT epochs—recovers standard Jtt. On the other hand, lowering KK makes the algorithm more similar to minimizing CVaR, since this more frequently updates the upweighted set to be the examples with higher loss under the current model, instead of the examples with higher loss under the static identification model.

Table 9 shows the results as we vary KK between 1010, 2020, 3030, and 5050 epochs on Waterbirds, re-tuning all hyperparameters for each value of KK. At high values of KK, where the error set remains fixed for many epochs, both average and worst-group accuracies are high. However, as KK decreases, the average and worst-group accuracies drop. These results show that at least on Waterbirds, holding the error set fixed appears to be critical for Jtt.

The analysis above suggests that the relatively poor worst-group performance of CVaR DRO might stem from how it dynamically computes which examples to upweight. We further study the behavior of CVaR DRO by analyzing the examples that it upweights. Concretely, throughout CVaR DRO training, we periodically identify the α\alpha-fraction of training examples with the highest loss and measure the worst-group precision and recall, where the worst group is defined as the group on which ERM achieves the lowest test accuracy.

Figure 4 shows the results using the value of α\alpha achieving the highest worst-group validation accuracy: α=0.2\alpha=0.2 on Waterbirds, α=0.00852\alpha=0.00852 on CelebA, and α=0.5\alpha=0.5 on MultiNLI. On Waterbirds, the worst-group examples (which comprise approximately 1% of the training set) make up 19% of the error set for Jtt, whereas they oscillate between 1% and 10% of the worst-α\alpha fraction for CVaR DRO. As a result, Jtt consistently upweights nearly 90% of the worst-group examples, whereas CVaR DRO oscillates between upweighting the worst group and the other groups, upweighting as little as 20% of the examples at some points during training. On CelebA, CVaR DRO upweights the worst-group examples with slightly higher precision than Jtt, but α\alpha is much smaller than the size of the error set; as a result, Jtt upweights nearly 95% of the worst-group examples, whereas CVaR DRO only upweights 13% of them. On MultiNLI, the worst group steadily gets less and less upweighted for CVaR DRO, whereas Jtt upweights it at a higher rate, though it still only comprises 2% of the error set for Jtt.

These results suggest that the CVaR objective might be overly conservative where the α\alpha-fraction of examples with highest loss often include many examples from other groups. Furthermore, the set of examples varies widely across different iterations of training. In contrast, Jtt upweights a fixed set of points. Empirically, we find that this allows Jtt to successfully use the worst-group accuracy on a small validation set to identify error sets that improve accuracy on groups we care about.

Discussion

In this work, we presented Just Train Twice (Jtt), a simple algorithm that substantially improves worst-group performance without requiring expensive group annotations during training. We conclude by discussing several directions for future work.

First, a better theoretical understanding of when and why Jtt works would help us to refine and further develop methods for training models that are less susceptible to spurious correlations. For example, it would be useful to understand why early-stopped ERM models (as in the identification models used by Jtt) seem to consistently latch onto the spurious correlations in our datasets, and why it seems to be important to fix the upweighted set instead of dynamically recomputing it, as in CVaR DRO.

Second, Jtt and many prior methods on robustness without group information all rely on a validation set that is representative of the distribution shift or annotated with group information. While these annotations are significantly cheaper that labeling the entire training set, it still requires the practitioner to be aware of any spurious correlations and define groups accordingly. Doing so may be notably difficult in real-world applications. Therefore, this leaves open the question of whether methods can perform well with mis-specified groups or no group annotations whatsoever.

Finally, while our experiments focus on group robustness in the presence of spurious correlations, Jtt is not specifically tailored to spurious correlations. Given Jtt’s simplicity, it would be straightforward to experiment with Jtt to see if it might improve performance under different types of distribution shifts, such as in domain generalization settings (Blanchard et al., 2011; Muandet et al., 2013).

Our code is publicly available at https://github.com/anniesch/jtt.

Acknowledgements

This work was supported by NSF Award Grant No. 1805310 and in part by Google. EL is supported by a National Science Foundation Graduate Research Fellowship under Grant No. DGE-1656518. AR is supported by a Google PhD Fellowship and Open Philanthropy Project AI Fellowship. SS is supported by a Herbert Kunzel Stanford Graduate Fellowship.

References

Appendix A Training Details

In this section, we detail the model architectures and hyperparameters used by each approach. Within each dataset, we used the same model architecture across all approaches: ResNet-50 (He et al., 2016) for Waterbirds and CelebA, and BERT for MultiNLI and CivilComments (Devlin et al., 2019). For ResNet-50, we used the PyTorch (Paszke et al., 2017) implementation of ResNet-50, starting from ImageNet-pretrained weights. For BERT, we used the the HuggingFace implementation (Wolf et al., 2019) of BERT, also starting from pretrained weights.

We use the LfF implementation released by Nam et al. (2020). We use the group DRO and ERM implementations released by Sagawa et al. (2020a) and also implement CVaR DRO and Jtt on top of this code base, with the CVaR DRO implementation adapted from Levy et al. (2020). For the group DRO experiments on Waterbirds, CelebA, and MultiNLI, we directly use the reported performance numbers from Sagawa et al. (2020a). We note that these numbers utilize group-specific loss adjustments that encourage the model to attain lower training losses on smaller groups, which was shown to improve worst-group generalization. We train our own group DRO model on CivilComments-WILDS as it was not included in Sagawa et al. (2020a); for this, we did not implement these group adjustments. We train our own models for all other algorithms.

For CVaR DRO, we tune the size of the worst-case subpopulation α∈{0.1,0.2,0.5}\alpha\in\{0.1,0.2,0.5\}. For CelebA, we additionally tried α=# smallest group examples# training examples=0.00852\alpha=\frac{\text{\# smallest group examples}}{\text{\# training examples}}=0.00852.

For LfF, we tune the hyperparameter qq by grid searching over q∈{0.1,0.3,0.5,0.7,0.9}q\in\{0.1,0.3,0.5,0.7,0.9\}. For CivilComments, we additionally sample two values log-uniformly from (0,0.1](0,0.1]. This hyperparameter was not tuned in the experiments in Nam et al. (2020).

For Jtt, we additionally tune the number of epochs of training the identification model TT and the upsampling factor λup\lambda_{\text{up}}. While developing Jtt, we tried the following values of TT and λup\lambda_{\text{up}} without looking at test results, though not necessarily all combinations: T∈{1,2,40,50,60}T\in\{1,2,40,50,60\} and λup∈{5,10,20,30,40,50,100,∣training set∣∣error set∣}\lambda_{\text{up}}\in\{5,10,20,30,40,50,100,\frac{|\text{training set}|}{|\text{error set}|}\}. For the final experiments we tune over λup∈{20,50,100}\lambda_{\text{up}}\in\{20,50,100\} for the vision datasets (Waterbirds and CelebA) and λup∈{4,5,6}\lambda_{\text{up}}\in\{4,5,6\} for the NLP datasets (MultiNLI and CivilComments). Additionally, Waterbirds requires more training epochs than the others, due to its much smaller training set size, so we tune over T∈{40,50,60}T\in\{40,50,60\} for Waterbirds and T∈{1,2}T\in\{1,2\} for all other datasets.

ERM has no additional algorithm-specific hyperparameters. For group DRO, we fixed the step size ηq\eta_{q} for updating group weights to its default value of 0.010.01 from Sagawa et al. (2020a), without tuning.

All approaches are optimized for up to 300300 epochs with batch size 6464, using batch normalization (Ioffe & Szegedy, 2015), and no data augmentation. We chose this smaller batch size (compared to the batch size of 128 used in Sagawa et al. (2020a)) for computational convenience. All approaches are optimized with stochastic gradient descent (SGD) with momentum 0.90.9.

Our grid search over α\alpha for CVaR DRO yields α=0.2\alpha=0.2. Our grid search over qq for LfF yields q=0.5q=0.5. Our grid search over TT and λup\lambda_{\text{up}} for Jtt yields T=60T=60 epochs and λup=100\lambda_{\text{up}}=100.

We train all approaches for up to 5050 epochs with batch size 128128, using batch normalization and no data augmentation. Like in Waterbirds, we optimize all approaches with SGD with momentum 0.90.9.

Our grid search over α\alpha for CVaR DRO yields α=0.00852\alpha=0.00852. Our grid search over qq for LfF yields q=0.5q=0.5. Our grid search over TT and λup\lambda_{\text{up}} for Jtt yields T=1T=1 epoch and λup=50\lambda_{\text{up}}=50.

Jtt achieves the highest validation worst-group accuracy using SGD optimization without clipping for the initial model, and using the AdamW optimizer with clipping for the final model. All other approaches achieve highest validation worst-group accuracy using AdamW with clipping. Our grid search over α\alpha for CVaR DRO yields α=0.5\alpha=0.5. Our grid search over qq for LfF yields q=0.1q=0.1. Our grid search over TT and λup\lambda_{\text{up}} for Jtt yields T=2T=2 epochs and λup=6\lambda_{\text{up}}=6.

Jtt achieves the highest validation worst-group accuracy using SGD optimization without clipping for the initial model, and using the AdamW optimizer with clipping for the final model. All other approaches achieve highest validation worst-group accuracy using AdamW with clipping. Our grid search over α\alpha for CVaR DRO yielded α=0.5\alpha=0.5. Our grid search over qq for LfF yielded q=0.00001q=0.00001. Our grid search over TT and λup\lambda_{\text{up}} for Jtt yields T=2T=2 epochs and λup=6\lambda_{\text{up}}=6.

We also note that our group DRO approach uses a different spurious attribute compared to the group DRO results reported in Koh et al. (2021). Our group DRO uses the spurious attribute of any demographic identity being mentioned, while the one in Koh et al. (2021) uses only mentions of the Black demographic. Both perform similarly: ours achieves 0.3% lower worst-group accuracy, but 0.5% higher average accuracy.

Appendix B Dataset Details

We use the Waterbirds dataset introduced by Sagawa et al. (2020a), which is constructed by cropping out images of birds from the CUB dataset (Wah et al., 2011) and pasting them on backgrounds from the Places dataset (Zhou et al., 2017). In this dataset, images of seabirds (albatross, auklet, cormorant, frigatebird, fulmar, gull, jaeger, kittiwake, pelican, puffin, or tern) and waterfowl (gadwall, grebe, mallard, merganser, guillemot, or Pacific loon) are labeled as waterbirds, and all other birds are labeled as landbirds.

Backgrounds from the ocean and natural lake categories in the Places dataset are considered to have spurious attribute a=a= water background, while backgrounds from the bamboo forest or broadleaf forest categories are considered to have spurious attribute a=a= land background.

There are two minority groups: (land background, waterbird) and (water background, landbird); and two majority groups: (land background, landbird) and (water background, waterbird). We use the same training / valid / test splits from Sagawa et al. (2020a). In the training data, 95% of the waterbirds appear on water backgrounds, and 95% of the landbirds appear on land backgrounds, so the minority groups contain far fewer examples than the majority groups. In the validation and test sets, both the landbirds and waterbirds are evenly split between the water and land backgrounds.

B.2 CelebA

We use the task setup from Sagawa et al. (2020a) on the CelebA celebrity face dataset (Liu et al., 2015). The label yy is set to be the Blond_Hair attribute, and the spurious attribute aa is set to be the Male attribute: being female spurious correlates with having blond hair. The minority groups are (blond, male) and (not blond, female), although the (blond, male) group is significantly smaller than the (not blond, female) group. The majority groups are (blond, female) and (not blond, male). We use the standard train / valid / test splits from Sagawa et al. (2020a).

B.3 MultiNLI

We use the task setup from Sagawa et al. (2020a) on the MultiNLI natural language inference dataset (Williams et al., 2018). Given two sentences, a premise and a hypothesis, the task is to predict whether the hypothesis is entailed by, neutral with, or contradicted by the premise. The spurious attribute aa is a binary indicator for when any of the negation words nobody, no, never, or nothing appear in the second sentence (the hypothesis), which spuriously correlates with the contradiction label. We use the standard train / valid / test splits from Sagawa et al. (2020a).

B.4 CivilComments-WILDS

We use the CivilComments-WILDS dataset from Koh et al. (2021), which is derived from the Jigsaw dataset (Borkan et al., 2019). Given a real online comment, the task is to predict whether the comment is toxic or not toxic. The spurious attribute aa is an 8-dimensional binary vector, where each entry is a binary indicator of whether the following 8 demographic identities are mentioned in the online comment: male, female, LGBTQ, Christian, Muslim, other religion, Black, and White.

Following Koh et al. (2021), we consider the 16 potentially overlapping groups equal to (identity, toxic) and (identity, not toxic) for all 8 identities. We use the standard train / valid / test splits from Koh et al. (2021).

Appendix C Additional Experimental Results

We include the CivilComments error set analysis in Table 12 for space constraints.

C.2 Additional analysis

Below, we present a series of analyses that involve partitioning the dataset into two groups: groups in which spurious correlation holds with y=ay=a, and groups in which spurious correlation does not h old with 9≠a9\neq a. For this investigation, we focus on Waterbirds and CelebA, where all groups can be clearly partitioned as above because we consider binary classification tasks with binary spurious attributes. In Waterbirds, the y=ay=a groups are waterbirds on water backgrounds and landbirds on land backgrounds; the y≠ay\neq a groups are waterbirds on land backgrounds and landbirds on water backgrounds. In CelebA, the y=ay=a groups are blond females and non-blond males; the y≠ay\neq a groups are non-blond females and blond males. In contrast, it is unclear how to partition the groups as above in MultiNLI, in which we consider a multi-class classification problem, and in CivilComments-WILDS, in which we have multiple spurious a ttributes corresponding to different demographic identities.

We first study how worst-group accuracy changes when we remove y=ay=a examples or y≠ay\neq a examples from the error set, as summarized in Figure 5. In both datasets, removing either the y=ay=a or y≠ay\neq a examples from the error set significantly decreases worst-group accuracy. While this reduction in worst-group accuracy could stem from the fact that we consider a fixed set of hyperparameters including the upweight factor (which was tuned for Jtt with the full error set), it is possible that both y=ay=a and y≠ay\neq a examples contribute to the improvement in worst-group accuracy. In particular, because both datasets have substantial label imbalance, it is expected that upweighting groups from rare labels is important to perform well on all groups, and in fact, groups with y≠ay\neq a and with rare yy are upweighted as discussed in Section 5.3.

Next, we explore if the particular y=ay=a or y≠ay\neq a examples that Jtt upsamples is important, or if upsampling any collection of examples in these groups yields high worst-group accuracy. To do this, we study how average and worst-group accuracies change when we replace examples in Jtt’s error set with randomly selected examples from specific groups. Specifically, we study what happens when we upsample the following four variants of Jtt’s error set:

Replace y≠ay\neq a: We replace the y≠ay\neq a examples in the error set with an equal number of randomly selected y≠ay\neq a examples, leaving the y=ay=a examples in the error set uncha nged.

Replace y=ay=a: We replace the y=ay=a examples in the error set with an equal number of randomly selected y=ay=a examples, leaving the y≠ay\neq a examples in the error set unchanged.

Replace both: We replace both the y≠ay\neq a examples in the error set with an equal number of randomly selected y≠ay\neq a examples, and the y=ay=a examples in the error set with an e qual number of randomly selected y=ay=a examples.

Random sample: We replace all examples in the error set with an equal number of randomly selected examples. This yields a different fraction of y≠ay\neq a examples in the error set compared to Replace both.

Figure 6 compares upsampling these variants of the error set with the original unmodified error set (unmodified error set). Compared to upsampling the original error set, upsampling Replace y≠ay\neq a slightly decreases worst-group accuracy and leaves average accuracy unchanged. This suggests that upsampling most y≠ay\neq a examples helps improve worst-group accuracy, though the particular y≠ay\neq a examples Jtt identifies in the error set are still slightly better than rand om. On the other hand, upsampling Replace y=ay=a significantly decreases worst-group accuracy, although it slightly improves average accuracy compared to the original error set. This suggests that the particular y=ay=a examples Jtt identifies in the error set are important for improving worst-group accuracy. This could be because the label balance within upsampled y=ay=a changes, or for other reasons. Finally, the low worst-group and average accuracies of both Replace both and Random sample show that merely upsampling random y=ay=a and y≠ay\neq a examples is insufficient to achieve high worst-group accuracy.

We present the performance of a simple baseline, in which we upweight y≠ay\neq a examples using ground-truth group annotations, in Table 13. While Upsample minority improves worst-group error over ERM, this baseline is limited in a few ways. First, y≠ay\neq a groups are not necessarily groups with the worst accuracies or smallest number of examples, for example due to label imbalance. So while we y≠ay\neq a examples are counter-examples to the spurious correlations, it’s not necessarily expected that they improve the worst-group performance well. Secondly, in the presence of ground-truth examples, it is possible to reweight each of the groups independently, rather than reweighting y=ay=a and y≠ay\neq a groups. Prior work has observed much higher worst-group performance by reweighting the groups than Upsample minority (Sagawa et al., 2020a).