执行摘要
- 一句话:修复 FSDP logits 温度缩放 in-place 操作崩溃
- 推荐动作:值得跟随阅读,了解 FSDP 引擎中温度缩放的实现细节以及 in-place 操作对视图张量的潜在问题。PR 提交者可以关注是否为类似场景添加测试覆盖。
功能与动机
PR body 指出,output.logits.squeeze(0) 返回的可能是一个视图(view),对于通过自定义 autograd 函数产生 logits 的模型,div_() 操作会因干扰反向传播而被 PyTorch 拒绝。
实现拆解
- rmpad 分支(移除填充路径):在
verl/workers/engine/fsdp/transformer_impl.py 约第 1138-1142 行,将 logits_rmpad.div_(...) 替换为 logits_rmpad = logits_rmpad / ...,并在除法前检查 logits_rmpad 是否为 DTensor,若是则调用 full_tensor() 收集完整张量。
- BSHD 分支(批处理序列路径):在约第 1249-1252 行,对
logits 张量也执行相同的 DTensor 检查和 out-of-place 除法。
- 两处修改均参考了 torchtitan 和 automodel 的实现模式,确保与 TP(张量并行)场景兼容。
关键文件:
verl/workers/engine/fsdp/transformer_impl.py(模块 引擎;类别 source;类型 core-logic): 核心变更文件,修改了 prepare_model_outputs 方法中两处温度缩放逻辑
关键符号:prepare_model_outputs
关键源码片段
verl/workers/engine/fsdp/transformer_impl.py
核心变更文件,修改了 prepare_model_outputs 方法中两处温度缩放逻辑
# verl/workers/engine/fsdp/transformer_impl.py ( 关键变更片段 )
# 在 rmpad 分支中(约第 1138-1142 行):
else:
logits_rmpad = output.logits.squeeze(0) # (total_nnz, vocab_size)
# With TP, logits are DTensors sharded on vocab dim; gather for log_softmax.
if isinstance(logits_rmpad, DTensor):
logits_rmpad = logits_rmpad.full_tensor() # 收集完整 logits
logits_rmpad = logits_rmpad / temperature_rmpad.clamp(min=1e-8).unsqueeze(-1).to(logits_rmpad.dtype) # out-of-place
# 在 BSHD 分支中(约第 1249-1252 行):
logits = output.logits # (bsz, response_length, vocab_size)
temperature = output_args["temperature"] # (bsz,)
temperature = temperature.unsqueeze(-1).unsqueeze(-1)
# With TP, logits are DTensors sharded on vocab dim; gather for log_softmax.
if isinstance(logits, DTensor):
logits = logits.full_tensor() # 收集完整 logits
logits = logits / temperature.clamp(min=1e-8).to(logits.dtype) # out-of-place
评论区精华
Reviewer wuxibin89 指出还需要修改 BSHD 分支,作者 zjchenn 确认后增加对应修改。
- BSHD 分支也需要同步修改 (correctness): 作者 zjchenn 接受建议并补充了 BSHD 分支的修改。
风险与影响
- 风险:低风险。改动仅在 FSDP 引擎的 eager logits 温度缩放路径中,将 in-place 操作改为 out-of-place,语义等价。新增的
DTensor 处理仅在启用 TP 时触发。未修改测试,但改动量小,逻辑清晰。
- 影响:影响范围局限于 FSDP 引擎中启用了 remove padding 且不使用 fused kernel 的场景,仅对使用 TP 且具有自定义 autograd 的模型有实质性影响。
- 风险标记:核心路径变更, 缺少测试覆盖
关联脉络
参与讨论