关于SFT泛化性的研究:一个带有奖励修正的强化学习视角

发表
Yongliang WuYongliang Wu 提交
作者: Yongliang Wu, Yizhou Zhou, Zhou Ziheng, Yingzhe Peng, Xinyu Ye, Xinting Hu, Wenbo Zhu, Lu Qi, Ming-Hsuan Yang, Xu Yang

摘要

AI 生成总结
动态微调(DFT)通过动态调整梯度,提高了大型语言模型(LLM)的泛化能力,其性能优于标准监督微调(SFT),并在离线强化学习中表现出竞争力。
我们提出了一种对大型语言模型(LLM)的监督微调(SFT)进行改进的简单但有理论依据的方法,以解决其与强化学习(RL)相比泛化能力有限的问题。通过数学分析,我们揭示了标准SFT梯度隐含地编码了一种有问题的奖励结构,这可能会严重限制模型的泛化能力。为了纠正这一点,我们提出了动态微调(DFT),通过动态地根据令牌的概率重新缩放目标函数来稳定每个令牌的梯度更新。值得注意的是,这一行代码的更改在多个具有挑战性的基准和基础模型上显著优于标准SFT,展示了极大的泛化能力提升。此外,我们的方法在离线RL设置中也显示出具有竞争力的结果,提供了一种有效而更简单的替代方案。这项工作将理论洞察与实践解决方案相结合,实质性地提升了SFT的性能。代码将发布在 https://github.com/yongliang-wu/DFT
查看 arXiv 页面查看 PDF

评论

Yongliang WuYongliang Wu
论文提交者

我们提出了一种简单但具有理论基础的改进方法,用于大型语言模型(LLM)的监督微调(SFT),以解决其与强化学习(RL)相比泛化能力有限的问题。通过数学分析,我们发现标准SFT梯度隐式编码了一种有问题的奖励结构,这可能会严重限制模型的泛化能力。为了纠正这一点,我们提出了动态微调(DFT),通过使用此令牌的概率动态重新缩放目标函数来稳定每个令牌的梯度更新。值得注意的是,这一行代码的更改在多个具有挑战性的基准测试和基础模型上显著优于标准SFT,表现出大大改进的泛化能力。此外,我们的方法在离线RL设置中也显示出有竞争力的结果,提供了一种有效但更简单的替代方案。这项工作将理论洞察与实际解决方案相结合,极大地提升了SFT的性能。代码将在 https://github.com/yongliang-wu/DFT 提供。

Chuan JianguoChuan Jianguo

你们在通用的LLM上测试过吗,比如通义千问2.5/3?增益仍然显著吗?

Yongliang WuYongliang Wu
论文提交者

我们尚未推出,但您可期待我们下一版本包含更多模型和任务。

Qing LiQing Li

干得好!

Yongliang WuYongliang Wu
论文提交者

谢谢~

future.lifuture.li

从现在起,全是DFT。

Li DongLi Dong

DFT 等同于最大化 p(y|x) 而非 log p(y|x) (SFT) 吗?

Yongliang WuYongliang Wu
论文提交者

不,我不这么认为。

Li DongLi Dong

由于 sg 阻断了梯度流,因此在微分过程中,sg(πθ(y⋆∣x)) 被视为常数。所以:
∇θLDFT=−(πθ(y⋆∣x)sg(πθ(y⋆∣x))​)∇θπθ(y⋆∣x)
但是 sg(πθ(y⋆∣x)) 在数值上近似于 πθ(y⋆∣x),所以:
πθ(y⋆∣x)sg(πθ(y⋆∣x))​≈1
因此:
∇θLDFT≈−∇θπθ(y⋆∣x)

Yongliang WuYongliang Wu
论文提交者

这看起来在数学上是等价的,但使用 logits 更具数值稳定性,直接应用 softmax 可能会导致溢出。我将尝试直接优化这个损失函数,看看效果如何。谢谢你的评论!

haoranhaoran

优化“p(yi | x, y{

haoranhaoran

“sg(p(y|x))**(n) p_{\theta}(y|x)”怎么样?或者在 p(y|x) 和 log p(y|x) 之间是否存在某种中间形式(或广义形式)?

Vadim KataevVadim Kataev

不错的方法!我更愿意把这个直接放到摘要中,以节省读者的时间:

DFT 是对标准 SFT 的一行代码更改:将每个 token 的损失与其预测概率(分离以避免梯度流动)进行缩放。
loss = loss * torch.softmax(shiftlogits, dim=-1).gather(1, shiftlabels.unsqueeze(-1)).squeeze(-1).detach()

Hermione GrangerHermione Granger

很棒的工作!SFT 到 RL 的连接令人惊叹(也很有直觉)。它表明 SFT 在隐式地使用一个不适定且极其稀疏的奖励函数来进行 RL 训练,而这项工作修复了它(尽管奖励仍然是“稀疏”的)。

顺便说一句,方程(8)似乎少了一个负号:)

RinRin

使用 DFT 损失在 ms-swift 上训练数百个 LLM 和 MLLM

DFT 与 SFT 的一些结果比较

训练曲线

(DFT:红色,SFT:绿色)
training-curves

MATH500 评估
脚本:https://github.com/yongliang-wu/DFT

模型(Qwen2.5-Math-1.5B) MATH500 (%) 原始 31.24 SFT 42.66 DFT 56.81

脚本可以在此处找到

beiqingbeiqing

7月8日,我独立得出了同样的结论,并推导出了与您完全相同的公式。当时,我立即在知乎上发布了这个想法。然而,由于我需要写一份专利申请,几天后我撤回了......当时,我证明了 SFT 和 RL 之间的等价性,以展示动态奖励的有效性。我通过证明动态策略熵可以集成离策略和在策略方法来实现这一点。

syzygysyzygy

有趣,

Kacper GłombiewskiKacper Głombiewski

很棒的工作,我可以在微调 gpt2/opt 时使用 DFT 而不是 SFT 吗?

Daniel van StrienDaniel van Strien

您现在可以在 SFT 的最新版本中使用 DFT:https://github.com/huggingface/trl/releases/tag/v0.23.0 (只需在 SFTConfig 中切换损失即可)