Rank-Constrained Adaptation for Reliable Real-World Performance¶
Conference: NeurIPS2026
arXiv: 2602.06924
Area: AI Safety
Keywords: group robustness, unknown subgroups, misclassification awareness, low-rank adaptation, weighted covariance
TL;DR¶
MARLA uses error probabilities from a frozen ERM model on a held-out adaptation set with task labels to construct a weighted feature subspace, learns a low-rank logit correction only within that subspace, and improves worst-group accuracy without subgroup labels while separating model-selection conditions with no, partial, and complete subgroup knowledge.
Background & Motivation¶
Average accuracy does not ensure reliable predictions for every subgroup: shortcuts, class imbalance, and attribute imbalance can steer a model toward prevalent examples. Group DRO optimizes worst-case risk over known groups, but deployment-relevant groups may not be completely specified in advance. Although JTT and AFR can avoid group labels in their training losses, standard evaluations often use complete validation-group information for hyperparameter selection and early stopping. Training without group labels and avoiding group labels throughout training and model selection are different commitments.
The paper primarily studies failure modes already present in the data but not identified or annotated as relevant groups, rather than entirely unseen domains. Section 3 provides an existence result: assuming positive probability for every group and attained optimization minima, there is a way to omit a group such that ERM remains optimal for the incomplete Group DRO objective, or some optimizer has greater risk on the omitted group than ERM. This is neither a theorem that every omission harms every unknown group nor an unconditional guarantee for MARLA. It establishes that the validation grouping itself is an assumption requiring scrutiny.
MARLA starts from error structure already encoded in frozen representations. If examples handled poorly by the base model share variation along a few directions, it may be unnecessary to name the groups and retrain the whole classifier. Core Idea: emphasize the geometry of failure examples through true-label probabilities, then strictly restrict the classifier correction to the corresponding low-dimensional subspace.
Method¶
Overall Architecture¶
The inputs are a classification model trained by empirical risk minimization (ERM) and an adaptation set separate from base-model training data. The adaptation set must contain inputs and ground-truth task labels, but does not require subgroup labels. MARLA applies “Error-probability weighting,” “Weighted subspace estimation,” and “Restricted logit correction” in sequence: it freezes the encoder and original head, computes weights and weighted covariance, fixes the leading eigenvectors, and trains only a small coefficient matrix.
At inference time, a new example passes through the original encoder; its projection supplies a correction added to the original logits before classification. Neither the example’s true label nor group identity is needed, and covariance is not dynamically recomputed. Task labels are used only for adaptation training. Dashed edges below indicate training supervision or transfer of learned parameters; solid edges show inference data flow.
%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
H["Held-out adaptation set<br/>Inputs + task labels"] -.-> W["Error-probability weighting"]
W -.-> S["Weighted subspace estimation"]
S -.->|Fixed basis and adaptation supervision| C["Restricted logit correction"]
X["Inference input<br/>No label or group identity"] --> E["Frozen ERM model<br/>Features and original logits"]
E --> C
C --> O["Corrected logits → classification"]
Key Designs¶
1. Error-probability weighting: find failure signals through task supervision rather than group inference
For each adaptation example, MARLA obtains the frozen model’s softmax probability assigned to the true class. The misclassification score is its complement, measuring probability mass assigned to incorrect classes. It is neither a hard indicator of whether the prediction is wrong nor a calibrated guarantee that the example will be misclassified. Correct but low-confidence predictions can therefore remain relevant to adaptation.
Here, \(\gamma>0\) controls concentration, and normalization is over adaptation examples. Low true-label probability receives greater weight, but the algorithm never generates pseudo-labels for race, gender, or other subgroup identities. It seeks failure-associated signals rather than treating demographic attributes as prediction targets. Main experiments do not apply an additional class-frequency rebalancing multiplier; Appendix A explicitly specifies rebalance=False.
Because the score lies in \([0,1]\), for fixed \(\gamma\) the ratio between two unnormalized weights is at most \(e^\gamma\). This differs from weighting directly by an unbounded loss, but does not make arbitrarily large \(\gamma\) safe: the text experiments use large numerical values, which can still produce highly concentrated weights. The appendix implementation subtracts the maximum log weight before exponentiation to mitigate overflow and detaches the weights from the gradient graph. Adaptation does not retrain the base model to change these scores.
2. Weighted subspace estimation: capture shared failure variation rather than maximum overall variance
High-variance directions from ordinary PCA may primarily reflect majority examples and irrelevant noise, obscuring the examples that need correction. MARLA uses its weights to compute the mean of frozen embeddings and the covariance of centered features, increasing the geometric contribution of examples with low true-label probability. The eigenvectors associated with the largest \(k\) eigenvalues form a fixed orthonormal basis \(V_k\).
Centering is used only for covariance estimation: the paper’s logit correction uses uncentered encoder features and must not silently be replaced by a mean-subtracted projection. \(V_k\) is neither a trainable LoRA factor nor a separate branch for each group. It is one shared set of directions estimated on the adaptation data and subsequently frozen.
The weighted spectrum guides rank candidates. Cumulative weighted variance (CWV) is the sum of the first \(k\) eigenvalues divided by the sum of all eigenvalues. Ranks corresponding to 50%–90% CWV define a search region, followed by selection on a separate validation set. The spectrum narrows the search; it neither proves that these directions cover every deployment failure nor replaces model selection. Validation criteria must differ across subgroup-knowledge settings.
Appendix F clarifies the conditions for this mechanism: the failure component must receive greater weight on average, its low-dimensional covariance signal must dominate nuisance variation and finite-sample error, and a sufficient eigengap must exist. Only under these conditions does eigenspace perturbation analysis support proximity to a failure-associated subspace. The connection between low confidence, shared geometry, and groups requiring improvement is thus supported empirically and by conditional theory, not directly observed by the algorithm.
3. Restricted logit correction: preserve the original predictor and learn an increment along selected directions
With the encoder and original head frozen, the only trainable parameter is \(A\in\mathbb R^{k\times C}\), where \(C\) is the number of classes. Features are projected into the fixed subspace, mapped by \(A\) to class-specific logit increments, and added to the original output:
Zero initialization of \(A\) exactly recovers ERM predictions at the start of adaptation without discarding the original head. The induced classifier update is \(A^\top V_k^\top\), has rank at most \(k\), and makes no update along the orthogonal complement of the selected subspace. This is an exact algebraic restriction, not a soft regularizer encouraging low rank. Adaptation optimizes only \(kC\) parameters, although the fixed basis must still be stored or merged into the head; trainable parameter count is not total storage cost.
The restriction addresses the possibility that unconstrained head retraining moves along irrelevant directions. When failure signals are concentrated, a targeted correction can be more controlled than full-space adjustment. Conversely, this head-level increment cannot recover predictive information absent from the representation, or capture dispersed failure directions when the selected subspace is too small. The method is not weight quantization or pruning, and does not claim to improve encoder representations for transfer to other tasks.
Loss & Training¶
Adaptation minimizes weighted cross-entropy on the same held-out set, updating only \(A\) while keeping weights, basis, and base model fixed. The central objective is:
The default training split allocates 80%/20% to ERM and MARLA, rather than obtaining adaptation labels from the test set. Image experiments use ImageNet-pretrained ResNet-50, text experiments use BERT, and the two tabular datasets use a residual MLP. Adaptation features can be cached, eliminating repeated encoder back-propagation. Appendix C.3 sets the additional regularization coefficient to 0 for the main dataset configurations. The synthetic mechanism experiment in Appendix E separately uses mild \(\ell_2\) regularization and should not be conflated with the main setting.
The no-subgroup-knowledge setting uses no subgroup labels for MARLA training or model selection; the matched-basis comparison in Appendix D.1 explicitly selects hyperparameters using validation worst-class accuracy. Partial knowledge uses WGA over groups defined by the known attribute for selection and early stopping; complete knowledge uses WGA over all evaluation-relevant groups. The known attribute limits the grouping available for selection, rather than removing examples from unknown groups. Test WGA still uses the full evaluation grouping, so absence of group labels during adaptation does not remove the evaluator’s need for audit labels.
Candidate scales for \(\gamma\) depend on the score distribution: images exhibit a more distinct high-score tail, whereas text scores are compressed. One numerical setting cannot be transferred indiscriminately across datasets. Some sensitivity curves in Appendix D.2 select rank by validation WGA and therefore cannot directly establish fully group-label-free tuning. Data-guided candidate generation and group-supervised selection must be described separately.
Key Experimental Results¶
Main Results¶
WGA is the lowest within-group accuracy among the specified test groups. Avg Acc is accuracy over all test examples, not an equally weighted average of group accuracies. Group definitions vary by dataset; clinical demographic groups and label-conditioned groups must not be conflated. The table selects no-subgroup-knowledge results from the paper’s Table 1, in %, reported as mean ± standard deviation over three runs. Gains are percentage-point differences in mean WGA relative to ERM.
| Dataset | ERM WGA | MARLA WGA | ERM Avg Acc | MARLA Avg Acc | WGA gain |
|---|---|---|---|---|---|
| GRACE | 93.3 ± 0.7 | 94.2 ± 1.6 | 95.1 ± 0.4 | 95.8 ± 1.3 | +0.9 |
| MIMIC-IV | 72.1 ± 0.0 | 78.2 ± 2.0 | 80.0 ± 0.0 | 81.0 ± 0.2 | +6.1 |
| Waterbirds | 69.1 ± 4.7 | 90.1 ± 0.1 | 84.1 ± 1.7 | 93.7 ± 0.7 | +21.0 |
| CelebA | 57.6 ± 0.8 | 82.8 ± 0.5 | 95.0 ± 0.1 | 95.2 ± 0.1 | +25.2 |
| CivilComments | 63.2 ± 1.2 | 71.6 ± 0.7 | 85.4 ± 0.2 | 91.4 ± 0.4 | +8.4 |
| MultiNLI | 66.4 ± 2.3 | 69.7 ± 0.1 | 81.0 ± 0.3 | 81.2 ± 0.2 | +3.3 |
| CheXpert | 41.7 ± 3.4 | 75.3 ± 0.4 | 88.6 ± 0.7 | 79.6 ± 0.2 | +33.6 |
Improvement over ERM does not imply dominance over every baseline. Waterbirds DPE reaches 91.0 ± 0.5 WGA, exceeding MARLA’s 90.1 ± 0.1; CheXpert DFR reaches 75.8 ± 0.3. GSR in Table 1 uses group-labeled validation examples, and the authors explicitly flag the supervision mismatch. It is not a strictly group-supervision-free comparator. CheXpert’s WGA gain also accompanies a decline in overall accuracy from 88.6 to 79.6, precluding a claim of simultaneous improvement on every metric.
In the partial-knowledge experiments of Table 2, MARLA obtains the highest WGA in six of nine attribute-availability configurations. With only age or sex known, GRACE WGA is 92.6 ± 3.7 or 95.0 ± 0.4; with only race or gender known, MIMIC-IV WGA is 78.5 ± 1.9 or 79.1 ± 0.6. Test audits still cover complete intersectional groups and imply no behavioral or capability assumptions about unknown attributes. CheXpert yields 73.4 ± 1.9 and 73.3 ± 1.7, below EA’s 75.8 ± 0.9 and 75.6 ± 0.8. Incomplete group knowledge does not necessarily make MARLA optimal.
Complete knowledge is not a guaranteed win either: Table 3 reports CelebA WGA of 85.0 ± 0.9 for MARLA versus 89.4 ± 0.2 for GIC; MultiNLI WGA is 69.6 ± 0.8 for MARLA versus 73.4 ± 0.6 for AFR. These models are selected using complete group information and cannot be compared with the table above as if supervision conditions were identical.
Ablation Study¶
The paper’s Table 4 compares low-rank and full-rank correction without subgroup knowledge. Both metrics are retained below to avoid equating greater robustness with monotonic improvement in overall prediction quality.
| Dataset | MARLA WGA | Full-rank WGA | MARLA Avg Acc | Full-rank Avg Acc |
|---|---|---|---|---|
| Waterbirds | 90.1 ± 0.1 | 82.0 ± 4.5 | 93.7 ± 0.7 | 92.0 ± 2.0 |
| MultiNLI | 69.7 ± 0.1 | 69.7 ± 1.7 | 81.2 ± 0.2 | 81.0 ± 0.5 |
| CelebA | 82.8 ± 0.5 | 81.3 ± 0.4 | 95.2 ± 0.1 | 90.2 ± 0.1 |
| CivilComments | 71.6 ± 0.7 | 68.7 ± 1.3 | 91.4 ± 0.4 | 91.0 ± 0.6 |
| CheXpert | 75.3 ± 0.4 | 56.4 ± 9.1 | 79.6 ± 0.2 | 79.1 ± 0.8 |
The main text describes the rank constraint as beneficial on all five datasets, but mean WGA is identical on MultiNLI, with differences in standard deviation and overall accuracy. This is not restated as a strict WGA increase on all five datasets.
Appendix Table 15 further fixes representations, rank, optimizer, and training budget, changing only the basis. Values below are WGA mean ± standard deviation in %, with subgroup-label-free validation selection.
| Dataset | Error-weighted basis | Unweighted PCA basis | Random orthonormal basis |
|---|---|---|---|
| Waterbirds | 90.1 ± 0.1 | 76.6 ± 6.9 | 64.7 ± 8.7 |
| CivilComments | 71.6 ± 0.7 | 54.2 ± 9.5 | 60.1 ± 0.4 |
| GRACE | 94.2 ± 1.6 | 91.2 ± 0.8 | 91.1 ± 1.0 |
| MIMIC-IV | 78.2 ± 2.0 | 76.9 ± 0.5 | 77.0 ± 0.5 |
Key Findings¶
- Low rank alone is insufficient: at the same rank budget, the error-weighted basis outperforms PCA and random bases on all four datasets. The WGA margin over PCA is 13.5 percentage points on Waterbirds and 17.4 on CivilComments.
- The synthetic mechanism experiments construct 1, 4, or 8 minority-relevant directions and recover WGA when the correction subspace is sufficiently large. Here, \(k\) counts available feature directions; it should not be equated with the algebraic rank of the final binary-classification weight matrix or a universally valid group-count estimator.
- Single-RTX8000 end-to-end timings in Table 12 include base-model training and feature caching: ERM/MARLA take 25.05/26.53 minutes on Waterbirds and 158.36/167.03 minutes on CelebA. These are not isolated adaptation-matrix training times and do not include a complete hyperparameter search.
- Clinical results contain unexplained numerical differences: Table 1 gives GRACE MARLA WGA as 94.2 ± 1.6, whereas age×sex WGA in Table 18 is 81.56 ± 2.19. MIMIC-IV gives 78.2 ± 2.0 in Table 1 versus 63.22 ± 1.33 in Table 19. Both sets are retained separately; identical checkpoints or evaluation definitions cannot be assumed, and the values must not be silently reconciled.
Highlights & Insights¶
- The method separately controls which examples merit correction and where the classifier may move. Weighting emphasizes failure signals while geometry restricts update freedom; the matched-basis ablation supports this mechanism more directly than trainable parameter counts alone.
- Information available during model selection becomes an explicit experimental variable. Other robust-learning evaluations can likewise record attributes available for training, early stopping, tuning, and test auditing instead of declaring a method group-label-free based only on its loss.
- A zero-initialized increment preserves the original predictor as the starting point. This is useful for classifiers with adequate representations that need inexpensive boundary correction, rather than as a replacement for learning missing features.
Limitations & Future Work¶
- Error scores depend on task labels and probability quality, and can be distorted by noisy labels, structured miscalibration, or data contamination. Temperature perturbations in Appendix Table 20 produce a maximum WGA change of 3.7 percentage points, supporting stability only within that tested range.
- The adaptation set must cover relevant failure examples with identifiable low-dimensional structure. The mechanism does not guarantee recovery for unseen groups or missing representation features; coverage diagnostics and criteria for declining head-level correction are useful directions.
- Clinical datasets exhibit substantial class imbalance, so high demographic-group accuracy cannot establish reliability on positive cases. AUROC, AUPRC, sensitivity, balanced accuracy, and label-conditioned WGA require joint inspection, and the main-text/appendix numerical differences require clarification before deployment reliability can be assessed.
- Some appendix hyperparameter configurations and summary tables are not fully aligned, and several baseline results come from original papers rather than uniform reruns. Reproduction should specify configurations, selection criteria, and search budgets for each information regime rather than reuse a single grid.
Related Work & Insights¶
- vs AFR: Both use difficult-example signals from a base model for second-stage adaptation. AFR can retrain the full head; MARLA fixes a data-estimated subspace and optimizes only an additive correction within it.
- vs JTT: JTT upweights hard misclassifications and trains again, whereas MARLA uses continuous true-label probabilities and freezes representations. Absence of group labels in training and tuning must still be checked separately.
- vs DFR / Group DRO: These methods can use group-balanced retraining or an explicit worst-group objective in their respective settings. MARLA constructs its update without group labels, but other methods can remain preferable with complete information.
- vs LoRA: Both limit trainable parameters, but one MARLA factor comes from the error-weighted spectrum of frozen features, only the coefficient matrix is trained, and the update acts on classification logits. It is not standard LoRA applied throughout the backbone.
Rating¶
- Novelty: 4/5. The combination of error-weighted spectral estimation and strictly restricted correction is well defined, and incomplete-group-knowledge evaluation is valuable.
- Experimental Thoroughness: 4/5. Seven datasets and matched-basis/rank ablations provide broad evidence, but clinical numerical discrepancies and configuration inconsistencies constrain reproducibility assessment.
- Writing Quality: 3/5. Method definitions and guarantee boundaries are clear, although some summary claims do not fully match tabulated means.
- Value: 4/5. Useful for inexpensive robustness correction of existing classifiers, but not evidence that clinical deployment safety has been validated.