论文

LeanGRPO:消除扩散强化学习中的冗余重新计算

LeanGRPO: Eliminating Redundant Recomputation in Diffusion RL

AI 基础设施分布式训练基础设施

摘要

扩散强化学习 (RL) 最近在 后训练 图像和视频生成模型中取得了巨大成功。然而,大多数扩散强化学习方法,包括 DanceGRPO 和 FlowGRPO,都会在轨迹采样后通过梯度跟踪重新计算选定的时间步长。在使用相同后端进行采样与更新的 同策略 训练下,这种重新计算在数学上是多余的。直观上,轨迹采样和策略更新步骤可以重用相同的前馈主干以避免冗余计算,但这样做可能会在轨迹采样期间产生大量内存开销。为了解决这个问题,我们通过重构数据并行布局并引入两种无需重新计算的轨迹-logprob扩散RL训练计划来提出LeanGRPO:(1)LeanGRPO-Retain在轨迹采样期间启用梯度跟踪,并在更新期间直接重用生成的计算图和保存的激活用于向后,不需要重新计算; (2)LeanGRPO-Reweight也在rollout期间启用梯度,但立即使用临时优势反向传播每个选定的步骤并延迟梯度同步,然后在轨迹完成后使用真实优势校正临时梯度。这些计划针对不同的模型规模和输入大小。在使用 FLUX.1-dev 和 Wan 的 FlowGRPO/DanceGRPO 中,LeanGRPO 实现了高达 1.83 倍的端到端加速,同时保留了原始优化目标。