Search papers, labs, and topics across Lattice.
This paper introduces JAGG (Jacobian-Aggregated Group Gradient), a method that significantly reduces the computational burden of Group Relative Policy Optimization (GRPO) in training diffusion models by minimizing the number of backward passes required. By leveraging the smooth and nearly linear variation of hidden states and velocity predictions along the sampling trajectory, JAGG reduces the need for full transformer backward passes from W to just 2 per group of W consecutive steps. Experimental results demonstrate that JAGG achieves approximately 2x speedup in backward computation without compromising model quality, making high-resolution text-to-image training more feasible.
JAGG slashes the backward pass cost in GRPO training of diffusion models by nearly 50%, enabling efficient high-resolution text-to-image generation.
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.