From 9fdb71732ae323ff590ecd69043dd3fd4cc6e510 Mon Sep 17 00:00:00 2001 From: Vedant V Jhaveri Date: Mon, 21 Sep 2026 17:04:18 -0700 Subject: [PATCH] Avoid materializing GDN QKV tensors during target verification (#33778) Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../layers/attention/linear/gdn_backend.py | 124 ++++++++++++++---- .../attention/linear/kernels/gdn_triton.py | 1 + .../linear/kernels/kernel_backend.py | 2 + .../gdn/test_gdn_replayssm_spec_fold.py | 36 +++++ .../attention/unittests/gdn/test_triton.py | 21 ++- .../test_gdn_prefill_backend_policy.py | 59 ++++++++- 6 files changed, 214 insertions(+), 29 deletions(-) diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py index 21148e6b0..15ac8abac 100644 --- a/python/sglang/srt/layers/attention/linear/gdn_backend.py +++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py @@ -476,13 +476,8 @@ class GDNKernelDispatcher: query_start_loc: torch.Tensor, **kwargs, ) -> torch.Tensor: - # FlashInfer verify supports a linear MTP chain. Tree-shaped drafts - # carry parent indices and must use Triton even when decode/prefill use - # FlashInfer. - verify_kernel = ( - self.tree_verify_kernel - if kwargs.get("retrieve_parent_token") is not None - else self.verify_kernel + verify_kernel = self._get_target_verify_kernel( + kwargs.get("retrieve_parent_token") ) return verify_kernel.target_verify( A_log=A_log, @@ -498,6 +493,22 @@ class GDNKernelDispatcher: **kwargs, ) + def target_verify_supports_strided_qkv( + self, retrieve_parent_token: Optional[torch.Tensor] + ) -> bool: + verify_kernel = self._get_target_verify_kernel(retrieve_parent_token) + return ( + getattr(verify_kernel, "supports_strided_target_verify_qkv", False) is True + ) + + def _get_target_verify_kernel(self, retrieve_parent_token: Optional[torch.Tensor]): + # Tree drafts use Triton even when linear MTP verification uses FlashInfer. + return ( + self.tree_verify_kernel + if retrieve_parent_token is not None + else self.verify_kernel + ) + class GDNAttnBackend(MambaAttnBackendBase): """Attention backend for GDN (Gated Delta Network) linear attention.""" @@ -531,9 +542,19 @@ class GDNAttnBackend(MambaAttnBackendBase): model_runner.device, ) ) + self._use_strided_target_verify_qkv = False + + def init_forward_metadata_out_graph( + self, + forward_batch: ForwardBatch, + in_capture: bool = False, + ): + super().init_forward_metadata_out_graph(forward_batch, in_capture=in_capture) + self._init_target_verify_qkv_routing(forward_batch) def init_forward_metadata(self, forward_batch: ForwardBatch): super().init_forward_metadata(forward_batch) + self._init_target_verify_qkv_routing(forward_batch) self.mis_metadata = None if forward_batch.multi_item_delimiter_indices is not None: if not self.enable_mis: @@ -558,6 +579,61 @@ class GDNAttnBackend(MambaAttnBackendBase): forward_batch, self.forward_metadata, self.device ) + def _init_target_verify_qkv_routing(self, forward_batch: ForwardBatch) -> None: + # CUDA-graph metadata leaves mode and draft count at their defaults. + if not forward_batch.forward_mode.is_target_verify(): + self._use_strided_target_verify_qkv = False + return + + metadata = self.forward_metadata + mamba_pool = self.req_to_token_pool.mamba_pool + mamba_cache = mamba_pool.mamba_cache + is_gdn_replayssm = not getattr(mamba_pool, "replayssm_is_kda", False) + use_replayssm_fold = ( + mamba_cache.replayssm_rawv is not None + and getattr(mamba_pool, "replayssm_spec_fold", False) + and is_gdn_replayssm + ) + use_replayssm_spec = ( + mamba_cache.replayssm_d is not None + and getattr(mamba_pool, "replayssm_cache_base", None) is not None + and is_gdn_replayssm + ) + self._use_strided_target_verify_qkv = self._target_verify_supports_strided_qkv( + retrieve_parent_token=metadata.retrieve_parent_token, + use_replayssm_fold=use_replayssm_fold, + use_replayssm_spec=use_replayssm_spec, + ssm_dtype=mamba_cache.temporal.dtype, + draft_token_num=forward_batch.spec_info.draft_token_num, + ) + + def _replayssm_fold_uses_cutedsl( + self, ssm_dtype: torch.dtype, draft_token_num: int + ) -> bool: + return ( + self.kernel_dispatcher.verify_kernel_is_flashinfer + and ssm_dtype == torch.bfloat16 + and draft_token_num >= 3 + ) + + def _target_verify_supports_strided_qkv( + self, + *, + retrieve_parent_token: Optional[torch.Tensor], + use_replayssm_fold: bool, + use_replayssm_spec: bool, + ssm_dtype: torch.dtype, + draft_token_num: int, + ) -> bool: + # ReplaySSM Triton routes accept strides; the CuTeDSL fold does not. + if use_replayssm_fold: + return not self._replayssm_fold_uses_cutedsl(ssm_dtype, draft_token_num) + if use_replayssm_spec: + return True + return self.kernel_dispatcher.target_verify_supports_strided_qkv( + retrieve_parent_token + ) + def forward_decode( self, layer: RadixLinearAttention, @@ -851,6 +927,17 @@ class GDNAttnBackend(MambaAttnBackendBase): mamba_cache_params.intermediate_conv_window[0] ) intermediate_state_indices = self.verify_intermediate_state_indices + mamba_pool = self.req_to_token_pool.mamba_pool + use_replayssm_fold = ( + mamba_cache_params.replayssm_rawv is not None + and getattr(mamba_pool, "replayssm_spec_fold", False) + and not getattr(mamba_pool, "replayssm_is_kda", False) + ) + use_replayssm_spec = ( + mamba_cache_params.replayssm_d is not None + and getattr(mamba_pool, "replayssm_cache_base", None) is not None + and not getattr(mamba_pool, "replayssm_is_kda", False) + ) else: has_initial_states = forward_batch.extend_prefix_lens > 0 @@ -927,7 +1014,11 @@ class GDNAttnBackend(MambaAttnBackendBase): actual_seq_len = mixed_qkv.shape[0] qkv_dim = layer.q_dim + layer.k_dim + layer.v_dim - if (is_cuda() or is_hip() or is_xpu()) and qkv_dim <= MAX_FUSED_QKV_SPLIT_DIM: + if ( + (is_cuda() or is_hip() or is_xpu()) + and qkv_dim <= MAX_FUSED_QKV_SPLIT_DIM + and not self._use_strided_target_verify_qkv + ): query, key, value = fused_qkv_split_gdn_prefill( mixed_qkv, layer.num_q_heads, @@ -951,17 +1042,6 @@ class GDNAttnBackend(MambaAttnBackendBase): # ReplaySSM verify protocols: fold-every-commit (ring-write during # verify, fold on commit), circular ring, or the snapshotting # fallback when neither ring is allocated. - mamba_pool = self.req_to_token_pool.mamba_pool - use_replayssm_fold = ( - mamba_cache_params.replayssm_rawv is not None - and getattr(mamba_pool, "replayssm_spec_fold", False) - and not getattr(mamba_pool, "replayssm_is_kda", False) - ) - use_replayssm_spec = ( - mamba_cache_params.replayssm_d is not None - and getattr(mamba_pool, "replayssm_cache_base", None) is not None - and not getattr(mamba_pool, "replayssm_is_kda", False) - ) if use_replayssm_fold: core_attn_out = self._replayssm_fold_target_verify( layer=layer, @@ -1219,11 +1299,7 @@ class GDNAttnBackend(MambaAttnBackendBase): seq_len = query.shape[1] batch_size = query_start_loc.shape[0] - 1 draft_token_num = seq_len // batch_size - if ( - self.kernel_dispatcher.verify_kernel_is_flashinfer - and ssm_states.dtype == torch.bfloat16 - and draft_token_num >= 3 - ): + if self._replayssm_fold_uses_cutedsl(ssm_states.dtype, draft_token_num): from sglang.kernels.ops.attention.cutedsl_gdn_mtp_ring import ( gated_delta_rule_mtp, ) diff --git a/python/sglang/srt/layers/attention/linear/kernels/gdn_triton.py b/python/sglang/srt/layers/attention/linear/kernels/gdn_triton.py index b16c41ca8..8660a2c23 100644 --- a/python/sglang/srt/layers/attention/linear/kernels/gdn_triton.py +++ b/python/sglang/srt/layers/attention/linear/kernels/gdn_triton.py @@ -42,6 +42,7 @@ class TritonGDNKernel(LinearAttnKernelBase): """Triton-based kernel for GDN (Gated Delta Network) linear attention.""" supports_packed_decode: bool = not is_cpu() and not is_npu() + supports_strided_target_verify_qkv: bool = True def packed_decode( self, diff --git a/python/sglang/srt/layers/attention/linear/kernels/kernel_backend.py b/python/sglang/srt/layers/attention/linear/kernels/kernel_backend.py index ec5ae6463..eb008a28b 100644 --- a/python/sglang/srt/layers/attention/linear/kernels/kernel_backend.py +++ b/python/sglang/srt/layers/attention/linear/kernels/kernel_backend.py @@ -12,6 +12,8 @@ class LinearAttnKernelBase(ABC): uses_state_checkpoints: bool = False supports_fused_chain_verify: bool = False + # Opt in only when target-verify kernels honor non-unit token strides. + supports_strided_target_verify_qkv: bool = False # True when extend() honors the fp32 track snapshot (track_state / # track_chunk_idx), natively or by routing tracked batches to a kernel diff --git a/test/registered/attention/unittests/gdn/test_gdn_replayssm_spec_fold.py b/test/registered/attention/unittests/gdn/test_gdn_replayssm_spec_fold.py index 51f253fe4..f5e531e22 100644 --- a/test/registered/attention/unittests/gdn/test_gdn_replayssm_spec_fold.py +++ b/test/registered/attention/unittests/gdn/test_gdn_replayssm_spec_fold.py @@ -137,6 +137,42 @@ class TestGdnReplayssmSpecFold(CustomTestCase): ) self.assertTrue(torch.equal(out_plain, out_ring), f"{dtype=}") + def test_ring_write_accepts_strided_qkv_views(self): + inputs = _make_window(12) + packed_qkv = torch.cat( + [inputs[name].reshape(B * T, -1) for name in ("q", "k", "v")], + dim=-1, + ) + q, k, v = packed_qkv.split([H * K, H * K, HV * V], dim=-1) + strided_qkv = { + "q": q.view(1, B * T, H, K), + "k": k.view(1, B * T, H, K), + "v": v.view(1, B * T, HV, V), + } + self.assertTrue( + all(not tensor.is_contiguous() for tensor in strided_qkv.values()) + ) + + def run(qkv): + state = self._state(torch.float32).unsqueeze(0).contiguous() + rings = _make_rings() + output = _run_verify( + {**inputs, **qkv}, + self.gating, + state[0], + self.slots, + rings=rings, + ) + _fold(state, rings, self.slots, self.accept_lens) + return {"output": output, "state": state, **rings} + + contiguous = run( + {name: tensor.contiguous() for name, tensor in strided_qkv.items()} + ) + strided = run(strided_qkv) + for name in contiguous: + self.assertTrue(torch.equal(contiguous[name], strided[name]), name) + def test_fold_matches_snapshot_baseline(self): for dtype in (torch.float32, torch.bfloat16): state = self._state(dtype) diff --git a/test/registered/attention/unittests/gdn/test_triton.py b/test/registered/attention/unittests/gdn/test_triton.py index 24be44741..129e76a5c 100644 --- a/test/registered/attention/unittests/gdn/test_triton.py +++ b/test/registered/attention/unittests/gdn/test_triton.py @@ -1,6 +1,6 @@ import unittest from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch import torch @@ -8,6 +8,7 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( HybridLinearAttnBackend, MambaAttnBackendBase, ) +from sglang.srt.layers.attention.linear import gdn_backend from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.utils import is_hip from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci @@ -316,6 +317,24 @@ class TestTritonGDNBackendCorrectness(CustomTestCase): ): run_gdn_eagle_verify_case(self, case, topk=topk, spec_kind=spec_kind) + def test_triton_target_verify_skips_prefill_qkv_materialization(self): + case, topk, spec_kind = self.EAGLE_VERIFY_CASES[0] + with patch.object( + gdn_backend, + "fused_qkv_split_gdn_prefill", + side_effect=AssertionError, + ): + run_gdn_eagle_verify_case(self, case, topk=topk, spec_kind=spec_kind) + + def test_triton_prefill_keeps_contiguous_qkv_materialization(self): + with patch.object( + gdn_backend, + "fused_qkv_split_gdn_prefill", + wraps=gdn_backend.fused_qkv_split_gdn_prefill, + ) as split_spy: + run_gdn_attention_case(self, self.CASES[0]) + split_spy.assert_called() + def test_runner_mode_eagle_verify_cuda_graph_cases(self): for case, topk, spec_kind in self.EAGLE_VERIFY_CUDA_GRAPH_CASES: with self.subTest( diff --git a/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py b/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py index 012bc8efe..182ee17a8 100644 --- a/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py +++ b/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py @@ -4,9 +4,7 @@ from unittest.mock import MagicMock, patch, sentinel import torch -from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( - MambaAttnBackendBase, -) +from sglang.srt.layers.attention.hybrid_linear_attn_backend import MambaAttnBackendBase from sglang.srt.layers.attention.linear import gdn_backend from sglang.srt.layers.attention.linear.gdn_backend import ( GDNAttnBackend, @@ -81,6 +79,16 @@ def make_runner( class TestFlashInferGDNPrefillBackendPolicy(CustomTestCase): + @staticmethod + def make_target_verify_routing_backend(): + backend = object.__new__(GDNAttnBackend) + nominal_capability = MagicMock(return_value=False) + backend.kernel_dispatcher = SimpleNamespace( + verify_kernel_is_flashinfer=True, + target_verify_supports_strided_qkv=nominal_capability, + ) + return backend, nominal_capability + def test_mis_requires_triton_prefill_backend(self): runner = make_runner(self, enable_mis=True) with self.assertRaisesRegex(ValueError, "Triton linear-attention prefill"): @@ -250,6 +258,7 @@ class TestFlashInferGDNPrefillBackendPolicy(CustomTestCase): backend.kernel_dispatcher = SimpleNamespace(extend_uses_state_checkpoints=True) metadata = SimpleNamespace(has_mamba_track_mask=True, track_ssm_h_src=None) forward_batch = SimpleNamespace( + forward_mode=SimpleNamespace(is_target_verify=lambda: False), multi_item_delimiter_indices=None, mamba_track_mask=torch.tensor([True]), mamba_track_indices=torch.tensor([7]), @@ -264,7 +273,10 @@ class TestFlashInferGDNPrefillBackendPolicy(CustomTestCase): torch.testing.assert_close(metadata.conv_states_mask_indices, torch.tensor([7])) def test_tree_verify_uses_triton_kernel(self): - flashinfer_kernel = MagicMock(supports_target_verify=True) + flashinfer_kernel = MagicMock( + supports_target_verify=True, + supports_strided_target_verify_qkv=False, + ) with ( patch.object(gdn_backend, "is_cuda", return_value=True), patch( @@ -279,6 +291,10 @@ class TestFlashInferGDNPrefillBackendPolicy(CustomTestCase): ) self.assertIsInstance(dispatcher.tree_verify_kernel, TritonGDNKernel) + self.assertFalse(dispatcher.target_verify_supports_strided_qkv(None)) + self.assertTrue( + dispatcher.target_verify_supports_strided_qkv(sentinel.parent_token) + ) tensor = sentinel.tensor with patch.object( @@ -295,6 +311,41 @@ class TestFlashInferGDNPrefillBackendPolicy(CustomTestCase): tree_verify.assert_called_once() flashinfer_kernel.target_verify.assert_not_called() + def test_target_verify_strided_input_capability_is_opt_in(self): + dispatcher = GDNKernelDispatcher( + LinearAttnKernelBackend.TRITON, + LinearAttnKernelBackend.TRITON, + ) + + self.assertTrue(dispatcher.target_verify_supports_strided_qkv(None)) + dispatcher.verify_kernel = SimpleNamespace() + self.assertFalse(dispatcher.target_verify_supports_strided_qkv(None)) + + def test_target_verify_strided_qkv_routing(self): + backend, nominal_capability = self.make_target_verify_routing_backend() + cases = ( + ("fp32_fold", True, False, torch.float32, 4, True, False), + ("short_bf16_fold", True, False, torch.bfloat16, 2, True, False), + ("cutedsl_fold", True, False, torch.bfloat16, 4, False, False), + ("circular", False, True, torch.bfloat16, 4, True, False), + ("nominal", False, False, torch.bfloat16, 4, False, True), + ) + for name, fold, circular, dtype, draft_tokens, expected, delegates in cases: + with self.subTest(name=name): + nominal_capability.reset_mock() + actual = backend._target_verify_supports_strided_qkv( + retrieve_parent_token=None, + use_replayssm_fold=fold, + use_replayssm_spec=circular, + ssm_dtype=dtype, + draft_token_num=draft_tokens, + ) + self.assertEqual(actual, expected) + if delegates: + nominal_capability.assert_called_once_with(None) + else: + nominal_capability.assert_not_called() + def test_helion_backend_reports_kda_only(self): cases = ( (LinearAttnKernelBackend.HELION, LinearAttnKernelBackend.TRITON),