From 72cac880223e8326f42abdd8dd50d4f0bea718e1 Mon Sep 17 00:00:00 2001 From: AndyLi429 <68410213+AndyLi429@users.noreply.github.com> Date: Thu, 25 Jun 2026 14:04:56 +0800 Subject: [PATCH] [NPU] perf: precompute mamba conv-state track indices once per batch (#29105) --- .../npu/attention/ascend_gdn_backend.py | 22 +++++++++++++------ 1 file changed, 15 insertions(+), 7 deletions(-) diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py index d62483c9a..7f506e7f1 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py @@ -47,6 +47,15 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase): prefill_backend = get_linear_attn_prefill_backend() self.kernel_dispatcher = GDNKernelDispatcher(decode_backend, prefill_backend) + def _prepare_mamba_track_metadata(self, forward_batch: ForwardBatch): + if self.forward_metadata.has_mamba_track_mask: + 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[mamba_track_mask_indices] + ) + def prepare_gdn_inputs( self, bs: int, @@ -85,6 +94,7 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase): forward_batch.forward_mode, forward_batch.spec_info, ) + self._prepare_mamba_track_metadata(forward_batch) self.graph_mode = True def init_forward_metadata(self, forward_batch: ForwardBatch): @@ -96,6 +106,7 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase): forward_batch.forward_mode, forward_batch.spec_info, ) + self._prepare_mamba_track_metadata(forward_batch) self.graph_mode = False def forward_decode( @@ -221,16 +232,13 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase): ).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.transpose(1, 2)[conv_dst[mask_indices]] = mixed_qkv_to_track + conv_states.transpose(1, 2)[ + forward_metadata.conv_states_mask_indices + ] = mixed_qkv_to_track kernel_size = layer.conv_weights.shape[-1] conv_states_for_prefill = conv_states[:, -(kernel_size - 1) :, :] conv_states_tmp = conv_states_for_prefill.transpose(1, 2).contiguous()