๐ค AI Summary
This work addresses the high computational cost of aligning diffusion models with human preferences via GRPO, which requires backpropagating gradients through the expensive DiT backbone at every sampling stepโparticularly prohibitive for high-resolution text-to-image generation. To mitigate this, the authors propose JAGG, the first method to introduce a trajectory linearity assumption into reinforcement learning training for diffusion models. Leveraging the near-linear relationship between DiT hidden states and velocity predictions, JAGG approximates intermediate-step gradients via interpolation of endpoint Jacobians and aggregates multi-step upstream signals into just two joint backward passes. Combined with intra-batch gradient sharing, cosine-similarity-based routing, and an adaptive activation strategy (jagg_frac), JAGG achieves approximately 2ร acceleration in backpropagation while preserving image generation quality on standard text-to-image benchmarks.
๐ Abstract
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.