From bfa4e4a57b4fa791cce40e717423d9623d7bc91d Mon Sep 17 00:00:00 2001 From: Alex Nails Date: Mon, 3 Aug 2026 23:58:28 -0700 Subject: [PATCH] [Nemotron] Hoist mamba track-mask host syncs out of the per-layer prefill path (#32589) Co-authored-by: Claude Opus 5 (1M context) --- .../attention/hybrid_linear_attn_backend.py | 26 ++++++++----------- .../srt/layers/attention/mamba/mamba.py | 9 +++---- .../layers/attention/mamba/mamba2_metadata.py | 12 +++++++++ .../runner/decode_cuda_graph_runner.py | 8 ++++-- 4 files changed, 32 insertions(+), 23 deletions(-) diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 0c6d9adb9..00226405b 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -99,6 +99,11 @@ class MambaAttnBackendBase(AttentionBackend): forward_batch.mamba_track_indices = self._translate_mamba_indices( forward_batch.mamba_track_indices ) + # Resolve the tracked-row selection once per forward + has_mamba_track_mask = bool( + forward_batch.mamba_track_mask is not None + and forward_batch.mamba_track_mask.any() + ) _real_bs = forward_batch._original_batch_size if _real_bs is not None and _real_bs < mamba_cache_indices.shape[0]: mamba_cache_indices = mamba_cache_indices.clone() @@ -198,10 +203,7 @@ class MambaAttnBackendBase(AttentionBackend): forward_batch.extend_start_loc[-1] + forward_batch.extend_seq_lens[-1] ) - if ( - forward_batch.mamba_track_mask is not None - and forward_batch.mamba_track_mask.any() - ): + if has_mamba_track_mask: track_conv_indices = self._init_track_conv_indices( query_start_loc, forward_batch ) @@ -215,11 +217,6 @@ class MambaAttnBackendBase(AttentionBackend): else: raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode=}") - has_mamba_track_mask = bool( - forward_batch.mamba_track_mask is not None - and forward_batch.mamba_track_mask.any() - ) - return ForwardMetadata( query_start_loc=query_start_loc, mamba_cache_indices=mamba_cache_indices, @@ -800,6 +797,10 @@ class Mamba2AttnBackend(MambaAttnBackendBase): is_target_verify=forward_batch.forward_mode.is_target_verify(), draft_token_num=draft_token_num, ) + # `forward` slices the track destinations from ([-num_decodes:]) + assert ( + self.forward_metadata.num_decodes == forward_batch.batch_size + ), f"{self.forward_metadata.num_decodes=} != {forward_batch.batch_size=}" def init_forward_metadata(self, forward_batch: ForwardBatch): metadata = self._forward_metadata(forward_batch) @@ -831,17 +832,12 @@ class Mamba2AttnBackend(MambaAttnBackendBase): output=output, layer_cache=layer_cache, metadata=self.forward_metadata, - forward_batch=forward_batch, 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 - and forward_batch.mamba_track_mask is not None - and forward_batch.mamba_track_mask.any() - ): + if intermediate_states is not None: self._track_mamba_state_extend( forward_batch, intermediate_states, diff --git a/python/sglang/srt/layers/attention/mamba/mamba.py b/python/sglang/srt/layers/attention/mamba/mamba.py index 632766518..4e21ae80a 100644 --- a/python/sglang/srt/layers/attention/mamba/mamba.py +++ b/python/sglang/srt/layers/attention/mamba/mamba.py @@ -27,7 +27,6 @@ from sglang.srt.layers.linear import ( ) from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.mem_cache.memory_pool import MambaPool -from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_loader.weight_utils import ( composed_weight_loader, sharded_weight_loader, @@ -445,7 +444,6 @@ class MambaMixer2(torch.nn.Module): output: Optional[torch.Tensor] = None, layer_cache: MambaPool.State, metadata: Mamba2Metadata, - forward_batch: ForwardBatch, mup_vector: Optional[torch.Tensor] = None, use_triton_causal_conv: bool = False, ): @@ -564,14 +562,13 @@ class MambaMixer2(torch.nn.Module): x = hidden_states_B_C_p.transpose( 0, 1 ) # this is the form that causal-conv see + # Runs once per mamba layer if ( - forward_batch.mamba_track_mask is not None - and forward_batch.mamba_track_mask.any() + metadata.has_mamba_track_mask and metadata.track_conv_indices is not None ): x_to_track = x[:, metadata.track_conv_indices].transpose(0, 1) - mask_indices = forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0] - conv_state[forward_batch.mamba_track_indices[mask_indices]] = x_to_track + conv_state[metadata.conv_states_mask_indices] = x_to_track ccfn = ( causal_conv1d_fn if not use_triton_causal_conv diff --git a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py index 977b7f7cc..42d8a0c97 100644 --- a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py +++ b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py @@ -275,6 +275,16 @@ class Mamba2Metadata(ForwardMetadata): if forward_batch.spec_info is not None else 1 ) + # Resolve the tracked-row selection once per 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] + 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, @@ -288,6 +298,8 @@ class Mamba2Metadata(ForwardMetadata): track_ssm_final_src=forward_metadata.track_ssm_final_src, track_ssm_final_dst=forward_metadata.track_ssm_final_dst, 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, diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index bbca3c05d..7323effe8 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -176,8 +176,12 @@ def build_replay_fb_view( # The mamba-track registry slot (VIRTUAL ids) is the v2p translate SOURCE # for the backend, which copies the result into its own static buffer and # reads THAT in the decode track-save — this slot is never mutated. None - # when mamba-track is disabled. - mamba_track_indices=getattr(buffers, "mamba_track_indices", None), + # when mamba-track is disabled, slice to [:bs] like every other buffer + mamba_track_indices=( + None + if buffers.mamba_track_indices is None + else buffers.mamba_track_indices[:bs] + ), spec_info=forward_batch.spec_info, )