train_grpo.py中torch+grpo路径下old_per_token_logps未detach导致GRPO优势项梯度抵消 #788
Closed
zeta172417-design
started this conversation in
General
Replies: 1 comment
|
代码早已改为“一律重新 forward + 显式 .detach()”,该 Bug 不复存在,你没有把fork或者本地的代码对齐版本 |
0 replies
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Uh oh!
There was an error while loading. Please reload this page.
问题描述
我在使用
train_grpo.py训练时,发现当使用如下配置:也就是
torch+dense+grpo路径时,old_per_token_logps和per_token_logps可能来自同一个未detach的张量,导致GRPO中的ratio虽然数值为1,但梯度被抵消,进而使reward/advantage项无法有效更新actor。相关代码
当前代码中大致是这样:
在
torch+dense+grpo路径下,会走else分支,于是:然后GRPO中计算:
等价于:
这时
ratio数值上等于1,但因为是同一个计算图张量相减,所以梯度也会被抵消:因此GRPO的reward/advantage项可能无法有效更新actor,训练主要只剩下KL项在起作用。
建议修改
可以将
old_per_token_logps显式detach:这样ratio变成:
数值上初始仍然是1,但梯度不会被抵消:
这样GRPO的advantage项才能正常通过当前策略的
per_token_logps更新actor。为什么这个问题主要出现在torch+dense+grpo路径
我理解是:
sglang路径下,old_per_token_logps来自外部推理服务,本身不连接本地PyTorch计算图;use_moe路径下,会重新forward计算当前per_token_logps,不再是同一个张量直接相减;torch+dense路径下,当前per_token_logps直接复用rollout_result.per_token_logps,如果old_per_token_logps不detach,就会出现exp(x - x)导致梯度抵消的问题。观察到的现象
修改前,
loss_type=grpo时模型似乎较难从reward/advantage中学习,kl_ref变化也较小。将
old_per_token_logps改为:之后,
kl_ref变化明显增大,说明actor相对ref model开始发生更明显的策略偏移,这也符合GRPO优势项梯度恢复后的预期。想确认的问题
请问这里
old_per_token_logps未detach是有意设计,还是一个可能的实现bug?从GRPO/PPO的理论上看,
old_per_token_logps应该表示采样时旧策略的固定logprob,因此似乎应当作为常数处理,不应该继续参与反向传播。All reactions