Prhub

#2158 Remove ctx.set_materialize_grads(False) which may cause issues

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

执行摘要

移除 set_materialize_grads 并修复梯度传播

根据 commit message "which may cause issues" 和 backward 中新增的 RuntimeError 检查,说明原设置可能导致 grad_log_prob 为 None 未被正确处理,引发静默错误。

值得精读。变更小但涉及梯度流的核心路径,移除高风险配置,建议在训练实验中验证梯度正确性与收敛效果。

讨论亮点

无 review 评论。

实现拆解

  1. slime/utils/ppo_utils.py_VocabParallelLogProbEntropy.forward 中,移除第198行的 ctx.set_materialize_grads(False) 调用;
  2. 调整 backward 返回类型注解,将 torch.Tensor | None 改为 torch.Tensor,并添加显式的 None 检查:若 grad_log_probNone 则抛出 RuntimeError
  3. 移除 ctx.cast_log_prob_grad_to_bfloat16 标志,简化梯度转换逻辑,将 log_prob 分支梯度直接计算而不做条件转换。
文件 模块 状态 重要度
slime/utils/ppo_utils.py PPO 工具 modified 6.86

关键符号

_VocabParallelLogProbEntropy.forward _VocabParallelLogProbEntropy.backward

关键源码片段

slime/utils/ppo_utils.py core-logic

核心变更文件,移除 `set_materialize_grads(False)` 并调整 backward 梯度处理逻辑,直接影响 PPO 训练中的 logprob/entropy 梯度计算。

# slime/utils/ppo_utils.pyclass _VocabParallelLogProbEntropy(torch.autograd.Function):
    @staticmethod
    def forward(
        ctx,
        vocab_parallel_logits: torch.Tensor,
        target: torch.Tensor,
        log_prob_keep_mask: torch.Tensor | None,
        process_group,
        with_entropy: bool,
        with_entropy_grad: bool,
    ) -> tuple[torch.Tensor, torch.Tensor]:
        # --- 关键修改:移除了 ctx.set_materialize_grads(False) ---
        # 原代码此处调用了 set_materialize_grads(False),
        # 这会导致 backward 中 grad_log_prob 为 None 时静默跳过梯度计算,
        # 可能引起训练中的梯度丢失。
        with_entropy_grad = with_entropy and with_entropy_grad
        # ... 后续 softmax 计算保持不变 ...
        ctx.with_entropy_grad = with_entropy_grad
        # --- 同时移除了 cast_log_prob_grad_to_bfloat16 标志 ---
        # 原代码:ctx.cast_log_prob_grad_to_bfloat16 = vocab_parallel_logits.is_cuda
        # 移除该标志以避免不必要的精度转换,简化梯度计算路径。
​
    @staticmethod
    def backward(ctx, grad_log_prob, grad_entropy):
        # --- 新增显式 None 检查 ---
        if grad_log_prob is None:
            raise RuntimeError(
                "_VocabParallelLogProbEntropy expected a materialized grad_log_prob. "
                "Do not call ctx.set_materialize_grads(False)."
            )
        # 后续计算:先处理 entropy 梯度(若 with_entropy_grad),
        # 再处理 log_prob 梯度(不再有条件 bfloat16 转换)。
        # 原代码的 cast_log_prob_grad_to_bfloat16 分支被移除,
        # 统一以 float32 计算梯度,与 Megatron 的 fused vocab-parallel CE backward 解耦。

评论区精华

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

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

风险与影响

移除 set_materialize_grads(False) 后,PyTorch 会为所有输入张量生成梯度,可能略微增加内存占用,但幅度很小;同时移除 cast_log_prob_grad_to_bfloat16 逻辑,可能改变梯度精度(bfloat16 -> float32),对 Megatron 训练收敛性的影响需验证。

直接影响所有使用 _VocabParallelLogProbEntropy 的 PPO 训练流程(特别是涉及 logprob 和 entropy 计算的路径),消除梯度不物化导致的 bug,但可能小幅增加显存开销并改变梯度数据类型。

核心路径变更 梯度精度变化

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论