[Fix] Don't free the multi-CTAs KV counter the decode graphs captured (#39175)
Co-authored-by: mmangkad <mohammad.angkad@radixark.ai> Co-authored-by: kpham-sgl <khoa.pham@radixark.ai> Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
co-authored by
mmangkad
kpham-sgl
Claude Fable 5.1
parent
44bdf225d8
commit
11e661fd45
@@ -308,6 +308,9 @@ class DeepseekSparseAttnBackend(
|
|||||||
# (page-table width) and never reads seq_lens_cpu / seq_lens_sum; opt out of
|
# (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.
|
# the D2H sync. The eager fallback derives lengths from GPU seq_lens.
|
||||||
needs_cpu_seq_lens: bool = False
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -549,7 +552,6 @@ class DeepseekSparseAttnBackend(
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.workspace_buffer = None
|
self.workspace_buffer = None
|
||||||
self._multi_ctas_kv_counter_buffer = None
|
|
||||||
|
|
||||||
def _make_aiter_dsa_decode_metadata_buffer(
|
def _make_aiter_dsa_decode_metadata_buffer(
|
||||||
self,
|
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(
|
def _build_forward_metadata_cuda_graph(
|
||||||
self,
|
self,
|
||||||
bs: int,
|
bs: int,
|
||||||
@@ -3426,14 +3461,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
batch_size = page_table_1.shape[0]
|
batch_size = page_table_1.shape[0]
|
||||||
_, num_heads, head_dim = q_all.shape
|
_, num_heads, head_dim = q_all.shape
|
||||||
|
|
||||||
self._multi_ctas_kv_counter_buffer = (
|
multi_ctas_kv_counter_buffer = self._multi_ctas_kv_counter_for(batch_size)
|
||||||
grow_multi_ctas_kv_counter_buffer_if_needed(
|
|
||||||
self._multi_ctas_kv_counter_buffer,
|
|
||||||
torch.device(self.device),
|
|
||||||
self.num_q_heads,
|
|
||||||
batch_size,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
q = q_all.view(batch_size, 1, num_heads, head_dim)
|
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)
|
kv = kv_cache.view(-1, 1, self.real_page_size, self.kv_cache_dim)
|
||||||
@@ -3455,7 +3483,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
backend="trtllm-gen",
|
backend="trtllm-gen",
|
||||||
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
|
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
|
||||||
sparse_mla_top_k_lens=sparse_mla_top_k_lens,
|
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
|
return out
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user