Spread Spurious Attribute: Improving Worst-group Accuracy with Spurious Attribute Estimation

Junhyun Nam, Jaehyung Kim, Jaeho Lee, Jinwoo Shin

Introduction

Machine learning models trained on datasets containing spurious correlation (also known as “shortcuts”) often end up learning such shortcuts instead of intended solutions (Geirhos et al., 2020). For example, consider an image classification dataset of ‘cows’ and ‘camels,’ in which most images of cows appear on grasslands and camels on deserts. When trained on such a dataset, models often learn to make predictions based on the landscape instead of the object (Beery et al., 2018). This phenomenon can lead to very low test accuracies on groups underrepresented in the training set.

Various approaches have been proposed to resolve this gap, and the idea of minimizing the worst-group loss—e.g., Sagawa et al. (2020)—has arisen as one of the most promising solutions. This approach forces high performances on both the majority group (e.g., cows standing on grasslands) and the minority group which contradicts the spurious correlation (e.g., cows standing on deserts). Despite its effectiveness, the worst-group loss minimization approach has a drawback: The learner requires supervision on which group each training sample belongs. For example, the learner needs to know that sample belongs to both the ‘cow’ and the ‘desert’ categories to utilize such group information to perform worst-group loss minimization. Even if we put aside the issue of identifying the spuriously correlated attributes (‘desert’ and ‘grass’) in the first place, one still needs to collect additional annotation on such spurious attributes. Acquiring such fine-grained annotation is presumably more expensive to collect, as the annotator needs a clear understanding of both the target attribute (‘cow’ and ‘camel’) and the spurious attribute (‘desert’ and ‘grass’).

Acknowledging this difficulty, recent works propose worst-group loss minimization algorithms that require a smaller number of group-labeled training samples (i.e., training samples with spurious attribute annotations) (Nam et al., 2020; Liu et al., 2021). At a high level, these works share a similar strategy (Fig. 1, Left): The methods first use a specialized mechanism to identify minority group samples among group-unlabeled training samples, and train a model in a way that puts more emphasis on identified-as-minority samples, e.g., by upweighting. A small set of group-labeled samples are used to tune hyperparameters of this procedure; as Liu et al. (2021) shows, the performance of trained models are very sensitive to these hyperparameters, indicating the high dependency of such algorithms on the availability of the group-labeled samples. Although these methods achieve higher worst-group accuracy than completely annotation-free approaches, e.g., Sohoni et al. (2020), they fail to perform comparably to the algorithms which use full annotations on spurious attributes, e.g., Sagawa et al. (2020). This performance gap gives rise to the following question: Can we closely achieve the full-annotation performance using a partially annotated dataset, if we use group-labeled samples more actively than hyperparameter-tuning?

Contribution. This paper proposes a worst-group loss minimization algorithm, coined Spread Spurious Attribute (SSA). At a high level, SSA consists of two phases—pseudo-labeling and robust training—that are designed to make a full use out of the group-labeled samples (Fig. 1, Right):

Pseudo-labeling: SSA trains a spurious attribute predictor, using both group-labeled and group-unlabeled samples (i.e., samples lacking spurious attribute annotations). We find that group imbalances underlying the dataset can render pseudo-labels to be biased towards the majority group, which can be detrimental to the performance (as in Kim et al. (2020)), and propose using group-wise adaptive thresholds to mitigate this problem. The predictions are used as pseudo-labels (Lee, 2013) on the group-unlabeled samples.

Robust training: Based on the pseudo-labeled training set, SSA trains a model which predicts the target attribute, using the worst-group loss minimization algorithms developed for fully supervised cases. We use Group DRO (Sagawa et al., 2020) as a default choice, and re-use group-labeled samples for the hyperparameter tuning.

Our experimental results suggest that SSA is a general yet effective framework for worst-group loss minimization. On various benchmark datasets, SSA consistently achieves performances comparable to fully supervised Group DRO and outperforms other baseline methods even when using only 5% of the group-labeled samples compared to the baselines. To be more specific, SSA requires less than 1.5% (as little as 0.6%) of group-labeled samples—993 out of 182637 on CelebA, and 4123 out of 288637 on MultiNLI—to achieve a worst-group accuracy similar to that of 100% usage. Moreover, we find that such benefits of SSA persist when combined with other robust training methods, such as Correct-N-Contrast (CNCWe use the term CNC to refer to its supervised contrastive learning procedure with contrastive batch sampling only, rather than the entire process which includes spurious attribute annotation estimation via ERM.; Zhang et al. (2021)), the supervised contrastive learning with contrastive batch sampling using group information. Finally, we also empirically observe that the group-wise adaptive thresholding technique—which we proposed for better pseudo-labeling—enjoy a broader usage for addressing the general problem of semi-supervised learning under class imbalance; existing pseudo-labeling techniques (Kim et al., 2020; Wei et al., 2021) require certain assumptions in datasets to work well for the task, while ours does not.

Related Works

Improving worst-group accuracy with group annotations. It has been known in various literature that machine learning models often perform significantly worse on the samples from groups with a relatively small number of training samples than on samples from the majority group. The class imbalance problem (Japkowicz, 2000; Johnson & Khoshgoftaar, 2019) is one of the representative cases of this phenomenon. Here, class labels naturally define group identities and thus do not require any additional group annotations. In this context, popular strategies to improve worst-group performances are using re-weighted loss for each group (Huang et al., 2016; Khan et al., 2017) and re-sampling the given dataset to balance the group distribution of the training dataset (Chawla et al., 2002; He & Garcia, 2009). This paper, in contrast, considers a setup where the group is defined as a combination of the label (target attribute) and the spuriously correlated attribute; unlike the class imbalance literature, the availability of the full group identity information is not guaranteed, as spurious attributes may not have proper annotations. Under the setting where spurious attribute annotations are available on all samples, Sagawa et al. (2020) gives an online optimization algorithm that shows promising results on minimizing the worst-group loss.

Improving worst-group accuracy without group annotations. To reduce additional annotation costs, recent works aim to train a robust model without requiring group annotations for all training samples. These works utilize models trained with standard procedure to identify samples that disagree with spurious correlations. Nam et al. (2020) train a model to be intentionally biased using the generalized cross entropy loss, and use it to identify-and-upweight high-loss samples for training another model. Liu et al. (2021) train a standard ERM model for a few epochs and upweight samples misclassified by this model to train the second model. Although these approaches do not use any group information for training, they still use a small number of group-labeled samples for hyperparameter tuning, which is critical to their worst-group performance. In this paper, we design a method to utilize group-labeled samples more efficiently. We note that, in the class imbalance literature, another line of works proposes to use group-labeled samples more actively; they use labeled samples for determining sample weights for the training set through meta-learning (Ren et al., 2018) or training a model to predict sample weights (Shu et al., 2019).

Problem setup

Consider the learning a classifier in the presence of spurious correlations in the training set. Following the prior work of Sagawa et al. (2020), we cast this problem as minimizing the worst-group loss, where the group identity is determined by the target attribute (that we want to predict) and the spurious attribute (that we want to ignore). We assume that we do not have spurious attribute annotations of the samples in the training set (group-unlabeled set), but have an access to additional set of samples with both spurioust attribute and target attribute annotations (group-labeled set).

More formally, we let each sample be a triplet consisting of an input x∈Xx\in\mathcal{X}, a target attribute y∈Yy\in\mathcal{Y}, and a spurious attribute a∈Aa\in\mathcal{A}. Our goal is to train a parameterized model fθ:X→Yf_{\theta}:\mathcal{X}\to\mathcal{Y} that minimizes the worst-group expected loss on test samples; for the purpose of avoiding learning spurious correlations, we define the group as an attribute pair g:=(y,a)∈Y×A=:Gg:=(y,a)\in\mathcal{Y}\times\mathcal{A}=:\mathcal{G}. In other words, we aim to minimize

Spread Spurious Attribute

We now describe the algorithm we propose, Spread Spurious Attribute (SSA). At a high level, SSA consists of two phases: pseudo-labeling and robust training. In the pseudo-labeling phase, SSA trains a model to predict the spurious attribute. In particular, we use both group-labeled and group-unlabeled samples to generate pseudo-labels on the group-unlabeled training samples (Section 4.1), with an adaptive thresholding technique for balancing the number of samples in each group (Section 4.2). Then, in the robust training phase, SSA uses generated pseudo-labels to train a model that predicts the target attribute with a small worst-group loss (Section 4.3). We use Group DRO (Sagawa et al., 2020) as our default robust training method, but SSA also performs well when combined with other existing algorithms that use full spurious label supervisions (see Section 5.3).

Using both group-labeled set DL\mathcal{D}_{L} and group-unlabeled set DU\mathcal{D}_{U}, we train a model to predict spurious attributes on the group-unlabeled samples. The predictions will be used to generate artificial spurious attribute labels on the group-unlabeled training samples, called pseudo-labels (Lee, 2013). We emphasize that we also use DU\mathcal{D}_{U} for training the model, instead of training solely based on DL\mathcal{D}_{L}. In fact, as we shall see in Section 5.2, the additional use of training set brings a significant performance boost. More specifically, we first partition both the group-labeled and group-unlabeled set into two:

We use DL∘,DU∘\mathcal{D}^{\circ}_{L},\mathcal{D}_{U}^{\circ} to train the spurious attribute predictor that make prediction on DU∙\mathcal{D}^{\bullet}_{U}, and validate the model with DL∙\mathcal{D}_{L}^{\bullet}.We write ∘\circ to denote that samples are “visible” during the training phase, and ∙\bullet to denote that they are not. To train this predictor, we use a loss function consisting of two terms: the supervised loss for the samples in DL∘\mathcal{D}_{L}^{\circ}, and the unsupervised loss for samples in DU∘\mathcal{D}_{U}^{\circ}. The supervised loss is simply the standard cross entropy loss. For the unsupervised loss, we use the cross entropy loss between the prediction and the pseudo-labels, i.e., labels generated by taking the arg max⁡\operatorname*{arg\,max} of predictions. Following prior works (e.g., Sohn et al. (2020)), we apply the loss only if the confidence of the prediction exceeds some threshold τ≥0\tau\geq 0. More formally, let us denote the class probability estimate of the predictor on the attribute aa given the input xx as p^(a∣x)\hat{p}(a|x), and the pseudo-label from this prediction as a^(x)=arg max⁡a∈Ap^(a∣x)\hat{a}(x)=\operatorname*{arg\,max}_{a\in\mathcal{A}}\hat{p}(a|x). Then, the supervised and unsupervised losses are

where CE(p^(⋅∣x),a)\textrm{CE}(\hat{p}(\cdot|x),a) denotes the cross-entropy loss between the prediction p^(⋅∣x)\hat{p}(\cdot|x) and label aa. The total loss is then given as the sum of the supervised and unsupervised loss

As we will describe in Section 4.2, we additionally use group-wise adaptive thresholds for pseudo-labeling (i.e., set different τ\tau for each group). The thresholds are determined in a way that balances the pseudo-group (i.e., the pair g^=(y,a^(x))\hat{g}=(y,\hat{a}(x))) population of the samples with confidence exceeding the group-wise threshold. This strategy helps the model to avoid making a pseudo-label prediction biased toward the majority group. Also, we note that we make predictions only on samples in DU∙\mathcal{D}_{U}^{\bullet} and not on samples in DU∘\mathcal{D}_{U}^{\circ}; empirically, we observe that this “splitting” of group-unlabeled samples is beneficial for performance comparing to using the whole DU\mathcal{D}_{U} for training (at the cost of running pseudo-labeling multiple times). For more discussions and ablation studies, see Section A.2.

2 Balancing generated pseudo-labels via adaptive thresholds

As we train the spurious attribute predictor with the loss (Eq. 4), the prediction confidence of the model on each sample gradually increases. Ideally, we want the number of highly-confident samples (i.e., with confidence greater than τ\tau) to increase uniformly over all pseudo-groups so that each pseudo-group contributes evenly to weight updates. However, if the underlying group population is severely imbalanced—as is common with spurious attributes—samples from the majority group often attains high prediction confidence significantly faster than samples from minority groups, by receiving more frequent gradient updates. This training imbalance leads to a severer imbalance in the pseudo-labels, resulting in detrimental effects on the downstream training.

To mitigate this majority group bias, we set different thresholds for each pseudo-group, so that the same number of samples from each group is used for training with Eq. 4. To do this, we first set a fixed threshold τgmin⁡\tau_{g_{\min}} for the group with the smallest population in the training split of the group-labeled set, i.e., the group defined as

Next, we count the number of samples in DU∘\mathcal{D}_{U}^{\circ} which (a) belongs to gmin⁡g_{\min} after pseudo-labeling, and (b) the prediction confidence exceeds τgmin⁡\tau_{g_{\min}}. Then, we set thresholds for other groups to have same number of samples from each group. Concretely, let DU∘(g,τg)\mathcal{D}_{U}^{\circ}(g,\tau_{g}) be the set of samples in the training split of the group-unlabeled set with pseudo-group gg and confidence greater than equal to τg\tau_{g}, i.e.,

Then, for each group g≠gmin⁡g\neq g_{\min}, we set τg\tau_{g} to be the smallest real number such that

With this group-wise adaptive threshold, we use the unsupervised loss revised as

We show the effectiveness of this group-wise adaptive threshold in Section 5.2, and further applicability to class-imbalance problem in Section 5.4.

3 Worst-group loss minimization with estimated pseudo-groups

After training the model using the revised loss (Section 4.2), we generate final pseudo-labels for all samples in the training set. In other words, we generate

We put pseudo-labels on all samples without applying any threshold, which allows us to utilize the whole training set. In the robust training phase, we use the pseudo-labeled dataset D~U\widetilde{\mathcal{D}}_{U} and the group-labeled dataset DL\mathcal{D}_{L} to perform worst-group loss minimization. We use Group DRO (Sagawa et al., 2020) as our default algorithm, using D~U\widetilde{\mathcal{D}}_{U} for training and DL\mathcal{D}_{L} for validation. We note that SSA can also use other robust training subroutines instead of group DRO; in Section 5.3, we show that our framework performs well when combined with CNC (Zhang et al., 2021).

Experiments

Here, we briefly describe the experiment setup that will be used throughout all experiments in this section, except for Section 5.4 where we consider a slightly different scenario.

Datasets. We evaluate SSA on two image classification datasets (Waterbirds, CelebA) and two natural language processing datasets (MultiNLI, CivilComments-WILDS) containing spurious correlations. For all datasets, we use the validation split of the dataset as the group-labeled set. Below, we briefly describe each dataset and the corresponding spurious correlations.

Waterbirds (Sagawa et al., 2020): Waterbirds is an artificial dataset generated by combining bird photographs in the Caltech-UCSD Birds dataset (Wah et al., 2011) with landscapes from Places (Zhou et al., 2017). The goal is to classify the target attributes Y={waterbird, landbird}\mathcal{Y}=\{\textrm{waterbird, landbird}\} given the spurious correlations with the background landscape A={water background, land background}\mathcal{A}=\{\textrm{water background, land background}\}.

CelebA (Liu et al., 2015): CelebA dataset consists of the face pictures of celebrities, with various annotations on facial/demographic features. We use the hair color as the target attribute Y={blond, non-blond}\mathcal{Y}=\{\textrm{blond, non-blond}\}, given the spurious correlations with the gender A={male, female}\mathcal{A}=\{\textrm{male, female}\}.

MultiNLI (Williams et al., 2018): MultiNLI is a multi-genre natural language corpus where each data instance consists of two sentences and a label indicating whether the second sentence is entailed by, contradicts, or neutral to the first. We use this label as the target attribute (i.e., Y={entailed, neutral, contradictory}\mathcal{Y}=\{\textrm{entailed, neutral, contradictory}\}), and use the existence of the negating words as the spurious attribute (i.e., A={negation, no negation}\mathcal{A}=\{\textrm{negation, no negation}\}).

CivilComments-WILDS (Borkan et al., 2019; Koh et al., 2021): CivilComments-WILDS consists of comments generated by online users, each of which are labeled with the toxicity indicator Y={toxic,non-toxic}\mathcal{Y}=\{\textrm{toxic},\textrm{non-toxic}\}. We use demographic attributes of the mentioned identity A={male, female, White, Black, LGBTQ, Muslim, Christian, other religion}\mathcal{A}=\{\textrm{male, female, White, Black, LGBTQ, Muslim, Christian, other religion}\} as a spurious attribute for evaluation purpose. We note that a comment can contain multiple such identities, so that groups defined by Y×A\mathcal{Y}\times\mathcal{A} can be overlapped. Therefore, we use A′={any identity, no identity}\mathcal{A}^{\prime}=\{\textrm{any identity, no identity}\} as a spurious attribute for training, following Liu et al. (2021).

Models. For the all experiments on image classification datasets, we use ResNet-50 (He et al., 2016) starting from ImageNet-pretrained weights. For experiments on language datasets, we use pretrained BERT (Devlin et al., 2019). We use the same architecture for predicting the spurious attribute (in the pseudo-labeling phase) and the target attribute (in the robust training phase).

We compare the average and worst-case performance of the proposed SSA against the standard empirical risk minimization (ERM) and recent methods that tackles spurious correlation without group annotation for the training, including CVaR DRO (Levy et al., 2020), LfF (Nam et al., 2020), EIIL (Creager et al., 2021), JTT (Liu et al., 2021), and Group DRO (Sagawa et al., 2020) requiring group annotation to minimize the worst-group loss.

In Table 1, 2, we report average accuracies and the worst-group accuracies on all datasets we consider. Our method consistently outperforms all the other approaches that use spurious attribute annotated dataset for validation while using the same amount of spurious attribute annotation. Notably, our method shows comparable performance to Group DRO which uses full amount of spurious attribute annotation for the training set, even outperforms on CelebA dataset. We also run our algorithm using only 5% of the default validation set to show efficiency of our algorithm to improve the worst-group accuracy. We further provide analysis on varying group-labeled set size below.

Effect of group-labeled set size. Although we focus on improving the worst-group performance with given amount of spurious annotation, reducing the amount of supervision is an important topic to discuss especially with high annotation cost. In the main results in Table 1, 2, we use the default validation sets provided by each dataset as DL\mathcal{D}_{L}. To further test whether SSA can achieve high worst-group accuracy with reduced amount of supervision, we run our algorithm with small fraction of the default validation sets as group-labeled sets. Following Liu et al. (2021), we run our method using 100%, 20%, 10%, 5% of the default validation set. In Table 3, we report the worst-group accuracy of JTT and our algorithm on Waterbirds and CelebA, with various group-labeled set size. Surprisingly, we find that our method maintains high worst-group accuracy even with the 10% of the original validation set. Most notably, the number of attribute annotated samples used for training spurious attribute predictor is 58 in total, 6 for the worst-group in Waterbirds when we only use 10% of the default validation set.

2 Detailed analysis on the pseudo-labeling phase of SSA

We now take a closer look at the pseudo-labeling phase of the proposed SSA. In particular, we focus on validating the effectiveness of the group-wise adaptive threshold we introduced in Section 4.2. The purpose of the adaptive threshold was to prevent the pseudo-labels from being biased towards the majority group during the pseudo-labeling phase; the prediction confidence of the majority group samples may exceed the threshold faster than samples of minority group, increasing the contribution of majority-group samples on the loss even further. In the first set of experiments (Table 4), we perform ablation studies on the pseudo-labeling and the adaptive threshold to see their effects on the spurious attribute prediction performance of the SSA. In the second set of experiments (Table 5), we validate if the pseudo-labels are biased towards the majority group and check that our adaptive threshold strategy successfully addresses the phenomenon.

Accuracy of the spurious attribute predictor. In Table 4, we report the group-wise spurious attribute prediction accuracies of SSA on CelebA dataset (using 5%5\% or 10%10\% of the validation split as the group-labeled set) in pseudo-labeling phase for following ablations: (1) Vanilla: Does not use any pseudo-labels and train the model using only the validation set samples. (2) Pseudo-labeling: Uses pseudo-labeling with a fixed threshold, and (3) +Group-wise threshold: Identical to SSA, using group-wise adaptive thresholds for pseudo-labeling. We find that using the group-wise threshold increases the spurious attribute prediction accuracy for the worst-group—(y,a)=(Blond,Male)(y,a)=(\text{Blond},\text{Male}), providing  7%~{}7\% boost in both cases. Interestingly, we observe that pseudo-labeling with fixed thresholds can even slightly degrade the performance, when the validation set is too small (5%).

Population comparison. In Table 5, we compare the populations of pseudo-labeled samples that contribute to the training, for pseudo-labeling using a fixed threshold (‘Pseudo-labeling’) and the adaptive threshold (‘+Group-wise threshold’); for each method, we trained the spurious attribute predictor until the worst-group spurious attribute prediction accuracy reaches the highest point. We used CelebA dataset with only 5% of the validation split. From the experimental results, we find that naïve pseudo-labeling with a fixed threshold indeed leads the pseudo-labels to be biased towards the majority group using only 0.68% of the minority group (male blond) for training while the true fraction of the group is over 0.8%0.8\%. On the other hand, using the adaptive threshold successfully addresses this problem, lifting the fraction of blond male samples to 0.88%0.88\%, which is close to the population level. More impressively, we observe that pseudo-labeling phase adaptive threshold uses relatively uniform number of samples from each group to training the spurious attribute predictor.

3 SSA combined with Supervised Contrastive Learning

We now examine whether the proposed SSA still remains to be beneficial when combined with other robust training procedures. As an example, we choose a recent robust training procedure (Zhang et al., 2021) proposed as an alternative to Group DRO. More specifically, CNC (Zhang et al., 2021) considers a following procedure base on the supervised contrastive loss (Khosla et al., 2020): As in JTT (Liu et al., 2021), we first train a standard ERM model. Then, we use the contrastive loss to maximize the representational similarity between samples have the same target label but different ERM prediction, while minimizing the representational similarity of samples with different target attribute but same ERM prediction. This procedure can be smoothly combined the SSA framework, by replacing the ERM predictions with the pseudo-labels generated in the first phase of SSA.

In Table 6, we compare the performance of SSA using Group DRO or CNC with the performance of the robust training methods using the full spurious attribute annotations on the training set. Both models learned with SSA achieved comparable performance to the fully supervised counterparts. This result suggests that the benefit of SSA may not be constrained on a specific robust training method, and may be combined with more general classes of robust training algorithms. Also, we find that none of the robust training method consistently outperform the other; CNC with group label slightly outperforms Group DRO on Waterbirds, and Group DRO does better on CelebA.

4 Application to general semi-supervised learning under class imbalance

In Section 4.2, we proposed to use adaptive thresholding to mitigate confirmation bias when pseudo-labeling group-imbalanced datasets. Interestingly, it turns out that the benefit of this idea also extends (without any modification) to a more general scenario of semi-supervised learning (SSL) on datasets with class imbalances. In fact, two settings are quite similar, except that SSA aims to put pseudo-labels on spurious attributes while SSL methods estimate target attributes. To demonstrate this point, we evaluate adaptive thresholds under the SSL setup, and compare it with the performances of baseline pseudo-labeling-based SSL algorithms. FixMatch (Sohn et al., 2020) is a recent SSL method proposed without considerations on the class imbalance issue. DARP (Kim et al., 2020) and CReST (Wei et al., 2021) build on FixMatch to handle class imbalance, but are primarily designed under the assumptions that there exists sufficiently many labeled data at hand (to estimate the class imbalance ratio), and that class distributions of labeled and unlabeled datasets are identical, respectively.In fact, in the previous spurious attribute setups, adopting DARP/CReST methods did not provide much gain over naïve pseudo-labeling; we suspect the reason to be the violation of these assumptions. In contrast, our method of adaptive thresholds does not rely on such assumptions to solve the same task.

We consider a slightly more challenging experimental setup than in Kim et al. (2020): We construct an artificial labeled dataset from the CIFAR-10 (Krizhevsky et al., 2009) by controlling the number of samples in the majority class mmajm_{\text{maj}} (i.e., the largest class) and the ratio between the largest and smallest class sizes γlab≥1\gamma_{\text{lab}}\geq 1.Larger γlab\gamma_{\text{lab}} thus indicates a more severe imbalance. In a similar manner, we construct an unlabeled dataset using some parameters nmajn_{\text{maj}} and γunlab\gamma_{\text{unlab}}. We select these parameters so that the number of labeled samples is small, and the imbalance ratios of labeled and unlabeled sets have bigger discrepancies (We provide further details in Section A.6). In Table 7, we observe that empirical gains from both DARP and CReST are limited comparing to supervised learning with full ground-truth labels (oracle), due to the violation of their inherent assumptions in the experimental setups of our choice, i.e., limited labeled data and different class distribution between labeled and unlabeled datasets. However, our method successfully reduces such gap by effectively constructing the pseudo-labels for minority classes. Overall, these results implies that the proposed method has a potential to provide a more robust semi-supervised learning solution (despite its simplicity) in more realistic scenarios, which we think is an interesting direction to explore further in the future.

Conclusion

In this work, we present Spread Spurious Attribute (SSA), an algorithm for improving worst-group accuracy in the presence of spurious correlation. SSA framework uses a small amount of spurious attribute annotated samples to estimate group identities of the training set samples. With the generated attribute annotated training set, we successfully train a robust model using existing worst-case loss minimization algorithms. SSA is highly effective given a limited amount of the spurious attribute annotated samples, but still does not completely remove the need for supervision on spurious attributes which is an important future direction.

This paper addresses the problem of mitigating the harmful effects of spurious correlations in the dataset, which is deeply intertwined with the topics of machine bias and fairness. By targeting demographic groups (e.g., gender, race) carefully, our method has a potential to be used for preventing the machine to make predictions on the basis of biases on such demographic identities. One possible pitfall, however, is that this group-based evaluation can provide a false sense of morality; as making a fair decision in terms of group cannot be equated with the individual fairness (even putting aside the issues of gerrymandering), our performance reports based on the worst-group performance should not be naïvely taken as a measure of justice. About the responsible research practice: We have used benchmark datasets, and thus believe that our research practice has not posed any foreseeable hazard in this regard.

Reproducibility statement

We provide descriptions of the experimental setup and implementation details in Sections A.2, A.3 and A.5. Also, we provide our source code as a part of the open-to-public supplementary materials.

Acknowledgments

This work was supported by Institute of Information & Communications Technology Planning & Evaluation (IITP) grant funded by the Korea government(MSIT) (No.2019-0-00075, Artificial Intelligence Graduate School Program (KAIST) and No.2019-0-01396, Development of framework for analyzing, detecting, mitigating of bias in AI model and training data)

References

Appendix A Experimental details

Algorithm 1 provides pseudocode for Spread Spread Attribute combined with Group DRO.

A.2 Discussions on splitting the group-labeled and group-unlabeled sets

Recall that in the pseudo-labeling phase of SSA (Section 4.1, we partition both the group-labeled set and the group-unlabeled set into two subsets:

SSA then trains a spurious attribute predictor based on DL∘,DU∘\mathcal{D}^{\circ}_{L},\mathcal{D}_{U}^{\circ}, with hyperparameters tuned using DL∙\mathcal{D}_{L}^{\bullet}. The trained model is then used to make predictions on the samples in DU∙\mathcal{D}_{U}^{\bullet}.

Here, it is easy to see that, if we want to make a best prediction on a particular data point x⋆∈DUx_{\star}\in\mathcal{D}_{U}, then the optimal split would be the one that uses the largest number of group-unlabeled samples for training the model, i.e.,

However, such partitioning requires training ∣DU∣|\mathcal{D}_{U}| different models to label all samples in the group-unlabeled set, which is computationally infeasible. Thus, we propose partitioning the group-unlabeled dataset into KK equally-sized subsets DU(1),…,DU(K)\mathcal{D}_{U}^{(1)},\ldots,\mathcal{D}_{U}^{(K)}, and run the algorithm KK times, using

for the ii-th training iteration. For all experiments appearing in this paper, we used K=3K=3 for a simple evaluation; the empirical performance of SSA may improve if we use a larger value of KK.

For the group-labeled set, we do not require such trick. We simply partitioned the group-labeled into two equally size subsets, and used one for training and another for validation. We kept the split fixed throughout the whole pseudo-labeling phase.

Discussion. The splitting procedure aims to mitigate the potential negative effect of “self-confirmation” that may take place in the pseudo-labeling procedure. During the training of pseudo-labeler, SSA utilizes the pseudo-attributes of the training samples whenever the prediction confidence exceeds a certain threshold. In other words, when a sample gains a high-enough confidence, pseudo-labeling may continually strengthen its own predictions, and can be very problematic in group-DRO-like scenarios where we expect a severe group imbalance. Data splitting helps mitigate this effect by separating the samples that we train on and the samples we make final predictions on (as a side note, we also propose group-adaptive thresholds to mitigate the same effect). Empirically, we indeed observe that data splitting helps improve the downstream worst-group performance. Tables 8 and 9 give a comparison of SSA with/without splitting on CelebA and Waterbirds dataset, respectively.

A.3 Training details

Models. As we briefly discussed in the main text, we use pretrained ResNet-50 and BERT for image and natural language dataset experiments, respectively. For ResNet-50, we use the torchvision implementation. For BERT, we use the huggingface implementation.

Hyperparameter Tuning - Overall. We separately tune the hyperparameters for the pseudo-labeling phase and the robust training phase. For the pseudo-labeling phase, the tuning criterion is the worst-group spurious label classification accuracy on DL∙\mathcal{D}_{L}^{\bullet}. For the robust training phase, the tuning criterion is the worst-group prediction accuracy of the trained target attribute classifier on the whole validation set DL\mathcal{D}_{L}. We fix the threshold τgmin\tau_{g_{\text{min}}} for the group with smallest population as 0.95 following (Sohn et al., 2020) in pseudo-labeling phase.

Baseline Implementation. For EIIL, we directly take results on Waterbirds and CivilComments from Creager et al. (2021), while the results on CelebA and MultiNLI are new. For environment inference on CelebA, we follow the same procedure as for Waterbirds. We use one epoch trained ERM as a reference classifier. We optimize the EI objective with a learning rate of 0.01 for 20k steps using the Adam optimizer. For environment inference on MultiNLI, we follow the same procedure as for CivilComments-WILDS. We train an ERM model for 5 epochs and choose the reference classifier using the best epoch based on the validation worst-group accuracy. We use the error split heuristic instead of optimizing EI objective as Creager et al. (2021) did. We then train the robust model with the same procedure as we did.

A.4 Runtime analysis

In Table 10, we provide the time required for the pseudo-labeling phase and the robust training phase on a single Nvidia Titan XP for each dataset. Compared with the vanilla Group DRO, SSA requires an additional pseudo-labeling phase. When we select K=3K=3, the overhead is as small as x0.5 on the CelebA dataset, and does not exceed x4 in the worst case (Waterbird). We make two additional remarks. First, other baseline methods for addressing the lack of full group annotation (e.g., JTT, LfF) also require additional computation and runtime. For instance, JTT requires a preliminary training run to estimate the spurious label. Second, the number of iterations we used for pseudo-labeling has not been optimized for achieving the best runtime-performance tradeoff, and thus can be further improved with a more careful search.

A.5 Details on supervised contrastive learning

Given an anchor (xi,yi)(x_{i},y_{i}), original supervised contrastive loss uses same class samples (y=yi)(y=y_{i}) as positive samples and different class samples (y≠yi)(y\neq y_{i}) as negative samples to maximize similarity of representation between samples belong to same class. Correct-N-Contrast (CNC; Zhang et al. (2021)) first trains ERM model as JTT and use ERM prediction as a surrogate to ground truth spurious attribute. With obtained ERM prediction, CNC uses supervised contrastive loss to maximize similarity of representation between samples having same target and different ERM prediction, while minimizing similarity of representation between samples having different target and same ERM prediction. Similar to the second stage of CNC, given an anchor (xi,yi,ai)(x_{i},y_{i},a_{i}), we use samples having same target and different spurious attribute as positive samples (y+=yi,a+≠ai)(y^{+}=y_{i},a^{+}\neq a_{i}) and samples having different target and same spurious attribute as negative samples (y−≠yi,a−=ai)(y^{-}\neq y_{i},a^{-}=a_{i}). To be specific, we sample MM samples from each group. Given an anchor (xi,yi,ai)(x_{i},y_{i},a_{i}), we use samples from group (yi,a+)(y_{i},a^{+}) for a+≠aia^{+}\neq a_{i} as a set of positive samples B+B^{+} and samples from group (y−,ai)(y^{-},a_{i}) for y−≠yiy^{-}\neq y_{i} as a set of negative samples B−B^{-}. In addition to standard cross entropy loss, we minimizes following supervised contrastive loss for each anchor xix_{i}:

We use M=16M=16 for both Waterbirds and CelebA. Except batch size, we follow Zhang et al. (2021) for other hyperparameters. For Waterbirds, we use temperature 0.1, contrastive weight 0.75, SGD optimizer with momentum 0.9, learning rate 1e-4, weight decay 1e-3 and use gradient accumulation to update model parameter every 32 batches. For CelebA, we use temperature 0.05, contrastive weight 0.75, SGD optimizer with momentum 0.9, learning rate 1e-5, weight decay 1e-1 and use gradient accumulation to update model parameter every 32 batches.

A.6 Details on class-imbalanced semi-supervised learning

We consider a classification problem with KK classes. In other words, our goal is to train a predictor X→Y\mathcal{X}\to\mathcal{Y} with Y={1,…,K}\mathcal{Y}=\{1,\ldots,K\}. We assume that we have access to two types of datasets: labeled, and unlabeled. We let

The number of samples in some class k∈Yk\in\mathcal{Y} will be denoted by mkm_{k} and nkn_{k}, respectively, i.e., ∑k=1Kmk=n\sum_{k=1}^{K}m_{k}=n and ∑k=1Kmk=m\sum_{k=1}^{K}m_{k}=m. Without loss of generality, we assume that the number of labeled data in each class is ordered in a descending order, i.e.,

We define the imbalance ratio of this labeled dataset as the ratio between the sample sizes of the largest class and the smallest class, i.e., γlab=m1/mK≥1\gamma_{\text{lab}}=m_{1}/m_{K}\geq 1. This quantity is used as a key parameter that controls the degree of class imbalance in the labeled dataset; higher the γlab\gamma_{\text{lab}}, more severe the imbalance. To determine the number of samples in other classes, we use an exponential decay function, i.e., mk=m1⋅γlab(−k−1)/(K−1)m_{k}=m_{1}\cdot\gamma_{\text{lab}}^{(-k-1)/(K-1)}. We assume that the same ordering holds for the unlabeled samples, i.e., n1≥⋯≥nKn_{1}\geq\cdots\geq n_{K}, and let define the imbalance ratio as γunlab=n1/nK\gamma_{\text{unlab}}=n_{1}/n_{K}. For simplicity, we let γunlab=1\gamma_{\text{unlab}}=1, i.e., all classes have the same number of unlabeled samples.

To evaluate the classification performance of models trained under the imbalanced dataset, we report two popular metrics: balanced accuracy (bACC) and geometric mean scores (GM), which are defined by the arithmetic and geometric mean over class-wise sensitivity, respectively. Mean and standard deviation are reported across three random trials, respectively.