Fix paged SWA retraction resume accounting (#33794)

Co-authored-by: zhisbug <1654062+zhisbug@users.noreply.github.com>
This commit is contained in:
Hao Zhang
2026-08-06 16:32:54 -07:00
committed by GitHub
co-authored by zhisbug
parent e0af47b03e
commit af7c62e337
2 changed files with 53 additions and 3 deletions
+8 -3
View File
@@ -88,7 +88,7 @@ from sglang.srt.observability.req_time_stats import (
set_time_batch,
)
from sglang.srt.runtime_context import get_disagg, get_parallel
from sglang.srt.utils import get_num_new_pages, is_npu
from sglang.srt.utils import ceil_align, get_num_new_pages, is_npu
from sglang.srt.utils.network import NetworkAddress
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
@@ -406,6 +406,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
def _prealloc_required_tokens(self, req: Req) -> Tuple[int, int]:
full_len, swa_len = self._prealloc_kv_lens(req)
page_size = self.token_to_kv_pool_allocator.page_size
if page_size > 1:
# Match the allocator, which charges whole pages for both pools.
full_len = ceil_align(full_len, page_size)
swa_len = ceil_align(swa_len, page_size)
swa_reserved = self.num_reserved_decode_tokens
if self.scheduler.server_args.disable_radix_cache:
swa_reserved = 0
@@ -1120,8 +1125,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
hicache_reserved_tokens=reserved_restore_tokens,
)
if uses_swa_tail_prealloc:
# SWA budget uses simple decrement (no radix cache eviction in
# the SWA pool, so page-rounding drift is negligible).
# SWA has no radix cache eviction, so decrement its
# page-aligned requirement directly.
swa_allocatable_tokens -= swa_required
decode_req.req.cache_protected_len = total_prefix_len
@@ -30,6 +30,51 @@ class FakeReceiver:
class TestDecodeQueueCleanup(CustomTestCase):
def test_paged_swa_retraction_resume_uses_physical_page_budget(self):
page_size = 128
fill_len = 574
physical_tokens_per_req = 5 * page_size
physical_available = 18 * page_size
reqs = [
SimpleNamespace(
rid=f"req-{i}",
origin_input_ids=[0] * fill_len,
output_ids=[],
is_retracted=True,
load_kv_cache=MagicMock(),
)
for i in range(4)
]
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
queue.retracted_queue = reqs.copy()
queue.num_reserved_decode_tokens = 0
queue.req_to_token_pool = SimpleNamespace(available_size=lambda: len(reqs))
queue.token_to_kv_pool_allocator = SimpleNamespace(page_size=page_size)
queue.scheduler = SimpleNamespace(
sliding_window_size=2047,
server_args=SimpleNamespace(disable_radix_cache=True),
)
queue._uses_swa_tail_prealloc = MagicMock(return_value=True)
queue._swa_aware_allocatable_token_budgets = MagicMock(
return_value=(physical_available, physical_available)
)
def pre_alloc(_req):
nonlocal physical_available
self.assertGreaterEqual(physical_available, physical_tokens_per_req)
physical_available -= physical_tokens_per_req
queue._pre_alloc = MagicMock(side_effect=pre_alloc)
resumed = queue.resume_retracted_reqs()
self.assertEqual(resumed, reqs[:3])
self.assertEqual(queue.retracted_queue, reqs[3:])
self.assertEqual(physical_available, 3 * page_size)
self.assertEqual(queue._pre_alloc.call_count, 3)
def test_prealloc_abort_clears_receiver_before_removing_request(self):
receiver = FakeReceiver()
req = SimpleNamespace(