Keep tiny CP cache-hit suffixes off async prefetch
The repeated-request hang reproduced on g0034 after the cached-prefix request created both MLA and index async prefetchers with prefix_lens=[40320] and extend_lens=[65]. The existing one-page default only blocked sub-page suffixes, so a barely-over-one-page suffix still entered the next-layer collective path before any later forward progress was logged. Raise the default async prefetch extend gate to one page per CP lane while keeping the env override. This only gates async prefetcher object creation; target partial-current reuse still uses the synchronous page-slot compose/current-splice path when no prefetcher exists. Constraint: cp_size=8,page_size=64 repeated prompt had extend_len=65 and hung immediately after has_mla=True has_index=True create_result logs Rejected: Disable partial-current reuse for short extends | that would lose the cache-hit benefit and regress current/full reuse Rejected: Disable all async prefetch by default | broader performance impact than the observed tiny-suffix failure Confidence: medium Scope-risk: moderate Directive: Do not lower the default below one page per CP lane without ETE proof that repeated cache-hit tiny suffixes no longer hang Tested: Remote g0034 container py_compile for touched runtime/prefetch/test files; targeted C22 tests passed 2 tests plus 2 subtests; full test_cp_shared_kv_runtime.py passed 73 tests plus 2 subtests Not-tested: Full multi-node ETE repeated-request run after this threshold change
This commit is contained in:
@@ -1404,7 +1404,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
32,
|
||||
)
|
||||
|
||||
def test_mla_prefetch_min_async_extend_tokens_defaults_to_one_page_and_can_override(
|
||||
def test_mla_prefetch_min_async_extend_tokens_defaults_to_one_page_per_cp_lane_and_can_override(
|
||||
self,
|
||||
):
|
||||
from sglang.srt.environ import envs
|
||||
@@ -1412,7 +1412,15 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
|
||||
envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_EXTEND_TOKENS.clear()
|
||||
self.assertEqual(
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(page_size=64),
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(
|
||||
cp_size=8, page_size=64
|
||||
),
|
||||
512,
|
||||
)
|
||||
self.assertEqual(
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(
|
||||
cp_size=None, page_size=64
|
||||
),
|
||||
64,
|
||||
)
|
||||
self.assertEqual(
|
||||
@@ -1422,13 +1430,17 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
|
||||
with envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_EXTEND_TOKENS.override(0):
|
||||
self.assertEqual(
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(page_size=64),
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(
|
||||
cp_size=8, page_size=64
|
||||
),
|
||||
0,
|
||||
)
|
||||
|
||||
with envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_EXTEND_TOKENS.override(128):
|
||||
self.assertEqual(
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(page_size=64),
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(
|
||||
cp_size=8, page_size=64
|
||||
),
|
||||
128,
|
||||
)
|
||||
|
||||
@@ -1439,16 +1451,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
def is_context_parallel_extend(self):
|
||||
return True
|
||||
|
||||
forward_batch = SimpleNamespace(
|
||||
uses_cp_shared_kv=True,
|
||||
hisparse_coordinator=None,
|
||||
forward_mode=Mode(),
|
||||
batch_size=1,
|
||||
token_to_kv_pool=SimpleNamespace(page_size=64, start_layer=0),
|
||||
cp_shared_kv_layout=SimpleNamespace(cp_size=8, cp_rank=0),
|
||||
extend_prefix_lens_cpu=[16320],
|
||||
extend_seq_lens_cpu=[16],
|
||||
)
|
||||
metadata = SimpleNamespace(
|
||||
real_page_table=torch.arange(256, dtype=torch.int64),
|
||||
page_table_1=torch.zeros((1, 16336), dtype=torch.int32),
|
||||
@@ -1461,45 +1463,62 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
dense_num_pages=257,
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
prefetch, "cp_shared_kv_mla_prefetch_enabled", return_value=True
|
||||
), patch.object(
|
||||
prefetch, "cp_shared_kv_debug_enabled", return_value=False
|
||||
), patch.object(
|
||||
prefetch.torch.cuda, "is_available", return_value=True
|
||||
), patch.object(
|
||||
prefetch, "_is_cuda_stream_capturing", return_value=False
|
||||
), patch.object(
|
||||
prefetch, "is_nsa_prefill_cp_in_seq_split", return_value=True
|
||||
), patch.object(
|
||||
prefetch, "get_attention_cp_group", return_value=SimpleNamespace(pynccl_comm=object())
|
||||
), patch.object(
|
||||
prefetch.torch.cuda, "Stream", return_value=stream
|
||||
), patch.object(
|
||||
prefetch, "_prefetch_pool_get_key_buffer", return_value=kv_cache
|
||||
) as mla_getter, patch.object(
|
||||
prefetch, "get_or_build_shared_token_kv_slot_remap", return_value=remap
|
||||
) as token_remap, patch.object(
|
||||
prefetch,
|
||||
"_prefetch_pool_get_index_buffer",
|
||||
side_effect=AssertionError("index getter should not be reached"),
|
||||
) as index_getter:
|
||||
mla_result = prefetch.CpSharedKVMlaPrefetcher.maybe_create(
|
||||
forward_batch=forward_batch,
|
||||
metadata=metadata,
|
||||
topk_transform_is_paged=True,
|
||||
)
|
||||
index_result = prefetch.CpSharedKVIndexPrefetcher.maybe_create(
|
||||
forward_batch=forward_batch,
|
||||
metadata=metadata,
|
||||
topk_transform_is_paged=True,
|
||||
)
|
||||
for extend_len in (16, 65):
|
||||
with self.subTest(extend_len=extend_len):
|
||||
forward_batch = SimpleNamespace(
|
||||
uses_cp_shared_kv=True,
|
||||
hisparse_coordinator=None,
|
||||
forward_mode=Mode(),
|
||||
batch_size=1,
|
||||
token_to_kv_pool=SimpleNamespace(page_size=64, start_layer=0),
|
||||
cp_shared_kv_layout=SimpleNamespace(cp_size=8, cp_rank=0),
|
||||
extend_prefix_lens_cpu=[16320],
|
||||
extend_seq_lens_cpu=[extend_len],
|
||||
)
|
||||
|
||||
self.assertIsNone(mla_result)
|
||||
self.assertIsNone(index_result)
|
||||
mla_getter.assert_not_called()
|
||||
token_remap.assert_not_called()
|
||||
index_getter.assert_not_called()
|
||||
with patch.object(
|
||||
prefetch, "cp_shared_kv_mla_prefetch_enabled", return_value=True
|
||||
), patch.object(
|
||||
prefetch, "cp_shared_kv_debug_enabled", return_value=False
|
||||
), patch.object(
|
||||
prefetch.torch.cuda, "is_available", return_value=True
|
||||
), patch.object(
|
||||
prefetch, "_is_cuda_stream_capturing", return_value=False
|
||||
), patch.object(
|
||||
prefetch, "is_nsa_prefill_cp_in_seq_split", return_value=True
|
||||
), patch.object(
|
||||
prefetch,
|
||||
"get_attention_cp_group",
|
||||
return_value=SimpleNamespace(pynccl_comm=object()),
|
||||
), patch.object(
|
||||
prefetch.torch.cuda, "Stream", return_value=stream
|
||||
), patch.object(
|
||||
prefetch, "_prefetch_pool_get_key_buffer", return_value=kv_cache
|
||||
) as mla_getter, patch.object(
|
||||
prefetch,
|
||||
"get_or_build_shared_token_kv_slot_remap",
|
||||
return_value=remap,
|
||||
) as token_remap, patch.object(
|
||||
prefetch,
|
||||
"_prefetch_pool_get_index_buffer",
|
||||
side_effect=AssertionError("index getter should not be reached"),
|
||||
) as index_getter:
|
||||
mla_result = prefetch.CpSharedKVMlaPrefetcher.maybe_create(
|
||||
forward_batch=forward_batch,
|
||||
metadata=metadata,
|
||||
topk_transform_is_paged=True,
|
||||
)
|
||||
index_result = prefetch.CpSharedKVIndexPrefetcher.maybe_create(
|
||||
forward_batch=forward_batch,
|
||||
metadata=metadata,
|
||||
topk_transform_is_paged=True,
|
||||
)
|
||||
|
||||
self.assertIsNone(mla_result)
|
||||
self.assertIsNone(index_result)
|
||||
mla_getter.assert_not_called()
|
||||
token_remap.assert_not_called()
|
||||
index_getter.assert_not_called()
|
||||
|
||||
def test_fused_mla_store_uses_tai_kernel_when_enabled(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
|
||||
Reference in New Issue
Block a user