跳转至

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 模型时应用这些补丁。