diff --git a/python/sglang/srt/hardware_backend/npu/allocator_npu.py b/python/sglang/srt/hardware_backend/npu/allocator_npu.py index bbc3e9dec..f425b85ce 100644 --- a/python/sglang/srt/hardware_backend/npu/allocator_npu.py +++ b/python/sglang/srt/hardware_backend/npu/allocator_npu.py @@ -48,7 +48,7 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator): num_new_pages_item = num_new_pages_tensor.item() else: num_new_pages_item = num_new_pages - if self.need_sort and num_new_pages_item > len(self.free_pages): + if num_new_pages_item > len(self.free_pages): self.merge_and_sort_free() if num_new_pages_item > len(self.free_pages): @@ -116,7 +116,6 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator): if num_new_pages > len(self.free_pages): self.merge_and_sort_free() - if num_new_pages > len(self.free_pages): return None @@ -144,11 +143,7 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator): if self.free_group is None: device = free_index.device free_page_indices = torch.unique(free_index.cpu() // self.page_size) - free_page_indices = free_page_indices.to(device) - if self.need_sort: - self.release_pages = torch.cat((free_page_indices, self.release_pages)) - else: - self.free_pages = torch.cat((free_page_indices, self.free_pages)) + self._release_page_ids(free_page_indices.to(device)) else: self.free_group.append(self._copy_for_free_group(free_index)) diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index e2af23cfb..e87269a72 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -171,14 +171,12 @@ class SchedulerInvariantChecker: self.req_to_token_pool.mamba_pool.size, ) if leak: - # Page-level leak diagnosis for mamba. Allocator flavors without - # page free-lists (free_pages is None) skip the page census — the - # dump must never crash the watchdog thread that calls it. - free_pages = self.token_to_kv_pool_allocator.free_pages - release_pages = self.token_to_kv_pool_allocator.release_pages - if free_pages is None or release_pages is None: + # Pools without a page free list return None; skip the census rather + # than crash the watchdog thread that runs this dump. + free_pages = self.token_to_kv_pool_allocator.get_all_free_pages() + if free_pages is None: return leak, msg - free_full_pages = set(free_pages.tolist() + release_pages.tolist()) + free_full_pages = set(free_pages.tolist()) cached_full_pages = set(self.tree_cache.all_values_flatten().tolist()) full_page_msg = "" if ( @@ -386,18 +384,9 @@ class SchedulerInvariantChecker: if not sub_allocs: return - def _free_pages(a): - free = a.free_pages - release = getattr(a, "release_pages", None) - return ( - torch.cat((free, release)) - if release is not None and len(release) > 0 - else free - ) - # Check B: every sub-pool's free set has no duplicate pages. for i, sub in enumerate(sub_allocs): - free = _free_pages(sub) + free = sub.get_all_free_pages() uniq = torch.unique(free) if uniq.numel() != free.numel(): raise_error_or_warn( @@ -409,7 +398,7 @@ class SchedulerInvariantChecker: # Check A: owner pages (full-pool indices) must not be in the full free # set (sub_allocs[0] is the full pool, even on hybrid-SWA). - full_unique = torch.unique(_free_pages(sub_allocs[0])) + full_unique = torch.unique(sub_allocs[0].get_all_free_pages()) stale = owner_pages[torch.isin(owner_pages, full_unique)] if stale.numel() > 0: raise_error_or_warn( diff --git a/python/sglang/srt/mem_cache/allocator/base.py b/python/sglang/srt/mem_cache/allocator/base.py index 1417c00b5..fc1cbe257 100644 --- a/python/sglang/srt/mem_cache/allocator/base.py +++ b/python/sglang/srt/mem_cache/allocator/base.py @@ -93,6 +93,14 @@ class BaseTokenToKVPoolAllocator(abc.ABC): def get_kvcache(self): return self._kvcache + def get_all_free_pages(self): + # Debug / invariant census; None when the pool has no page free list. + if self.free_pages is None: + return None + if self.release_pages is None or len(self.release_pages) == 0: + return self.free_pages + return torch.cat((self.free_pages, self.release_pages)) + def free_group_begin(self): assert self.free_group is None, "free groups cannot be nested" self.free_group = [] diff --git a/python/sglang/srt/mem_cache/allocator/paged.py b/python/sglang/srt/mem_cache/allocator/paged.py index f21c077f6..8ad1d2329 100755 --- a/python/sglang/srt/mem_cache/allocator/paged.py +++ b/python/sglang/srt/mem_cache/allocator/paged.py @@ -146,6 +146,19 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): pass self.clear() + def available_size(self): + return (len(self.free_pages) + self.num_staged_pages) * self.page_size + + def get_all_free_pages(self): + return torch.cat((self.free_pages, *self.staged_pages)) + + def merge_and_sort_free(self): + if not self.staged_pages: + return + self.free_pages, _ = torch.sort(self.get_all_free_pages()) + self.staged_pages = [] + self.num_staged_pages = 0 + def alloc(self, need_size: int): # page-aligned allocation, returning contiguous indices of pages if self.debug_mode: @@ -154,7 +167,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): ), "The allocation size should be page-aligned" num_pages = need_size // self.page_size - if self.need_sort and num_pages > len(self.free_pages): + if num_pages > len(self.free_pages): self.merge_and_sort_free() if num_pages > len(self.free_pages): return None @@ -185,9 +198,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): ) bs = len(prefix_lens) - if self.need_sort and extend_num_tokens // self.page_size + bs + 1 > len( - self.free_pages - ): + if extend_num_tokens // self.page_size + bs + 1 > len(self.free_pages): self.merge_and_sort_free() out_indices = torch.empty( @@ -231,7 +242,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): ) bs = len(seq_lens) - if self.need_sort and bs > len(self.free_pages): + if bs > len(self.free_pages): self.merge_and_sort_free() out_indices = torch.empty((bs,), dtype=torch.int64, device=self.device) @@ -303,13 +314,13 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): ) def _debug_check_no_duplicate_pages(self): - # span both containers: need_sort (PD disagg) routes frees into release_pages - pages = torch.cat((self.free_pages, self.release_pages)) + pages = self.get_all_free_pages() assert len(torch.unique(pages)) == len(pages) def _release_page_ids(self, *page_ids: torch.Tensor): if self.need_sort: - self.release_pages = torch.cat((*page_ids, self.release_pages)) + self.staged_pages.extend(page_ids) + self.num_staged_pages += sum(ids.numel() for ids in page_ids) else: self.free_pages = torch.cat((*page_ids, self.free_pages)) @@ -335,7 +346,9 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): ) self.free_group = None self.free_page_reps_group = [] - self.release_pages = torch.empty((0,), dtype=torch.int64, device=self.device) + # need_sort only: freed pages wait here, unsorted, until an alloc runs short. + self.staged_pages: list[torch.Tensor] = [] + self.num_staged_pages = 0 def get_cpu_copy(self, indices, mamba_indices=None): return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices) diff --git a/test/registered/unit/layers/test_minicpm_sparse_cache.py b/test/registered/unit/layers/test_minicpm_sparse_cache.py index d42fd0f72..c34c39d46 100644 --- a/test/registered/unit/layers/test_minicpm_sparse_cache.py +++ b/test/registered/unit/layers/test_minicpm_sparse_cache.py @@ -53,6 +53,9 @@ class RecordingAllocator(BaseTokenToKVPoolAllocator): def available_size(self): return self.capacity - len(self.live) + def get_all_free_pages(self): + return self.free_pages + def clear(self): self.next_slot = 1 self.live.clear() @@ -244,7 +247,6 @@ def test_streaming_session_release_frees_compressed_slots(): def test_mamba_leak_diagnostic_does_not_report_reserved_slots(): pool, _, _, allocator = make_pool_and_req(capacity=69) allocator.free_pages = torch.arange(6, 70, dtype=torch.int64) - allocator.release_pages = torch.empty(0, dtype=torch.int64) pool.mamba_pool = SimpleNamespace(size=1) pool.mamba_allocator = SimpleNamespace( size=1, diff --git a/test/registered/unit/managers/test_kv_page_invariants.py b/test/registered/unit/managers/test_kv_page_invariants.py index 83a670274..2d2109c69 100644 --- a/test/registered/unit/managers/test_kv_page_invariants.py +++ b/test/registered/unit/managers/test_kv_page_invariants.py @@ -22,7 +22,7 @@ def _make_checker(page_size=_PAGE_SIZE, row_width=4096, num_reqs=8, free_pages=N alloc = SimpleNamespace( page_size=page_size, free_pages=free_pages, - release_pages=torch.empty(0, dtype=torch.int64), + get_all_free_pages=lambda: free_pages, ) tc = SimpleNamespace(slots={}) _ps, _rtp, _alloc, _tc = page_size, rtp, alloc, tc diff --git a/test/registered/unit/mem_cache/test_paged_allocator_lazy_release.py b/test/registered/unit/mem_cache/test_paged_allocator_lazy_release.py new file mode 100644 index 000000000..ff1eb5c3d --- /dev/null +++ b/test/registered/unit/mem_cache/test_paged_allocator_lazy_release.py @@ -0,0 +1,70 @@ +"""need_sort allocators stage released pages and merge only under allocation pressure.""" + +import unittest +from unittest.mock import patch + +import torch + +from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + +PAGE_SIZE = 2 +NUM_PAGES = 16 + + +def _make_allocator(*, need_sort: bool) -> PagedTokenToKVPoolAllocator: + return PagedTokenToKVPoolAllocator( + size=NUM_PAGES * PAGE_SIZE, + page_size=PAGE_SIZE, + dtype=torch.float16, + device="cpu", + kvcache=None, + need_sort=need_sort, + ) + + +def _page_ids(indices: torch.Tensor) -> torch.Tensor: + return indices[::PAGE_SIZE] // PAGE_SIZE + + +class TestPagedAllocatorLazyRelease(CustomTestCase): + def test_release_stages_until_free_pages_run_dry(self): + allocator = _make_allocator(need_sort=True) + allocated = allocator.alloc(24) + free_pages_before = allocator.free_pages + + with patch.object(torch, "cat", wraps=torch.cat) as cat_mock: + allocator.free(allocated[:2]) + allocator.free(allocated[4:6]) + self.assertIs(allocator.free_pages, free_pages_before) + self.assertEqual(allocator.num_staged_pages, 2) + self.assertEqual(allocator.available_size(), 12) + cat_mock.assert_not_called() + + # free_pages still covers this: no merge. + primary = allocator.alloc(8) + self.assertTrue(torch.equal(_page_ids(primary), torch.arange(13, 17))) + cat_mock.assert_not_called() + + # Pressure: one merge, staged pages come back sorted. + reused = allocator.alloc(4) + self.assertTrue(torch.equal(_page_ids(reused), torch.tensor([1, 3]))) + self.assertEqual(cat_mock.call_count, 1) + + self.assertEqual(allocator.num_staged_pages, 0) + self.assertEqual(allocator.staged_pages, []) + + def test_clear_drops_staged_pages(self): + allocator = _make_allocator(need_sort=True) + allocator.free(allocator.alloc(4)) + + allocator.clear() + self.assertEqual(allocator.available_size(), allocator.size) + self.assertEqual(allocator.staged_pages, []) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_paged_free_segment.py b/test/registered/unit/mem_cache/test_paged_free_segment.py index 2efc0f407..4f0b125fc 100644 --- a/test/registered/unit/mem_cache/test_paged_free_segment.py +++ b/test/registered/unit/mem_cache/test_paged_free_segment.py @@ -67,11 +67,12 @@ class TestFreeSegment(unittest.TestCase): alloc.free_segment(row[:0], start_pos=0) self.assertEqual(len(alloc.free_pages), before) - def test_need_sort_routes_to_release_pages(self): + def test_need_sort_defers_released_pages(self): alloc = _make_allocator(need_sort=True) row = _make_kv_row(alloc, 2 * PAGE_SIZE) alloc.free_segment(row, start_pos=0) - self.assertEqual(len(alloc.release_pages), 2) + self.assertEqual(len(alloc.staged_pages), 1) + self.assertEqual(alloc.num_staged_pages, 2) def test_group_defers_until_group_end(self): alloc = _make_allocator() @@ -108,8 +109,8 @@ class TestFreeSegment(unittest.TestCase): with self.assertRaises(AssertionError): alloc.free_group_end() - def test_group_end_debug_assert_covers_release_pages(self): - # need_sort routes frees into release_pages; the duplicate check must + def test_group_end_debug_assert_covers_staged_releases(self): + # need_sort stages frees in chunks; the duplicate check must # not go vacuous there (PD disaggregation runs with need_sort=True). alloc = _make_allocator(need_sort=True) alloc.debug_mode = True