# PR #27202 完整报告

- 仓库：`sgl-project/sglang`
- 标题：[NPU][Bugfix] fix MTP accuracy regression on Qwen3 hybrid models
- 合并时间：2026-06-09 15:47
- 原文链接：http://prhub.com.cn/sgl-project/sglang/pull/27202

---

# 执行摘要

- 一句话：修复 NPU MTP 精度回归，涉及三个底层 bug
- 推荐动作：值得精读，特别是 `update_mamba_state_after_mtp_verify` 中 NPU 特有的 conv 状态重建与回滚策略，体现了硬件差异下的适配设计。建议关注后续 PR 是否在 CUDA 与 NPU 间推进公共抽象层。

# 功能与动机

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

# 实现拆解

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 注意力；类别 source；类型 core-logic；符号 _replay_metadata, update_mamba_state_after_mtp_verify）: 核心修复文件，包含 SSM 索引别名修复和 mamba 状态跟踪添加，直接影响 MTP 精度。
- `python/sglang/srt/hardware_backend/npu/memory_pool_npu.py`（模块 NPU 内存；类别 source；类型 core-logic；符号 _init_kv_copy_and_warmup）: 添加 NPU 特有的 KV 复制预热方法，避免基类中不兼容的初始化导致错误。

关键符号：_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`

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

```python
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

```

# 评论区精华

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

- NPU verify 阶段缺少 mamba_track_indices 处理 (correctness): 已在 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_num`（`init_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 复制预热被跳过

# 关联脉络

- 暂无明显关联 PR