Avoid materializing GDN QKV tensors during target verification (#33778)

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
Vedant V Jhaveri
2026-09-21 17:04:18 -07:00
committed by GitHub
co-authored by Copilot
parent 506698761d
commit 9fdb71732a
6 changed files with 214 additions and 29 deletions
@@ -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)
@@ -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(
@@ -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),