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 模型上是一个未解决的问题。
两个关键挑战
-
确定性与随机性的矛盾:Flow Matching 模型基于确定性 ODE 采样——给定相同的初始噪声和 prompt,输出完全相同。而 RL 的核心机制需要随机采样来探索环境并从奖励中学习。没有随机性,就没有探索;没有探索,RL 就无法工作。
-
采样效率:在线 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 定义噪声插值:
模型训练回归速度场 \(v_\theta(x_t, t)\),目标速度 \(v = x_1 - x_0\):
这是 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 正则化的期望累积奖励:
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)\}\),优势函数:
Flow-GRPO 目标:
其中重要性比率 \(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 信号可以传播到整个去噪过程。

2.4 ODE-to-SDE 转换——核心创新
问题:确定性 ODE \(dx_t = v_t dt\) 无法支持 GRPO,原因有二:
- 重要性比率 \(r_t^i(\theta) = p_\theta(x_{t-1}|x_t, c)\) 在确定性动力学下难以计算(需要估计概率密度,涉及散度计算)
- 确定性采样缺乏探索性,严重降低 RL 训练效率
解决方案:将确定性 Flow-ODE 转换为等价的随机 SDE,保持所有时间步的边际概率密度不变。
推导过程
第一步:正向 SDE
考虑一般 SDE:\(dx_t = f_{SDE}(x_t, t)dt + \sigma_t dw\)
其边际密度满足 Fokker-Planck 方程:
ODE 的边际演化:
令两者相等(保持边际分布不变),利用恒等式 \(\nabla^2[\sigma_t^2 p_t] = \sigma_t^2 \nabla \cdot (p_t \nabla \log p_t)\),解得漂移项:
正向 SDE 为:
第二步:逆向 SDE
由 Anderson (1982) 的逆向时间 SDE 公式:若正向 SDE 为 \(dx_t = f(x_t, t)dt + g(t)dw\),则逆向 SDE 为:
代入 \(g(t) = \sigma_t\),逆向 SDE 变为:
第三步:计算 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 论文):
解出 score:
第四步:最终逆向 SDE
将 score 代入逆向 SDE,得到最终形式:
Euler-Maruyama 离散化:
其中 \(\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 散度有闭式解:
这进一步简化为:
关键洞察: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),同时图像质量指标不降反升

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. 创新与局限性
核心创新
-
ODE-to-SDE 转换:这是方法最核心的贡献。通过 Fokker-Planck 方程的约束,在确定性 ODE 上添加精确校准的随机项,使得边际分布完全不变但获得探索能力。这不是简单的"加噪声"——噪声的强度和漂移项的修正都是精确推导的,确保数学上的等价性。
-
Denoising Reduction:一个简单但实用的发现——RL 训练不需要高精度采样。4 倍加速对工程实践的意义巨大。
-
首次将 GRPO 引入 Flow Matching:打通了在线 RL 与 Flow Matching 之间的桥梁,证明了 Flow Matching 模型同样可以从在线 RL 中获益。
-
KL 正则化的深入分析:证明了 KL 不仅防止奖励黑客,而且不是早停的等价物——适当 KL 可以在高奖励和质量保持之间取得最优平衡。
局限性
-
条件 OT 的限制:ODE-to-SDE 转换中的 score 计算 \(\nabla \log p_t(x)\) 使用了条件 OT 的线性插值假设。对于非 Rectified Flow 的其他 Flow Matching 变体,推导需要修改。
-
噪声调度的启发式选择:\(\sigma_t = a\sqrt{t/(1-t)}\) 的选择有一定启发式成分。\(a\) 的最优值(0.7)需要调参,且可能依赖具体模型和任务。理论上最优的噪声调度未被讨论。
-
仅验证了 T2I 任务:方法在视频生成、3D 生成等其他 Flow Matching 应用场景的适用性未验证。论文自己也在 Limitations 中提到了视频生成的挑战。
-
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 可能有偏差。论文未讨论这种近似误差的影响。
-
与并发工作的差异:论文 [56] 通过将速度预测重新参数化为高斯分布来实现随机性(需要重新训练),而 Flow-GRPO 通过 ODE-to-SDE 转换避免重新训练。两种方法的相对优劣缺乏直接比较。
-
奖励设计的局限性:当前使用的奖励(目标检测器、编辑距离、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_0\),\(x_1 \sim \mathcal{N}(0, I)\),因此 \(x_t\) 的条件分布为:
其中:
- 均值:\(\mu_t = (1-t)x_0\)
- 方差:\(\Sigma_t = t^2 I\)
第二步:写出高斯对数概率密度
多维高斯分布的对数概率密度:
代入条件分布:
化简二次项:
第三步:对 \(x_t\) 求梯度
利用向量微积分恒等式 \(\nabla_x \|x - a\|^2 = 2(x - a)\):
第四步:用 \(x_1\) 表示
根据线性插值公式 \(x_t - (1-t)x_0 = tx_1\),代入得:
最终结果:
A.2 边际 score 的推导
目标:计算边际分布 \(p_t(x_t)\) 的对数梯度 \(\nabla \log p_t(x_t)\)。
第一步:关键恒等式
对于任意两个随机变量 \(a\) 和 \(b\),有对数导数恒等式:
证明:
第二步:应用到我们的问题
选择以 \(x_0\) 为条件(因为我们已经知道条件 score):
第三步:代入条件 score
将 A.1 中得到的条件 score \(\nabla \log p_{t|0}(x_t|x_0) = -\frac{x_1}{t}\) 代入:
第四步:转换为关于 \(x_1\) 的期望
注意到给定 \(x_t\) 时,\(x_1\) 是 \(x_0\) 的确定性函数:
因此我们可以将期望转换为关于 \(x_1\) 的后验分布:
最终结果:
另一种等价推导
也可以通过联合分布的梯度来推导:
这与前面的结果一致。
A.3 Score 与速度场的关系推导
目标:推导 \(\nabla \log p_t(x)\) 与 \(v_t(x)\) 的数学关联。
-
基础定义
- 线性插值条件分布:\(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]\)
-
核心推导
边际得分的期望形式:\(\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) \] -
工程意义
将难以计算的边际得分转换为模型可学习的速度场\(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) 在文件末尾触发。初始化顺序:
-
Accelerator 构造(
:345):DeepSpeed ZeRO-2、bf16,梯度累积被放大为gradient_accumulation_steps(24) × num_train_timesteps(5) = 120,因为每个时间步都要累积梯度。 -
加载 FluxPipeline(
:370):FluxPipeline.from_pretrained("black-forest-labs/FLUX.1-dev"),冻结 vae/text_encoder/text_encoder_2。 -
LoRA(
:409):config.use_lora=True->get_peft_model注入 LoRA(target_modules 含 attn/ff 等模块),仅 LoRA 参数requires_grad。 -
EMA + 优化器(
:441,:461):EMAModuleWrapper(decay=0.9,每 8 步更新)+ AdamW。 -
数据集(
:469):prompt_fn=="general_ocr"->TextPromptDataset(读dataset/pickscore/train.txt/test.txt)+DistributedKRepeatSampler(每个 prompt 重复k=24次,保证各卡同步采样)。 -
奖励函数(
:549):reward_fn = flow_grpo.rewards.multi_score(device, {"pickscore": 1.0})-> 内部用pickscore_score(PickScoreScorer 打分)。 -
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 用的是当前(可训练、带梯度)的参数,不是采样时冻结的旧参数。
所以代入后得到的是:
即「当前策略下,从 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。

浙公网安备 33010602011771号