CAR-MIL: Counterfactual Attention Regularization for Multiple Instance Learning¶
Conference: ECCV 2026
Paper: ECCV 2026 Poster
PDF: ECCV 2026 PDF
Area: Medical Imaging
Keywords: Multiple Instance Learning / Counterfactual Explanations / Attention Regularization / Digital Pathology / Whole Slide Imaging
TL;DR¶
To address the long-standing challenges of unfaithful attention attribution and lack of contradictory evidence in Multiple Instance Learning (MIL), CAR-MIL introduces a dual-branch counterfactual attention regularization framework that enforces prediction divergence under minimal, structured attention perturbations, decoupling complementary supporting and refuting evidence distributions while boosting classification performance across five digital pathology benchmarks and synthetic tasks.
Background & Motivation¶
Multiple Instance Learning (MIL) has emerged as the standard paradigm for weakly-supervised learning in computational pathology, where gigapixel whole-slide images (WSIs) are divided into thousands of patch instances with only slide-level diagnostic labels. Standard attention-based MIL architectures (such as ABMIL, CLAM, and TransMIL) learn attention scores to aggregate instance embeddings into a global bag representation. Despite their impressive classification accuracy, existing models suffer from two critical limitations: first, attention weights are learned merely as an implicit byproduct of classification loss minimization without explicit structural guidance, often resulting in attention collapse or overfitting to spurious correlations in non-diagnostic stroma; second, standard attention only captures evidence supporting the predicted class, offering no insight into which contradictory evidence would challenge or overturn the decision, severely impairing clinical reliability and interpretability.
Counterfactual Explanations (CE) offer a principled lens for reasoning about such evidence by identifying minimal input perturbations that alter a model's prediction, naturally differentiating supporting evidence from contradictory evidence. However, conventional CE methods are strictly post-hoc, iteratively optimizing input-space perturbations on already-trained models at prohibitive computational cost without ever refining the underlying attention representations during training. Meanwhile, existing attempts to integrate counterfactual interventions into training (such as CIA-MIL) rely on uncontrolled random attention perturbations, which act merely as a diagnostic heuristic and often degrade classification performance.
This paper shifts counterfactual reasoning from the input space to the attention logit space and embeds it directly into training dynamics. By constructing a lightweight counterfactual attention branch that shares the backbone feature extractor and downstream classifier with the factual branch, the model is regularized to produce an alternative prediction under minimal attention deviation. Core idea: introduce counterfactual attention regularization into MIL training via a lightweight dual-branch architecture and a dual 'prediction divergence under minimal attention proximity' constraint, forcing decision shifts to arise exclusively from sparse, structured redistributions of instance attention and thereby decoupling complementary supporting and refuting attention maps.
Method¶
Overall Architecture¶
CAR-MIL builds upon standard attention-based MIL. Given a bag of \(N\) instance patches \(B = \{x_1, \dots, x_N\}\), a frozen or pre-trained feature extractor \(E\) (such as the pathology foundation model UNI or ResNet) independently extracts instance embeddings \(z_j = E(x_j) \in \mathbb{R}^d\). The system then forks into a parallel dual-branch attention formulation: the Factual Branch computes raw instance attention logits \(u\) via attention module \(\psi\), normalizes them into attention weights \(a\) via Softmax, pools instance embeddings into a bag representation \(\hat{Z}\), and predicts bag-level class logits through classifier \(\varphi\); the Counterfactual Branch shares the exact same encoder \(E\) and downstream classifier \(\varphi\), but employs an independent, lightweight attention head \(\psi_{\text{cf}}\) to generate counterfactual attention logits \(u^{\text{cf}}\), pooled bag representation \(\hat{Z}^{\text{cf}}\), and counterfactual class logits.
During training, the framework is optimized end-to-end under three joint objectives: the factual classification loss guarantees accurate slide-level diagnosis; the counterfactual evidence differential loss forces a label-consistent divergence on the ground-truth class logits; and the attention logit proximity regularization penalizes large deviations between factual and counterfactual logit vectors, ensuring that decision shifts stem strictly from minimal, targeted redistributions of attention across critical diagnostic instances.
%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
A["Input WSI Bag<br/>N Image Patches"] --> B["Feature Extractor E<br/>Extract Instance Features zj"]
B --> C["Dual-Branch Evidence Decoupling Architecture<br/>Shared Encoder and Classifier"]
C --> D["Counterfactual Evidence Differential Constraint<br/>Ldiff: Force Ground-Truth Logit Divergence"]
C --> E["Attention Logit Proximity Regularization<br/>Ldiv: L1/Cosine Minimal Perturbation"]
D --> F["Joint Optimization & Dual Output<br/>Prediction + Supporting/Refuting Heatmaps"]
E --> F
Key Designs¶
1. Dual-Branch Evidence Decoupling Architecture: lightweight symmetric counterfactual head sharing the backbone and classifier
To avoid the intractable computational overhead of post-hoc counterfactual search, CAR-MIL internalizes counterfactual reasoning into a parallel attention branch. For extracted instance features \(\{z_j\}_{j=1}^N\), the factual and counterfactual heads produce unnormalized attention logits: $\(u_j = \psi(z_j), \quad u_j^{\text{cf}} = \psi_{\text{cf}}(z_j)\)$ After Softmax normalization, instance features are aggregated into bag embeddings \(\hat{Z}\) and \(\hat{Z}^{\text{cf}}\), which are then passed through the identical, shared downstream classifier \(\varphi\) to yield factual logits \(F(u)\) and counterfactual logits \(F(u^{\text{cf}})\). Because \(E\) and \(\varphi\) are strictly shared, the counterfactual branch adds negligible parameter overhead (a lightweight single or two-layer MLP), and at test time it can either be queried for complementary counterfactual heatmaps or dropped entirely for standard forward inference. Crucially, all prediction discrepancies between the two branches are guaranteed to originate exclusively from variations in attention allocation.
2. Counterfactual Evidence Differential Constraint: margin-based loss driving counterfactual prediction degradation on ground-truth classes
Allowing the counterfactual branch unconstrained freedom risks learning arbitrary, unrelated representations. To enforce meaningful counterfactual divergence, CAR-MIL defines the evidence differential vector in logit space: $\(\Delta F(u, u^{\text{cf}}) = F(u) - F(u^{\text{cf}}) \in \mathbb{R}^K\)$ For ground-truth class \(y \in \{1, \dots, K\}\), \(\Delta F_y\) reflects the factual branch's relative evidential advantage over the counterfactual branch. The evidence differential loss \(\mathcal{L}_{\text{diff}}\) is defined as: $\(\mathcal{L}_{\text{diff}} = \log \left( 1 + \sum_{k \neq y} \exp \left( -(\Delta F_y - \Delta F_k) \right) \right)\)$ Analytically, whenever \(\mathcal{L}_{\text{diff}} \le \varepsilon\), the following margin bound holds for every non-ground-truth class \(k \neq y\): $\(\Delta F_y - \Delta F_k \ge -\log(e^\varepsilon - 1)\)$ Minimizing \(\mathcal{L}_{\text{diff}}\) explicitly forces the factual branch to provide significantly stronger evidence on the true class \(y\) compared to the counterfactual branch than on any competing class \(k\). Consequently, the counterfactual branch is steered to selectively suppress supporting positive evidence, shifting its prediction away from the target class.
3. Attention Logit Proximity Regularization: geometric distance penalties enforcing sparse and localized attention shifts
Without an explicit proximity constraint, the counterfactual head could easily diverge into an inverted or uniform distribution that bears no relation to the factual decision boundary. Guided by the principle of minimal perturbation, CAR-MIL imposes an attention proximity loss \(\mathcal{L}_{\text{div}} = D(u, u^{\text{cf}})\) directly in logit space.
Two geometric metrics are formulated: first, a normalized \(L_1\) distance: $\(D_{L_1}(u, u^{\text{cf}}) = \frac{1}{N} \sum_{j=1}^N \left| u_j - u_j^{\text{cf}} \right|\)$ Owing to the sparsity-inducing nature of \(L_1\) regularization, this constraint forces the counterfactual branch to alter the attention logits of only a small subset of critical instances while leaving the vast majority untouched. Second, a scale-invariant cosine dissimilarity metric: $\(D_{\cos}(u, u^{\text{cf}}) = 1 - \frac{\sum_{j=1}^N u_j u_j^{\text{cf}}}{\|u\|_2 \|u^{\text{cf}}\|_2}\)$ Cosine dissimilarity focuses on the directional reallocation of attention across instance sub-populations, proving especially effective in complex multi-class tasks where relative attention orientation matters more than absolute logit magnitudes. Under this proximity penalty, the factual branch is incentivized to concentrate its attention on the most indispensable diagnostic patterns, while the counterfactual branch naturally directs its attention toward refuting or non-malignant tissue regions.
Loss & Training¶
The complete CAR-MIL framework is optimized end-to-end via a composite objective: $\(\mathcal{L} = \mathcal{L}_{\text{cls}} + \alpha \mathcal{L}_{\text{diff}} + \lambda \mathcal{L}_{\text{div}}\)$ where \(\mathcal{L}_{\text{cls}} = \text{CE}(\hat{y}(u), y)\) is the standard cross-entropy classification loss on the factual branch, and \(\alpha, \lambda \ge 0\) control the contribution of counterfactual differential and proximity regularizers. Training uses the Adam optimizer with a learning rate of 0.0002 for TCGA and CAMELYON16 datasets (0.0001 for BRACS) and weight decay of 1e-5. Moderate regularization weights (\(\alpha, \lambda \in [0.2, 1.0]\)) yield robust improvements across tasks.
Key Experimental Results¶
Main Results¶
CAR-MIL was evaluated on five diverse digital pathology WSI datasets covering cancer subtyping, TP53 mutation prediction, multi-class tissue grading, and lymph node metastasis detection using UNI foundation model features:
| Dataset | Task Type | Metric | CAR-MIL (Best) | ABMIL Baseline | Prev. SOTA (Model) | Gain |
|---|---|---|---|---|---|---|
| TCGA-BRCA | Breast Cancer Subtyping | AUC (%) / F1 (%) | 95.5±2.2 / 89.6±1.4 | 95.3±1.7 / 88.9±1.3 | 95.4±1.5 / 88.0±1.1 (MaxMIL) | AUC +0.2% / F1 +0.7% |
| TCGA-NSCLC | Lung Cancer Subtyping | AUC (%) / F1 (%) | 98.0±0.5 / 95.0±1.8 | 97.6±1.0 / 94.2±1.2 | 97.9±0.8 / 94.1±1.7 (CLAM) | AUC +0.4% / F1 +0.8% |
| TCGA-LUAD | TP53 Mutation Prediction (Hard) | AUC (%) / F1 (%) | 78.1±4.1 / 72.6±4.5 | 74.0±4.3 / 72.1±5.4 | 77.5±2.5 / 71.1±5.6 (CIA-MIL) | AUC +4.1% / F1 +1.5% |
| BRACS | 7-class Tissue Subtyping (Hard) | BACC (%) / F1 (%) | 43.1±2.0 / 40.4±2.4 | 38.2±4.1 / 36.0±4.5 | 42.3±2.0 / 40.0±2.0 (DSMIL) | BACC +4.9% / F1 +4.4% |
| CAMELYON16 | Lymph Node Metastasis Detection | AUC (%) / AUPRC+ (%) | 99.9±0.1 / 94.8±1.0 | 98.7±0.3 / 93.2±1.2 | 99.7±0.3 / 94.4±0.3 (CLAM) | AUC +1.2% / AUPRC+ +1.6% |
Note: On BRCA, the Cosine variant achieved F1 89.6% while L1 achieved AUC 95.5%; on LUAD and BRACS, Cosine regularization demonstrated substantial superiority; on CAMELYON16, L1 reached 99.9% AUC and 94.8% patch-level AUPRC+.
Ablation Study¶
1. Backbone Generalization: Integrating CAR into Diverse MIL Aggregators
| Model | Variant | TCGA-BRCA AUC | TCGA-LUAD AUC | BRACS BACC | CAMELYON16 AUPRC+ |
|---|---|---|---|---|---|
| ABMIL | Original Baseline | 95.3 ± 1.7 | 74.0 ± 4.3 | 38.2 ± 4.1 | 93.2 ± 1.2 |
| ABMIL | + CAR (CAR-MIL) | 95.0 ± 1.7 | 78.1 ± 4.1 | 43.1 ± 2.0 | 94.8 ± 1.0 |
| DSMIL | Original Baseline | 94.1 ± 1.6 | 66.9 ± 4.7 | 42.3 ± 2.0 | 80.6 ± 11.2 |
| DSMIL | + CAR (CAR-DSMIL) | 94.5 ± 2.2 | 71.2 ± 4.1 | 42.5 ± 5.1 | 87.8 ± 2.9 |
| TransMIL | Original Baseline | 93.2 ± 2.7 | 71.7 ± 5.4 | 40.0 ± 3.6 | 28.2 ± 3.6 |
| TransMIL | + CAR (CAR-TransMIL) | 94.0 ± 1.3 | 73.6 ± 4.4 | 42.6 ± 6.1 | 32.6 ± 0.3 |
2. Synthetic MNIST-bags Benchmark: Instance Evidence Alignment
| Dataset Variant | Model | Evidence Proxy | Bag AUC (%) | Instance AUPRC+ (%) | Instance AUPRC± (%) |
|---|---|---|---|---|---|
| Adjacent Pairs (Contextual) | ABMIL | Attention Weights | 85.7 ± 5.9 | 76.9 ± 6.1 | 61.5 ± 0.9 |
| Adjacent Pairs (Contextual) | TransMIL | Rollout Attention | 94.6 ± 4.4 | 81.5 ± 6.9 | 61.7 ± 2.1 |
| Adjacent Pairs (Contextual) | CAR-MIL | Attention Weights | 94.1 ± 5.9 | 85.7 ± 6.5 | 84.6 ± 5.2 |
| Four Bags (Interaction) | ABMIL | Attention Weights | 99.1 ± 0.1 | 86.8 ± 0.2 | 53.1 ± 0.1 |
| Four Bags (Interaction) | TransMIL | Rollout Attention | 99.5 ± 0.1 | 81.6 ± 6.3 | 51.9 ± 0.8 |
| Four Bags (Interaction) | CAR-MIL | Attention Weights | 99.6 ± 0.1 | 86.8 ± 0.6 | 89.7 ± 0.9 |
Key Findings¶
- Gains scale with task difficulty: On well-saturated binary tasks (BRCA, NSCLC), CAR-MIL matches or marginally outperforms strong baselines. However, on difficult clinical tasks such as TP53 mutation prediction (LUAD) and 7-class subtyping (BRACS), CAR-MIL delivers massive gains (+4.1% AUC and +4.9% BACC), proving that regularizing attention learning dynamics is especially vital under noisy, ambiguous label regimes.
- Superior attention-evidence fidelity: On the synthetic MNIST-bags benchmark with ground-truth instance roles, CAR-MIL boosts two-class positive/negative evidence AUPRC± from ~53-61% to 84.6% and 89.7%, confirming that the learned attention accurately captures contextual evidence.
- Strictly monotonic degradation in MORF perturbation analysis: Progressively removing top-attended instances leads to a smooth, monotonic drop in prediction confidence for CAR-MIL, whereas ABMIL exhibits erratic, non-monotonic curves, demonstrating that CAR-MIL attention rankings faithfully reflect true instance importance.
Highlights & Insights¶
- From Post-Hoc Diagnostic to In-Training Attention Regularizer: Instead of using counterfactuals merely as an expensive post-training explanation tool, CAR-MIL embeds counterfactual reasoning into training dynamics, rectifying attention misallocation at the root.
- Emergence of Complementary Positive and Negative Evidence Maps: Qualitative inspection on BRACS WSIs reveals that the factual branch accurately localizes ductal carcinoma in situ (DCIS) ducts, while the counterfactual branch highlights surrounding benign tissue with an anti-correlated attention pattern (\(r = -0.67\)), providing clinicians with intuitive contrastive evidence.
- Plug-and-Play Architectural Versatility: The CAR mechanism seamlessly plugs into DSMIL and TransMIL without altering their inference pipelines, and produces even greater relative gains when paired with general convolutional encoders (ResNet50) compared to domain-specific foundation models (UNI).
Limitations & Future Work¶
- Author-admitted limitation: In multi-class scenarios, the counterfactual branch shifts to alternative classes determined automatically by the evidence differential loss (e.g., class 1 naturally shifts to class 3 in Four Bags); explicit user steering toward a predetermined counterfactual target class is not yet supported.
- Identified limitation: Hyperparameters \(\alpha\) and \(\lambda\) require balanced tuning, and the geometric choice between L1 and Cosine distances depends on class complexity (L1 favors sparse localized binary detection, Cosine favors diffuse multi-class reallocation).
- Future directions: The learned counterfactual maps could serve as auxiliary pseudo-labels for active learning, weakly supervised lesion segmentation, or interactive clinical debugging systems.
Related Work & Insights¶
- vs ABMIL [Ilse et al., ICML 2018]: ABMIL learns attention implicitly via cross-entropy loss, making it prone to spurious correlations; CAR-MIL enforces explicit minimal-perturbation counterfactual regularization, yielding robust attention that faithfully reflects diagnostic evidence.
- vs CIA-MIL [Chraki et al., 2024]: CIA-MIL uses random attention interventions to diagnose causal impact, which often trades off classification performance; CAR-MIL employs an adversarial, learnable counterfactual branch that preserves and enhances classification accuracy while improving explainability.
- vs CLAM [Lu et al., Nat. Biomed. Eng. 2021]: CLAM relies on instance-level pseudo-clustering with SVM loss requiring complex hyperparameter tuning; CAR-MIL achieves sharper lesion localization through an elegant geometric logit distance loss on a shared backbone.
Rating¶
- Novelty: ⭐⭐⭐⭐⭐ [Pioneering translation of counterfactual explanation principles into an in-training MIL attention regularizer]
- Experimental Thoroughness: ⭐⭐⭐⭐⭐ [Extensive validation across synthetic interaction benchmarks and five public pathology WSI datasets with MORF and heatmap analyses]
- Writing Quality: ⭐⭐⭐⭐⭐ [Mathematically clear formulation, well-structured arguments, and self-contained experimental reporting]
- Value: ⭐⭐⭐⭐⭐ [Delivers both state-of-the-art predictive performance and trustworthy, dual-perspective interpretability for clinical computational pathology]