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 91194c494..5be0b7dc8 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -217,6 +217,11 @@ 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, @@ -228,6 +233,7 @@ class MambaAttnBackendBase(AttentionBackend): track_ssm_h_dst=track_ssm_h_dst, track_ssm_final_src=track_ssm_final_src, track_ssm_final_dst=track_ssm_final_dst, + has_mamba_track_mask=has_mamba_track_mask, ) def init_forward_metadata(self, forward_batch: ForwardBatch): @@ -613,10 +619,7 @@ class MambaAttnBackendBase(AttentionBackend): Note: Conv state tracking for extend is handled separately via gather operations using indices computed by `_init_track_conv_indices`. """ - if ( - forward_batch.mamba_track_mask is not None - and forward_batch.mamba_track_mask.any() - ): + if forward_metadata.has_mamba_track_mask: h = h.squeeze(0) if forward_metadata.track_ssm_h_src.numel() > 0: diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py index 700ccfdf6..4dad415b1 100644 --- a/python/sglang/srt/layers/attention/linear/gdn_backend.py +++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py @@ -256,6 +256,18 @@ class GDNAttnBackend(MambaAttnBackendBase): prefill_backend = get_linear_attn_prefill_backend() self.kernel_dispatcher = GDNKernelDispatcher(decode_backend, prefill_backend) + def init_forward_metadata(self, forward_batch: ForwardBatch): + super().init_forward_metadata(forward_batch) + if self.forward_metadata.has_mamba_track_mask: + self.forward_metadata.mamba_track_mask_indices = ( + forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0] + ) + self.forward_metadata.conv_states_mask_indices = ( + forward_batch.mamba_track_indices[ + self.forward_metadata.mamba_track_mask_indices + ] + ) + def forward_decode( self, layer: RadixLinearAttention, @@ -394,16 +406,13 @@ class GDNAttnBackend(MambaAttnBackendBase): mixed_qkv = mixed_qkv_processed.transpose(1, 2).view(seq_len, -1) else: mixed_qkv = mixed_qkv.transpose(0, 1) - if ( - forward_batch.mamba_track_mask is not None - and forward_batch.mamba_track_mask.any() - ): - conv_dst = forward_batch.mamba_track_indices + if forward_metadata.has_mamba_track_mask: mixed_qkv_to_track = mixed_qkv[ :, forward_metadata.track_conv_indices ].transpose(0, 1) - mask_indices = forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0] - conv_states[conv_dst[mask_indices]] = mixed_qkv_to_track + conv_states[forward_metadata.conv_states_mask_indices] = ( + mixed_qkv_to_track + ) mixed_qkv = causal_conv1d_fn( mixed_qkv, diff --git a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py index 5eeb2b65e..3c1548e9b 100644 --- a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py +++ b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py @@ -41,6 +41,10 @@ class ForwardMetadata: is_target_verify: bool = False draft_token_num: int = 1 + has_mamba_track_mask: bool = False + mamba_track_mask_indices: Optional[torch.Tensor] = None + conv_states_mask_indices: Optional[torch.Tensor] = None + @dataclass(kw_only=True) class Mamba2Metadata(ForwardMetadata):