patch.py — mcore 兼容性补丁¶
文件路径¶
verl/models/mcore/patch.py
文件概述¶
修复 Megatron-Core 0.12 版本中 MLA (Multi-Latent Attention) 的两个 bug,并提供 mbridge 库的兼容性补丁。
关键代码讲解¶
apply_patch() -- MLA Bug 修复¶
Bug 1: get_query_key_value_tensors 在 packed_seq_params 时行为异常
def patch_get_query_key_value_tensors(self, hidden_states, ...):
# 修复:当 packed_seq_params 不为 None 时,
# 需要对 q_pos_emb, k_pos_emb 等做 squeeze(1) 去掉 batch 维度
if packed_seq_params is not None:
q_pos_emb = q_pos_emb.squeeze(1)
k_pos_emb = k_pos_emb.squeeze(1)
q_no_pe = q_no_pe.squeeze(1)
k_no_pe = k_no_pe.squeeze(1)
value = value.squeeze(1)
# 修复:k_pos_emb expand 时的维度处理
if packed_seq_params is not None:
k_pos_emb = k_pos_emb.expand(-1, self.num_attention_heads_per_partition, -1)
else:
k_pos_emb = k_pos_emb.expand(-1, -1, self.num_attention_heads_per_partition, -1)
Bug 2: forward 中 THD 格式的输出 reshape
def patch_forward(self, hidden_states, attention_mask, ...):
# THD 格式下,value 需要 padding 到 query 的 head_dim
if non_dsa_thd_qkv_format and query.shape[-1] != v_dim:
value = F.pad(value, [0, query.shape[-1] - v_dim])
# 输出需要截断回 v_dim 并 reshape
if non_dsa_thd_qkv_format:
if core_attn_out.ndim == 2:
core_attn_out = core_attn_out.reshape(*core_attn_out.shape[:-1], -1, value.shape[-1])
if query.shape[-1] != v_dim:
core_attn_out = core_attn_out[..., :v_dim]
core_attn_out = core_attn_out.reshape(core_attn_out.size(0), 1, -1)
版本控制¶
mcore_ge_013 = version.parse(megatron.core.__version__) >= version.parse("0.13.0")
# 只对 mcore < 0.13 应用 get_query_key_value_tensors 补丁
if not mcore_ge_013:
MLASelfAttention.get_query_key_value_tensors = patch_get_query_key_value_tensors
# forward 补丁始终应用
MultiLatentAttention.forward = patch_forward
apply_patch_mbridge() -- mbridge 兼容性补丁¶
def apply_patch_mbridge():
"""为旧版 mcore 补充缺失的 get_tensor_model_parallel_group_if_none 函数"""
try:
from megatron.core.utils import get_tensor_model_parallel_group_if_none
except ImportError:
# 自行实现并注入
megatron.core.utils.get_tensor_model_parallel_group_if_none = get_tensor_model_parallel_group_if_none
核心函数列表¶
| 函数名 | 作用 |
|---|---|
apply_patch() |
修复 mcore 0.12 的 MLA bug |
apply_patch_mbridge() |
mbridge 兼容性补丁 |
与其他模块的关系¶
- 被
config_converter.py的hf_to_mcore_config_dpskv3()调用 - 被
mbridge.py在导入时调用
小结¶
这个文件是典型的版本兼容性补丁。MLA 在 packed sequence 场景下的 bug 会导致训练失败,因此必须在使用 DeepSeek-V3 等 MLA 模型时应用这些补丁。