Prhub

#2123 Fix training stuck on all-gather cp

原始 PR 作者 zhuzilin 合并时间 2026-07-01 14:21 文件变更 1 提交数 1 评论 0 代码增减 +6 / -5

执行摘要

修复 all-gather CP 模式下训练卡住的问题

在 all-gather context parallel 训练场景下,模型会卡住停止训练。作者通过排查发现,_allgather_cp_redistribute_extract_per_sample 中创建零张量的方式存在梯度图不一致问题,以及 policy_loss_function 中 logprob 比较使用了错误的变量,这些都会导致训练卡死。

建议合并该 PR,它解决了严重的训练卡死问题。修改小而精确,值得精读以理解 all-gather CP 模式下梯度图构建的细节。特别是 requires_grad 继承和切片取代 torch.zeros 的做法,可作为保持计算图连续性的参考模式。

讨论亮点

该 PR 没有评论或 review 讨论。

实现拆解

  1. _allgather_cp_redistribute 中零张量梯度传递修复loss.py:202-205):将 torch.zeros(response_length, ... requires_grad=True) 改为 requires_grad=ref_value.requires_grad,确保零张量的 requires_grad 属性与原始参考值一致,避免在特定 CP 分区下梯度图断裂。

  2. _extract_per_sample 中空切片创建方式优化loss.py:444-449):将 torch.zeros((0,), ...) 替换为切片操作 log_prob_full[:0]entropy_full[:0],保持与已有计算图的一致性和连续性,消除不必要的 tensor 创建。

  3. policy_loss_function 中 logprob 对比变量修复loss.py:1073-1075):新增 use_rollout_logprobs 配置(args.use_rollout_logprobs),当为 True 时使用 old_log_probs,否则使用当前 log_probs,避免统一使用 old_log_probs 导致不匹配和梯度问题。

文件 模块 状态 重要度
slime/backends/megatron_utils/loss.py 损失计算 modified 6.23

关键符号

_allgather_cp_redistribute _extract_per_sample policy_loss_function

关键源码片段

slime/backends/megatron_utils/loss.py core-logic

核心修改文件,包含三处关键 bug 修复:梯度图连续性修复、空切片创建优化、logprob 对比逻辑修复。

# _allgather_cp_redistribute 中零张量 requires_grad 继承修复
if value is None or e <= s:
    full_resp = torch.zeros(
        response_length,
        dtype=ref_dtype,
        device=ref_device,
        requires_grad=ref_value.requires_grad, # 原为 True,改为继承 ref_value
    )# _extract_per_sample 中空切片方法优化
if e <= s:
    log_probs_list.append(log_prob_full[:0]) # 原为 torch.zeros((0,), ...)
    if entropy_full is not None:
        entropy_list.append(entropy_full[:0]) # 原为 torch.zeros((0,), ...)# policy_loss_function 中 logprob 对比变量的选择
log_probs_to_compare = log_probs if args.use_rollout_logprobs else old_log_probs
# 根据配置选择正确的 log_probs,避免错误变量导致梯度不一致
train_rollout_logprob_abs_diff = sum_of_sample_mean(
    (log_probs_to_compare - rollout_log_probs).abs()
)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险较低:三个修改都是针对特定分支场景的边界条件修复,且修改范围小(只有 6 行增加、5 行删除)。主要风险在于 use_rollout_logprobs 配置若未正确设置,可能影响 rollout logprob 对比的计算逻辑,但不会导致崩溃。需要确保该配置在实际训练脚本中正确传递。

直接影响所有使用 Megatron 后端且启用 all-gather context parallel 的训练任务,修复了训练卡死的问题,提升了训练稳定性。对非 all-gather CP 模式无影响。use_rollout_logprobs 配置的引入增加了灵活性,但需要用户明确设置。

核心路径变更 配置键新增

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论