Fix paged SWA retraction resume accounting (#33794)
Co-authored-by: zhisbug <1654062+zhisbug@users.noreply.github.com>
This commit is contained in:
@@ -88,7 +88,7 @@ from sglang.srt.observability.req_time_stats import (
|
|||||||
set_time_batch,
|
set_time_batch,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_disagg, get_parallel
|
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.network import NetworkAddress
|
||||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
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]:
|
def _prealloc_required_tokens(self, req: Req) -> Tuple[int, int]:
|
||||||
full_len, swa_len = self._prealloc_kv_lens(req)
|
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
|
swa_reserved = self.num_reserved_decode_tokens
|
||||||
if self.scheduler.server_args.disable_radix_cache:
|
if self.scheduler.server_args.disable_radix_cache:
|
||||||
swa_reserved = 0
|
swa_reserved = 0
|
||||||
@@ -1120,8 +1125,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
hicache_reserved_tokens=reserved_restore_tokens,
|
hicache_reserved_tokens=reserved_restore_tokens,
|
||||||
)
|
)
|
||||||
if uses_swa_tail_prealloc:
|
if uses_swa_tail_prealloc:
|
||||||
# SWA budget uses simple decrement (no radix cache eviction in
|
# SWA has no radix cache eviction, so decrement its
|
||||||
# the SWA pool, so page-rounding drift is negligible).
|
# page-aligned requirement directly.
|
||||||
swa_allocatable_tokens -= swa_required
|
swa_allocatable_tokens -= swa_required
|
||||||
decode_req.req.cache_protected_len = total_prefix_len
|
decode_req.req.cache_protected_len = total_prefix_len
|
||||||
|
|
||||||
|
|||||||
@@ -30,6 +30,51 @@ class FakeReceiver:
|
|||||||
|
|
||||||
|
|
||||||
class TestDecodeQueueCleanup(CustomTestCase):
|
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):
|
def test_prealloc_abort_clears_receiver_before_removing_request(self):
|
||||||
receiver = FakeReceiver()
|
receiver = FakeReceiver()
|
||||||
req = SimpleNamespace(
|
req = SimpleNamespace(
|
||||||
|
|||||||
Reference in New Issue
Block a user