执行摘要
移除 set_materialize_grads 并修复梯度传播
根据 commit message "which may cause issues" 和 backward 中新增的 RuntimeError 检查,说明原设置可能导致 grad_log_prob 为 None 未被正确处理,引发静默错误。
值得精读。变更小但涉及梯度流的核心路径,移除高风险配置,建议在训练实验中验证梯度正确性与收敛效果。
无 review 评论。
根据 commit message "which may cause issues" 和 backward 中新增的 RuntimeError 检查,说明原设置可能导致 grad_log_prob 为 None 未被正确处理,引发静默错误。
值得精读。变更小但涉及梯度流的核心路径,移除高风险配置,建议在训练实验中验证梯度正确性与收敛效果。
无 review 评论。
slime/utils/ppo_utils.py 的 _VocabParallelLogProbEntropy.forward 中,移除第198行的 ctx.set_materialize_grads(False) 调用;torch.Tensor | None 改为 torch.Tensor,并添加显式的 None 检查:若 grad_log_prob 为 None 则抛出 RuntimeError;ctx.cast_log_prob_grad_to_bfloat16 标志,简化梯度转换逻辑,将 log_prob 分支梯度直接计算而不做条件转换。| 文件 | 模块 | 状态 | 重要度 |
|---|---|---|---|
slime/utils/ppo_utils.py |
PPO 工具 | modified | 6.86 |
slime/utils/ppo_utils.py
core-logic
核心变更文件,移除 `set_materialize_grads(False)` 并调整 backward 梯度处理逻辑,直接影响 PPO 训练中的 logprob/entropy 梯度计算。
# slime/utils/ppo_utils.py
class _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 链接,后续同步到相关引用后会出现在这里。
参与讨论