AGG: Jacobian-Aggregated Group Gradient for Efficient GRPO Training of Diffusion Models

📄 arXiv: 2607.17572v1 📥 PDF

作者: Ruiyi Ding, Jie Li, He Kang, Ziyan Liu, Chengru Song, Yuan chen

分类: cs.LG, cs.CV, eess.SY

发布日期: 2026-07-20

备注: 21 pages


💡 一句话要点

提出JAGG以解决扩散模型GRPO训练的计算瓶颈问题

🎯 匹配领域: 支柱二:RL算法与架构 (RL & Architecture) 支柱九:具身大模型 (Embodied Foundation Models)

关键词: 扩散模型 群体相对策略优化 雅可比聚合 反向传播 文本到图像生成 计算效率 生成模型

📋 核心要点

  1. 现有的GRPO算法在扩散模型训练中面临计算瓶颈,导致高分辨率文本到图像训练成本过高。
  2. 本文提出JAGG,通过对雅可比矩阵的加权插值和聚合,显著减少反向传播的计算量。
  3. 实验结果显示,JAGG在文本到图像生成任务中实现了约2倍的反向传播速度提升,且几乎没有质量损失。

📝 摘要(中文)

群体相对策略优化(GRPO)是一种强大的强化学习算法,用于将生成模型与人类偏好对齐。然而,在扩散和流匹配模型中扩展GRPO会引入严重的计算瓶颈:在采样轨迹的每个时间步都必须通过高容量的DiT骨干网络反向传播梯度。为了解决这一问题,本文提出了JAGG(雅可比聚合组梯度),通过对端点雅可比的加权插值来近似中间步骤的雅可比,从而将完整的变换器反向传播次数从W减少到每组W个连续步骤的2次。实验结果表明,JAGG在文本到图像基准测试中实现了约2倍的反向传播加速,且质量损失微乎其微。

🔬 方法详解

问题定义:本文旨在解决扩散模型训练中GRPO算法的计算瓶颈问题,现有方法在每个时间步都需进行高成本的反向传播,导致训练效率低下。

核心思路:JAGG通过对雅可比矩阵进行加权插值,减少了反向传播的次数,从而降低了计算成本。该方法利用了DiT隐藏状态和速度预测在轨迹上的平滑性和近线性变化。

技术框架:JAGG的整体架构包括两个主要模块:首先,通过端点雅可比的加权插值来近似中间步骤的雅可比;其次,将每步的上游信号聚合为两个复合梯度,通过一次联合反向传播应用。

关键创新:JAGG的主要创新在于将完整的反向传播次数从W减少到每组W个连续步骤的2次,这一设计显著提高了训练效率。与现有方法相比,JAGG在保持生成质量的同时,显著降低了计算成本。

关键设计:在JAGG中,采用了基于余弦相似度的路由规则(jagg_frac),仅在假设成立的地方部署JAGG,从而确保插值的准确性。

📊 实验亮点

实验结果表明,JAGG在文本到图像基准测试中实现了约2倍的反向传播速度提升,且生成质量几乎没有下降。这一结果显示了JAGG在实际应用中的巨大潜力。

🎯 应用场景

该研究的潜在应用领域包括高效的文本到图像生成、增强现实和虚拟现实等需要实时生成内容的场景。通过提高训练效率,JAGG能够推动生成模型在实际应用中的广泛部署,提升用户体验。

📄 摘要(原文)

Group Relative Policy Optimization (GRPO) is a powerful reinforcement learning algorithm for aligning generative models with human preferences. While successful in large language models~\cite{shao2024deepseekmathpushinglimitsmathematical}, its extension to diffusion and flow matching models introduces a severe computational bottleneck: gradients must be back-propagated through the high-capacity DiT backbone at \emph{every} timestep of the sampling trajectory, making high-resolution text-to-image (T2I) training prohibitively expensive. Training-free DiT inference acceleration methods (e.g., $Δ$-DiT, ScalingCache) exploit the fact that DiT hidden states and velocity predictions vary \emph{smoothly and nearly linearly} along the trajectory. We ask whether the same linearity can reduce the backward-pass cost of DiT RL training, and answer affirmatively with \textbf{JAGG} (\textbf{J}acobian-\textbf{A}ggregated \textbf{G}roup \textbf{G}radient), which reduces full transformer backward passes from $W$ to $2$ per group of $W$ consecutive steps. JAGG approximates intermediate-step Jacobians via $t$-weighted interpolation of the endpoint Jacobians, then aggregates per-step upstream signals into two composite gradients applied through a single joint backward pass. We prove this interpolation is \emph{exact} when the velocity is linear in $(z,t)$, and a cosine-similarity routing rule (\texttt{jagg_frac}) deploys JAGG only where the assumption holds. Experiments on T2I benchmarks show JAGG delivers $\sim$2$\times$ backward speedup with negligible quality degradation.