[P/D disagg] Decode-side radix cache for SWA hybrid models (unified radix tree) (#27770)
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
Shangming Cai
parent
aa3f766799
commit
978244d671
@@ -48,11 +48,13 @@ def _has_mooncake():
|
||||
class DisaggregationDecodeRadixCacheTestMixin:
|
||||
extra_decode_args = ["--disaggregation-decode-enable-radix-cache"]
|
||||
transfer_backend_name = None
|
||||
model_name = DEFAULT_MODEL_NAME_FOR_TEST
|
||||
gsm8k_min_score = 0.80
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
cls.model = try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST)
|
||||
cls.model = try_cached_model(cls.model_name)
|
||||
cls.transfer_backend = [
|
||||
"--disaggregation-transfer-backend",
|
||||
cls.transfer_backend_name,
|
||||
@@ -117,8 +119,8 @@ class DisaggregationDecodeRadixCacheTestMixin:
|
||||
metrics_second = run_eval(args)
|
||||
print(f"Second run metrics: {metrics_second}")
|
||||
|
||||
self.assertGreater(metrics_first["score"], 0.80)
|
||||
self.assertGreater(metrics_second["score"], 0.80)
|
||||
self.assertGreater(metrics_first["score"], self.gsm8k_min_score)
|
||||
self.assertGreater(metrics_second["score"], self.gsm8k_min_score)
|
||||
|
||||
accuracy_drop = metrics_first["score"] - metrics_second["score"]
|
||||
self.assertLessEqual(
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
"""SWA coverage for decode-side radix cache on gpt-oss-20b.
|
||||
|
||||
The decode worker reuses full-attention prefix KV while transferring the SWA
|
||||
window fresh per request. This path requires the unified radix tree and validates
|
||||
both multi-turn cache hits and two-pass GSM8K accuracy.
|
||||
"""
|
||||
|
||||
import unittest
|
||||
|
||||
from test_disaggregation_decode_radix_cache import (
|
||||
DisaggregationDecodeRadixCacheTestMixin,
|
||||
_has_nixl,
|
||||
)
|
||||
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||
PDDisaggregationServerBase,
|
||||
)
|
||||
from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST_MXFP4_WITH_MOE, is_in_ci
|
||||
|
||||
register_cuda_ci(est_time=600, stage="extra-b", runner_config="8-gpu-h200")
|
||||
|
||||
SWA_SERVER_ARGS = ["--page-size", "64", "--attention-backend", "triton"]
|
||||
|
||||
|
||||
@unittest.skipUnless(
|
||||
is_in_ci() or _has_nixl(),
|
||||
"NIXL is required for decode radix cache disaggregation coverage.",
|
||||
)
|
||||
class TestDisaggregationDecodeRadixCacheSWANixl(
|
||||
DisaggregationDecodeRadixCacheTestMixin, PDDisaggregationServerBase
|
||||
):
|
||||
transfer_backend_name = "nixl"
|
||||
model_name = DEFAULT_MODEL_NAME_FOR_TEST_MXFP4_WITH_MOE
|
||||
# The 512-token eval cap truncates mxfp4 gpt-oss reasoning. On the fixed
|
||||
# 500-question H200 sample, the score has a roughly 2-point standard error,
|
||||
# so keep the original 0.45 absolute floor and rely on the two-pass
|
||||
# non-regression check below to catch decode-cache corruption.
|
||||
gsm8k_min_score = 0.45
|
||||
# SWA + decode-side radix cache is gated to the unified radix tree.
|
||||
extra_prefill_env = {"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"}
|
||||
extra_decode_env = {"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"}
|
||||
extra_prefill_args = SWA_SERVER_ARGS
|
||||
extra_decode_args = [
|
||||
"--disaggregation-decode-enable-radix-cache",
|
||||
*SWA_SERVER_ARGS,
|
||||
]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -37,10 +37,10 @@ class TestMamba2ExtraBufferKL(KLDivergenceMixin, DefaultServerBase):
|
||||
# Decode-seeded reuse is the regression trigger (the graphed decode
|
||||
# track-save); the broken path fails at KL ~1.5, so 0.005 discriminates
|
||||
# cleanly while absorbing bf16 reuse noise. Prefill reuse (chunk-aligned
|
||||
# intermediate h states) is inherently looser; threshold matches the manual
|
||||
# TestNvidiaNemotronNanoV2BF16ExtraBuffer calibration.
|
||||
# intermediate h states) is inherently looser; 0.012 retains a wide margin
|
||||
# below the broken path while covering the observed bf16 calibration noise.
|
||||
kl_div_thres = 0.005
|
||||
kl_div_thres_prefill = 0.01
|
||||
kl_div_thres_prefill = 0.012
|
||||
kl_div_max_samples = 16
|
||||
|
||||
other_args = [
|
||||
|
||||
@@ -72,6 +72,9 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
queue._swa_aware_allocatable_token_budgets = MagicMock(
|
||||
return_value=(physical_available, physical_available)
|
||||
)
|
||||
queue._swa_tail_allocatable_token_budget = MagicMock(
|
||||
side_effect=lambda **_: physical_available
|
||||
)
|
||||
|
||||
def pre_alloc(_req):
|
||||
nonlocal physical_available
|
||||
@@ -170,6 +173,72 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
self.assertTrue(all(r is not decode_req for r in queue.pending_reqs))
|
||||
self.assertIsNone(decode_req.kv_receiver)
|
||||
|
||||
def test_swa_reclaim_failure_rejects_only_request(self):
|
||||
receiver = FakeReceiver()
|
||||
req = SimpleNamespace(
|
||||
rid="swa-reclaim-failed",
|
||||
origin_input_ids=[1, 2, 3],
|
||||
output_ids=[],
|
||||
finished_reason=None,
|
||||
return_logprob=False,
|
||||
sampling_params=SimpleNamespace(max_new_tokens=1),
|
||||
)
|
||||
decode_req = SimpleNamespace(
|
||||
req=req,
|
||||
kv_receiver=receiver,
|
||||
waiting_for_input=True,
|
||||
is_rebootstrap=False,
|
||||
)
|
||||
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue.pp_size = 1
|
||||
queue.queue = [decode_req]
|
||||
queue.pending_reqs = [decode_req]
|
||||
queue.retracted_queue = []
|
||||
queue.num_reserved_decode_tokens = 0
|
||||
queue._resolve_pending_reqs = MagicMock()
|
||||
queue._update_handshake_waiters = MagicMock()
|
||||
queue._uses_swa_tail_prealloc = MagicMock(return_value=True)
|
||||
queue._swa_aware_allocatable_token_budgets = MagicMock(
|
||||
return_value=(1024, 1024)
|
||||
)
|
||||
queue._prealloc_required_tokens = MagicMock(return_value=(3, 3))
|
||||
queue._prealloc_kv_lens = MagicMock(return_value=(3, 3))
|
||||
queue._reclaim_swa_tail_capacity = MagicMock(
|
||||
return_value=(
|
||||
"SWA eviction insufficient: needed=64, available=0, "
|
||||
"req=swa-reclaim-failed"
|
||||
)
|
||||
)
|
||||
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
|
||||
queue._pre_alloc = MagicMock()
|
||||
queue.req_to_token_pool = MagicMock()
|
||||
queue.req_to_token_pool.available_size.return_value = 1
|
||||
queue.req_to_metadata_buffer_idx_allocator = MagicMock()
|
||||
queue.req_to_metadata_buffer_idx_allocator.available_size.return_value = 1
|
||||
|
||||
scheduler = MagicMock()
|
||||
scheduler.running_batch.reqs = []
|
||||
scheduler.enable_priority_scheduling = False
|
||||
scheduler.enable_hisparse = False
|
||||
scheduler.server_args.disaggregation_decode_enable_radix_cache = False
|
||||
scheduler.output_streamer = MagicMock()
|
||||
queue.scheduler = scheduler
|
||||
|
||||
preallocated, failed = queue.pop_preallocated()
|
||||
|
||||
self.assertEqual(preallocated, [])
|
||||
self.assertEqual(failed, [decode_req])
|
||||
self.assertEqual(queue.queue, [])
|
||||
self.assertEqual(queue.pending_reqs, [])
|
||||
self.assertTrue(receiver.clear_called)
|
||||
self.assertIsNone(decode_req.kv_receiver)
|
||||
self.assertIsInstance(req.finished_reason, FINISH_ABORT)
|
||||
queue._pre_alloc.assert_not_called()
|
||||
scheduler.output_streamer.stream_output.assert_called_once_with(
|
||||
[req], req.return_logprob
|
||||
)
|
||||
|
||||
def test_ensure_prefill_info_tolerates_cleared_receiver(self):
|
||||
# A req whose kv_receiver was already cleared must not crash on .abort().
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
|
||||
@@ -34,6 +34,7 @@ import torch
|
||||
from sglang.srt.disaggregation.decode import DecodePreallocQueue
|
||||
from sglang.srt.disaggregation.decode_hicache_mixin import DecodePrefixMatch
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
DecLockRefParams,
|
||||
InsertParams,
|
||||
MatchPrefixParams,
|
||||
)
|
||||
@@ -95,6 +96,71 @@ def _make_req(fill_ids, req_pool_idx=0, cache_protected_len=0, last_node=None):
|
||||
class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
"""Test lock_ref balance across decode transfer scenarios."""
|
||||
|
||||
def test_swa_tail_len_keeps_page_aligned_matchable_window(self):
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue._uses_swa_tail_prealloc = MagicMock(return_value=True)
|
||||
queue.scheduler = SimpleNamespace(
|
||||
sliding_window_size=127,
|
||||
server_args=SimpleNamespace(disaggregation_decode_enable_radix_cache=True),
|
||||
)
|
||||
queue.token_to_kv_pool_allocator = MagicMock(page_size=64)
|
||||
|
||||
tail_len = queue._swa_tail_len(895)
|
||||
|
||||
self.assertEqual(tail_len, 191)
|
||||
swa_start = 895 - tail_len
|
||||
radix_key_len = (895 // 64) * 64
|
||||
self.assertGreaterEqual(radix_key_len - swa_start, 127)
|
||||
|
||||
def test_swa_admission_counts_evictable_capacity(self):
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue.scheduler = MagicMock()
|
||||
queue.scheduler.running_batch.reqs = []
|
||||
queue.scheduler.sliding_window_size = 128
|
||||
queue.scheduler.last_batch = None
|
||||
queue.retracted_queue = []
|
||||
queue._need_space_for_single_req = MagicMock(return_value=0)
|
||||
queue._active_req_count = MagicMock(return_value=1)
|
||||
queue.token_to_kv_pool_allocator = MagicMock()
|
||||
queue.token_to_kv_pool_allocator.size_swa = 256
|
||||
queue.token_to_kv_pool_allocator.swa_available_size.return_value = 0
|
||||
queue.tree_cache = MagicMock()
|
||||
queue.tree_cache.swa_evictable_size.return_value = 192
|
||||
|
||||
budget = queue._swa_tail_allocatable_token_budget(
|
||||
count_retracted=False,
|
||||
reserved_tokens=64,
|
||||
)
|
||||
|
||||
# 192 reclaimable tokens minus 64 reserved for active-request growth.
|
||||
self.assertEqual(budget, 128)
|
||||
|
||||
def test_reclaim_swa_tail_capacity_page_rounds(self):
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue.token_to_kv_pool_allocator = MagicMock(page_size=64)
|
||||
queue.token_to_kv_pool_allocator.swa_available_size.side_effect = [64, 192]
|
||||
queue.tree_cache = MagicMock()
|
||||
|
||||
error = queue._reclaim_swa_tail_capacity(129, "req-1")
|
||||
|
||||
self.assertIsNone(error)
|
||||
params = queue.tree_cache.evict.call_args.args[0]
|
||||
self.assertEqual(params.num_tokens, 0)
|
||||
self.assertEqual(params.swa_num_tokens, 128)
|
||||
|
||||
def test_reclaim_swa_tail_capacity_fails_before_allocation(self):
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue.token_to_kv_pool_allocator = MagicMock(page_size=64)
|
||||
queue.token_to_kv_pool_allocator.swa_available_size.side_effect = [64, 128]
|
||||
queue.tree_cache = MagicMock()
|
||||
|
||||
error = queue._reclaim_swa_tail_capacity(129, "req-1")
|
||||
|
||||
self.assertEqual(
|
||||
error,
|
||||
"SWA eviction insufficient: needed=192, available=128, req=req-1",
|
||||
)
|
||||
|
||||
def _populate_prefix(self, cache, prefix_ids, prefix_values):
|
||||
"""Insert a prefix into the tree so future requests can match it."""
|
||||
cache.insert(
|
||||
@@ -303,6 +369,9 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
req.last_node = object()
|
||||
req.finished_reason = None
|
||||
req.cache_protected_len = 0
|
||||
req.swa_uuid_for_lock = 123
|
||||
req.swa_prefix_lock_released = False
|
||||
req.pd_rebootstrap_in_progress = False
|
||||
req.sampling_params.max_new_tokens = 16
|
||||
|
||||
decode_req = MagicMock()
|
||||
@@ -319,6 +388,10 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
queue.num_reserved_decode_tokens = 0
|
||||
queue._resolve_pending_reqs = MagicMock()
|
||||
queue._update_handshake_waiters = MagicMock()
|
||||
queue._uses_swa_tail_prealloc = MagicMock(return_value=True)
|
||||
queue._swa_tail_len = MagicMock(return_value=8)
|
||||
queue._swa_aware_allocatable_token_budgets = MagicMock(return_value=(8, 8))
|
||||
queue._swa_tail_allocatable_token_budget = MagicMock(return_value=8)
|
||||
queue._match_prefix_and_lock = MagicMock(
|
||||
return_value=DecodePrefixMatch(
|
||||
prefix_indices=torch.arange(4, dtype=torch.int64),
|
||||
@@ -349,21 +422,32 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
||||
scheduler.running_batch = running_batch
|
||||
scheduler.server_args = server_args
|
||||
scheduler.enable_hisparse = False
|
||||
scheduler.enable_decode_hicache = False
|
||||
scheduler.enable_priority_scheduling = False
|
||||
scheduler.waiting_queue = []
|
||||
scheduler.last_batch = None
|
||||
scheduler.output_streamer = MagicMock()
|
||||
queue.scheduler = scheduler
|
||||
|
||||
# Initial budget says the request fits; post-lock budget says it does not.
|
||||
queue._allocatable_token_budgets = MagicMock(side_effect=[8, 3])
|
||||
# The 4-token match is locked, then capped to zero because the whole
|
||||
# 8-token request is inside the SWA window. Admission rejection must
|
||||
# still release the original matched-node lock.
|
||||
queue._allocatable_token_budgets = MagicMock(return_value=3)
|
||||
|
||||
preallocated, failed = queue.pop_preallocated()
|
||||
|
||||
self.assertEqual(preallocated, [])
|
||||
self.assertEqual(failed, [])
|
||||
queue._pre_alloc.assert_not_called()
|
||||
queue.tree_cache.dec_lock_ref.assert_called_once_with(req.last_node)
|
||||
self.assertEqual(queue._allocatable_token_budgets.call_count, 2)
|
||||
queue.tree_cache.dec_swa_lock_only.assert_called_once_with(req.last_node, 123)
|
||||
queue.tree_cache.dec_lock_ref.assert_called_once_with(
|
||||
req.last_node,
|
||||
DecLockRefParams(swa_uuid_for_lock=123),
|
||||
skip_swa=True,
|
||||
)
|
||||
self.assertFalse(req.swa_prefix_lock_released)
|
||||
queue._swa_tail_len.assert_called_once_with(8)
|
||||
queue._allocatable_token_budgets.assert_called_once()
|
||||
|
||||
def test_repeated_incremental_no_leak(self):
|
||||
"""Multiple incremental transfers shouldn't leak lock_refs."""
|
||||
|
||||
@@ -116,6 +116,7 @@ class TestDeepSeekV4HiSparseAllocator(CustomTestCase):
|
||||
device=torch.device("cpu"),
|
||||
page_size=256,
|
||||
available_size=MagicMock(return_value=fill_len),
|
||||
swa_available_size=MagicMock(return_value=swa_tail_len),
|
||||
alloc_extend_swa_tail=MagicMock(return_value=kv_loc),
|
||||
alloc_logical_only=MagicMock(return_value=kv_loc),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user