[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:
co-authored by
Claude Opus 5
parent
23ea7b6481
commit
bfa4e4a57b
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user