Prhub

#45240 [XPU][DeepSeek-V4] Fix MTP: sync with upstream fixes #44821 and #43746

原始 PR 作者 majian4work 合并时间 2026-06-12 15:41 文件变更 1 提交数 3 评论 1 代码增减 +29 / -18

执行摘要

修复 XPU 上 DeepSeek V4 MTP 实现的多项 bug

XPU 平台的 DeepSeek V4 MTP 实现需要与 NV/AMD 版本保持同步,以修复多个已知 bug:e_proj/h_proj 缺少 prefix 导致 compressed-tensors 无法正确应用量化的 ignore/target 规则;缺失 mhc_post 调用导致残差混合错误;权重映射中引用了不存在的 norm_gate 键。

值得合入,修复了明确的 bug,且经过 review 验证。建议关注是否有对应的测试覆盖,特别是 mhc_post 和权重映射相关的回归测试。

讨论亮点

PR 仅有两位 reviewer 批准(LGTM),无深入技术讨论。

实现拆解

  1. 传递 prefix 参数:在 DeepSeekV4MultiTokenPredictorLayer.__init__ 中,为 e_projh_projReplicatedLinear 构造函数添加 prefix 参数,使得 compressed-tensors 能够根据权重名正确匹配量化的排除或目标规则。

  2. 移除 torch.compile 依赖并引入融合 kernel:删除 from vllm.compilation.decorators import support_torch_compile 导入和 @support_torch_compile 装饰器。从 vllm.models.deepseek_v4.common.ops 导入 fused_mtp_input_rmsnormmtp_shared_head_rmsnorm 两个融合 Triton kernel,分别替代原先分离的 mask + enorm + hnorm 操作和 SharedHead.forward() 中的 norm 步骤。

  3. 修复 forward 逻辑:在 forward 中,使用 fused_mtp_input_rmsnorm 一步完成输入掩码、enorm 和 hnorm,同时调整 previous_hidden_states 的 reshape 时机以配合融合 kernel 的输入格式。在 mtp_block 前向之后添加缺失的 mhc_post 调用,修复了残差混合的 bug。

  4. 修复 compute_logits 和权重映射:在 compute_logits 中使用 mtp_shared_head_rmsnorm 替换单独的 shared_head 前向和 norm;移除 DeepSeekV4MTP 类上的 @support_torch_compile 装饰器。修复了 WEIGHT_NAME_REMAPPING 中不存在的 norm_gate 条目,并修正了 gate.bias 的映射路径。

文件 模块 状态 重要度
vllm/models/deepseek_v4/xpu/mtp.py 模型层 modified 6.92

关键符号

DeepSeekV4MultiTokenPredictorLayer.__init__ DeepSeekV4MultiTokenPredictorLayer.forward DeepSeekV4MultiTokenPredictorLayer.compute_logits DeepSeekV4MTP.__init__

关键源码片段

vllm/models/deepseek_v4/xpu/mtp.py core-logic

所有变更均在此文件,包括导入调整、prefix 传递、融合 kernel 替换、mhc_post 修复和权重映射修复。

# vllm/models/deepseek_v4/xpu/mtp.py ( 关键变更片段 )class DeepSeekV4MultiTokenPredictorLayer(nn.Module):
    def __init__(self, vllm_config, topk_indices_buffer, prefix, aux_stream_list=None):
        # ...
        # [ 修复 #44821] 为 e_proj 和 h_proj 传递 prefix,
        # 使 compressed-tensors 能正确匹配权重名
        self.e_proj = ReplicatedLinear(
            config.hidden_size, config.hidden_size, bias=False, return_bias=False,
            quant_config=quant_config,
            prefix=f"{prefix}.e_proj", # 新增 prefix
        )
        self.h_proj = ReplicatedLinear(
            config.hidden_size, config.hidden_size, bias=False, return_bias=False,
            quant_config=quant_config,
            prefix=f"{prefix}.h_proj", # 新增 prefix
        )
        # ...
​
    def forward(self, input_ids, positions, previous_hidden_states, inputs_embeds, ...):
        # 调整 reshape 顺序以适应融合 kernel
        previous_hidden_states = previous_hidden_states.view(
            -1, self.hc_mult, self.config.hidden_size
        )
        # [ 修复 #43746] 使用融合 Triton kernel 替代分离的 mask + enorm + hnorm
        inputs_embeds, previous_hidden_states = fused_mtp_input_rmsnorm(
            inputs_embeds, positions, previous_hidden_states,
            self.enorm.weight.data, self.hnorm.weight.data,
            self.enorm.variance_epsilon, self.hc_mult,
        )
        hidden_states = self.h_proj(previous_hidden_states) + self.e_proj(
            inputs_embeds
        ).unsqueeze(-2)
        hidden_states, residual, post_mix, res_mix = self.mtp_block(
            positions=positions, x=hidden_states, input_ids=None
        )
        # [ 修复 ] 缺失的 mhc_post 调用,修复残差混合 bug
        hidden_states = self.mtp_block.hc_post(
            hidden_states, residual, post_mix, res_mix
        )
        return hidden_states
​
    def compute_logits(self, mtp_layer, hidden_states):
        # [ 修复 ] 使用融合 kernel 替换 shared_head.forward()
        hidden_states = mtp_shared_head_rmsnorm(
            hidden_states,
            mtp_layer.shared_head.norm.weight.data,
            mtp_layer.shared_head.norm.variance_epsilon,
        )
        logits = self.logits_processor(mtp_layer.shared_head.head, hidden_states)
        return logits

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险较低:该 PR 主要同步上游修复,且改动集中在一个文件,并已验证与 NV/AMD 版本的差异仅为平台适配。但缺乏直接对应的测试文件(如 tests/ 下的新增或修改),回归风险需要通过已有集成测试覆盖。

直接影响 XPU(Intel GPU)上 DeepSeek V4 的 MTP(Multi-Token Prediction)推理路径,修复了几个可能影响正确性和性能的 bug。对用户而言,使用 XPU 运行 DeepSeek V4 模型时,推理结果会更准确,性能因融合 kernel 可能提升。不涉及其他平台或模型。

缺少测试覆盖

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论