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

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Alex Nails
2026-08-03 23:58:28 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 23ea7b6481
commit bfa4e4a57b
4 changed files with 32 additions and 23 deletions
@@ -99,6 +99,11 @@ class MambaAttnBackendBase(AttentionBackend):
forward_batch.mamba_track_indices = self._translate_mamba_indices( forward_batch.mamba_track_indices = self._translate_mamba_indices(
forward_batch.mamba_track_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 _real_bs = forward_batch._original_batch_size
if _real_bs is not None and _real_bs < mamba_cache_indices.shape[0]: if _real_bs is not None and _real_bs < mamba_cache_indices.shape[0]:
mamba_cache_indices = mamba_cache_indices.clone() mamba_cache_indices = mamba_cache_indices.clone()
@@ -198,10 +203,7 @@ class MambaAttnBackendBase(AttentionBackend):
forward_batch.extend_start_loc[-1] forward_batch.extend_start_loc[-1]
+ forward_batch.extend_seq_lens[-1] + forward_batch.extend_seq_lens[-1]
) )
if ( if has_mamba_track_mask:
forward_batch.mamba_track_mask is not None
and forward_batch.mamba_track_mask.any()
):
track_conv_indices = self._init_track_conv_indices( track_conv_indices = self._init_track_conv_indices(
query_start_loc, forward_batch query_start_loc, forward_batch
) )
@@ -215,11 +217,6 @@ class MambaAttnBackendBase(AttentionBackend):
else: else:
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode=}") 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( return ForwardMetadata(
query_start_loc=query_start_loc, query_start_loc=query_start_loc,
mamba_cache_indices=mamba_cache_indices, mamba_cache_indices=mamba_cache_indices,
@@ -800,6 +797,10 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
is_target_verify=forward_batch.forward_mode.is_target_verify(), is_target_verify=forward_batch.forward_mode.is_target_verify(),
draft_token_num=draft_token_num, 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): def init_forward_metadata(self, forward_batch: ForwardBatch):
metadata = self._forward_metadata(forward_batch) metadata = self._forward_metadata(forward_batch)
@@ -831,17 +832,12 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
output=output, output=output,
layer_cache=layer_cache, layer_cache=layer_cache,
metadata=self.forward_metadata, metadata=self.forward_metadata,
forward_batch=forward_batch,
mup_vector=mup_vector, mup_vector=mup_vector,
use_triton_causal_conv=use_triton_causal_conv, use_triton_causal_conv=use_triton_causal_conv,
) )
if forward_batch.mamba_track_mask is not None: if forward_batch.mamba_track_mask is not None:
if ( if intermediate_states is not None:
intermediate_states is not None
and forward_batch.mamba_track_mask is not None
and forward_batch.mamba_track_mask.any()
):
self._track_mamba_state_extend( self._track_mamba_state_extend(
forward_batch, forward_batch,
intermediate_states, intermediate_states,
@@ -27,7 +27,6 @@ from sglang.srt.layers.linear import (
) )
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.mem_cache.memory_pool import MambaPool 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 ( from sglang.srt.model_loader.weight_utils import (
composed_weight_loader, composed_weight_loader,
sharded_weight_loader, sharded_weight_loader,
@@ -445,7 +444,6 @@ class MambaMixer2(torch.nn.Module):
output: Optional[torch.Tensor] = None, output: Optional[torch.Tensor] = None,
layer_cache: MambaPool.State, layer_cache: MambaPool.State,
metadata: Mamba2Metadata, metadata: Mamba2Metadata,
forward_batch: ForwardBatch,
mup_vector: Optional[torch.Tensor] = None, mup_vector: Optional[torch.Tensor] = None,
use_triton_causal_conv: bool = False, use_triton_causal_conv: bool = False,
): ):
@@ -564,14 +562,13 @@ class MambaMixer2(torch.nn.Module):
x = hidden_states_B_C_p.transpose( x = hidden_states_B_C_p.transpose(
0, 1 0, 1
) # this is the form that causal-conv see ) # this is the form that causal-conv see
# Runs once per mamba layer
if ( if (
forward_batch.mamba_track_mask is not None metadata.has_mamba_track_mask
and forward_batch.mamba_track_mask.any()
and metadata.track_conv_indices is not None and metadata.track_conv_indices is not None
): ):
x_to_track = x[:, metadata.track_conv_indices].transpose(0, 1) 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[metadata.conv_states_mask_indices] = x_to_track
conv_state[forward_batch.mamba_track_indices[mask_indices]] = x_to_track
ccfn = ( ccfn = (
causal_conv1d_fn causal_conv1d_fn
if not use_triton_causal_conv if not use_triton_causal_conv
@@ -275,6 +275,16 @@ class Mamba2Metadata(ForwardMetadata):
if forward_batch.spec_info is not None if forward_batch.spec_info is not None
else 1 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( return Mamba2Metadata(
query_start_loc=query_start_loc, query_start_loc=query_start_loc,
mamba_cache_indices=forward_metadata.mamba_cache_indices, 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_src=forward_metadata.track_ssm_final_src,
track_ssm_final_dst=forward_metadata.track_ssm_final_dst, track_ssm_final_dst=forward_metadata.track_ssm_final_dst,
has_mamba_track_mask=forward_metadata.has_mamba_track_mask, 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_prefills=num_prefills,
num_prefill_tokens=num_prefill_tokens, num_prefill_tokens=num_prefill_tokens,
num_decodes=num_decodes, num_decodes=num_decodes,
@@ -176,8 +176,12 @@ def build_replay_fb_view(
# The mamba-track registry slot (VIRTUAL ids) is the v2p translate SOURCE # 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 # 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 # reads THAT in the decode track-save — this slot is never mutated. None
# when mamba-track is disabled. # when mamba-track is disabled, slice to [:bs] like every other buffer
mamba_track_indices=getattr(buffers, "mamba_track_indices", None), mamba_track_indices=(
None
if buffers.mamba_track_indices is None
else buffers.mamba_track_indices[:bs]
),
spec_info=forward_batch.spec_info, spec_info=forward_batch.spec_info,
) )