mtp_patch.py — MTP (Multi-Token Prediction) 补丁¶
文件路径¶
verl/models/mcore/mtp_patch.py
文件概述¶
为 Megatron-Core 的 MTP (Multi-Token Prediction) 功能提供 PPO 训练适配补丁。核心问题是:mcore 默认的 MTP 后处理直接计算 loss,但 PPO 需要的是 logits(用于后续计算 log_probs)。
关键代码讲解¶
patch_postprocess() / unpatch_postprocess()¶
def patch_postprocess(model):
"""替换 GPTModel._postprocess 为 PPO 兼容版本"""
model._postprocess_backup = model._postprocess
model._postprocess = _megatron_gptmodel_postprocess.__get__(model, model.__class__)
def unpatch_postprocess(model):
"""恢复原始 _postprocess"""
model._postprocess = model._postprocess_backup
PPO 兼容的 _postprocess¶
def _megatron_gptmodel_postprocess(self, hidden_states, labels, ...):
# 推理路径:委托给原始实现
if labels is None:
return self._postprocess_backup(...)
# 训练路径:MTP 处理 + 返回 logits(而非 loss)
if self.config.mtp_num_layers:
hidden_states_list = torch.chunk(hidden_states, 1 + self.config.mtp_num_layers, dim=0)
hidden_states = hidden_states_list[0] # 主 hidden_states
for mtp_layer_number in range(self.config.mtp_num_layers):
# 对每个 MTP 层计算 loss
mtp_labels = roll_tensor(mtp_labels, shifts=-1, ...)
mtp_loss = self.compute_output_layer_and_language_model_loss(
hidden_states_list[mtp_layer_number + 1], labels=mtp_labels, ...)
# 通过 MTPLossAutoScaler 将 MTP loss 附加到主梯度
hidden_states = MTPLossAutoScaler.apply(hidden_states, mtp_loss_scale * mtp_loss)
# 关键区别:返回 logits 而非 loss
logits, _ = self.output_layer(hidden_states, weight=output_weight)
return logits.transpose(0, 1).contiguous() # [s,b,h] -> [b,s,h]
核心区别:标准 mcore 在最后直接计算 cross-entropy loss 并返回。但 PPO 需要 logits 来计算 log_probs、entropy 等量,因此这个补丁改为返回 logits。
patch_mtp_layer_get_embeddings()¶
def patch_mtp_layer_get_embeddings(model):
"""修改 MTP 层的 _get_embeddings 方法,detach embedding"""
for layer in target_layers:
layer._get_embeddings = _patched_get_embeddings_for_detach.__get__(layer, ...)
def _patched_get_embeddings_for_detach(self, input_ids, position_ids, embedding, hidden_states, ...):
# ... 正常计算 embedding
decoder_input = embedding(input_ids=input_ids, position_ids=position_ids)
# 关键:detach decoder_input 和 hidden_states
decoder_input = decoder_input.detach()
hidden_states = hidden_states.detach()
return input_ids, position_ids, decoder_input, hidden_states
为什么 detach:在 PPO 训练中,MTP 层的梯度不应该回传到主 decoder,因此需要 detach 切断梯度流。
核心函数列表¶
| 函数名 | 作用 |
|---|---|
patch_postprocess() |
替换 _postprocess 为 PPO 版本 |
unpatch_postprocess() |
恢复原始 _postprocess |
_megatron_gptmodel_postprocess() |
PPO 兼容的后处理(返回 logits) |
patch_mtp_layer_get_embeddings() |
Detach MTP 层的 embedding |
unpatch_mtp_layer_get_embeddings() |
恢复 MTP 层的 embedding |
与其他模块的关系¶
- 被训练流程中的 actor 模型前后处理调用
- 使用 mcore 的
MTPLossAutoScaler、roll_tensor等工具 - 与
model_forward.py配合使用
小结¶
MTP 补丁解决了 mcore MTP 功能与 PPO 训练的不兼容问题。两个核心修改:1) 后处理返回 logits 而非 loss;2) MTP 层的 embedding detach 防止梯度泄漏。