From 0668a7f51ac5b88dd8406a832941a3af64d4d2d3 Mon Sep 17 00:00:00 2001 From: Jincong Chen Date: Fri, 10 Apr 2026 17:53:57 +0800 Subject: [PATCH] [Perf] Remove two operations in gdn_backend extend verify path (#22444) --- .../srt/layers/attention/linear/gdn_backend.py | 12 ++++-------- 1 file changed, 4 insertions(+), 8 deletions(-) diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py index 4dad415b1..1f463430e 100644 --- a/python/sglang/srt/layers/attention/linear/gdn_backend.py +++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py @@ -255,6 +255,9 @@ class GDNAttnBackend(MambaAttnBackendBase): decode_backend = get_linear_attn_decode_backend() prefill_backend = get_linear_attn_prefill_backend() self.kernel_dispatcher = GDNKernelDispatcher(decode_backend, prefill_backend) + self.verify_intermediate_state_indices = torch.arange( + self.req_to_token_pool.size, dtype=torch.int32, device=model_runner.device + ) def init_forward_metadata(self, forward_batch: ForwardBatch): super().init_forward_metadata(forward_batch) @@ -373,14 +376,7 @@ class GDNAttnBackend(MambaAttnBackendBase): intermediate_conv_window_cache = ( mamba_cache_params.intermediate_conv_window[0] ) - has_initial_states = torch.ones( - seq_len // forward_batch.spec_info.draft_token_num, - dtype=torch.bool, - device=forward_batch.input_ids.device, - ) - intermediate_state_indices = torch.arange( - cache_indices.shape[0], dtype=torch.int32, device=cache_indices.device - ) + intermediate_state_indices = self.verify_intermediate_state_indices else: has_initial_states = forward_batch.extend_prefix_lens > 0