Prhub

#26351 [bugfix] commit Mamba states after NGRAM target verify

原始 PR 作者 xbfs 合并时间 2026-06-12 06:11 文件变更 1 提交数 3 评论 5 代码增减 +81 / -2

执行摘要

修复 NGRAM 推测解码下 Mamba 状态未提交导致的输出重复

修复 hybrid GDN 模型(如 Qwen3.5)在 NGRAM 推测解码下的输出损坏/重复 bug。PR body 描述用户可见问题:响应中出现 token 循环、重复片段和格式损坏,例如输出重复的 'I am Qwen' 模式。根本原因是 target verify 后接受的每请求推测状态未提交回持久 Mamba 缓存。

建议精读 _mamba_verify_update 的实现,理解 Mamba 状态提交机制和 track 点更新逻辑,这对于维护 hybrid 推测解码非常重要。此修复填补了 NGRAM 路径与 eagle_worker_v2 路径之间的行为差距。

讨论亮点

作者 xbfs 使用 AI 解决合并冲突后,reviewer hnyls2002 要求跑 NGRAM 相关测试(registered/spec/test_spec_ngram.pyregistered/spec/test_spec_ngram_extra.py),测试结果通过。没有其他讨论。

实现拆解

  1. 导入依赖调整:在 python/sglang/srt/speculative/ngram_worker.py 中新增 set_mamba_track_indices_from_reqsget_global_server_args 导入,为后续启用 Mamba 状态跟踪做准备。
  2. verify 前刷新 Mamba track indices:在 _prepare_for_speculative_decoding 方法中,当 enable_mamba_extra_buffer() 为真时,调用 set_mamba_track_indices_from_reqs(batch) 重建 track indices,并清除 mamba_track_maskmamba_track_seqlens,防止过期的 extend 阶段 mask 在 TARGET_VERIFY 期间触发错误的跟踪。
  3. 新增 _mamba_verify_update 方法:该方法在 forward_batch_generation 中的 verify(...) 调用之后立即执行。它计算每个请求最后一个正确步骤的索引(last_correct_step_indices),然后调用 attn_backend.update_mamba_state_after_mtp_verify(...) 将接受的 Mamba 状态提交到持久缓存。同时,如果启用了 mamba track 点更新,还根据 verify 前后的序列长度变化,计算是否需要更新 track 点并调用 set_mamba_track_indices_from_reqs
  4. 集成到 forward_batch_generation:在 target verify 分支中,verify_input.verify(...) 返回 accept_lensaccept_index 后,直接调用 self._mamba_verify_update(batch, accept_lens, accept_index, bs)
文件 模块 状态 重要度
python/sglang/srt/speculative/ngram_worker.py 推测解码 modified 7.59

关键符号

_mamba_verify_update

关键源码片段

python/sglang/srt/speculative/ngram_worker.py core-logic

唯一修改的文件,新增 `_mamba_verify_update` 方法和相关导入与前置刷新逻辑,是修复的核心。

# python/sglang/srt/speculative/ngram_worker.py
# 在 forward_batch_generation 的 target verify 分支中调用# 新增导入
from sglang.srt.managers.schedule_batch import (
    ScheduleBatch,
    set_mamba_track_indices_from_reqs, # 用于刷新 track indices
)
from sglang.srt.server_args import ServerArgs, get_global_server_args# 在 _prepare_for_speculative_decoding 中添加前置刷新
if get_global_server_args().enable_mamba_extra_buffer():
    # 从 reqs 重建 mamba track indices,避免使用 extend 阶段过期的 mask
    set_mamba_track_indices_from_reqs(batch)
    # 清除 mask/seqlens,防止在 TARGET_VERIFY 期间错误触发 in-forward tracking
    batch.mamba_track_mask = None
    batch.mamba_track_seqlens = None# 新增的 _mamba_verify_update 方法
def _mamba_verify_update(
    self,
    batch: ScheduleBatch,
    accept_lens: torch.Tensor,
    accept_index: torch.Tensor,
    bs: int,
) -> None:
    """Commit accepted speculative states for hybrid linear attention backends."""
    attn_backend = self.target_worker.model_runner.attn_backend
    # 如果后端不支持 Mamba 状态更新,直接返回
    if not hasattr(attn_backend, "update_mamba_state_after_mtp_verify"):
        return
    # idle 状态或没有接受 token 时无需操作
    if batch.forward_mode.is_idle() or accept_index.numel() == 0:
        return
​
    # 计算每个请求的最后一个正确步骤索引
    accept_indices_offset = torch.arange(
        0,
        bs * self.draft_token_num,
        step=self.draft_token_num,
        dtype=accept_lens.dtype,
        device=accept_lens.device,
    )
    req_idx = torch.arange(bs, dtype=torch.int64, device=accept_lens.device)
    last_correct_step_indices = (
        accept_index[req_idx, (accept_lens - 1).to(torch.int64)]
        - accept_indices_offset
    )
​
    # 如果启用了 mamba track 点更新,则准备
    if batch.mamba_track_indices is not None:
        seq_lens_pre_verify = batch.seq_lens
        seq_lens_post_verify = batch.seq_lens + accept_lens
        mamba_track_interval = self.server_args.mamba_track_interval
        to_track_mask = (
            seq_lens_pre_verify // mamba_track_interval
            != seq_lens_post_verify // mamba_track_interval
        )
        tracking_point = (
            seq_lens_post_verify // mamba_track_interval
        )
    else:
        to_track_mask = None
        tracking_point = None
​
    # 调用 attention backend 提交状态
    attn_backend.update_mamba_state_after_mtp_verify(
        batch,
        last_correct_step_indices,
        accept_lens,
        to_track_mask,
        tracking_point,
    )
​
    # 如果启用了 track 点更新,刷新 batch 中的 track indices
    if batch.mamba_track_indices is not None and to_track_mask is not None:
        set_mamba_track_indices_from_reqs(batch)

评论区精华

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

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

风险与影响

风险较低。变更仅限于 NGRAM 推测解码路径中的 hybrid GDN 模型,且通过条件 enable_mamba_extra_buffer() 保护,不影响非 hybrid 或非推测解码路径。但缺少直接针对此修复的测试用例覆盖,可能遗漏边缘情况。

直接影响:修复 hybrid GDN 模型(如 Qwen3.5)在 NGRAM 推测解码下的输出质量,消除 token 循环和重复。影响范围小,仅涉及特定模型和推测模式组合。对系统性能无显著影响,因为 Mamba 状态提交仅在 verify 后执行,且通过条件判断避免不必要的调用。

缺少测试覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论