arXiv:2506.07378cs.LGstat.ML2025-06被引 2

统一梯度与海森匹配,提升模型跨域泛化能力

Moment Alignment: Unifying Gradient and Hessian Matching for Domain Generalization

  • 通过理论分析统一三种跨域泛化方法
  • 新算法CMA避免反复反向传播,计算更高效
  • 适合追求高效可靠的跨域学习研究者

领域泛化(DG)旨在构建对未见目标域具有良好泛化性能的模型,以应对真实场景中普遍存在的分布偏移问题。现有方法尝试对齐不同领域的梯度与海森矩阵,但存在计算效率低、理论机制不清晰的问题。本文建立领域泛化的矩对齐理论,基于传递测度框架,扩展其定义至多源域情形,并给出目标误差上界。证明在特征提取器诱导出跨域不变最优预测器或不满足该条件时,对齐各领域导数均能提升传递测度。值得注意的是,矩对齐统一了不变风险最小化、梯度匹配与海森匹配三种此前分离的方法。进一步揭示特征矩与分类器头导数间的对偶关系,建立特征学习与分类器拟合的理论联系。基于此,提出闭式矩对齐(CMA)算法,以闭式方式对齐领域级梯度与海森矩阵,无需重复反向传播或采样估计海森矩阵,显著降低计算开销。在线性探测与全微调两个实验设置下,CMA均优于经验风险最小化及现有先进算法。

原文摘要 · Abstract (English)

Domain generalization (DG) seeks to develop models that generalize well to unseen target domains, addressing the prevalent issue of distribution shifts in real-world applications. One line of research in DG focuses on aligning domain-level gradients and Hessians to enhance generalization. However, existing methods are computationally inefficient and the underlying principles of these approaches are not well understood. In this paper, we develop the theory of moment alignment for DG. Grounded in \textit{transfer measure}, a principled framework for quantifying generalizability between two domains, we first extend the definition of transfer measure to domain generalization that includes multiple source domains and establish a target error bound. Then, we prove that aligning derivatives across domains improves transfer measure both when the feature extractor induces an invariant optimal predictor across domains and when it does not. Notably, moment alignment provides a unifying understanding of Invariant Risk Minimization, gradient matching, and Hessian matching, three previously disconnected approaches to DG. We further connect feature moments and derivatives of the classifier head, and establish the duality between feature learning and classifier fitting. Building upon our theory, we introduce \textbf{C}losed-Form \textbf{M}oment \textbf{A}lignment (CMA), a novel DG algorithm that aligns domain-level gradients and Hessians in closed-form. Our method overcomes the computational inefficiencies of existing gradient and Hessian-based techniques by eliminating the need for repeated backpropagation or sampling-based Hessian estimation. We validate the efficacy of our approach through two sets of experiments: linear probing and full fine-tuning. CMA demonstrates superior performance in both settings compared to Empirical Risk Minimization and state-of-the-art algorithms.

领域泛化梯度对齐海森矩阵优化理论

Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。