[论文解读] Representation via Representations: Domain Generalization via Adversarially Learned Invariant Representations
该论文提出了一种使用对抗性学习不变表示的域泛化方法,通过将域视为敏感属性以在多样化群体间强制实现不变性。它形式化了对抗性损失在域数量增加时的极限行为,证明了更多域能改善泛化性能,并首次为对抗性不变域泛化提供了非渐近界和最坏情况性能的充分条件理论保证。
We investigate the power of censoring techniques, first developed for learning {\em fair representations}, to address domain generalization. We examine {\em adversarial} censoring techniques for learning invariant representations from multiple "studies" (or domains), where each study is drawn according to a distribution on domains. The mapping is used at test time to classify instances from a new domain. In many contexts, such as medical forecasting, domain generalization from studies in populous areas (where data are plentiful), to geographically remote populations (for which no training data exist) provides fairness of a different flavor, not anticipated in previous work on algorithmic fairness. We study an adversarial loss function for $k$ domains and precisely characterize its limiting behavior as $k$ grows, formalizing and proving the intuition, backed by experiments, that observing data from a larger number of domains helps. The limiting results are accompanied by non-asymptotic learning-theoretic bounds. Furthermore, we obtain sufficient conditions for good worst-case prediction performance of our algorithm on previously unseen domains. Finally, we decompose our mappings into two components and provide a complete characterization of invariance in terms of this decomposition. To our knowledge, our results provide the first formal guarantees of these kinds for adversarial invariant domain generalization.
研究动机与目标
- 解决高维生物医学数据中的域泛化问题,其中训练数据来自多个多样化群体,但测试数据可能来自未见过的、地理位置遥远或非合作机构。
- 形式化对抗性不变表示学习在域数量增加时的理论行为,特别是当 k 趋于无穷大时的极限情况。
- 为所提出方法的泛化误差提供非渐近学习理论界和在未见域上实现鲁棒性能的充分条件。
- 通过该分解方法对表示映射进行分解,并完全表征不变性,从而更深入理解所学表示。
- 将公平性的概念从人口统计偏差扩展到地理和机构公平性,确保在多样化群体中实现医疗人工智能的公平模型性能。
提出的方法
- 该方法使用层次贝叶斯模型,其中域是从数据分布分布中独立同分布抽取的,训练数据从每个域中采样。
- 采用对抗性屏蔽学习框架:表示编码器 φ 将输入映射到潜在空间 Z,判别器 ψk 尝试预测编码表示的源域,分类器 f 预测结果。
- 最小化经验对抗性损失函数,该函数同时惩罚分类器 f 的误分类和判别器 ψk 对域的正确预测,从而促使表示 φ 对域特定特征保持不变。
- 该方法改编自公平表示学习,将域视为敏感属性以抑制虚假相关性。
- 理论分析刻画了对抗性损失在 k → ∞ 时的极限行为,并推导出泛化误差的非渐近界。
- 对表示映射进行分解,通过该分解完全表征不变性,从而实现对所学不变性的精确分析。
实验结果
研究问题
- RQ1随着训练域数量 k 增加,对抗性不变表示学习在域泛化中的性能如何提升?
- RQ2当 k 增大时,对抗性损失函数的极限行为是什么?它是否会收敛到一个有意义的值?
- RQ3可以为所提出方法的泛化误差建立哪些非渐近学习理论界?
- RQ4在什么条件下,未见域上的最坏情况预测性能可保证良好?
- RQ5如何通过表示映射的分解,完全表征所学表示中的不变性?
主要发现
- 在具有 4 个域的合成数据上,RVR(所提方法)达到 90.6% 的测试准确率,当域数增至 10 个时提升至 95.6%,优于逻辑回归和随机森林。
- 在具有 6 个域的彩色 MNIST 数据上,RVR 达到 97.7% 的测试准确率,显著优于 IRM(94.7%)和 CIDDG(96.9%),尤其在颜色-标签相关性不均等的设置下。
- 当仅使用 6 个域中的 3 个进行训练时,RVR 仍达到 86.1% 的准确率,使用全部 6 个域时提升至 97.7%,表明其具备强大的数据效率和可扩展性。
- 在 PACS 数据集上,RVR 在 Sketch 域上达到 80.8% 的准确率,优于 IRM(75.0%)和 CIDDG(73.8%),且在所有域上均保持强劲性能。
- 随着可见域数量的增加,测试准确率持续提升,验证了更多样化的域能带来更好不变表示的直觉。
- 理论分析首次为对抗性不变域泛化提供了正式保证,包括损失的极限行为和最坏情况性能的充分条件。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。