[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
)
# 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,
@@ -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
@@ -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,
@@ -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,
)