Improving Out-of-Distribution Robustness via Selective Augmentation

Huaxiu YaoYu WangSai LiLinjun ZhangWeixin LiangJames ZouChelsea Finn

article2022ICML271 citations

Develops LISA, a selective data interpolation method that pairs samples sharing labels across different domains or domains across different labels, achieving superior out-of-distribution generalization and theoretically reduced worst-group error across diverse shift benchmarks.

Listen

Machine learning systems frequently experience severe performance drops when deployed in real-world settings where operating conditions deviate from training conditions. This challenge typically arises from subpopulation shifts, such as class imbalances where models exploit misleading spurious correlations, or domain shifts, where models encounter entirely new environments like new hospitals or camera sensors. Standard techniques attempt to enforce stability across environments by adding complex mathematical penalties during training, but these regularizers often restrict model flexibility, prove difficult to optimize, and behave inconsistently across different applications.

The article evaluates a straightforward data augmentation method called LISA (Learning Invariant Predictors with Selective Augmentation) designed to improve model robustness under distribution shifts. The objective is to demonstrate that selectively interpolating data samples enables neural networks to ignore misleading domain attributes and rely solely on robust predictive features, without requiring artificial constraints on the model's internal structure.

To evaluate this framework, the authors conducted empirical benchmarks across nine standard image and natural language datasets, including the multi-domain WILDS and MetaShift benchmarks, alongside theoretical analysis using linear discriminant models. The approach linearly blends input features and target labels under two distinct strategies: intra-label augmentation, which pairs data points sharing the same class but coming from different domains to cancel domain-specific noise, and intra-domain augmentation, which pairs samples within the same domain possessing different labels to force the model to look beyond environmental artifacts.

The findings establish three primary outcomes. First, LISA consistently matched or outperformed seven leading robustness baselines across all nine datasets. For instance, in worst-case subpopulation accuracy, LISA achieved 89.3% on facial attribute recognition and 72.6% on toxic comment detection, clearly surpassing prior invariant-learning baselines. Second, in domain-shift environments, LISA raised accuracy on medical tumor identification to 77.1% compared to 70.3% for standard training, while also delivering leading performance on satellite imagery and cellular perturbation benchmarks. Third, controlled tests confirmed that these improvements stemmed directly from neutralizing spurious correlations rather than generic data expansion, showing the largest relative performance gains when the gap between training and testing distributions was widest.

These results carry significant practical implications for deploying reliable machine learning in safety-critical and high-risk domains such as medical diagnosis and content moderation. LISA mitigates the risk of catastrophic failures caused by environmental shifts while avoiding the hyperparameter fragility and training instability associated with prior regularization algorithms. By simply adjusting how training batches are constructed, teams can achieve superior reliability with minimal overhead.

Organizations developing machine learning models for shifting environments should consider integrating selective data blending into existing training pipelines as an alternative to complex loss regularizers. For tasks characterized by strong spurious correlations, practitioners should balance intra-label and intra-domain strategies, whereas intra-label mixing alone provides the greatest benefit when domain shifts are natural and widespread.

A primary operational limitation of LISA is its reliance on pairing samples with identical labels, which restricts its immediate utility in complex tasks such as object detection or generative modeling. Additionally, the approach assumes access to domain or group metadata during training, although empirical evidence indicates intra-label blending can succeed even without explicit domain tags when domain correlations are weak. Confidence in the underlying method remains high due to consistent empirical gains paired with formal theoretical validation.

Cover for Improving Out-of-Distribution Robustness via Selective Augmentation

Abstract

Machine learning algorithms typically assume that training and test examples are drawn from the same distribution. However, distribution shift is a common problem in real-world applications and can cause models to perform dramatically worse at test time. In this paper, we specifically consider the problems of subpopulation shifts (e.g., imbalanced data) and domain shifts. While prior works often seek to explicitly regularize internal representations or predictors of the model to be domain invariant, we instead aim to learn invariant predictors without restricting the model's internal representations or predictors. This leads to a simple mixup-based technique which learns invariant predictors via selective augmentation called LISA. LISA selectively interpolates samples either with the same labels but different domains or with the same domain but different labels. Empirically, we study the effectiveness of LISA on nine benchmarks ranging from subpopulation shifts to domain shifts, and we find that LISA consistently outperforms other state-of-the-art methods and leads to more invariant predictors. We further analyze a linear setting and theoretically show how LISA leads to a smaller worst-group error.

Table of Contents

  • 1 Introduction
  • 2 Preliminaries
  • 3 Learning Invariant Predictors with Selective Augmentation
  • 4 Experiments
  • 4.1 Evaluating Robustness to Subpopulation Shifts
  • 4.2 Evaluating Robustness to Domain Shifts
  • 4.3 Are the Performance Gains of LISA from Data Augmentation?
  • 4.4 Does LISA Lead to More Invariant Predictors?
  • 4.5 Effect of the Degree of Distribution Shifts
  • 5 Theoretical Analysis
  • 6 Related Work and Discussion
  • 7 Conclusion
  • References
  • A Additional Experiments
  • A.1 Additional Experiments on Subpopulation Shifts
  • A.1.1 Dataset Details
  • A.1.2 Training Details
  • A.1.3 Additional Results
  • A.2 Additional Experimental Settings on Domain Shifts
  • A.2.1 Dataset Details
  • A.2.2 Training Details
  • A.3 Strength of Spurious Correlation
  • A.4 Results on Datasets without Spurious Correlations
  • A.5 Additional Invariance Analysis
  • A.5.1 Additional Metrics of Invariant Predictor Analysis
  • A.5.2 Analysis of Learned Invariant Representations
  • A.6 Full Results of WILDS data
  • B Proofs of Theorem and Theorem
  • B.1 Decomposing the loss function
  • B.2 Classification errors of four methods with infinite training samples
  • B.2.1 Baseline method: ERM
  • B.2.2 Baseline method: Vanilla mixup
  • B.3 Intra-label LISA (LISA-L): mixup across domain
  • B.4 Intra-domain LISA (LISA-D): mixup within each domain
  • B.5 Finite sample analysis
  • B.6 A ξ\xi-dependent lower bound for EERM(w​s​t)−ELL(w​s​t)E^{(wst)}_{\textup{ERM}}-E^{(wst)}_{\textup{LL}}
  • B.7 Domain shifts: Proof of Theorem

Knowls

  1. Knowl 1 — Learning Invariant Predictors with Selective Augmentation

    model/method

    Learning Invariant Predictors with Selective Augmentation (LISA) is a data augmentation framework designed to learn domain-invariant predictors under subpopulation shifts and domain shifts without adding explicit regularization terms to the objective function. Given training samples (xi,yi,di)(x_i, y_i, d_i) and (xj,yj,dj)(x_j, y_j, d_j) from input space X\mathcal{X}, label space Y\mathcal{Y} (represented as one-hot vectors), and domain index set D\mathcal{D}, LISA performs linear interpolation:

    xmix=λxi+(1−λ)xj,ymix=λyi+(1−λ)yjx_{\text{mix}} = \lambda x_i + (1 - \lambda) x_j, \quad y_{\text{mix}} = \lambda y_i + (1 - \lambda) y_j

    where λ∈[0,1]\lambda \in [0, 1] is sampled from a Beta distribution Beta(α,β)\text{Beta}(\alpha, \beta). LISA restricts the pairing of samples according to two selective augmentation strategies:

    1. Intra-label LISA (LISA-L): Pairs samples with identical task labels but different domains (yi=yjy_i = y_j and di≠djd_i \neq d_j). This creates synthetic examples where the domain features are blended while the task label remains constant, neutralizing spurious associations between domain characteristics and labels.
    2. Intra-domain LISA (LISA-D): Pairs samples with identical domains but different task labels (di=djd_i = d_j and yi≠yjy_i \neq y_j). Because the domain remains fixed while the label varies continuously with λ\lambda, the model is forced to make predictions based on label-defining features rather than domain context.

    During training, LISA stochastically chooses between intra-label and intra-domain augmentation with probability psel∈[0,1]p_{\text{sel}} \in [0, 1] and 1−psel1 - p_{\text{sel}}, respectively. The standard linear interpolation can be replaced by CutMix for image tasks or Manifold Mixup on pre-trained representations for text tasks.

  2. Knowl 2 — Training Procedure of LISA

    algorithm

    The LISA optimization algorithm updates model parameters θ\theta by sampling minibatches and constructing interpolated batches using either intra-label or intra-domain sampling based on a selection probability pselp_{\text{sel}}.

    Input: Training dataset DD, learning rate γ\gamma, Beta distribution shape parameters α,β\alpha, \beta, selection probability psel∈[0,1]p_{sel} \in [0, 1]
    Output: Trained model parameters θ\theta
    while model not converged do
        Sample interpolation ratio λ∼Beta(α,β)\lambda \sim \text{Beta}(\alpha, \beta)
        Sample primary minibatch B1={(xi,yi,di)}i=1BB_1 = \{(x_i, y_i, d_i)\}_{i=1}^B from DD
        Initialize secondary minibatch B2←{}B_2 \leftarrow \{\}
        Sample strategy indicator s∼Bernoulli(psel)s \sim \text{Bernoulli}(p_{sel})
        if s=1s = 1 then
            for each (xi,yi,di)∈B1(x_i, y_i, d_i) \in B_1 do
                Randomly sample (xj,yj,dj)∈D(x_j, y_j, d_j) \in D satisfying (yj=yi)(y_j = y_i) and (dj≠di)(d_j \neq d_i)
                Append (xj,yj,dj)(x_j, y_j, d_j) to B2B_2
            end for
        else
            for each (xi,yi,di)∈B1(x_i, y_i, d_i) \in B_1 do
                Randomly sample (xj,yj,dj)∈D(x_j, y_j, d_j) \in D satisfying (yj≠yi)(y_j \neq y_i) and (dj=di)(d_j = d_i)
                Append (xj,yj,dj)(x_j, y_j, d_j) to B2B_2
            end for
        end if
        Construct mixed batch Bmix={(λxi+(1−λ)xj,λyi+(1−λ)yj)}i=1BB_{mix} = \{(\lambda x_i + (1 - \lambda) x_j, \lambda y_i + (1 - \lambda) y_j)\}_{i=1}^B
        Update parameters θ←θ−γ∇θ1B∑(x,y)∈Bmixℓ(fθ(x),y)\theta \leftarrow \theta - \gamma \nabla_\theta \frac{1}{B} \sum_{(x, y) \in B_{mix}} \ell(f_\theta(x), y)
    end while
    return θ\theta
  3. Knowl 3 — Worst-Group Misclassification Error Bound under Subpopulation Shifts

    theoretical result

    Consider a Gaussian mixture data model with binary label y∈{0,1}y \in \{0, 1\} and binary domain d∈{R,G}d \in \{R, G\}, where xi∣yi=y,di=d∼N(μ(y,d),Σ)x_i \mid y_i = y, d_i = d \sim \mathcal{N}(\mu^{(y,d)}, \Sigma) for positive definite covariance matrix Σ∈Rp×p\Sigma \in \mathbb{R}^{p \times p} and conditional means μ(y,d)∈Rp\mu^{(y,d)} \in \mathbb{R}^p. Assume cross-domain label shift invariance such that μ(1,R)−μ(0,R)=μ(1,G)−μ(0,G)=Δ\mu^{(1,R)} - \mu^{(0,R)} = \mu^{(1,G)} - \mu^{(0,G)} = \Delta, balanced marginals π(R)=π(1)=1/2\pi^{(R)} = \pi^{(1)} = 1/2, and minority group proportion π(0,R)=π(1,G)=α<1/4\pi^{(0,R)} = \pi^{(1,G)} = \alpha < 1/4. Let Δ~=E[xi∣yi=1]−E[xi∣yi=0]\tilde{\Delta} = \mathbb{E}[x_i \mid y_i = 1] - \mathbb{E}[x_i \mid y_i = 0], ∥v∥Σ=v⊤Σ−1v\|v\|_\Sigma = \sqrt{v^\top \Sigma^{-1} v}, and define the correlation parameter:

    ξ=Δ⊤Σ−1Δ~∥Δ∥Σ∥Δ~∥Σ\xi = \frac{\Delta^\top \Sigma^{-1} \tilde{\Delta}}{\|\Delta\|_\Sigma \|\tilde{\Delta}\|_\Sigma}

    where small ξ\xi indicates strong spurious correlation between domain and label. Let E^A(wst)\hat{E}^{(\text{wst})}_A denote the worst-group classification error on finite samples of size nn using method AA.

    If max⁡y,d∥μ(y,d)∥2≤C\max_{y,d} \|\mu^{(y,d)}\|_2 \le C, and (ξ,α)(\xi, \alpha) satisfies:

    ξ<min⁡{∥Δ~∥Σ∥Δ∥Σ,∥Δ∥Σ∥Δ~∥Σ}−Cα\xi < \min\left\{ \frac{\|\tilde{\Delta}\|_\Sigma}{\|\Delta\|_\Sigma}, \frac{\|\Delta\|_\Sigma}{\|\tilde{\Delta}\|_\Sigma} \right\} - C\alpha

    and E[λi2]/max⁡{var(λi),1/4}≥∥Δ~∥Σ2+∥Δ~∥Σ∥Δ∥Σ\mathbb{E}[\lambda_i^2] / \max\{\text{var}(\lambda_i), 1/4\} \ge \|\tilde{\Delta}\|_\Sigma^2 + \|\tilde{\Delta}\|_\Sigma \|\Delta\|_\Sigma, then for any selection probability psel∈[0,1]p_{\text{sel}} \in [0, 1]:

    E^LISA(wst)<min⁡{E^ERM(wst),E^mix(wst)}+OP(plog⁡nn+pαn)\hat{E}^{(\text{wst})}_{\text{LISA}} < \min\left\{ \hat{E}^{(\text{wst})}_{\text{ERM}}, \hat{E}^{(\text{wst})}_{\text{mix}} \right\} + O_P\left( \sqrt{\frac{p \log n}{n}} + \sqrt{\frac{p}{\alpha n}} \right)

    This proves that when spurious correlations exist (small ξ\xi) and p=o(αn)p = o(\alpha n), LISA achieves strictly smaller worst-group error than both empirical risk minimization (ERM) and vanilla mixup.

  4. Knowl 4 — Worst-Group Misclassification Error Bound under Domain Shifts

    theoretical result

    Consider the Gaussian mixture model with training data from domains d∈{R,G}d \in \{R, G\} and a test distribution drawn from a new, unseen domain:

    xi(0,∗)∼N(μ(0,∗),Σ),xi(1,∗)∼N(μ(1,∗),Σ)x_i^{(0,*)} \sim \mathcal{N}(\mu^{(0,*)}, \Sigma), \quad x_i^{(1,*)} \sim \mathcal{N}(\mu^{(1,*)}, \Sigma)

    where μ(1,∗)−μ(0,∗)=Δ\mu^{(1,*)} - \mu^{(0,*)} = \Delta. Define Δ~∗=2(μ(0,∗)−E[xi])\tilde{\Delta}^* = 2(\mu^{(0,*)} - \mathbb{E}[x_i]), ξ∗=Δ~⊤Σ−1Δ~∗∥Δ~∥Σ∥Δ∥Σ\xi^* = \frac{\tilde{\Delta}^\top \Sigma^{-1} \tilde{\Delta}^*}{\|\tilde{\Delta}\|_\Sigma \|\Delta\|_\Sigma}, and γ=Δ⊤Σ−1Δ~∗∥Δ~∥Σ∥Δ∥Σ\gamma = \frac{\Delta^\top \Sigma^{-1} \tilde{\Delta}^*}{\|\tilde{\Delta}\|_\Sigma \|\Delta\|_\Sigma}, with ∥v∥Σ=v⊤Σ−1v\|v\|_\Sigma = \sqrt{v^\top \Sigma^{-1} v}.

    Let E^A(wst∗)=max⁡y∈{0,1}E(y,∗)(bA,b0,A)\hat{E}^{(\text{wst}*)}_{A} = \max_{y \in \{0, 1\}} E^{(y,*)}(b_A, b_{0,A}) be the finite-sample worst-group error in the unseen domain. If max⁡y,d∥μ(y,d)∥2≤C\max_{y,d} \|\mu^{(y,d)}\|_2 \le C, 0≤ξ∗≤γξ0 \le \xi^* \le \gamma \xi, and:

    ξ<min⁡{γ2∥Δ~∥Σ∥Δ∥Σ,∥Δ∥Σ∥Δ~∥Σ}−Cα\xi < \min\left\{ \frac{\gamma}{2} \frac{\|\tilde{\Delta}\|_\Sigma}{\|\Delta\|_\Sigma}, \frac{\|\Delta\|_\Sigma}{\|\tilde{\Delta}\|_\Sigma} \right\} - C\alpha

    with E[λi2]/max⁡{var(λi),1/4}≥∥Δ~∥Σ2+∥Δ~∥Σ∥Δ∥Σ\mathbb{E}[\lambda_i^2] / \max\{\text{var}(\lambda_i), 1/4\} \ge \|\tilde{\Delta}\|_\Sigma^2 + \|\tilde{\Delta}\|_\Sigma \|\Delta\|_\Sigma, then for any psel∈[0,1]p_{\text{sel}} \in [0, 1]:

    E^LISA(wst∗)<min⁡{E^ERM(wst∗),E^mix(wst∗)}+OP(plog⁡nn+pαn)\hat{E}^{(\text{wst}*)}_{\text{LISA}} < \min\left\{ \hat{E}^{(\text{wst}*)}_{\text{ERM}}, \hat{E}^{(\text{wst}*)}_{\text{mix}} \right\} + O_P\left( \sqrt{\frac{p \log n}{n}} + \sqrt{\frac{p}{\alpha n}} \right)

    This shows that selective augmentation allows the linear classifier to generalise to unseen test domains with lower error than ERM or vanilla mixup when spurious correlations are present in the training environments.

  5. Knowl 5 — Empirical Performance on Subpopulation Shift Benchmarks

    data/table

    LISA was evaluated on four subpopulation shift benchmarks: Colored MNIST (CMNIST), Waterbirds, CelebA, and CivilComments. Evaluation focuses on worst-group accuracy (Worst) alongside average accuracy (Avg).

    Method CMNIST Waterbirds CelebA CivilComments
    Avg. Worst Avg. Worst Avg. Worst Avg. Worst
    ERM 27.8% 0.0% 97.0% 63.7% 94.9% 47.8% 92.2% 56.0%
    UW 72.2% 66.0% 95.1% 88.0% 92.9% 83.3% 89.8% 69.2%
    IRM 72.1% 70.3% 87.5% 75.6% 94.0% 77.8% 88.8% 66.3%
    IB-IRM 72.2% 70.7% 88.5% 76.5% 93.6% 85.0% 89.1% 65.3%
    V-REx 71.7% 70.2% 88.0% 73.6% 92.2% 86.7% 90.2% 64.9%
    CORAL 71.8% 69.5% 90.3% 79.8% 93.8% 76.9% 88.7% 65.6%
    GroupDRO 72.3% 68.6% 91.8% 90.6% 92.1% 87.2% 89.9% 70.0%
    DomainMix 51.4% 48.0% 76.4% 53.0% 93.4% 65.6% 90.9% 63.6%
    Fish 46.9% 35.6% 85.6% 64.0% 93.1% 61.2% 89.8% 71.1%
    LISA (ours) 74.0% 73.3% 91.8% 89.2% 92.4% 89.3% 89.2% 72.6%

    LISA achieves the top worst-group performance on CMNIST (73.3%73.3\%), CelebA (89.3%89.3\%), and CivilComments (72.6%72.6\%), and is competitive with GroupDRO on Waterbirds (89.2%89.2\% vs 90.6%90.6\%), while remaining consistently superior across diverse datasets compared to explicit invariance regularizers (such as IRM and V-REx).

  6. Knowl 6 — Empirical Performance on Domain Shift Benchmarks

    data/table

    LISA was evaluated on five domain shift benchmarks from WILDS and MetaShift covering medical imaging (Camelyon17), satellite imagery (FMoW), cellular microscopy (RxRx1), text review sentiment (Amazon), and natural scenes (MetaShift).

    Method Camelyon17 FMoW RxRx1 Amazon MetaShift
    Avg. Acc. Worst Acc. Avg. Acc. 10-th Per. Acc. Worst Acc.
    ERM 70.3 ±\pm 6.4% 32.3 ±\pm 1.25% 29.9 ±\pm 0.4% 53.8 ±\pm 0.8% 52.1 ±\pm 0.4%
    IRM 64.2 ±\pm 8.1% 30.0 ±\pm 1.37% 8.2 ±\pm 1.1% 52.4 ±\pm 0.8% 51.8 ±\pm 0.8%
    IB-IRM 68.9 ±\pm 6.1% 28.4 ±\pm 0.90% 6.4 ±\pm 0.6% 53.8 ±\pm 0.7% 52.3 ±\pm 1.0%
    V-REx 71.5 ±\pm 8.3% 27.2 ±\pm 0.78% 7.5 ±\pm 0.8% 53.3 ±\pm 0.0% 51.6 ±\pm 1.8%
    CORAL 59.5 ±\pm 7.7% 31.7 ±\pm 1.24% 28.4 ±\pm 0.3% 52.9 ±\pm 0.8% 47.6 ±\pm 1.9%
    GroupDRO 68.4 ±\pm 7.3% 30.8 ±\pm 0.81% 23.0 ±\pm 0.3% 53.3 ±\pm 0.0% 51.9 ±\pm 0.7%
    DomainMix 69.7 ±\pm 5.5% 34.2 ±\pm 0.76% 30.8 ±\pm 0.4% 53.3 ±\pm 0.0% 51.3 ±\pm 0.5%
    Fish 74.7 ±\pm 7.1% 34.6 ±\pm 0.18% 10.1 ±\pm 1.5% 53.3 ±\pm 0.0% 49.2 ±\pm 2.1%
    LISA (ours) 77.1 ±\pm 6.5% 35.5 ±\pm 0.65% 31.9 ±\pm 0.8% 54.7 ±\pm 0.0% 54.2 ±\pm 0.7%

    For these domain shift tasks, intra-label LISA (psel=1.0p_{\text{sel}} = 1.0) is used. LISA outperforms all baselines across every benchmark dataset, including both vision and NLP modalities.

  7. Knowl 7 — Ablation of Selective Augmentation against Substitute Mixup Strategies

    empirical result

    To evaluate whether LISA's performance stems from learning domain-invariant predictors rather than general data augmentation regularization, LISA was compared against two substitute mixup strategies:

    1. Vanilla Mixup: Pairs randomly chosen samples without group or domain restrictions.
    2. In-Group Mixup: Pairs samples belonging to the same label and the same domain (yi=yj,di=djy_i = y_j, d_i = d_j).

    In subpopulation shifts, Vanilla Mixup achieves only 3.1%3.1\% worst-group accuracy on CMNIST, 56.2%56.2\% on Waterbirds, and 46.4%46.4\% on CelebA. Combining Vanilla Mixup with group upweighting (UW) improves worst-group accuracy to 71.8%71.8\% on CMNIST and 85.6%85.6\% on Waterbirds, but LISA reaches 73.3%73.3\% and 89.2%89.2\%, respectively. Similarly, In-group Mixup + UW attains 71.6%71.6\% on CMNIST and 87.1%87.1\% on Waterbirds.

    On datasets modified to remove spurious correlations (balanced groups), Vanilla Mixup outperforms ERM and LISA (e.g., 74.28%74.28\% vs 73.18%73.18\% on CMNIST, 88.89%88.89\% vs 87.22%87.22\% on CelebA). However, under spurious correlations, Vanilla Mixup fails while LISA achieves the best performance. This confirms that LISA's gains come specifically from selectively eliminating spurious associations across domains.

  8. Knowl 8 — Predictor and Representation Invariance Quantitative Metrics

    data/table

    Model invariance is evaluated using unscaled model outputs (logits) gi,dg_{i,d} via two predictor invariance metrics and one representation invariance metric:

    • Accuracy of domain prediction (IPadpIP_{\text{adp}}): Accuracy of a logistic regression model trained on frozen logits to predict the domain label dd. Lower values indicate that logits contain less domain-identifying information.
    • Pairwise prediction KL divergence (IPklIP_{\text{kl}}): Average pairwise KL divergence between kernel-density estimated logit distributions across domains for a given class: 1∣Y∣∣D∣2∑y∈Y∑d,d′∈DKL(P(gDy∣D=d)∥P(gDy∣D=d′))\frac{1}{|\mathcal{Y}||\mathcal{D}|^2} \sum_{y \in \mathcal{Y}} \sum_{d, d' \in \mathcal{D}} \text{KL}(P(g_D^y \mid D=d) \parallel P(g_D^y \mid D=d')).
    • Pairwise representation KL divergence (IRklIR_{\text{kl}}): Pairwise KL divergence computed on pre-classifier latent representations hi,dh_{i,d}.
    Method IPadp↓IP_{\text{adp}} \downarrow IPkl↓IP_{\text{kl}} \downarrow
    CMNIST Waterbirds Camelyon17 MetaShift CMNIST Waterbirds Camelyon17 MetaShift
    ERM 82.85% 94.99% 49.43% 67.98% 6.286 1.888 1.536 1.205
    Vanilla mixup 92.34% 94.49% 52.79% 69.36% 4.737 2.912 0.790 1.171
    IRM 69.42% 95.12% 47.96% 67.59% 7.755 1.122 0.875 1.148
    IB-IRM 74.72% 94.78% 48.37% 67.39% 1.004 3.563 0.756 1.115
    V-REx 63.58% 93.32% 61.38% 68.38% 3.190 3.791 1.281 1.094
    LISA (ours) 58.42% 90.28% 45.15% 66.01% 0.567 0.134 0.723 1.001

    LISA achieves the lowest domain prediction accuracy and pairwise divergence across all datasets, confirming that it yields substantially more invariant output logits and latent representations (IRkl=0.421×10−8IR_{\text{kl}} = 0.421 \times 10^{-8} on CMNIST vs 1.683×10−81.683 \times 10^{-8} for ERM).

  9. Knowl 9 — Robustness Gain of LISA Scales with Distribution Shift Severity

    empirical result

    On the MetaShift dataset, distribution shift distance between training and test sets is quantitatively varied by altering background contexts based on meta-graph node distances (distances: 0.44,0.71,1.12,1.430.44, 0.71, 1.12, 1.43). As the distance increases:

    • At distance 0.440.44: ERM gets 80.1%80.1\%, IB-IRM gets 79.7%79.7\%, and LISA gets 81.3%81.3\% (+1.5%+1.5\% gain over best baseline).
    • At distance 0.710.71: ERM gets 68.4%68.4\%, GroupDRO gets 68.9%68.9\%, and LISA gets 69.7%69.7\% (+1.2%+1.2\% gain over best baseline).
    • At distance 1.121.12: ERM gets 52.1%52.1\%, IB-IRM gets 52.3%52.3\%, and LISA gets 54.2%54.2\% (+3.6%+3.6\% gain over best baseline).
    • At distance 1.431.43: ERM gets 33.2%33.2\%, GroupDRO gets 34.2%34.2\%, and LISA gets 37.5%37.5\% (+9.6%+9.6\% gain over best baseline).

    The performance improvement of LISA over the strongest baseline widens monotonically as the distribution distance grows, indicating that selective domain interpolation provides greater benefits under severe domain shifts.

  10. Knowl 10 — Limitations of LISA

    limitation

    LISA exhibits two main limitations:

    1. Dependence on Sample Matching by Exact Label: Intra-label LISA requires sampling pairs of data points sharing the same task label yy. In machine learning settings where labels are continuous, high-dimensional, unique, or complex (such as dense object detection, bounding box regression, and generative modeling), finding or matching samples with identical labels is non-trivial or impossible without task-specific approximations.
    2. Heuristic Application in Domain-Free Scenarios: On datasets with weak spurious correlations (e.g., Camelyon17, FMoW, and RxRx1 where Cramér's V is near zero), empirical performance is highest when applying intra-label interpolation without conditioning on domain annotations. However, domain-free intra-label interpolation lacks a formal theoretical guarantee and general selection principle.

Coverage note — None was omitted; all contributed algorithms, theoretical bounds, empirical benchmarks across subpopulation and domain shifts, ablation studies, invariance analyses, and stated limitations are fully represented.

References

  1. 1.Ahuja, K., Caballero, E., Zhang, D., Bengio, Y., Mitliagkas, I., and Rish, I. Invariance principle meets information bottleneck for out-of-distribution generalization. 2021.
  2. 2.Albuquerque, I., Monteiro, J., Darvishi, M., Falk, T. H., and Mitliagkas, I. Generalizing to unseen domains via distribution matching. arXiv preprint arXiv:1911.00804, 2019.
  3. 3.Anderson, T. W. An introduction to multivariate statistical analysis. Technical report, Wiley New York, 1962.
  4. 4.Arjovsky, M., Bottou, L., Gulrajani, I., and Lopez-Paz, D. Invariant risk minimization. arXiv preprint arXiv:1907.02893, 2019.
  5. 5.Bandi, P., Geessink, O., Manson, Q., Van Dijk, M., Balkenhol, M., Hermsen, M., Bejnordi, B. E., Lee, B., Paeng, K., Zhong, A., et al. From detection of individual metastases to classification of lymph node status at the patient level: the camelyon17 challenge. IEEE Transactions on Medical Imaging, 2018.
  6. 6.Ben-David, S., Blitzer, J., Crammer, K., Kulesza, A., Pereira, F., and Vaughan, J. W. A theory of learning from different domains. Machine learning, 79(1):151–175, 2010.
  7. 7.Borkan, D., Dixon, L., Sorensen, J., Thain, N., and Vasserman, L. Nuanced metrics for measuring unintended bias with real data for text classification. In Companion proceedings of the 2019 world wide web conference, pp. 491–500, 2019.
  8. 8.Cai, T. T. and Zhang, L. A convex optimization approach to high-dimensional sparse quadratic discriminant analysis. The Annals of Statistics, 49(3):1537–1568, 2021.
  9. 9.Cao, K., Wei, C., Gaidon, A., Arechiga, N., and Ma, T. Learning imbalanced datasets with label-distribution-aware margin loss. NeurIPS, 2019.
  10. 10.Cao, K., Chen, Y., Lu, J., Arechiga, N., Gaidon, A., and Ma, T. Heteroskedastic and imbalanced deep learning with adaptive regularization. arXiv preprint arXiv:2006.15766, 2020.
  11. 11.Chang, S., Zhang, Y., Yu, M., and Jaakkola, T. Invariant rationalization. In International Conference on Machine Learning, pp. 1448–1458. PMLR, 2020.
  12. 12.Christie, G., Fendley, N., Wilson, J., and Mukherjee, R. Functional map of the world. In Proceedings of the IEEE Conference on Computer Vision and Pattern Recognition, 2018.
  13. 13.Chuang, C.-Y. and Mroueh, Y. Fair mixup: Fairness via interpolation. ICLR, 2021.
  14. 14.Cramér, H. Mathematical Methods of Statistics (PMS-9), Volume 9. Princeton university press, 2016.
  15. 15.Creager, E., Jacobsen, J.-H., and Zemel, R. Environment inference for invariant learning. In International Conference on Machine Learning, pp. 2189–2200. PMLR, 2021.
  16. 16.Devlin, J., Chang, M.-W., Lee, K., and Toutanova, K. Bert: Pre-training of deep bidirectional transformers for language understanding. 2019.
  17. 17.Ganin, Y. and Lempitsky, V. Unsupervised domain adaptation by backpropagation. In International conference on machine learning, pp. 1180–1189. PMLR, 2015.
  18. 18.Ganin, Y., Ustinova, E., Ajakan, H., Germain, P., Larochelle, H., Laviolette, F., Marchand, M., and Lempitsky, V. Domain-adversarial training of neural networks. The journal of machine learning research, 17(1):2096–2030, 2016.
  19. 19.Goel, K., Gu, A., Li, Y., and Ré, C. Model patching: Closing the subgroup performance gap with data augmentation. In ICLR, 2021.
  20. 20.Guo, R., Zhang, P., Liu, H., and Kiciman, E. Out-of-distribution prediction with invariant risk minimization: The limitation and an effective fix. arXiv preprint arXiv:2101.07732, 2021.
  21. 21.He, K., Zhang, X., Ren, S., and Sun, J. Deep residual learning for image recognition. In Proceedings of the IEEE conference on computer vision and pattern recognition, pp. 770–778, 2016.
  22. 22.Huang, G., Liu, Z., Van Der Maaten, L., and Weinberger, K. Q. Densely connected convolutional networks. In Proceedings of the IEEE conference on computer vision and pattern recognition, pp. 4700–4708, 2017.
  23. 23.Khezeli, K., Blaas, A., Soboczenski, F., Chia, N., and Kalantari, J. On invariance penalties for risk minimization. arXiv preprint arXiv:2106.09777, 2021.
  24. 24.Koh, P. W., Sagawa, S., Xie, S. M., Zhang, M., Balsubramani, A., Hu, W., Yasunaga, M., Phillips, R. L., Gao, I., Lee, T., et al. Wilds: A benchmark of in-the-wild distribution shifts. In International Conference on Machine Learning, pp. 5637–5664. PMLR, 2021.
  25. 25.Koyama, M. and Yamaguchi, S. Out-of-distribution generalization with maximal invariant predictor. arXiv preprint arXiv:2008.01883, 2020.
  26. 26.Krishna, R., Zhu, Y., Groth, O., Johnson, J., Hata, K., Kravitz, J., Chen, S., Kalantidis, Y., Li, L.-J., Shamma, D. A., Bernstein, M., and Fei-Fei, L. Visual genome: Connecting language and vision using crowdsourced dense image annotations. 2016. URL https://arxiv.org/abs/1602.07332.
  27. 27.Krueger, D., Caballero, E., Jacobsen, J.-H., Zhang, A., Binas, J., Zhang, D., Le Priol, R., and Courville, A. Out-of-distribution generalization via risk extrapolation (rex). In International Conference on Machine Learning, pp. 5815–5826. PMLR, 2021.
  28. 28.Lee, H. B., Nam, T., Yang, E., and Hwang, S. J. Meta dropout: Learning to perturb latent features for generalization. In International Conference on Learning Representations, 2019.
  29. 29.Lee, Y., Yao, H., and Finn, C. Diversify and disambiguate: Learning from underspecified data. arXiv preprint arXiv:2202.03418, 2022.
  30. 30.Li, H., Pan, S. J., Wang, S., and Kot, A. C. Domain generalization with adversarial feature learning. In Proceedings of the IEEE Conference on Computer Vision and Pattern Recognition, pp. 5400–5409, 2018.
  31. 31.Liang, W. and Zou, J. Metadataset: A dataset of datasets for evaluating distribution shifts and training conflicts. In ICML2021 ML4data Workshop, 2021.
  32. 32.Liu, E. Z., Haghgoo, B., Chen, A. S., Raghunathan, A., Koh, P. W., Sagawa, S., Liang, P., and Finn, C. Just train twice: Improving group robustness without training group information. In ICML, pp. 6781–6792. PMLR, 2021a.
  33. 33.Liu, H., HaoChen, J. Z., Gaidon, A., and Ma, T. Self-supervised learning is more robust to dataset imbalance. arXiv preprint arXiv:2110.05025, 2021b.
  34. 34.Liu, Z., Luo, P., Wang, X., and Tang, X. Deep learning face attributes in the wild. In ICCV, 2015.
  35. 35.Long, M., Cao, Y., Wang, J., and Jordan, M. Learning transferable features with deep adaptation networks. In International conference on machine learning, pp. 97–105. PMLR, 2015.
  36. 36.Montanari, A., Ruan, F., Sohn, Y., and Yan, J. The generalization error of max-margin linear classifiers: High-dimensional asymptotics in the overparametrized regime. arXiv preprint arXiv:1911.01544, 2019.
  37. 37.Muandet, K., Balduzzi, D., and Schölkopf, B. Domain generalization via invariant feature representation. In International Conference on Machine Learning, pp. 10–18. PMLR, 2013.
  38. 38.Nam, J., Cha, H., Ahn, S.-S., Lee, J., and Shin, J. Learning from failure: De-biasing classifier from biased classifier. Advances in Neural Information Processing Systems, 33, 2020.
  39. 39.Ni, J., Li, J., and McAuley, J. Justifying recommendations using distantly-labeled reviews and fine-grained aspects. In Proceedings of the 2019 Conference on Empirical Methods in Natural Language Processing and the 9th International Joint Conference on Natural Language Processing (EMNLP-IJCNLP), 2019.
  40. 40.Qiao, F., Zhao, L., and Peng, X. Learning to learn single domain generalization. In Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition, pp. 12556–12565, 2020.
  41. 41.Rosenfeld, E., Ravikumar, P., and Risteski, A. The risks of invariant risk minimization. In ICLR, 2021.
  42. 42.Sagawa, S., Koh, P. W., Hashimoto, T. B., and Liang, P. Distributionally robust neural networks for group shifts: On the importance of regularization for worst-case generalization. In ICLR, 2020a.
  43. 43.Sagawa, S., Raghunathan, A., Koh, P. W., and Liang, P. An investigation of why overparameterization exacerbates spurious correlations. In ICML, pp. 8346–8356. PMLR, 2020b.
  44. 44.Sanh, V., Debut, L., Chaumond, J., and Wolf, T. Distilbert, a distilled version of bert: smaller, faster, cheaper and lighter. arXiv preprint arXiv:1910.01108, 2019.
  45. 45.Shi, Y., Seely, J., Torr, P. H., Siddharth, N., Hannun, A., Usunier, N., and Synnaeve, G. Gradient matching for domain generalization. arXiv preprint arXiv:2104.09937, 2021.
  46. 46.Shu, Y., Cao, Z., Wang, C., Wang, J., and Long, M. Open domain generalization with domain-augmented meta-learning. In Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition, pp. 9624–9633, 2021.
  47. 47.Sun, B. and Saenko, K. Deep coral: Correlation alignment for deep domain adaptation. In European conference on computer vision, pp. 443–450. Springer, 2016.
  48. 48.Taylor, J., Earnshaw, B., Mabey, B., Victors, M., and Yosinski, J. Rxrx1: An image set for cellular morphological variation across many experimental batches. In International Conference on Learning Representations (ICLR), 2019.
  49. 49.Tony Cai, T. and Zhang, L. High dimensional linear discriminant analysis: optimality, adaptive algorithm and missing data. Journal of the Royal Statistical Society: Series B (Statistical Methodology), 81(4):675–705, 2019.
  50. 50.Tzeng, E., Hoffman, J., Zhang, N., Saenko, K., and Darrell, T. Deep domain confusion: Maximizing for domain invariance. arXiv preprint arXiv:1412.3474, 2014.
  51. 51.Verma, V., Lamb, A., Beckham, C., Najafi, A., Mitliagkas, I., Lopez-Paz, D., and Bengio, Y. Manifold mixup: Better representations by interpolating hidden states. In International Conference on Machine Learning, pp. 6438–6447. PMLR, 2019.
  52. 52.Volpi, R., Namkoong, H., Sener, O., Duchi, J., Murino, V., and Savarese, S. Generalizing to unseen domains via adversarial data augmentation. arXiv preprint arXiv:1805.12018, 2018.
  53. 53.Wah, C., Branson, S., Welinder, P., Perona, P., and Belongie, S. The Caltech-UCSD Birds-200-2011 Dataset. Technical Report CNS-TR-2011-001, California Institute of Technology, 2011.
  54. 54.Wang, Y., Li, H., and Kot, A. C. Heterogeneous domain generalization via domain mixup. In ICASSP 2020-2020 IEEE International Conference on Acoustics, Speech and Signal Processing (ICASSP), pp. 3622–3626. IEEE, 2020.
  55. 55.Xu, M., Zhang, J., Ni, B., Li, T., Wang, C., Tian, Q., and Zhang, W. Adversarial domain adaptation with domain mixup. In Proceedings of the AAAI Conference on Artificial Intelligence, volume 34, pp. 6502–6509, 2020.
  56. 56.Yan, S., Song, H., Li, N., Zou, L., and Ren, L. Improve unsupervised domain adaptation with mixup training. arXiv preprint arXiv:2001.00677, 2020.
  57. 57.Yao, H., Zhang, L., and Finn, C. Meta-learning with fewer tasks through task interpolation. In International Conference on Learning Representations, 2021.
  58. 58.Yue, X., Zhang, Y., Zhao, S., Sangiovanni-Vincentelli, A., Keutzer, K., and Gong, B. Domain randomization and pyramid consistency: Simulation-to-real generalization without accessing target domain data. In Proceedings of the IEEE/CVF International Conference on Computer Vision, pp. 2100–2110, 2019.
  59. 59.Yun, S., Han, D., Oh, S. J., Chun, S., Choe, J., and Yoo, Y. Cutmix: Regularization strategy to train strong classifiers with localizable features. In Proceedings of the IEEE/CVF International Conference on Computer Vision, pp. 6023–6032, 2019.
  60. 60.Zhang, H., Cisse, M., Dauphin, Y. N., and Lopez-Paz, D. mixup: Beyond empirical risk minimization. 2018.
  61. 61.Zhang, J., Menon, A., Veit, A., Bhojanapalli, S., Kumar, S., and Sra, S. Coping with label shift via distributionally robust optimisation. In ICLR, 2021a.
  62. 62.Zhang, L., Deng, Z., Kawaguchi, K., Ghorbani, A., and Zou, J. How does mixup help with robustness and generalization? In ICLR, 2021b.
  63. 63.Zhang, L., Deng, Z., Kawaguchi, K., and Zou, J. When and how mixup improves calibration. arXiv preprint arXiv:2102.06289, 2021c.
  64. 64.Zhang, M., Sohoni, N. S., Zhang, H. R., Finn, C., and Ré, C. Correct-n-contrast: A contrastive approach for improving robustness to spurious correlations. In NeurIPS 2021 Workshop on Distribution Shifts: Connecting Methods and Applications, 2021d.
  65. 65.Zhao, L., Liu, T., Peng, X., and Metaxas, D. Maximum-entropy adversarial data augmentation for improved generalization and robustness. arXiv preprint arXiv:2010.08001, 2020.
  66. 66.Zhou, B., Lapedriza, A., Khosla, A., Oliva, A., and Torralba, A. Places: A 10 million image database for scene recognition. IEEE transactions on pattern analysis and machine intelligence, 40(6):1452–1464, 2017.
  67. 67.Zhou, C., Ma, X., Michel, P., and Neubig, G. Examining and combating spurious features under distribution shift. In ICML, 2021.
  68. 68.Zhou, F., Jiang, Z., Shui, C., Wang, B., and Chaib-draa, B. Domain generalization with optimal transport and metric learning. arXiv preprint arXiv:2007.10573, 2020a.
  69. 69.Zhou, K., Yang, Y., Hospedales, T., and Xiang, T. Deep domain-adversarial image generation for domain generalisation. In Proceedings of the AAAI Conference on Artificial Intelligence, volume 34, pp. 13025–13032, 2020b.

Citation

MLA
Yao, H., et al. “Improving Out-of-Distribution Robustness via Selective Augmentation”. arXiv, 2022, http://arxiv.org/abs/2201.00299v3.
APA
Yao, H., Wang, Y., Li, S., Zhang, L., Liang, W., Zou, J., & Finn, C. (2022). Improving Out-of-Distribution Robustness via Selective Augmentation. arXiv. http://arxiv.org/abs/2201.00299v3
Chicago
Yao, H., Y. Wang, S. Li, et al. 2022. “Improving Out-of-Distribution Robustness via Selective Augmentation”. arXiv. http://arxiv.org/abs/2201.00299v3.
Harvard
Yao, H. et al. (2022) “Improving Out-of-Distribution Robustness via Selective Augmentation”, arXiv [Preprint]. Available at: http://arxiv.org/abs/2201.00299v3.
Vancouver
1. Yao H, Wang Y, Li S, Zhang L, Liang W, Zou J, Finn C (2022) Improving Out-of-Distribution Robustness via Selective Augmentation. arXiv

BibTeX

@article{yao2022improving,
  title = {Improving Out-of-Distribution Robustness via Selective Augmentation},
  author = {Yao, Huaxiu and Wang, Yu and Li, Sai and Zhang, Linjun and Liang, Weixin and Zou, James and Finn, Chelsea},
  year = {2022},
  journal = {arXiv},
  url = {http://arxiv.org/abs/2201.00299v3},
  eprint = {2201.00299}
}
Metadata:arXiv

Source Code

This paper has an official code repository available. Click below to access the source code.

View Repository

Access the Paper

This paper is available from its original source. Click below to access the PDF.

Open PDF
License: https://creativecommons.org/licenses/by/4.0/