AGG:用于扩散模型高效GRPO训练的雅可比聚合组梯度
JAGG: Jacobian-Aggregated Group Gradient for Efficient GRPO Training of Diffusion Models
浏览论文内容
中文总结 AI 辅助
研究GRPO扩展到扩散模型的计算瓶颈问题,提出JAGG方法,通过特定插值和聚合梯度,减少反向传播次数,实验表明该方法能在质量损失小的情况下显著加速T2I训练中的DiT RL训练。
中文摘要 AI 辅助
组相对策略优化(GRPO)是一种强大的强化学习算法,可使生成模型与人类偏好保持一致。虽然在大语言模型中取得成功,但其扩展到扩散和流匹配模型时引入了严重的计算瓶颈:在采样轨迹的每个时间步都必须通过高容量的DiT主干进行梯度反向传播,使得高分辨率文本到图像(T2I)训练成本过高。无训练的DiT推理加速方法利用DiT隐藏状态和速度预测沿轨迹平滑且近似线性变化的事实。本文提出JAGG方法,通过端点雅可比的t加权插值近似中间步雅可比,将每步上游信号聚合为两个复合梯度,通过单次联合反向传播应用,证明了在速度为线性时插值是精确的,并通过实验验证该方法可在质量下降可忽略的情况下实现约2倍的反向加速。
英文摘要
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. The code for this work can be accessed through https://github.com/SchumiDing/JAGG.