[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 = 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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user