[NPU] perf: precompute mamba conv-state track indices once per batch (#29105)

This commit is contained in:
AndyLi429
2026-06-25 14:04:56 +08:00
committed by GitHub
parent 7c9804ef21
commit 72cac88022
@@ -47,6 +47,15 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
prefill_backend = get_linear_attn_prefill_backend() prefill_backend = get_linear_attn_prefill_backend()
self.kernel_dispatcher = GDNKernelDispatcher(decode_backend, 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( def prepare_gdn_inputs(
self, self,
bs: int, bs: int,
@@ -85,6 +94,7 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
forward_batch.forward_mode, forward_batch.forward_mode,
forward_batch.spec_info, forward_batch.spec_info,
) )
self._prepare_mamba_track_metadata(forward_batch)
self.graph_mode = True self.graph_mode = True
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
@@ -96,6 +106,7 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
forward_batch.forward_mode, forward_batch.forward_mode,
forward_batch.spec_info, forward_batch.spec_info,
) )
self._prepare_mamba_track_metadata(forward_batch)
self.graph_mode = False self.graph_mode = False
def forward_decode( def forward_decode(
@@ -221,16 +232,13 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
).view(seq_len, -1) ).view(seq_len, -1)
else: else:
mixed_qkv = mixed_qkv.transpose(0, 1) mixed_qkv = mixed_qkv.transpose(0, 1)
if ( if forward_metadata.has_mamba_track_mask:
forward_batch.mamba_track_mask is not None
and forward_batch.mamba_track_mask.any()
):
conv_dst = forward_batch.mamba_track_indices
mixed_qkv_to_track = mixed_qkv[ mixed_qkv_to_track = mixed_qkv[
:, forward_metadata.track_conv_indices :, forward_metadata.track_conv_indices
].transpose(0, 1) ].transpose(0, 1)
mask_indices = forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0] conv_states.transpose(1, 2)[
conv_states.transpose(1, 2)[conv_dst[mask_indices]] = mixed_qkv_to_track forward_metadata.conv_states_mask_indices
] = mixed_qkv_to_track
kernel_size = layer.conv_weights.shape[-1] kernel_size = layer.conv_weights.shape[-1]
conv_states_for_prefill = conv_states[:, -(kernel_size - 1) :, :] conv_states_for_prefill = conv_states[:, -(kernel_size - 1) :, :]
conv_states_tmp = conv_states_for_prefill.transpose(1, 2).contiguous() conv_states_tmp = conv_states_for_prefill.transpose(1, 2).contiguous()