diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index c40c08572..49165d6ed 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -308,6 +308,9 @@ class DeepseekSparseAttnBackend( # (page-table width) and never reads seq_lens_cpu / seq_lens_sum; opt out of # the D2H sync. The eager fallback derives lengths from GPU seq_lens. needs_cpu_seq_lens: bool = False + # init_cuda_graph_state sizes this for every backend, but only the TRT-LLM + # branch of __init__ allocates one. + _multi_ctas_kv_counter_buffer: Optional[torch.Tensor] = None def __init__( self, @@ -549,7 +552,6 @@ class DeepseekSparseAttnBackend( ) else: self.workspace_buffer = None - self._multi_ctas_kv_counter_buffer = None def _make_aiter_dsa_decode_metadata_buffer( self, @@ -1304,6 +1306,39 @@ class DeepseekSparseAttnBackend( ), } + # Sized by query rows, not requests: target verify captures + # speculative_num_draft_tokens rows per request. + self._ensure_multi_ctas_kv_counter_capacity(max(max_bs, max_num_tokens)) + + def _multi_ctas_kv_counter_for(self, num_query_rows: int) -> Optional[torch.Tensor]: + # A prefill batch wider than TRTLLM_MLA_MAX_BATCH_SIZE takes a temporary; + # rebinding would free the allocation the decode graphs captured. + counter = grow_multi_ctas_kv_counter_buffer_if_needed( + buffer=self._multi_ctas_kv_counter_buffer, + device=torch.device(self.device), + num_q_heads=self.num_q_heads, + batch_size=num_query_rows, + ) + # Capacity is set before capture, so a grow here is a broken invariant. + assert ( + counter is self._multi_ctas_kv_counter_buffer + or not torch.cuda.is_current_stream_capturing() + ), "multi_ctas_kv_counter_buffer grew during CUDA graph capture" + return counter + + def _ensure_multi_ctas_kv_counter_capacity(self, num_query_rows: int) -> None: + if self._multi_ctas_kv_counter_buffer is None: + return + # Must run before any capture: a later rebind frees what a graph replays. + self._multi_ctas_kv_counter_buffer = ( + grow_multi_ctas_kv_counter_buffer_if_needed( + buffer=self._multi_ctas_kv_counter_buffer, + device=torch.device(self.device), + num_q_heads=self.num_q_heads, + batch_size=num_query_rows, + ) + ) + def _build_forward_metadata_cuda_graph( self, bs: int, @@ -3426,14 +3461,7 @@ class DeepseekSparseAttnBackend( batch_size = page_table_1.shape[0] _, num_heads, head_dim = q_all.shape - self._multi_ctas_kv_counter_buffer = ( - grow_multi_ctas_kv_counter_buffer_if_needed( - self._multi_ctas_kv_counter_buffer, - torch.device(self.device), - self.num_q_heads, - batch_size, - ) - ) + multi_ctas_kv_counter_buffer = self._multi_ctas_kv_counter_for(batch_size) q = q_all.view(batch_size, 1, num_heads, head_dim) kv = kv_cache.view(-1, 1, self.real_page_size, self.kv_cache_dim) @@ -3455,7 +3483,7 @@ class DeepseekSparseAttnBackend( backend="trtllm-gen", skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), sparse_mla_top_k_lens=sparse_mla_top_k_lens, - multi_ctas_kv_counter_buffer=self._multi_ctas_kv_counter_buffer, + multi_ctas_kv_counter_buffer=multi_ctas_kv_counter_buffer, ) return out diff --git a/test/registered/kernels/ops/attention/test_dsa_multi_ctas_counter.py b/test/registered/kernels/ops/attention/test_dsa_multi_ctas_counter.py new file mode 100644 index 000000000..f7415a24f --- /dev/null +++ b/test/registered/kernels/ops/attention/test_dsa_multi_ctas_counter.py @@ -0,0 +1,139 @@ +"""Lifetime of the DSA multi-CTAs KV counter across CUDA graph capture. + +The decode graphs record this buffer's address, so nothing after capture may +reallocate it. _forward_trtllm, the production caller, needs a live FlashInfer +kernel and is not covered here. +""" + +import unittest +import weakref +from types import SimpleNamespace + +import torch + +from sglang.srt.layers.attention.dsa_backend import DeepseekSparseAttnBackend +from sglang.srt.layers.attention.trtllm_mla_backend import ( + TRTLLM_MLA_MAX_BATCH_SIZE, + grow_multi_ctas_kv_counter_buffer_if_needed, + make_persistent_multi_ctas_kv_counter_buffer, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="1-gpu-large") + +_NUM_Q_HEADS = 128 +_MAX_CTX_LEN = 64 +# 1024 captured requests x 9 draft tokens: a supported speculative capture whose +# query-row count exceeds TRTLLM_MLA_MAX_BATCH_SIZE. +_CAPTURED_BS = 1024 +_NUM_DRAFT_TOKENS = 9 +_NUM_CAPTURED_ROWS = _CAPTURED_BS * _NUM_DRAFT_TOKENS + + +def _make_backend(*, allocate_counter: bool = True): + backend = object.__new__(DeepseekSparseAttnBackend) + backend.device = "cuda" + backend.num_q_heads = _NUM_Q_HEADS + backend.real_page_size = 64 + backend.hisparse_coordinator = None + backend.speculative_num_draft_tokens = _NUM_DRAFT_TOKENS + backend.dsa_index_kpool = 1 + backend.use_fused_topk = False + backend.dsa_topk_backend = SimpleNamespace(should_use_topk_v2=lambda: False) + backend.dsa_index_topk = 2048 + backend.dsa_decode_impl = "trtllm" + backend.req_to_token = torch.zeros( + 8, _MAX_CTX_LEN, dtype=torch.int32, device="cuda" + ) + backend._multi_ctas_kv_counter_buffer = ( + make_persistent_multi_ctas_kv_counter_buffer( + device=torch.device("cuda"), + num_q_heads=_NUM_Q_HEADS, + max_batch_size=48, + ) + if allocate_counter + else None + ) + return backend + + +def _would_grow(backend, rows: int) -> bool: + return ( + grow_multi_ctas_kv_counter_buffer_if_needed( + buffer=backend._multi_ctas_kv_counter_buffer, + device=torch.device("cuda"), + num_q_heads=backend.num_q_heads, + batch_size=rows, + ) + is not backend._multi_ctas_kv_counter_buffer + ) + + +@unittest.skipUnless(torch.cuda.is_available(), "needs a CUDA device") +class TestMultiCtasKvCounterLifetime(CustomTestCase): + def test_request_sized_counter_would_grow_at_capture(self): + """The premise: sizing by requests undercounts captured query rows.""" + self.assertGreater(_NUM_CAPTURED_ROWS, TRTLLM_MLA_MAX_BATCH_SIZE) + self.assertTrue(_would_grow(_make_backend(), _NUM_CAPTURED_ROWS)) + + def test_init_cuda_graph_state_sizes_for_query_rows(self): + backend = _make_backend() + backend.init_cuda_graph_state( + max_bs=_CAPTURED_BS, max_num_tokens=_NUM_CAPTURED_ROWS + ) + self.assertFalse(_would_grow(backend, _NUM_CAPTURED_ROWS)) + + def test_init_cuda_graph_state_is_grow_only(self): + """A later, smaller graph must not discard an earlier allocation.""" + backend = _make_backend() + backend.init_cuda_graph_state( + max_bs=_CAPTURED_BS, max_num_tokens=_NUM_CAPTURED_ROWS + ) + sized = backend._multi_ctas_kv_counter_buffer + backend.init_cuda_graph_state(max_bs=8, max_num_tokens=64) + self.assertIs(backend._multi_ctas_kv_counter_buffer, sized) + + def test_init_cuda_graph_state_tolerates_backends_without_a_counter(self): + """Non-TRT-LLM branches leave the counter None; sizing must not read it.""" + backend = _make_backend(allocate_counter=False) + backend.init_cuda_graph_state( + max_bs=_CAPTURED_BS, max_num_tokens=_NUM_CAPTURED_ROWS + ) + self.assertIsNone(backend._multi_ctas_kv_counter_buffer) + + def test_counter_field_defaults_to_none_on_the_class(self): + """The sizing hook is unconditional, so every branch must leave it readable.""" + self.assertIsNone(DeepseekSparseAttnBackend._multi_ctas_kv_counter_buffer) + + def test_oversized_eager_call_keeps_the_captured_allocation(self): + """Holds only a weakref and an address, so a rebinding implementation + drops the last strong reference and the assertions see it.""" + backend = _make_backend() + backend.init_cuda_graph_state( + max_bs=_CAPTURED_BS, max_num_tokens=_NUM_CAPTURED_ROWS + ) + captured_ref = weakref.ref(backend._multi_ctas_kv_counter_buffer) + captured_ptr = backend._multi_ctas_kv_counter_buffer.data_ptr() + + counter = backend._multi_ctas_kv_counter_for(_NUM_CAPTURED_ROWS * 2) + self.assertIsNot(counter, backend._multi_ctas_kv_counter_buffer) + del counter + + self.assertIsNotNone(captured_ref()) + self.assertIs(backend._multi_ctas_kv_counter_buffer, captured_ref()) + self.assertEqual(backend._multi_ctas_kv_counter_buffer.data_ptr(), captured_ptr) + + def test_within_capacity_eager_call_reuses_the_captured_allocation(self): + backend = _make_backend() + backend.init_cuda_graph_state( + max_bs=_CAPTURED_BS, max_num_tokens=_NUM_CAPTURED_ROWS + ) + self.assertIs( + backend._multi_ctas_kv_counter_for(_NUM_CAPTURED_ROWS), + backend._multi_ctas_kv_counter_buffer, + ) + + +if __name__ == "__main__": + unittest.main()