JAGG:雅可比聚合群梯度,用于扩散模型的高效 GRPO 训练
JAGG: Jacobian-Aggregated Group Gradient for Efficient GRPO Training of Diffusion Models
摘要
组相对策略优化 (GRPO) 是一种强大的强化学习算法,用于使生成模型与人类偏好保持一致。虽然在 大语言模型~\cite{shao2024deepseekmathpushinglimitsmathematical} 中取得了成功,但其对扩散和流匹配模型的扩展引入了严重的计算瓶颈:梯度必须在采样轨迹的每个时间步长处通过高容量 DiT 主干进行反向传播,使得高分辨率文本到图像 (T2I) 的训练成本过高。 无需训练 DiT 推理加速方法(例如 $Δ$-DiT、ScalingCache)利用了 DiT 隐藏状态和速度预测沿轨迹 \emph{平滑且接近线性}变化的事实。我们询问相同的线性度是否可以降低 DiT RL 训练的后向传递成本,并用 \textbf{JAGG} (\textbf{J}acobian-\textbf{A}ggreated \textbf{G}roup \textbf{G}radient)给出肯定的回答,这将每组 $W$ 连续步骤的完整 Transformer 后向传递从 $W$ 减少到 $2$。 JAGG 通过端点雅可比行列式的 $t$ 加权插值来近似中间步骤雅可比行列式,然后将每步上游信号聚合为通过单个联合向后传递应用的两个复合梯度。我们证明,当速度在 $(z,t)$ 中呈线性时,这种插值是 \emph{exact},并且余弦相似性路由规则 (\texttt{jagg\_frac}) 仅在假设成立的情况下部署 JAGG。 T2I 基准测试表明,JAGG 提供了 $\sim$2$\times$ 向后加速,而质量下降可以忽略不计。这项工作的代码可以通过 https://github.com/SchumiDing/JAGG 访问。