PaperReading: DiffusionNFT
Yu Jiang

Link: DiffusionNFT: Online Diffusion Reinforcement with Forward Process.

DiffusionNFT 提出了一种在 forward process 中用 RL 训练 Diffusion Model 的方法, 在除了数据收集阶段, 不需要任何 reverse process. 所以不依赖任何 Solver 求解 ODE. 由于本人对于 Diffusion Model 的了解还非常的浅层, 也希望能够在未来回过头来继续补充本篇.

Diffusion 的加噪过程可以写为 , 并且可以重参数化为 .

而 Flow Model 的训练目标是训练一个速度场 , 并最小化误差的期望

那么训练得到的最优速度场为后验期望

下面回到正题, 对于 Online RL, 我们从 prompt 集合中 采样 张清晰图像 , 并用给定的 reward model 给出分数, 表示它质量最优的概率 . 通过 可以把采样得到的图像分为两个虚构的子集, 每张图像 会有 的概率进入 并有 的概率进入 . 那么在这两个数据集上的数据分布为:

RL 的每一步要求 . 我们根据定义可以得到 .

这是因为 可以表示为 的凸组合:

我们只需要说明 .

我们记 , 那么:

继而 . 那么很显然, 一个简单可行的微调策略就是每次用 生成若干张图片, 根据 的打分情况构造 , 并引导模型走向这个分布. 例如 Aligning Text-to-Image Models using Human Feedback 中也明确提到, 用这种拒绝采样的方法效果提升很大, 具体实现方法是从生成的 张图片中选取 top- 进行训练. 这种方法也别称作 Rejection FineTuning (RFT).

然而作者 argue, 并希望能够充分利用 内的样本. 为此, 我们需要更进一步地了解 的性质.

我们已经利用了它们的凸组合性质, 下面要证明它们的后验分布也是凸组合.

这一步的目的是为了和最优速度场相联系, 从而得到训练目标.

替换式 (1), 得到

从而得到了

在式 (1) 的两边同时作用 , 我们得到中间状态的凸组合:

并结合 符号, 我们也就得到的后验分布的凸组合以及 的表达式.

有了这一步, 并结合 flow model 的最优速度场, 我们可以得到在 下训练的得到的速度场可以凸组合得到 .

从而, 为了实现 , 我们可以同时利用 :

且此目标函数训练下的最优速度场为 .

证明如下:

首先对 进行变形, 用 替代 , 这是为了和速度场联系起来.

对内部的期望进行改造:

同理 . 因此内部的期望可以重写为

同理, 另一部分可以写为 . 而利用 的定义, 可以继续写为

因此, 最后 可以写为:

因此, 式 (2) 提出了一种 off-policy 的强化学习方法, 并于监督学习相结合. 下面是一些实践中的细节.

首先是 reward 的计算, 借鉴了 GRPO 的思路, 对 reward 做了归一化处理:

其中 是一个和 有关的量, 有点类似于 temperature.

其二是关于 , 做了 soft EMA: . 并且明确指出了 的取值是训练稳定性和收敛速度的平衡.

另一个小改动是启发式地设定 time weight function . 基模使用 rectifed flow model, 所以可以直接用 得到 .

训练采用 的 LoRA 微调, 每个 epoch 采样 48 个 prompts, 每个 prompt 采样 24 张图片.

Powered by Hexo & Theme Keep
Total words 29.6k