Prhub

#27202 [NPU][Bugfix] fix MTP accuracy regression on Qwen3 hybrid models

原始 PR 作者 AndyLi429 合并时间 2026-06-09 15:47 文件变更 2 提交数 14 评论 6 代码增减 +38 / -7

执行摘要

修复 NPU MTP 精度回归,涉及三个底层 bug

PR body 指出,在 NPU 上运行 Qwen3.6-35B-A3(w4a8, MTP 启用)时,ceval 准确率低于基线,开启 thinking 模式后准确率进一步下降且产生重复推理。分析确认三个 bug 导致。修复后 thinking 关闭模式下 ceval 从 87.5% 提升至 89.75%,开启模式下从 86% 提升至 90.71%。

值得精读,特别是 update_mamba_state_after_mtp_verify 中 NPU 特有的 conv 状态重建与回滚策略,体现了硬件差异下的适配设计。建议关注后续 PR 是否在 CUDA 与 NPU 间推进公共抽象层。

讨论亮点

Codex 审查机器人提出 P1 级问题:在 update_mamba_state_after_mtp_verify 中,当 mamba_track_indices 非空时,NPU 实现未像 CUDA 那样把中间状态散列到跟踪槽,可能导致前缀缓存状态错误。该评论已由作者在后续提交中解决,添加了 move_intermediate_cacheconv_state_rollback 的跟踪分支。

实现拆解

  1. SSM 状态索引别名修复ascend_hybrid_linear_attn_backend.py):在 _replay_metadata 的 target_verify 分支中,ssm_state_indices 构造改为 bs * draft_token_num(全 batch 覆盖),并直接 copy_ 整个 state_indices_list_gdn 张量,避免填充行使用 slot 0 造成数据别名。

  2. MTP verify 后 mamba 状态跟踪ascend_hybrid_linear_attn_backend.py):在 update_mamba_state_after_mtp_verify 中添加 mamba_track_indices / mamba_steps_to_track 散列逻辑。SSM 状态通过 move_intermediate_cache 写入跟踪槽;conv 状态因 NPU 无 per-step 中间缓存,采用从工作槽复制当前状态再各自回滚的方式处理。这是 thinking 模式长序列触发 prefix 重用时的关键修复。

  3. NPU KV 复制预热覆盖memory_pool_npu.py):新增 _init_kv_copy_and_warmup 方法,将 _kv_copy_config 设为 None,避免基类中依赖 data_strides / data_ptrs 的初始化在 NPU 上失败。

文件 模块 状态 重要度
python/sglang/srt/hardware_backend/npu/attention/ascend_hybrid_linear_attn_backend.py NPU 注意力 modified 6.8
python/sglang/srt/hardware_backend/npu/memory_pool_npu.py NPU 内存 modified 5.53

关键符号

_replay_metadata update_mamba_state_after_mtp_verify _init_kv_copy_and_warmup

关键源码片段

python/sglang/srt/hardware_backend/npu/attention/ascend_hybrid_linear_attn_backend.py core-logic

核心修复文件,包含 SSM 索引别名修复和 mamba 状态跟踪添加,直接影响 MTP 精度。

def update_mamba_state_after_mtp_verify(
    self,
    request_number: int,
    last_correct_step_indices: torch.Tensor,
    mamba_track_indices: Optional[torch.Tensor],
    mamba_steps_to_track: Optional[torch.Tensor],
):
    # 获取当前请求的 mamba 缓存索引和中间状态缓存
    state_indices_tensor = (
        self.linear_attn_backend.forward_metadata.mamba_cache_indices[:request_number]
    )
    mamba_caches = (
        self.linear_attn_backend.req_to_token_pool.get_speculative_mamba2_params_all_layers()
    )
    conv_states = mamba_caches.conv[0]
    ssm_states = mamba_caches.temporal
    intermediate_state_cache = mamba_caches.intermediate_ssm
​
    dst_indices_tensor = state_indices_tensor.to(torch.int64)
    src_indices_tensor = torch.arange(
        dst_indices_tensor.shape[0],
        device=dst_indices_tensor.device,
        dtype=torch.int64,
    )
    last_steps = last_correct_step_indices.to(torch.int64)
​
    # 1) 将 verify 后的 SSM 中间状态从 intermediate_state_cache 复制回 ssm_states 的工作槽
    move_intermediate_cache(
        ssm_states, intermediate_state_cache,
        dst_indices_tensor, src_indices_tensor, last_steps,
    )
​
    draft_token_num = intermediate_state_cache.shape[2]
​
    # 2) 如果存在需要跟踪的 mamba 状态(thinking 模式长序列触发),执行散列
    if mamba_track_indices is not None:
        assert mamba_steps_to_track is not None
        mamba_track_indices = mamba_track_indices.to(torch.int64)
        mamba_steps_to_track = mamba_steps_to_track.to(torch.int64)
​
        # 将 SSM 中间状态散列到跟踪槽
        move_intermediate_cache(
            ssm_states, intermediate_state_cache,
            mamba_track_indices, src_indices_tensor, mamba_steps_to_track,
        )
​
        # NPU 特有:conv 状态重建(无 per-step 中间缓存)
        track_mask = mamba_steps_to_track >= 0
        track_indices = mamba_track_indices[track_mask]
        if track_indices.numel() > 0:
            # 从 verify 时的工作槽复制当前 conv 状态到跟踪槽
            conv_states[:, track_indices] = conv_states[:, dst_indices_tensor[track_mask]]
​
    # 3) 回滚工作槽的 conv 状态
    if dst_indices_tensor.numel() > 0:
        conv_state_rollback(conv_states, dst_indices_tensor, last_steps, draft_token_num)
​
    # 4) 回滚跟踪槽的 conv 状态(如果有)
    if mamba_track_indices is not None and mamba_track_indices.numel() > 0:
        conv_state_rollback(conv_states, mamba_track_indices, mamba_steps_to_track, draft_token_num)
​
    return

评论区精华

NPU verify 阶段缺少 mamba_track_indices 处理 正确性

Codex 审查机器人提出现有 NPU 实现未处理 mamba_track_indices,可能导致前缀缓存状态错误(P1 级)。

结论:已在 PR 中添加了 move_intermediate_cache 和 conv_state_rollback 的跟踪分支,问题已解决。 · 已解决

风险与影响

  1. SSM 索引范围ssm_state_indices(bs-padding)*draft_token_num 改为 bs*draft_token_num,增加 padded 行的独立 slot。需要确保 state_indices_list_gdn 缓冲区大小始终是 bs*draft_token_numinit_cuda_graph_state 已保证),且 padded 槽的数据不会被真实请求读取。当前修复后 padded 槽不回读,但后续若引入读取 padded 槽的路径则会出错。
  2. conv 状态跟踪:NPU 特有的 conv 状态重建依赖 conv_states[:, track_indices] = conv_states[:, dst_indices_tensor[track_mask]],这假设工作槽与跟踪槽完全分离,且 conv_state_rollback 行为一致。如果 NPU 硬件升级改动 conv 布局,此逻辑需同步更新。
  3. KV 复制预热:将 _kv_copy_config 设为 None 跳过了预热,若将来 NPU 需要 KV 复制优化,需重新实现此方法。

影响范围:仅限于 Ascend NPU 后端的 MTP 推测解码路径,主要影响 Qwen3 hybrid 系列模型(如 Qwen3.6-35B-A3)。
影响程度:修复前 thinking 模式准确率低且重复推理,影响用户体验;修复后准确率恢复正常,模型生成质量显著提升。对其他后端(CUDA、ROCm、XPU)无影响。

SSM 索引范围扩大 conv 状态跟踪依赖硬件特性 KV 复制预热被跳过

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论