Avoid materializing GDN QKV tensors during target verification (#33778)
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
co-authored by
Copilot
parent
506698761d
commit
9fdb71732a
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
Reference in New Issue
Block a user