Prhub

#32589 [Nemotron] Hoist mamba track-mask host syncs out of the per-layer prefill path

原始 PR 作者 alexnails 合并时间 2026-08-04 14:58 文件变更 4 提交数 6 评论 3 代码增减 +32 / -23

执行摘要

消除 Mamba2 逐层 host sync,prefill 吞吐提升 1.5-3%

在 extra_buffer mamba radix 策略下,Mamba2 架构在 MambaMixer2.forward 内逐层解析 tracked-row 选择:forward_batch.mamba_track_mask.any() 经 Tensor.bool -> .item() 触发 stream sync,nonzero(as_tuple=True) 因输出尺寸数据依赖再触发一次同步。由于 prefill 分支总是 return_intermediate_states=True,.any() 对每个 mamba 层都会被求值,backend 中同一谓词的副本也无法短路。对 NVIDIA-Nemotron-Nano-9B-v2(hybrid_override_pattern = 27 个 M 层),每次 prefill forward 就是 56 次 host sync(27 次 mixer + 27 次 backend + 2 次 _forward_metadata),每一次都在层起始处排空 pipeline,CPU 无法预跑下一层 kernel。而 GDN / KDA / ascend-gdn 已把同样的两个张量提升进元数据,Mamba2 是唯一仍逐层推导的后端,且 ForwardMetadata 上所需的 mamba_track_mask_indices / conv_states_mask_indices 字段早已声明却从未填充。

值得精读,尤其是性能分析方法论和 metadata 提升模式。关注点:

  • 用 profiler 的 sync 计数(57 -> 2)验证代码推导的机制,而非只看端到端数字;
  • 对“sync 阻塞时间 vs CPU run-ahead 损失”的区分,纠正了常见的性能误读;
  • benchmark 设计(interleaved rounds、drift control、warmed ranges、disjoint ranges)可作为引擎侧性能 PR 的模板;
  • 提升 + 断言组合:把每层推导提升到每 forward,同时用 num_decodes == bs 断言保护切片约定。
讨论亮点

review 层面没有实质交锋:sshleifer 直接 APPROVED(LGTM),0 条 review 评论。有价值的讨论集中在 issue 评论中,作者 alexnails 发布了两项补测:

  1. sync 计数验证:Torch profiler 显示 per-request 的 aten::any 从 57.0 降到 2.0、cudaStreamSynchronize 从 99.0 降到 18.0、3 个请求的 host-sync wall time 从 30,407 us 降到 250 us,与代码推导的 56 次/forward 完全吻合。作者特别澄清:The 30,407 us -> 250 us row is not the speedup, and should not be read as one——sync 阻塞时间取决于已入队的 GPU 工作,真实收益是恢复的 CPU run-ahead,可信的收益数字是端到端 A/B(+3.0% / +1.5%)。

  2. decode no-regression:由于第 3、4 项改动触碰 decode replay 路径,作者补充了 decode-only A/B 确认无回归。

另两条评论来自 gemini-code-assist bot,是服务日落公告,无实质内容。

实现拆解

1. 谓词提升到 _forward_metadata

python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py 中,has_mamba_track_mask = bool(mask is not None and mask.any())_forward_metadata 末尾移到开头,只求值一次,并让 extend 分支的 track_conv_indices 初始化直接复用该值。这样每个 forward 只触发一次 .item() 同步,ForwardMetadata 携带的谓词成为唯一事实来源。

2. prepare_mixed 一次性解析 tracked-row 索引

python/sglang/srt/layers/attention/mamba/mamba2_metadata.py 中,Mamba2Metadata.prepare_mixed 在构造阶段(仅当 has_mamba_track_mask 为真)执行一次 nonzero(as_tuple=True) 得到 mamba_track_mask_indices,并在已翻译的 mamba_track_indices 上 gather 出 conv_states_mask_indices,作为不可变字段随 Mamba2Metadata 构造传入。这两个字段在 ForwardMetadata 上早已声明但从未填充,本次首次打通;相比 GDN 事后赋值,构造时传入保持了 metadata 不可变。

3. mixer 改为纯 metadata 驱动

python/sglang/srt/layers/attention/mamba/mamba.py 中,MambaMixer2.forward 删除 forward_batch 参数(keyword-only 参数,唯一调用点为 Mamba2AttnBackend.forward,nemotron_h / falcon_h1 / granitemoehybrid 均经 backend 路由),track-conv 写入改用 metadata.conv_states_mask_indices,谓词判断改用 metadata.has_mamba_track_mask,per-layer 的两处 host sync 归零。

4. backend 去冗余与 decode 路径配套

Mamba2AttnBackend.forward 删除对 mask.any() 的冗余再推导(_track_mamba_state_extend 已由 metadata.has_mamba_track_mask 门控);python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.pybuild_replay_fb_viewmamba_track_indices 从整缓冲改为 [:bs] 切片,与其它 buffer 一致,使 decode 的 translate + copy 从 O(max_bs) 降为 O(bs);init_forward_metadata_out_graph 新增 num_decodes == batch_size 断言,防止 eager mixed 与图路径切片约定不一致时静默损坏缓存状态。

5. 验证配套与测试缺口

PR 未新增单元测试(Checklist 未勾选),以 H200 基准 + GSM8K 精度验证覆盖:吞吐 +1.5%~3.0%、TTFT -4.4%~10.5%、GSM8K 0.890 持平(invalid 从 0.005 降为 0.000);profiler 显示 per-request sync 从 56 次降到 1 次。

文件 模块 状态 重要度
python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py 注意力后端 modified 6.51
python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py 解码图执行 modified 6.05
python/sglang/srt/layers/attention/mamba/mamba2_metadata.py 状态元数据 modified 5.97
python/sglang/srt/layers/attention/mamba/mamba.py 状态混合器 modified 5.48

关键符号

Mamba2AttnBackend._forward_metadata Mamba2AttnBackend.forward Mamba2AttnBackend.init_forward_metadata_out_graph Mamba2Metadata.prepare_mixed MambaMixer2.forward build_replay_fb_view

关键源码片段

python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py core-logic

核心改动所在:`_forward_metadata` 把 `has_mamba_track_mask` 谓词提升为每 forward 一次;`Mamba2AttnBackend.forward` 删除逐层冗余的 `mask.any()` 推导;`init_forward_metadata_out_graph` 新增 `num_decodes == bs` 断言保护图路径切片约定。

def forward(
    self,
    mixer: MambaMixer2,
    hidden_states: torch.Tensor,
    output: Optional[torch.Tensor],
    layer_id: int,
    forward_batch: ForwardBatch,
    mup_vector: Optional[torch.Tensor] = None,
    use_triton_causal_conv: bool = False,
):
    assert isinstance(self.forward_metadata, Mamba2Metadata)
    # Page-major 布局下状态存储是跨步的,只有识别 stride 的 Triton causal-conv
    # 能正确读取(CUDA 的 causal_conv1d 会读错);模型也可强制走 Triton。
    use_triton_causal_conv = (
        use_triton_causal_conv or get_memory().enable_page_major_kv_layout
    )
    layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
    # 不再向 mixer 传 forward_batch:tracked-row 选择已全部迁移到 metadata,
    # 逐层调用 mask.any() / nonzero() 触发的 host sync 已消除
    mixer_out, intermediate_states = mixer.forward(
        hidden_states=hidden_states,
        output=output,
        layer_cache=layer_cache,
        metadata=self.forward_metadata,
        mup_vector=mup_vector,
        use_triton_causal_conv=use_triton_causal_conv,
    )
​
    if forward_batch.mamba_track_mask is not None:
        if intermediate_states is not None:
            # has_mamba_track_mask 每 forward 只在 _forward_metadata 求值一次,
            # 这里直接沿用,不再触发 Tensor.__bool__ 的 stream sync
            self._track_mamba_state_extend(
                forward_batch,
                intermediate_states,
                layer_cache.temporal,
                self.forward_metadata,
            )
​
        if self.forward_metadata.num_decodes > 0:
            num_decodes = self.forward_metadata.num_decodes
            # decode 的 track-save 从尾部切片:eager mixed batch 的 decode 行在最后,
            # 图路径则在静态 buffer 头部(配合 build_replay_fb_view 的 [:bs] 切片与
            # init_forward_metadata_out_graph 的 num_decodes == bs 断言)
            track_mamba_states_if_needed(
                layer_cache.conv[0],
                layer_cache.temporal,
                self.forward_metadata.mamba_cache_indices[-num_decodes:],
                forward_batch.mamba_track_mask[-num_decodes:],
                self.forward_metadata.mamba_track_indices[-num_decodes:],
                num_decodes,
                check_freed_slots=self.enable_unified_memory,
            )
​
    return mixer_out
python/sglang/srt/layers/attention/mamba/mamba2_metadata.py core-logic

`Mamba2Metadata.prepare_mixed` 在构造阶段一次性解析 `mamba_track_mask_indices` / `conv_states_mask_indices`,激活 ForwardMetadata 上长期空置的字段,是“每层两次 sync -> 每 forward 零次”的关键。

        # ... 上方已完成 chunked prefill 的 seq_idx / chunk_offsets 等元数据计算 ...
        draft_token_num = (
            getattr(forward_batch.spec_info, "draft_token_num", 1)
            if forward_batch.spec_info is not None
            else 1
        )
        # Resolve the tracked-row selection once per forward
        # nonzero 的输出尺寸数据依赖,逐层调用会同步 host;这里每 forward 只做一次
        mamba_track_mask_indices = None
        conv_states_mask_indices = None
        if forward_metadata.has_mamba_track_mask:
            mamba_track_mask_indices = forward_batch.mamba_track_mask.nonzero(
                as_tuple=True
            )[0]
            # 在已翻译(virtual -> physical)的 track 索引上 gather,
            # MambaMixer2.forward 据此直接写入 conv_state,避免逐层重复推导
            conv_states_mask_indices = forward_batch.mamba_track_indices[
                mamba_track_mask_indices
            ]
        return Mamba2Metadata(
            query_start_loc=query_start_loc,
            mamba_cache_indices=forward_metadata.mamba_cache_indices,
            mamba_track_indices=forward_metadata.mamba_track_indices,
            has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
            mamba_track_mask_indices=mamba_track_mask_indices,
            conv_states_mask_indices=conv_states_mask_indices,
            num_prefills=num_prefills,
            num_prefill_tokens=num_prefill_tokens,
            num_decodes=num_decodes,
            is_target_verify=forward_batch.forward_mode.is_target_verify(),
            draft_token_num=draft_token_num,
            mixed_metadata=cls.MixedMetadata(
                has_initial_states=has_initial_states,
                prep_initial_states=prep_initial_states,
                chunk_size=chunk_size,
                seq_idx=seq_idx,
                chunk_indices=chunk_indices,
                chunk_offsets=chunk_offsets,
                extend_seq_lens_cpu=extend_seq_lens_cpu,
            ),
        )

评论区精华

sync 计数与 decode 回归的补充验证 性能

作者在 issue 评论中发布 Torch profiler 对比:per-request 的 aten::any 从 57.0 降到 2.0、cudaStreamSynchronize 从 99.0 降到 18.0、3 个请求 host-sync wall time 从 30,407 us 降至 250 us,与代码推导的 56 次 /forward 吻合;并澄清该 wall time 不是速度收益,真实代价是 CPU run-ahead 丢失。

结论:机制验证成立:per-layer 推导确实消失;decode-only A/B 确认第 3、4 项改动(切片与断言)无回归。 · 已解决

与 #32555 的边界划分 设计

PR body 说明与 #32555 互补不重叠:源头切片(build_replay_fb_view)不修复 _replay_metadata 分发完整 track buffer 导致的 stale tail 问题,#32555 仍必须合并;两者触碰不同 hunk(约 200 行间隔),合并无冲突。

结论:两个修复保持独立,本 PR 不替代 #32555。 · 已解决

风险与影响

  1. 核心路径变更:MambaMixer2.forward 删除 forward_batch 参数。所有形参均为 keyword-only,提交信息论证唯一调用点是 Mamba2AttnBackend.forward(nemotron_h、falcon_h1、granitemoehybrid 均经 backend 路由);但任何绕过 backend 的直接调用或外部扩展都会因签名变化而失败。
  2. 新增运行期断言:init_forward_metadata_out_graph 断言 num_decodes == batch_size,设计意图是 fail-fast 而非静默损坏图路径下的 mamba 缓存状态,但需确认不存在合法的 num_decodes < bs 图回放场景被误杀;PR 未针对该断言新增单元测试。
  3. decode 回放路径变更:build_replay_fb_viewmamba_track_indices 由整缓冲改为 [:bs] 切片,影响 decode track-save 的输入范围;作者补充的 decode-only A/B 未显示回归,但改动没有对应单测。
  4. 正确性等价性:has_mamba_track_mask 由构造保证等价于 bool(mask is not None and mask.any())conv_states_mask_indices 是在同一批已翻译索引上的同构 gather,语义无变化;GSM8K 精度持平(0.890)。
  5. 收益范围有限:加速仅出现在 extra_buffer 策略 + Mamba2 架构 + prefill 场景,decode(CUDA graph 下 per-layer Python 不执行)与其他策略无收益,但理论上无回归。

影响模型:NVIDIA-Nemotron-Nano-9B-v2 等 hybrid_override_pattern 全 Mamba 层的混合架构,以及 Falcon-H1、GraniteMoE-Hybrid。影响场景:extra_buffer + radix cache 下的 prefill——1x H200 上 prefill-heavy(output len 1)实测输入 512 时吞吐 +3.0%、TTFT -10.5%,输入 2048 时 +1.5%、-4.4%;decode 吞吐不受影响。对团队/代码库:激活了 ForwardMetadata 上长期空置的 mamba_track_mask_indices / conv_states_mask_indices 字段,使 Mamba2 后端与 GDN / KDA 的 metadata 驱动模式对齐,为后续 mamba 状态跟踪优化确立了范式。

核心路径变更 缺少测试覆盖 新增运行期断言

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论