执行摘要
- 一句话:修复 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%。
实现拆解
-
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 造成数据别名。
-
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 重用时的关键修复。
-
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 精度。
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 的跟踪分支,问题已解决。
风险与影响
关联脉络
参与讨论