flow_grpo

2505.05470_Flow-GRPO

Flow-GRPO: Training Flow Matching Models via Online RL 论文分析

论文: Flow-GRPO: Training Flow Matching Models via Online RL
作者: Jie Liu, Gongye Liu, Jiajun Liang, Yangguang Li, Jiaheng Liu, Xintao Wang, Pengfei Wan, Di Zhang, Wanli Ouyang (CUHK MMLab / Kuaishou / Shanghai AI Lab / NJU)
发表: arXiv 2025
arXiv: 2505.05470


1. 问题与动机

核心问题

Flow Matching 模型(如 Stable Diffusion 3、FLUX)已成为图像生成的主流范式,但在组合生成(如精确控制物体数量、空间位置、属性绑定)和文字渲染等需要精确对齐文本指令的任务上表现不佳。在线强化学习 (Online RL) 已在大语言模型中证明了其强大的对齐能力(DeepSeek-R1, OpenAI-o1),但如何将在线 RL 应用到 Flow Matching 模型上是一个未解决的问题。

两个关键挑战

  1. 确定性与随机性的矛盾:Flow Matching 模型基于确定性 ODE 采样——给定相同的初始噪声和 prompt,输出完全相同。而 RL 的核心机制需要随机采样来探索环境并从奖励中学习。没有随机性,就没有探索;没有探索,RL 就无法工作。

  2. 采样效率:在线 RL 需要大量采样来估计策略梯度,但 Flow Matching 模型每次采样需要多步迭代去噪(SD3.5-M 默认 40 步),对大模型来说代价极高。

核心假设

  • Flow Matching 的确定性 ODE 可以被转化为等价的随机 SDE,且两者保持相同的边际分布——从而在保持模型行为不变的前提下引入 RL 所需的随机性
  • 在 RL 训练中,采样使用的去噪步数可以大幅减少而不影响最终性能——训练时"粗略"采样,推理时"精细"采样

2. 核心方法

2.1 Flow Matching 基础 (Rectified Flow)

数据样本 \(x_0 \sim X_0\),噪声样本 \(x_1 \sim X_1\)。Rectified Flow 定义噪声插值:

\[x_t = (1-t)x_0 + tx_1, \quad t \in [0,1] \]

模型训练回归速度场 \(v_\theta(x_t, t)\),目标速度 \(v = x_1 - x_0\):

\[\mathcal{L}(\theta) = \mathbb{E}_{t, x_0, x_1}[\|v - v_\theta(x_t, t)\|^2] \]

这是 Flow Matching 框架中 OT-CFM 的具体实现——条件路径是线性插值,向量场是常数方向。

2.2 去噪即 MDP

将迭代去噪过程建模为马尔可夫决策过程 \((S, A, \rho_0, P, R)\):

  • 状态 \(s_t = (c, t, x_t)\):条件文本、时间步、当前样本
  • 动作 \(a_t = x_{t-1}\):模型预测的去噪结果
  • 策略 \(\pi(a_t|s_t) = p_\theta(x_{t-1}|x_t, c)\)
  • 转移 \(P(s_{t+1}|s_t, a_t)\):确定性转移
  • 初始状态 \(\rho_0(s_0) = (p(c), \delta_T, \mathcal{N}(0, I))\)
  • 奖励 \(R(s_t, a_t) = r(x_0, c)\) 当 \(t=0\),否则为 \(0\)

设计要点:奖励只在最终步 (\(t=0\)) 给出——只有完整的生成结果才被评估。这是一个稀疏奖励设定,但去噪过程的每一步都通过策略梯度获得信号(通过重要性采样的方式将最终奖励分配到每一步)。

2.3 GRPO 在 Flow Matching 上的应用

RL 目标是最大化带 KL 正则化的期望累积奖励:

\[\max_\theta \mathbb{E}_{\pi_\theta}\left[\sum_{t=0}^T \left(R(s_t, a_t) - \beta \cdot D_{KL}(\pi_\theta(\cdot|s_t) \| \pi_{ref}(\cdot|s_t))\right)\right] \]

GRPO (Group Relative Policy Optimization) 相比 PPO 的优势是无需价值网络——通过组内相对比较估计优势函数,大幅减少内存开销。

给定 prompt \(c\),采样一组 \(G\) 张图像 \(\{x_0^i\}_{i=1}^G\) 及其轨迹 \(\{(x_T^i, x_{T-1}^i, \ldots, x_0^i)\}\),优势函数:

\[\hat{A}_t^i = \frac{R(x_0^i, c) - \text{mean}(\{R(x_0^i, c)\}_{i=1}^G)}{\text{std}(\{R(x_0^i, c)\}_{i=1}^G)} \]

Flow-GRPO 目标:

\[J_{Flow-GRPO}(\theta) = \mathbb{E}_{c, \{x^i\} \sim \pi_{\theta_{old}}}\left[\frac{1}{G}\sum_{i=1}^G \frac{1}{T}\sum_{t=0}^{T-1}\left[\min(r_t^i(\theta)\hat{A}_t^i, \text{clip}(r_t^i(\theta), 1-\varepsilon, 1+\varepsilon)\hat{A}_i^t) - \beta D_{KL}(\pi_\theta \| \pi_{ref})\right]\right] \]

其中重要性比率 \(r_t^i(\theta) = p_\theta(x_{t-1}^i|x_t^i, c) / p_{\theta_{old}}(x_{t-1}^i|x_t^i, c)\)。

关键设计:每一步的去噪决策都参与策略梯度更新,而不仅仅是最终输出。这使得 RL 信号可以传播到整个去噪过程。

_attachments/flow_grpo/file-20260705010110460.png

2.4 ODE-to-SDE 转换——核心创新

问题:确定性 ODE \(dx_t = v_t dt\) 无法支持 GRPO,原因有二:

  1. 重要性比率 \(r_t^i(\theta) = p_\theta(x_{t-1}|x_t, c)\) 在确定性动力学下难以计算(需要估计概率密度,涉及散度计算)
  2. 确定性采样缺乏探索性,严重降低 RL 训练效率

解决方案:将确定性 Flow-ODE 转换为等价的随机 SDE,保持所有时间步的边际概率密度不变。

推导过程

第一步:正向 SDE

考虑一般 SDE:\(dx_t = f_{SDE}(x_t, t)dt + \sigma_t dw\)

其边际密度满足 Fokker-Planck 方程:

\[\partial_t p_t(x) = -\nabla \cdot [f_{SDE} p_t(x)] + \frac{1}{2}\nabla^2[\sigma_t^2 p_t(x)] \]

ODE 的边际演化:

\[\partial_t p_t(x) = -\nabla \cdot [v_t(x_t, t) p_t(x)] \]

令两者相等(保持边际分布不变),利用恒等式 \(\nabla^2[\sigma_t^2 p_t] = \sigma_t^2 \nabla \cdot (p_t \nabla \log p_t)\),解得漂移项:

\[f_{SDE} = v_t(x_t, t) + \frac{\sigma_t^2}{2}\nabla \log p_t(x) \]

正向 SDE 为:

\[dx_t = \left(v_t(x_t) + \frac{\sigma_t^2}{2}\nabla \log p_t(x_t)\right)dt + \sigma_t dw \]

第二步:逆向 SDE

由 Anderson (1982) 的逆向时间 SDE 公式:若正向 SDE 为 \(dx_t = f(x_t, t)dt + g(t)dw\),则逆向 SDE 为:

\[dx_t = [f(x_t, t) - g^2(t)\nabla \log p_t(x_t)]dt + g(t)d\bar{w} \]

代入 \(g(t) = \sigma_t\),逆向 SDE 变为:

\[dx_t = \left(v_t(x_t) - \frac{\sigma_t^2}{2}\nabla \log p_t(x_t)\right)dt + \sigma_t dw \]

第三步:计算 Score

对于 Rectified Flow 的线性插值 \(x_t = (1-t)x_0 + tx_1\):
(注意,由于以上已经保证了SDE和ODE的边际概率密度相同,因此可以通过ODE推理分布,与SDE的分布完全等价)

条件分布:\(p_{t|0}(x_t|x_0) = \mathcal{N}(x_t | (1-t)x_0, t^2 I)\)

条件 score:\(\nabla \log p_{t|0}(x_t|x_0) = -\frac{x_1}{t}\)
(通过带入高斯分布的公式推导得到)

边际 score:\(\nabla \log p_t(x_t) = -\frac{1}{t}\mathbb{E}[x_1|x_t]\)

利用速度场与 score 的关系(由 Theorem 3 of Flow Matching 论文):

\[v_t(x) = -\frac{x}{1-t} - \frac{t}{1-t}\nabla \log p_t(x) \]

解出 score:

\[\nabla \log p_t(x) = -\frac{x}{t} - \frac{1-t}{t}v_t(x) \]

第四步:最终逆向 SDE

将 score 代入逆向 SDE,得到最终形式:

\[dx_t = \left[v_t(x_t) + \frac{\sigma_t^2}{2t}(x_t + (1-t)v_t(x_t))\right]dt + \sigma_t dw \]

Euler-Maruyama 离散化:

\[x_{t+\Delta t} = x_t + \left[v_\theta(x_t, t) + \frac{\sigma_t^2}{2t}(x_t + (1-t)v_\theta(x_t, t))\right]\Delta t + \sigma_t\sqrt{\Delta t} \cdot \varepsilon \]

其中 \(\varepsilon \sim \mathcal{N}(0, I)\) 注入随机性。

噪声调度:\(\sigma_t = a\sqrt{t/(1-t)}\),其中 \(a\) 是控制噪声水平的超参数。

直觉理解:ODE-to-SDE 转换的本质是在原始确定性速度场上添加两个修正项:(1) 一个依赖于当前样本位置的偏移项 \(\frac{\sigma_t^2}{2t}(x_t + (1-t)v_t)\),它确保添加噪声后边际分布不变;(2) 一个随机扰动项 \(\sigma_t \sqrt{\Delta t}\varepsilon\),它为 RL 提供探索能力。这两者的精确形式由 Fokker-Planck 方程保证——只要 drift 和 diffusion 满足特定关系,边际分布就不变。

策略分布与 KL 散度

在 SDE 采样下,策略 \(\pi_\theta(x_{t-1}|x_t, c)\) 是各向同性高斯分布。KL 散度有闭式解:

\[D_{KL}(\pi_\theta \| \pi_{ref}) = \frac{\|\bar{x}_{t+\Delta t, \theta} - \bar{x}_{t+\Delta t, ref}\|^2}{2\sigma_t^2 \Delta t} \]

这进一步简化为:

\[D_{KL}(\pi_\theta \| \pi_{ref}) = \frac{\Delta t}{2}\left(\frac{\sigma_t(1-t)}{2t} + \frac{1}{\sigma_t}\right)^2 \|v_\theta(x_t, t) - v_{ref}(x_t, t)\|^2 \]

关键洞察:KL 散度直接正比于两个速度场的差异 \(\|v_\theta - v_{ref}\|^2\)。这意味着 KL 正则化的本质是约束 RL 训练不要偏离预训练速度场太远——越大的速度场修改,惩罚越重。系数中的时间依赖性 \(\left(\frac{\sigma_t(1-t)}{2t} + \frac{1}{\sigma_t}\right)^2\) 意味着不同时间步对 KL 的敏感度不同。

2.5 Denoising Reduction——效率创新

观察:Flow Matching 的在线 RL 训练中,生成样本使用的去噪步数不必与推理时相同。

策略:

  • 训练时:\(T = 10\) 步(快速采样用于数据收集)
  • 推理时:\(T = 40\) 步(原始默认步数,保证质量)

为什么可行:RL 优化的是奖励信号的方向,而非像素级精确。训练时只需要"足够好"的样本来区分优劣、估计优势函数,不需要最高质量的样本。减少步数带来 4 倍加速,且最终推理质量不受影响。

消融验证:进一步减少到 \(T=5\) 步时训练不稳定——步数太少导致采样质量过低,无法提供有效的奖励信号。


3. 实验设计逻辑

3.1 三个任务验证三个能力维度

任务 奖励函数 验证能力
组合生成 (GenEval) 目标检测 + 属性匹配 精确遵循复杂文本指令
文字渲染 (OCR) 编辑前后文字的距离 细粒度文本-像素对齐
人类偏好 (PickScore) PickScore 模型 整体视觉质量和偏好

3.2 关键结果

组合生成 (GenEval):

模型 Overall Counting Position Attr. Binding
SD3.5-M 0.63 0.50 0.24 0.52
SD3.5-M + Flow-GRPO 0.95 0.95 0.99 0.86
GPT-4o 0.84 0.85 0.75 0.61

SD3.5-M 经 Flow-GRPO 后在所有维度大幅超越 GPT-4o,尤其是位置关系从 0.24 提升到 0.99。

文字渲染:准确率从 59% → 92%

人类偏好:PickScore 从 21.72 → 23.31(带 KL),同时图像质量指标不降反升

_attachments/flow_grpo/file-20260705010110459.png

3.3 Reward hacking问题的系统分析

四维质量监控:Aesthetic Score、DeQA(图像质量)、ImageReward、UnifiedReward

设置 GenEval Aesthetic DeQA ImageReward
SD3.5-M 0.63 5.39 4.07 0.87
无 KL 0.95 4.93 2.77 0.44
有 KL 0.95 5.25 4.01 1.03

关键发现:

  • 无 KL:任务奖励上升但图像质量显著下降(DeQA 从 4.07 降到 2.77),偏好奖励也下降
  • 有 KL:任务奖励同样达到 0.95,但图像质量几乎不降(DeQA 4.01 vs 4.07),偏好奖励反而上升
  • KL 正则化不等价于早停——适当的 KL 可以在长时间训练后达到与无 KL 相同的高奖励,同时保持质量

3.4 消融研究

噪声水平 \(a\):

  • \(a=0.1\):探索不足,学习慢
  • \(a=0.7\):最佳,充分探索
  • \(a=1.0\):噪声过大,图像质量差,奖励为零,训练失败

组大小 \(G\):

  • \(G=6\) 或 \(12\):训练不稳定,最终崩溃(优势估计方差过大)
  • \(G=24\):稳定训练

初始噪声:每条轨迹使用不同随机噪声比共享初始噪声获得更高奖励——增加探索多样性有效。

3.5 泛化能力

  • 未见类别:在 T2I-CompBench++ 上,空间关系从 0.29 提升到 0.54,数值能力从 0.59 到 0.68
  • 未见数量:训练时只用了 2-4 个物体的 prompt,但泛化到 5-6 个物体(0.48 vs 基线 0.13)甚至 12 个物体(0.12 vs 0.02)

4. 创新与局限性

核心创新

  1. ODE-to-SDE 转换:这是方法最核心的贡献。通过 Fokker-Planck 方程的约束,在确定性 ODE 上添加精确校准的随机项,使得边际分布完全不变但获得探索能力。这不是简单的"加噪声"——噪声的强度和漂移项的修正都是精确推导的,确保数学上的等价性。

  2. Denoising Reduction:一个简单但实用的发现——RL 训练不需要高精度采样。4 倍加速对工程实践的意义巨大。

  3. 首次将 GRPO 引入 Flow Matching:打通了在线 RL 与 Flow Matching 之间的桥梁,证明了 Flow Matching 模型同样可以从在线 RL 中获益。

  4. KL 正则化的深入分析:证明了 KL 不仅防止奖励黑客,而且不是早停的等价物——适当 KL 可以在高奖励和质量保持之间取得最优平衡。

局限性

  1. 条件 OT 的限制:ODE-to-SDE 转换中的 score 计算 \(\nabla \log p_t(x)\) 使用了条件 OT 的线性插值假设。对于非 Rectified Flow 的其他 Flow Matching 变体,推导需要修改。

  2. 噪声调度的启发式选择:\(\sigma_t = a\sqrt{t/(1-t)}\) 的选择有一定启发式成分。\(a\) 的最优值(0.7)需要调参,且可能依赖具体模型和任务。理论上最优的噪声调度未被讨论。

  3. 仅验证了 T2I 任务:方法在视频生成、3D 生成等其他 Flow Matching 应用场景的适用性未验证。论文自己也在 Limitations 中提到了视频生成的挑战。

  4. Score 近似:边际 score \(\nabla \log p_t(x) = -\frac{x}{t} - \frac{1-t}{t}v_t(x)\) 是精确的,但实际使用的是模型估计的 \(v_\theta\) 而非真实的 \(v_t\)。当 \(v_\theta\) 不完美时,SDE 的边际分布与原始 ODE 可能有偏差。论文未讨论这种近似误差的影响。

  5. 与并发工作的差异:论文 [56] 通过将速度预测重新参数化为高斯分布来实现随机性(需要重新训练),而 Flow-GRPO 通过 ODE-to-SDE 转换避免重新训练。两种方法的相对优劣缺乏直接比较。

  6. 奖励设计的局限性:当前使用的奖励(目标检测器、编辑距离、PickScore)相对简单。更复杂的审美或语义奖励可能导致不同的奖励黑客模式。

开放方向

  • 多奖励联合优化:同时优化组合性、文字渲染和人类偏好
  • 视频生成:需要考虑时序一致性,奖励设计更复杂
  • 更鲁棒的奖励黑客防护:KL 正则化有效但不完美,某些 prompt 仍有轻微黑客
  • 其他 RL 算法的适配:PPO、REINFORCE 等是否也能通过 ODE-to-SDE 框架引入?
  • 与其他对齐方法的融合:DPO + GRPO 的混合策略

5. 总结

Flow-GRPO 的核心贡献是解决了"确定性 Flow Matching 模型如何进行在线 RL"这一看似矛盾的问题。答案优雅而精确:通过 Fokker-Planck 方程的约束,将 ODE 转换为边际等价的 SDE,在不改变模型原有行为的前提下注入 RL 所需的随机性。Denoising Reduction 则进一步解决了采样效率问题。两者结合,使得 Flow Matching 模型首次能够从在线 RL 中获益——在组合生成上超越 GPT-4o、文字渲染准确率提升 33 个百分点、且几乎不发生奖励黑客。这篇工作也揭示了一个更广泛的可能:Flow Matching 的数学结构(明确的边际分布、闭式 score)使其比扩散模型更适合精确的 RL 适配。


附录:详细数学推导

A.1 条件 score 的推导

目标:给定 \(x_0\),计算 \(p_{t|0}(x_t|x_0)\) 的对数梯度。

第一步:线性插值的条件分布

Rectified Flow 的条件路径是线性插值:

\[x_t = (1-t)x_0 + tx_1 \]

给定 \(x_0\),\(x_1 \sim \mathcal{N}(0, I)\),因此 \(x_t\) 的条件分布为:

\[x_t | x_0 \sim \mathcal{N}\left((1-t)x_0, t^2 I\right) \]

其中:

  • 均值:\(\mu_t = (1-t)x_0\)
  • 方差:\(\Sigma_t = t^2 I\)

第二步:写出高斯对数概率密度

多维高斯分布的对数概率密度:

\[\log p(x) = -\frac{d}{2}\log(2\pi) - \frac{1}{2}\log|\Sigma| - \frac{1}{2}(x-\mu)^T \Sigma^{-1} (x-\mu) \]

代入条件分布:

\[\log p_{t|0}(x_t|x_0) = \text{常数} - \frac{1}{2}\left(x_t - (1-t)x_0\right)^T (t^2 I)^{-1} \left(x_t - (1-t)x_0\right) \]

化简二次项:

\[\log p_{t|0}(x_t|x_0) = \text{常数} - \frac{\|x_t - (1-t)x_0\|^2}{2t^2} \]

第三步:对 \(x_t\) 求梯度

利用向量微积分恒等式 \(\nabla_x \|x - a\|^2 = 2(x - a)\):

\[\begin{align*} \nabla_{x_t} \log p_{t|0}(x_t|x_0) &= -\frac{1}{2t^2} \cdot 2\left(x_t - (1-t)x_0\right) \\ &= -\frac{x_t - (1-t)x_0}{t^2} \end{align*} \]

第四步:用 \(x_1\) 表示

根据线性插值公式 \(x_t - (1-t)x_0 = tx_1\),代入得:

\[\nabla_{x_t} \log p_{t|0}(x_t|x_0) = -\frac{tx_1}{t^2} = -\frac{x_1}{t} \]

最终结果:

\[\boxed{\nabla \log p_{t|0}(x_t|x_0) = -\frac{x_1}{t}} \]


A.2 边际 score 的推导

目标:计算边际分布 \(p_t(x_t)\) 的对数梯度 \(\nabla \log p_t(x_t)\)。

第一步:关键恒等式

对于任意两个随机变量 \(a\) 和 \(b\),有对数导数恒等式:

\[\nabla \log p(a) = \mathbb{E}_{b|a}\left[\nabla \log p(a|b)\right] \]

证明:

\[\begin{align*} \nabla \log p(a) &= \frac{\nabla p(a)}{p(a)} \\ &= \frac{1}{p(a)} \nabla \int p(a|b) p(b) db \\ &= \frac{1}{p(a)} \int \nabla p(a|b) p(b) db \\ &= \frac{1}{p(a)} \int \frac{\nabla p(a|b)}{p(a|b)} p(a|b) p(b) db \\ &= \int \nabla \log p(a|b) \cdot \frac{p(a|b) p(b)}{p(a)} db \\ &= \mathbb{E}_{b|a}\left[\nabla \log p(a|b)\right] \end{align*} \]

第二步:应用到我们的问题

选择以 \(x_0\) 为条件(因为我们已经知道条件 score):

\[\nabla \log p_t(x_t) = \mathbb{E}_{x_0|x_t}\left[\nabla \log p_{t|0}(x_t|x_0)\right] \]

第三步:代入条件 score

将 A.1 中得到的条件 score \(\nabla \log p_{t|0}(x_t|x_0) = -\frac{x_1}{t}\) 代入:

\[\begin{align*} \nabla \log p_t(x_t) &= \mathbb{E}_{x_0|x_t}\left[-\frac{x_1}{t}\right] \\ &= -\frac{1}{t} \mathbb{E}_{x_0|x_t}\left[x_1\right] \end{align*} \]

第四步:转换为关于 \(x_1\) 的期望

注意到给定 \(x_t\) 时,\(x_1\) 是 \(x_0\) 的确定性函数:

\[x_1 = \frac{x_t - (1-t)x_0}{t} \]

因此我们可以将期望转换为关于 \(x_1\) 的后验分布:

\[\mathbb{E}_{x_0|x_t}\left[x_1(x_0)\right] = \mathbb{E}_{x_1|x_t}\left[x_1\right] \]

最终结果:

\[\boxed{\nabla \log p_t(x_t) = -\frac{1}{t}\mathbb{E}[x_1|x_t]} \]

另一种等价推导

也可以通过联合分布的梯度来推导:

\[\begin{align*} p_t(x_t) &= \int p_{t|0}(x_t|x_0) p_0(x_0) dx_0 \\ \nabla \log p_t(x_t) &= \frac{1}{p_t(x_t)} \int \nabla p_{t|0}(x_t|x_0) p_0(x_0) dx_0 \\ &= \int \nabla \log p_{t|0}(x_t|x_0) \cdot \frac{p_{t|0}(x_t|x_0) p_0(x_0)}{p_t(x_t)} dx_0 \\ &= \int \nabla \log p_{t|0}(x_t|x_0) \cdot p_{0|t}(x_0|x_t) dx_0 \\ &= \mathbb{E}_{x_0|x_t}\left[\nabla \log p_{t|0}(x_t|x_0)\right] \end{align*} \]

这与前面的结果一致。


A.3 Score 与速度场的关系推导

目标:推导 \(\nabla \log p_t(x)\) 与 \(v_t(x)\) 的数学关联。

  1. 基础定义

    • 线性插值条件分布:\(x_t=(1-t)x_0+tx_1\),其中 \(x_0\sim p_0, x_1\sim\mathcal{N}(0,I)\)
    • 条件得分:\(\nabla\log p_{t|0}(x_t|x_0)=-\frac{x_1}{t}\)
    • 速度场:\(v_t(x)=\mathbb{E}_{x_0|x_t}[x_1-x_0]\)
  2. 核心推导
    边际得分的期望形式:\(\nabla\log p_t(x)=\mathbb{E}_{x_0|x_t}\left[\nabla\log p_{t|0}(x_t|x_0)\right]\)
    结合线性插值关系联立求解,最终得到:

    \[v_t(x) = -\frac{x}{1-t} - \frac{t}{1-t}\nabla\log p_t(x) \]

    整理为得分形式:

    \[\nabla\log p_t(x) = -\frac{x}{t} - \frac{1-t}{t}v_t(x) \]

  3. 工程意义
    将难以计算的边际得分转换为模型可学习的速度场\(v_t(x)\)的函数,无需额外训练即可估计得分,是Flow-GRPO中ODE-to-SDE转换的核心数学基础。


A.4 flow-grpo代码解析

Flow GRPO 训练流程分析(grpo_flux.sh)

本节梳理 sh scripts/single_node/grpo_flux.sh 的完整调用链路,并深入剖析 GRPO/PPO 中 sample["log_probs"]节 的计算方式与重要性采样(importance sampling)的原理。

1. Shell 入口

scripts/single_node/grpo_flux.sh 只有 3 行核心内容:


accelerate launch \
--config_file scripts/accelerate_configs/deepspeed_zero2.yaml \
--num_processes=8 --main_process_port 29501 \
scripts/train_flux.py \
--config config/grpo.py:pickscore_flux_8gpu

  • 通过 accelerate launch 拉起 8 个进程(单机 8 GPU);

  • 分布式策略由 scripts/accelerate_configs/deepspeed_zero2.yaml 决定:distributed_type: DEEPSPEED、zero_stage: 2、downcast_bf16、单机 8 卡;

  • 入口脚本是 scripts/train_flux.py,配置参数 --config config/grpo.py:pickscore_flux_8gpu。

2. 配置加载

scripts/train_flux.py:42 用 ml_collections 注册 flag:


config_flags.DEFINE_config_file("config", "config/base.py", "Training configuration.")

config/grpo.py:pickscore_flux_8gpu 的冒号语法含义:加载 config/grpo.py 模块,并调用其 get_config("pickscore_flux_8gpu")(grpo.py:1018),即 globals()["pickscore_flux_8gpu"]()。

该函数继承自 compressibility()(调用 config/base.py 的 get_config()),关键覆盖项:

字段 值
pretrained.model black-forest-labs/FLUX.1-dev
sample.num_steps / eval_num_steps 6 / 28
sample.guidance_scale 3.5
sample.train_batch_size 3
sample.num_image_per_prompt 24
sample.num_batches_per_epoch int(48/(8*3/24)) = 48
train.gradient_accumulation_steps 48//2 = 24
train.timestep_fraction 0.99 -> num_train_timesteps = int(6*0.99) = 5
train.beta 0(不启用 KL loss)
sample.noise_level / global_std / ema 0.9 / True / True
reward_fn {"pickscore": 1.0}
prompt_fn general_ocr
3. main() 初始化阶段

scripts/train_flux.py:326 main() 启动,app.run(main) 在文件末尾触发。初始化顺序:

  1. Accelerator 构造(:345):DeepSpeed ZeRO-2、bf16,梯度累积被放大为 gradient_accumulation_steps(24) × num_train_timesteps(5) = 120,因为每个时间步都要累积梯度。

  2. 加载 FluxPipeline(:370):FluxPipeline.from_pretrained("black-forest-labs/FLUX.1-dev"),冻结 vae/text_encoder/text_encoder_2。

  3. LoRA(:409):config.use_lora=True -> get_peft_model 注入 LoRA(target_modules 含 attn/ff 等模块),仅 LoRA 参数 requires_grad。

  4. EMA + 优化器(:441, :461):EMAModuleWrapper(decay=0.9,每 8 步更新)+ AdamW。

  5. 数据集(:469):prompt_fn=="general_ocr" -> TextPromptDataset(读 dataset/pickscore/train.txt / test.txt)+ DistributedKRepeatSampler(每个 prompt 重复 k=24 次,保证各卡同步采样)。

  6. 奖励函数(:549):reward_fn = flow_grpo.rewards.multi_score(device, {"pickscore": 1.0}) -> 内部用 pickscore_score(PickScoreScorer 打分)。

  7. accelerator.prepare(:557)包裹 transformer/optimizer/dataloader;起 ThreadPoolExecutor(max_workers=8) 异步算奖励。

4. 训练主循环

scripts/train_flux.py:597 while True:,每个 epoch 分三个阶段:

① EVAL(每 eval_freq=30 epoch)

eval()(:219):EMA 权重拷贝过去 -> 对测试集跑 pipeline_with_logprob(28 步)-> 算奖励 -> wandb 记录图像与 eval_reward_* -> EMA 还原。

② SAVE(每 save_freq=30 epoch,主进程)

save_ckpt()(:315):把 EMA 权重临时拷过去,unwrap_model(...).save_pretrained() 存 LoRA 到 logs/pickscore/flux-group24-8gpu/checkpoints/checkpoint-<step>/lora。

③ SAMPLING(:605)

循环 num_batches_per_epoch=48 次:

  • DistributedKRepeatSampler 产出 prompts -> compute_text_embeddings 调 encode_prompt;

  • pipeline_with_logprob(flow_grpo/diffusers_patch/flux_pipeline_with_logprob.py:22):6 步去噪循环,每步

  • transformer 前向预测 noise_pred;

  • sde_step_with_logprob(flow_grpo/diffusers_patch/sd3_sde_with_logprob.py:11):把 flow-matching 的 ODE 步改造成 SDE 步,记录每步的 prev_sample / log_prob / prev_sample_mean / std_dev_t(noise_level=0.9);

  • 累积 all_latents(每步潜变量)、all_log_probs(每步对数概率);

  • VAE decode 出图 -> 异步提交 reward_fn(images, prompts, metadata, only_strict=True)。

采完后汇总所有 batch:torch.cat 拼成 (48*3, ...);reward["avg"] 沿时间维 repeat 成 (N, 5);accelerator.gather 跨卡收集奖励。

④ 优势计算(:754)

per_prompt_stat_tracking=True:

  • PerPromptStatTracker.update(flow_grpo/stat_tracking.py:11):global_std=True -> 同一 prompt 的组内均值 + 全局 std 归一化得 GRPO 优势;

  • calculate_zero_std_ratio 统计零方差比例;advantages 再按卡切回本进程。

⑤ TRAINING(:799)

num_inner_epochs=1,对 batch 打乱重排后,逐 sample 逐 timestep:

  • compute_log_prob(:186):用当前 transformer 重新前向 + sde_step_with_logprob(传 prev_sample=sample["next_latents"]),得到新的 log_prob、prev_sample_mean、std_dev_t;

  • PPO 核心(:847):


ratio = exp(log_prob - sample["log_probs"][:, j]) # 新/旧策略概率比
unclipped_loss = -advantages * ratio
clipped_loss = -advantages * clamp(ratio, 1±clip_range)
policy_loss = mean(max(unclipped, clipped))

beta=0 故无 KL loss;loss.backward() -> accelerator.clip_grad_norm_ -> optimizer.step();

  • 在 accelerator.sync_gradients(累积到第 120 步)时 reduce 指标、wandb 记录、global_step += 1、ema.step(每 8 步更新 EMA)。

每个 epoch 的 2 次梯度更新 = samples_per_epoch(1152) // total_train_batch_size(576),与脚本注释「每 epoch 更新两次」一致。


5. sample["log_probs"] 是如何计算的

sample["log_probs"] 是采样阶段由 pipeline_with_logprob 在去噪过程中逐步记录下来的「旧策略」对数概率,它表示:在当前潜变量 x_{σ_j} 下,转移到下一步 x_{σ_{j+1}} 的 SDE 转移概率的对数密度。

  • 5.1 在采样循环中收集

scripts/train_flux.py:640 调用 pipeline_with_logprob,返回的第 5 个值就是每步的 log_prob:


images, latents, image_ids, text_ids, log_probs = pipeline_with_logprob(pipeline, ...)

log_probs = torch.stack(log_probs, dim=1) # (batch_size, num_steps=6)

随后存进 sample 字典(:678):


"latents": latents[:, :-1], # 每步的输入 x_{σ_j}
"next_latents": latents[:, 1:], # 每步的输出 x_{σ_{j+1}}
"log_probs": log_probs, # 每步转移的 log p(x_{σ_{j+1}} | x_{σ_j})

注意 latents 由 flux_pipeline_with_logprob.py:139 的 all_latents 累积得到(初始噪声 + 每步输出,共 num_steps+1 个),切片后 latents[:, j] 与 next_latents[:, j] 配对,log_probs[:, j] 就是这一对的转移对数概率。

  • 5.2 每一步的去噪与 log_prob 计算

去噪循环在 flux_pipeline_with_logprob.py:144-173,每步做两件事:


# (a) transformer 预测 velocity
noise_pred = self.transformer(
hidden_states=latents, # x_{σ_j}
timestep=timestep / 1000,
guidance=guidance, pooled_projections=...,
encoder_hidden_states=prompt_embeds, ...
)[0]

# (b) 用 SDE 步推进,同时算 log_prob
latents, log_prob, prev_latents_mean, std_dev_t = sde_step_with_logprob(
self.scheduler, noise_pred.float(), t..., latents.float(),
noise_level=noise_level, # 0.9
)
all_latents.append(latents)
all_log_probs.append(log_prob)

关键点:这一步没有传 prev_sample,所以 sde_step_with_logprob 内部 prev_sample is None,会从高斯中采样出下一步 latent —— 即把 Flux 原本确定性的 flow-matching Euler 步改成了随机 SDE 步,这是能算出 log_prob 的前提。

5.3 sde_step_with_logprob 的数学

pickscore_flux_8gpu 没有设置 config.sample.sde_type,走默认的 'sde' 分支(sd3_sde_with_logprob.py:49-68)。

记 sample = x_{σ_j}、model_output = v_θ(预测的 velocity)、dt = σ_{j+1} - σ_j < 0(sigmas 递减,噪声->数据):

① SDE 噪声尺度:

std_dev_t = sqrt(σ / (1 - σ)) * noise_level

② 转移分布的均值(确定性部分):

prev_sample_mean = sample * (1 + std_dev_t**2/(2*σ)*dt)
+ model_output * (1 + std_dev_t**2*(1-σ)/(2*σ)*dt)

③ 采样下一步(因为 prev_sample is None):

prev_sample = prev_sample_mean + std_dev_t * sqrt(-dt) * ε # ε ~ N(0, I)

即下一步服从 N(prev_sample_mean, (std_dev_t·√(-dt))²)。

④ 该步的 log_prob —— 标准高斯对数密度:

log_prob = -((prev_sample.detach() - prev_sample_mean)**2) / (2 * (std_dev_t*sqrt(-dt))**2)
- log(std_dev_t * sqrt(-dt))
- log(sqrt(2π))

三行分别对应 -(x-μ)²/(2σ²)、-log σ、-log√(2π),即 log N(prev_sample; prev_sample_mean, (std_dev_t·√(-dt))²)。

⑤ 跨维度平均(:89):

log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))

把通道维、空间维全部平均,得到每个样本一个标量。所以 log_probs 最终形状是 (batch, num_steps)。

关键细节:prev_sample.detach() —— 在计算 log_prob 时把「实际采到的下一步」当作常数,梯度只通过 prev_sample_mean(即 model_output)流动。采样阶段在 torch.no_grad() 下,所以这里只是记录数值;真正需要梯度是在训练阶段。

  • 5.4 小结
去噪第 j 步:
x_{σ_j} ──transformer──> v_θ(x_{σ_j}, t_j)
		│
		└─ sde_step_with_logprob (sde 分支, noise_level=0.9)
			├─ μ_j = f(x_{σ_j}, v_θ, σ_j, σ_{j+1}) # 转移均值
			├─ x_{σ_{j+1}} ~ μ_j + std_dev_t·√(-dt)·ε # 采样(无 prev_sample)
			└─ log p = log N(x_{σ_{j+1}}; μ_j, (std_dev_t·√(-dt))²)
						↓ mean over C/H/W
			-> 标量, 存入 all_log_probs[j]
			
最终: sample["log_probs"] = stack(all_log_probs) # (B, 6)
训练只用到前 5 步(num_train_timesteps=int(6*0.99)=5)

它本质上是对 SDE 化的 flow-matching 采样轨迹逐步求出的高斯转移对数密度,作为策略梯度公式里的「旧策略 log π_old」。


6. 为什么固定 next_latents 得到的是「当前策略」log_prob

澄清一个常见误解:next_latents 不是「从多个样本中挑出的最优样本」。它只是采样阶段从旧策略的高斯转移分布里实际抽到的那一个样本(一条确定下来的轨迹)。

  • 6.1 next_latents 是什么

采样阶段 flux_pipeline_with_logprob.py:163:

prev_sample = prev_sample_mean + std_dev_t * sqrt(-dt) * ε # ε ~ N(0,I)

因为当时 prev_sample=None,函数从 N(prev_sample_mean, (std_dev_t·√(-dt))²) 里随机抽出一个 prev_sample,这就是 next_latents。它是一个已经发生、固定下来的数值(采样在 torch.no_grad() 下完成),里面没有任何「选最优」的成分 —— 抽到什么就是什么。

「多个样本」其实是另一个机制:num_image_per_prompt=24 让同一个 prompt 生成 24 条不同的轨迹(每条都有自己的 next_latents 序列),这 24 条组成一个 group,用于 GRPO 算优势(组内归一化)。这和单步 log_prob 的计算是两回事 —— 优势回答的是「这条轨迹相对组内其他轨迹好不好」,log_prob 回答的是「策略产生这一步转移的概率多大」。

  • 6.2 为什么传 prev_sample=next_latents 就得到「当前策略」log_prob
    看 log_prob 公式(sd3_sde_with_logprob.py:64):
log_prob = -((prev_sample - prev_sample_mean)**2) / (2*σ²) - log σ - log√(2π)

这是标准高斯对数密度 log N(prev_sample; prev_sample_mean, σ²)。这个密度由两部分决定:

  • prev_sample(代入点):你把 next_latents 代进去,意思是「在这个固定的转移结果上求密度」。next_latents.detach() 固定了它,梯度不从这里走。这一步保证了你求的是「产生同一条轨迹」的概率,而不是重新采样的概率。
  • prev_sample_mean(分布参数):它由 model_pred = transformer(x_{σ_j}, ...) 算出来(compute_log_prob:197)。这里的 transformer 用的是当前(可训练、带梯度)的参数,不是采样时冻结的旧参数。

所以代入后得到的是:

\[\log \pi_{\theta_{当前}}(x_{\sigma_{j+1}} \mid x_{\sigma_j}) \]

即「当前策略下,从 x_{σ_j} 转移到这个固定的 x_{σ_{j+1}} 的对数概率」。它是「当前策略的 log prob」是因为分布的均值/方差由当前模型决定,而不是因为 next_latents 是什么最优样本。

  • 6.3 这正是 PPO / 重要性采样的标准做法

把采样和训练对照:

采样阶段 训练阶段 compute_log_prob
模型参数 旧策略 θ_old(冻结) 当前策略 θ(可训练)
prev_sample None -> 重新采样得到 x_{σ_{j+1}} next_latents -> 固定不采样,代入求密度
得到 log π_{θ_old}(x_{σ_{j+1}}|x_{σ_j}) = sample["log_probs"] log π_{θ}(x_{σ_{j+1}}|x_{σ_j}) = 新的 log_prob
  • 采样时我们收集经验:用旧策略走了一条随机轨迹,记录每步的概率 log π_old;
  • 训练时我们重放这条轨迹:同一个 x_{σ_j} -> x_{σ_{j+1}},但用当前模型重新评估它发生的概率 log π_θ。
  • 一个非常重要的点:采样旧轨迹的时候,采样了n个sample,更新新模型的时候,会逐一按照这n个sample更新,对于第一个sample,旧策略和新策略相同,对于后续的sample,新策略已经进行了backward更新,因此不再是旧策略,但轨迹仍然是旧轨迹。

于是比值:

ratio = exp(log_prob - sample["log_probs"][:, j]) # = π_θ / π_old

就是重要性采样比,用来把「在旧策略采到的样本上」估计的梯度,修正成「对当前策略有效」的梯度。这是 PPO 的核心,也是 next_latents 必须固定的原因 —— 如果训练时重新采样,那两条轨迹就不是同一条了,比值就失去意义。

  • 6.4 总结
  • next_latents = 采样阶段抽到的那个转移结果,固定不变,代表「已发生的轨迹」,与「最优」无关;
  • 把它代入用当前模型算出的高斯密度,得到的自然就是「当前策略产生这条轨迹的 log prob」;
  • 与采样时记录的旧 log_prob 相减,得到 PPO 的 importance ratio;
  • 「选最优/组内比较」是 GRPO 的优势计算(PerPromptStatTracker),和 log_prob 是正交的两件事:优势决定梯度的方向和大小(这条轨迹多好),ratio 决定梯度的校正系数(新旧策略对该轨迹的相对概率)。

7. 调用链总览
grpo_flux.sh
└─ accelerate launch (8 GPU, DeepSpeed ZeRO-2, bf16)
	└─ train_flux.py --config config/grpo.py:pickscore_flux_8gpu
		└─ app.run(main)
			├─ Accelerator(混合精度, grad_accum=24×5=120)
			├─ FluxPipeline.from_pretrained("FLUX.1-dev") -> get_peft_model(LoRA)
			├─ TextPromptDataset + DistributedKRepeatSampler(k=24)
			├─ reward_fn = rewards.multi_score({"pickscore":1.0})
			└─ while True:
				├─ EVAL : eval() -> pipeline_with_logprob(28步) -> reward_fn
				├─ SAVE : save_ckpt() -> LoRA.save_pretrained
				├─ SAMPLE : pipeline_with_logprob(6步)
				│     └─ transformer -> sde_step_with_logprob (记 log_prob)
				│     └─ vae.decode -> reward_fn(异步)
				├─ ADV : PerPromptStatTracker.update (GRPO 优势, global_std)
				└─ TRAIN : compute_log_prob(重算 log_prob)
						└─ ratio = exp(new - old)
						└─ PPO clipped policy loss -> backward -> optimizer.step -> ema.step

核心思想:用 SDE 改造的 flow-matching 采样保留每步 log_prob,训练时重算当前策略 log_prob,二者之比构成 PPO 的 importance ratio,配合 group-relative advantage(同 prompt 内归一化)做策略梯度更新 —— 即 Flow GRPO。

posted @ 2026-07-16 11:02  kiyoxi  阅读(43)  评论(0)    收藏  举报