跳转至

Rank-Constrained Adaptation for Reliable Real-World Performance

会议: NeurIPS2026
arXiv: 2602.06924
领域: AI 安全
关键词: 群体鲁棒性、未知子群、误分类感知、低秩适配、加权协方差

一句话总结

MARLA 用冻结 ERM 模型在有任务标签的留出适配集上的错误概率构建加权特征子空间,仅在该子空间内学习低秩 logit 修正,在不使用子群标签的条件下改善最差群体准确率,同时明确区分无、部分和完整子群知识的模型选择条件。

研究背景与动机

平均准确率并不能保证每个子群都得到可靠预测:训练分布中的捷径、类别不平衡和属性不平衡可能使模型主要服务于占比高的样本。Group DRO 通过已知群体的最差风险进行优化,但部署时需要关注的群体未必事先完整列出。JTT、AFR 等方法虽然可以不在训练损失中使用群体标签,常见评测仍用完整的验证群体信息选超参数和早停。因此,“训练不需要群体标签”与“从训练到模型选择都不需要群体标签”是两个不同承诺。

本文重点研究已出现在数据中、却未被识别或标注为相关群体的失败模式,而不是完全未出现的新域。第 3 节给出一个存在性结论:在所有群体均有正概率、优化问题取得最小值等条件下,存在一种遗漏群体的方式,使 ERM 仍是该不完整 Group DRO 目标的最优解,或某个最优解在遗漏群体上的风险高于 ERM。这不是“每次遗漏都会伤害每个未知群体”的定理,也不是 MARLA 的无条件性能保证;它说明验证分组本身就是需要审查的假设。

MARLA 的切入点是冻结表示空间中已有的错误结构。若基模型难以正确处理的样本在少数方向上具有共同变化,就不必先命名群体、再重训整个分类器。核心 idea:用真实标签概率把失败样本的几何结构凸显出来,再把分类器修正严格限制在这一结构对应的低维子空间中。

方法详解

整体框架

输入是一个完成经验风险最小化(ERM)训练的分类模型,以及与基模型训练数据分离的适配集;适配集必须包含输入和任务真实标签,但无需子群标签。MARLA 依次执行“错误概率加权”“加权子空间估计”“受限 logit 修正”:冻结编码器和原分类头,计算样本权重及加权协方差,固定其前若干特征向量,再只训练一个小矩阵。

推理时,新样本经过原编码器,其投影产生修正项,与原 logits 相加后分类。无需推理样本的真实标签、群体身份或动态重算协方差;任务标签只用于适配训练。下图虚线表示训练监督或学得参数的传递,实线保留推理数据流。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    H["留出适配集<br/>输入 + 任务标签"] -.-> W["错误概率加权"]
    W -.-> S["加权子空间估计"]
    S -.->|固定基底与适配监督| C["受限 logit 修正"]
    X["推理输入<br/>无需标签或群体身份"] --> E["冻结 ERM 模型<br/>特征与原 logits"]
    E --> C
    C --> O["修正 logits → 分类"]

关键设计

1. 错误概率加权:用任务监督寻找失败信号,而不是推断群体身份

对适配样本,先取冻结基模型分配给真实类别的 softmax 概率。误分类分数为该概率的补数,衡量模型把多少概率质量分给错误类别;它不是一次预测是否错误的硬指示量,也不是经校准的“该样本一定会被误分”的保证。这样,预测正确但把握不足的样本仍能进入适配关注范围。

\[ \mu_i=\frac{\exp\!\big(\gamma(1-p_{y_i})\big)}{\sum_j\exp\!\big(\gamma(1-p_{y_j})\big)}. \]

其中 \(\gamma>0\) 控制权重的集中程度,归一化在适配样本间进行。低真实标签概率获得较大权重,但算法始终没有为样本生成种族、性别或其他子群伪标签;它寻找的是与失败有关的信号,而非把人口属性当作预测目标。主实验也没有额外乘类别频率的重平衡因子,附录 A 明确对应 rebalance=False。

分数位于 \([0,1]\),所以固定 \(\gamma\) 时,两样本的未归一化权重比最多为 \(e^\gamma\)。这与让无界损失直接支配加权不同,但不意味着任意大 \(\gamma\) 都安全:文本实验使用较大的数值,仍可能让权重高度集中。附录实现先减去最大对数权重再取指数,缓解数值溢出,并把权重从梯度图中分离;适配期间不重新训练基模型来改变分数。

2. 加权子空间估计:寻找失败样本的共同变化,而非全体数据的最大方差

普通 PCA 的大方差方向可能主要反映多数样本及无关噪声,恰好掩盖需要修补的样本。MARLA 用上述权重求冻结嵌入的加权均值,再对中心化特征形成协方差;低真实标签概率样本对几何结构的贡献随之增大。取其最大特征值对应的前 \(k\) 个特征向量,组成固定正交基底 \(V_k\)。

\[ \bar z=\sum_i\mu_i z_i,\qquad \Sigma_{\mathrm w}=\sum_i\mu_i(z_i-\bar z)(z_i-\bar z)^\top. \]

这里的中心化只用于协方差估计,原文的 logit 修正使用未中心化的编码器特征,不能悄悄替换成减均值后的投影。\(V_k\) 不是可训练 LoRA 因子,也不是为每个群体单独建立的分支。它是在适配集上一次估计得到的统一方向集合,随后固定。

秩候选由加权谱引导:累计加权方差 CWV 定义为前 \(k\) 个特征值之和除以全部特征值之和,作者用达到 50%–90% CWV 对应的秩确定搜索区域,再在独立验证集上选择。谱只是缩小搜索范围,并不自动证明这些方向就是所有部署失败方向,也不替代模型选择;不同群体知识条件下,验证指标必须不同。

附录 F 将这一机制的有效条件讲得更清楚:失败成分需要平均获得更高权重,其低维协方差信号需要强于干扰和有限样本误差,且存在足够的特征值间隙。只有这些条件成立时,谱扰动分析才支持估计子空间接近失败相关子空间。因而“低置信度”“共同几何结构”和“需要改善的群体”之间的对应关系是经验与条件性理论支持的假设,不是算法直接观察到的事实。

3. 受限 logit 修正:保留原预测器,只学习选定方向上的增量

冻结原编码器和分类头后,唯一可训练参数是 \(A\in\mathbb R^{k\times C}\),其中 \(C\) 为类别数。先把特征投影到固定子空间,再由 \(A\) 转成各类别的 logit 增量,与原输出相加:

\[ f_{\mathrm{MARLA}}(x)=f(x)+A^\top V_k^\top e_{\hat\psi}(x). \]

\(A\) 从零初始化,所以适配开始时完整恢复 ERM 预测,不需要先抛弃原分类头。诱导的分类器更新为 \(A^\top V_k^\top\),秩至多为 \(k\),且在所选子空间的正交补上不产生更新。这是严格的代数约束,而不只是让更新“倾向于”低秩的软正则化;适配只优化 \(kC\) 个参数,不过固定基底仍需要存储或合并进分类头,不能把可训练参数量等同于全部存储成本。

这种限制针对的是不受约束的分类头重训可能沿无关方向移动的问题:当失败信号集中时,小幅定向修补比全空间调整更可控。相反,如果错误需要表示空间里不存在的特征,或相关方向分散且所选秩不足,这个头部增量无法凭空恢复信息。本文不是权重量化或模型剪枝方法,也没有声称改善编码器对其他任务的迁移表示。

损失函数 / 训练策略

适配阶段对同一留出集合最小化加权交叉熵,只更新 \(A\);权重、基底、原模型都固定。核心目标是:

\[ \min_A\sum_i\mu_i\,\ell\!\left(f(x_i)+A^\top V_k^\top e_{\hat\psi}(x_i),y_i\right). \]

默认把训练集按 80%/20% 分给 ERM 与 MARLA,而非从测试集取适配标签。图像使用 ImageNet 预训练 ResNet-50,文本使用 BERT,两个表格数据集使用残差 MLP。适配特征可缓存,优化因此不必反复执行编码器反向传播;附录 C.3 的主数据集配置将额外正则系数设为 0。附录 E 的合成机制实验另用轻度 \(\ell_2\) 正则,不能混为主实验设置。

无子群知识条件不使用子群标签进行 MARLA 训练或模型选择;附录 D.1 的匹配基底比较明确按验证集最差类别准确率选择超参数。部分知识条件按已知属性定义的群体 WGA 选择与早停,完整知识条件按全部评测相关群体的 WGA 选择。已知属性只限定可用于选择的分组,不意味着删除未知群体样本。测试时仍用完整评测分组计算 WGA,因此“无群体标签”并不表示评估者也不需要审计标签。

\(\gamma\) 的候选尺度取决于分数分布:图像中高错误分数尾部较分明,文本分数较压缩,因而不能跨数据集照搬同一数值。附录 D.2 的部分敏感性曲线使用验证 WGA 选择秩,这些曲线不能直接充当完全无群体标签调参的证据。数据引导候选搜索与群体监督选择应分开描述。

实验关键数据

主实验

WGA 是指定测试分组中最低的组内准确率;Avg Acc 是全部测试样本的总体准确率,而非各群体准确率的等权平均。具体分组依数据集而异,临床人口分组与标签条件分组不可混用。下表选取原文表 1 的无子群知识结果,单位为 %,数值是三次运行的均值 ± 标准差;提升是 WGA 均值相对 ERM 的百分点差。

数据集 ERM WGA MARLA WGA ERM Avg Acc MARLA Avg Acc WGA 提升
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

相对 ERM 改善不等于全面超过所有基线:Waterbirds 的 DPE WGA 为 91.0 ± 0.5,超过 MARLA 的 90.1 ± 0.1;CheXpert 的 DFR 为 75.8 ± 0.3。表 1 的 GSR 使用群体标注验证样本,作者明确注明监督条件不匹配,不能把它当成严格无群体监督对照。CheXpert 的 WGA 提升伴随总体准确率从 88.6 降至 79.6,也不能概括成所有指标同时改善。

部分知识实验的表 2 中,MARLA 在九种属性可用配置的六种达到最高 WGA。GRACE 只知道年龄或性别时分别为 92.6 ± 3.7、95.0 ± 0.4;MIMIC-IV 只知道种族或性别时为 78.5 ± 1.9、79.1 ± 0.6。测试审计仍覆盖完整交叉属性群体,不对未知属性作行为或能力推断。CheXpert 两种配置分别为 73.4 ± 1.9、73.3 ± 1.7,低于 EA 的 75.8 ± 0.9、75.6 ± 0.8,说明不完整群体知识并不必然使 MARLA 最优。

完整知识也不是 MARLA 必胜的场景:表 3 的 CelebA 上 MARLA WGA 为 85.0 ± 0.9,而 GIC 为 89.4 ± 0.2;MultiNLI 上 MARLA 为 69.6 ± 0.8,AFR 为 73.4 ± 0.6。这些结果使用完整群体信息选模型,不能与上表无信息结果当作同一监督条件比较。

消融实验

原文表 4 比较无子群知识下的低秩与全秩修正。下表保留两个指标,避免把更稳健误读为总体预测能力单调上升。

数据集 MARLA WGA 全秩 WGA MARLA Avg Acc 全秩 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

原文正文称秩约束在五个数据集都“有益”,但 MultiNLI 的 WGA 均值实际相同,只是标准差与总体准确率不同;这里不改写为五个数据集 WGA 都严格提高。

附录表 15 进一步固定表示、秩、优化器和训练预算,只替换基底;下表为 WGA 均值 ± 标准差,单位为 %,验证选择不使用子群标签。

数据集 错误加权基底 未加权 PCA 基底 随机正交基底
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

关键发现

  • 低秩本身不够:同样的秩预算下,错误加权基底在四个数据集均优于 PCA 和随机基底。Waterbirds 相对 PCA 的 WGA 差为 13.5 个百分点,CivilComments 为 17.4 个百分点。
  • 合成机制实验人为设置 1、4、8 个少数群体相关方向,WGA 在修正子空间维数足够时恢复。这里的 \(k\) 是可使用的特征方向数,不应等同于最终二分类权重矩阵的代数秩或普遍成立的群体计数估计器。
  • 表 12 的单 RTX8000 端到端时间包含基模型训练与特征缓存:Waterbirds 为 ERM 25.05 分钟、MARLA 26.53 分钟;CelebA 为 158.36、167.03 分钟。这不是只训练适配矩阵的计时,也不包含完整超参数搜索成本。
  • 临床指标存在尚未解释的原文数值差异:表 1 的 GRACE MARLA WGA 为 94.2 ± 1.6,而表 18 的年龄×性别 WGA 为 81.56 ± 2.19;MIMIC-IV 对应表 1 为 78.2 ± 2.0,表 19 为 63.22 ± 1.33。两组数值分别保留,不能假定是同一检查点或同一评测口径,也不能自行合并。

亮点与洞察

  • 把“哪些样本值得修”与“允许分类器往哪里移动”分别控制。权重凸显失败信号,几何约束限制更新自由度,匹配基底消融比单纯报告低秩参数量更能支撑机制。
  • 把模型选择中的信息泄漏变成正式实验变量。该评测思路可迁移到其他鲁棒学习方法:分别记录训练、早停、调参和测试审计可用的属性,避免仅凭训练损失宣布无需群体标签。
  • 零初始化增量保留原预测器作为起点。适用于已有较好表示、需要低成本修补决策边界的分类器,而不是替代重新学习缺失特征的方案。

局限与展望

  • 错误分数依赖任务标签与概率质量,可能被噪声标签、结构化失校准或数据污染误导。附录表 20 的温度扰动最大 WGA 变化为 3.7 个百分点,仅支持该有限扰动范围内的稳定性。
  • 适配集需要覆盖相关失败样本,且它们要在低维方向上形成可识别结构。未出现的新群体或缺失的表示特征不受该机制保证,可进一步研究覆盖诊断与何时拒绝执行头部修补。
  • 临床数据存在显著类别不平衡,高人口群体准确率不能代替阳性类别上的可靠性。应联合审查 AUROC、AUPRC、敏感度、平衡准确率和标签条件 WGA;当前主文与附录数值差异也需要作者澄清后再讨论实际部署可靠性。
  • 附录中的部分超参数配置与汇总表未完全统一,且若干基线取自原论文而非统一重跑。公开复现应报告每个信息条件下的准确配置、选择指标与搜索预算,而非只复用一份网格。

相关工作与启发

  • vs AFR:二者都用基模型的困难样本信号进行二阶段适配;AFR 可重训完整分类头,MARLA 固定数据估计的子空间,只优化其中的加性修正。
  • vs JTT:JTT 放大硬误分类样本并进行再次训练,MARLA 使用连续真实标签概率,同时冻结表示。训练不看群体标签与调参不看群体标签仍须分别检验。
  • vs DFR / Group DRO:这些方法在相应设置中可借助群体平衡重训或显式最差群体目标;MARLA 不需要群体标签构造更新,但完整信息条件下其他方法仍可能更好。
  • vs LoRA:都限制可训练参数量,但 MARLA 的一个方向因子来自冻结特征的错误加权谱,只有系数矩阵可训练,而且作用于分类 logits。不能把它描述成对整个骨干网络应用标准 LoRA。

评分

  • 新颖性: 4/5。错误加权谱与严格受限修正的组合明确,不完整群体知识评测也有价值。
  • 实验充分度: 4/5。七数据集、匹配基底与秩消融较全面,但临床数值差异和配置不一致限制复现判断。
  • 写作质量: 3/5。方法定义清晰且标注保证边界,部分正文概括与表格均值仍不完全一致。
  • 价值: 4/5。适合已有分类器的低成本鲁棒修补,但不能据此宣称临床部署安全性已获验证。