[PD] Optimize paged allocator free-list release (#37146)
Co-authored-by: wangwenming.41 <wangwenming.41@jd.com> Co-authored-by: hnyls2002 <lsyincs@gmail.com> Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
co-authored by
wangwenming.41
hnyls2002
Liangsheng Yin
parent
4bc34117f1
commit
c05f8ae830
@@ -48,7 +48,7 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator):
|
|||||||
num_new_pages_item = num_new_pages_tensor.item()
|
num_new_pages_item = num_new_pages_tensor.item()
|
||||||
else:
|
else:
|
||||||
num_new_pages_item = num_new_pages
|
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()
|
self.merge_and_sort_free()
|
||||||
|
|
||||||
if num_new_pages_item > len(self.free_pages):
|
if num_new_pages_item > len(self.free_pages):
|
||||||
@@ -116,7 +116,6 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator):
|
|||||||
|
|
||||||
if num_new_pages > len(self.free_pages):
|
if num_new_pages > len(self.free_pages):
|
||||||
self.merge_and_sort_free()
|
self.merge_and_sort_free()
|
||||||
|
|
||||||
if num_new_pages > len(self.free_pages):
|
if num_new_pages > len(self.free_pages):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -144,11 +143,7 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator):
|
|||||||
if self.free_group is None:
|
if self.free_group is None:
|
||||||
device = free_index.device
|
device = free_index.device
|
||||||
free_page_indices = torch.unique(free_index.cpu() // self.page_size)
|
free_page_indices = torch.unique(free_index.cpu() // self.page_size)
|
||||||
free_page_indices = free_page_indices.to(device)
|
self._release_page_ids(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))
|
|
||||||
else:
|
else:
|
||||||
self.free_group.append(self._copy_for_free_group(free_index))
|
self.free_group.append(self._copy_for_free_group(free_index))
|
||||||
|
|
||||||
|
|||||||
@@ -171,14 +171,12 @@ class SchedulerInvariantChecker:
|
|||||||
self.req_to_token_pool.mamba_pool.size,
|
self.req_to_token_pool.mamba_pool.size,
|
||||||
)
|
)
|
||||||
if leak:
|
if leak:
|
||||||
# Page-level leak diagnosis for mamba. Allocator flavors without
|
# Pools without a page free list return None; skip the census rather
|
||||||
# page free-lists (free_pages is None) skip the page census — the
|
# than crash the watchdog thread that runs this dump.
|
||||||
# dump must never crash the watchdog thread that calls it.
|
free_pages = self.token_to_kv_pool_allocator.get_all_free_pages()
|
||||||
free_pages = self.token_to_kv_pool_allocator.free_pages
|
if free_pages is None:
|
||||||
release_pages = self.token_to_kv_pool_allocator.release_pages
|
|
||||||
if free_pages is None or release_pages is None:
|
|
||||||
return leak, msg
|
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())
|
cached_full_pages = set(self.tree_cache.all_values_flatten().tolist())
|
||||||
full_page_msg = ""
|
full_page_msg = ""
|
||||||
if (
|
if (
|
||||||
@@ -386,18 +384,9 @@ class SchedulerInvariantChecker:
|
|||||||
if not sub_allocs:
|
if not sub_allocs:
|
||||||
return
|
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.
|
# Check B: every sub-pool's free set has no duplicate pages.
|
||||||
for i, sub in enumerate(sub_allocs):
|
for i, sub in enumerate(sub_allocs):
|
||||||
free = _free_pages(sub)
|
free = sub.get_all_free_pages()
|
||||||
uniq = torch.unique(free)
|
uniq = torch.unique(free)
|
||||||
if uniq.numel() != free.numel():
|
if uniq.numel() != free.numel():
|
||||||
raise_error_or_warn(
|
raise_error_or_warn(
|
||||||
@@ -409,7 +398,7 @@ class SchedulerInvariantChecker:
|
|||||||
|
|
||||||
# Check A: owner pages (full-pool indices) must not be in the full free
|
# 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).
|
# 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)]
|
stale = owner_pages[torch.isin(owner_pages, full_unique)]
|
||||||
if stale.numel() > 0:
|
if stale.numel() > 0:
|
||||||
raise_error_or_warn(
|
raise_error_or_warn(
|
||||||
|
|||||||
@@ -93,6 +93,14 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
|||||||
def get_kvcache(self):
|
def get_kvcache(self):
|
||||||
return self._kvcache
|
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):
|
def free_group_begin(self):
|
||||||
assert self.free_group is None, "free groups cannot be nested"
|
assert self.free_group is None, "free groups cannot be nested"
|
||||||
self.free_group = []
|
self.free_group = []
|
||||||
|
|||||||
@@ -146,6 +146,19 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
pass
|
pass
|
||||||
self.clear()
|
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):
|
def alloc(self, need_size: int):
|
||||||
# page-aligned allocation, returning contiguous indices of pages
|
# page-aligned allocation, returning contiguous indices of pages
|
||||||
if self.debug_mode:
|
if self.debug_mode:
|
||||||
@@ -154,7 +167,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
), "The allocation size should be page-aligned"
|
), "The allocation size should be page-aligned"
|
||||||
|
|
||||||
num_pages = need_size // self.page_size
|
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()
|
self.merge_and_sort_free()
|
||||||
if num_pages > len(self.free_pages):
|
if num_pages > len(self.free_pages):
|
||||||
return None
|
return None
|
||||||
@@ -185,9 +198,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
)
|
)
|
||||||
|
|
||||||
bs = len(prefix_lens)
|
bs = len(prefix_lens)
|
||||||
if self.need_sort and extend_num_tokens // self.page_size + bs + 1 > len(
|
if extend_num_tokens // self.page_size + bs + 1 > len(self.free_pages):
|
||||||
self.free_pages
|
|
||||||
):
|
|
||||||
self.merge_and_sort_free()
|
self.merge_and_sort_free()
|
||||||
|
|
||||||
out_indices = torch.empty(
|
out_indices = torch.empty(
|
||||||
@@ -231,7 +242,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
)
|
)
|
||||||
|
|
||||||
bs = len(seq_lens)
|
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()
|
self.merge_and_sort_free()
|
||||||
|
|
||||||
out_indices = torch.empty((bs,), dtype=torch.int64, device=self.device)
|
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):
|
def _debug_check_no_duplicate_pages(self):
|
||||||
# span both containers: need_sort (PD disagg) routes frees into release_pages
|
pages = self.get_all_free_pages()
|
||||||
pages = torch.cat((self.free_pages, self.release_pages))
|
|
||||||
assert len(torch.unique(pages)) == len(pages)
|
assert len(torch.unique(pages)) == len(pages)
|
||||||
|
|
||||||
def _release_page_ids(self, *page_ids: torch.Tensor):
|
def _release_page_ids(self, *page_ids: torch.Tensor):
|
||||||
if self.need_sort:
|
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:
|
else:
|
||||||
self.free_pages = torch.cat((*page_ids, self.free_pages))
|
self.free_pages = torch.cat((*page_ids, self.free_pages))
|
||||||
|
|
||||||
@@ -335,7 +346,9 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
)
|
)
|
||||||
self.free_group = None
|
self.free_group = None
|
||||||
self.free_page_reps_group = []
|
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):
|
def get_cpu_copy(self, indices, mamba_indices=None):
|
||||||
return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices)
|
return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices)
|
||||||
|
|||||||
@@ -53,6 +53,9 @@ class RecordingAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
def available_size(self):
|
def available_size(self):
|
||||||
return self.capacity - len(self.live)
|
return self.capacity - len(self.live)
|
||||||
|
|
||||||
|
def get_all_free_pages(self):
|
||||||
|
return self.free_pages
|
||||||
|
|
||||||
def clear(self):
|
def clear(self):
|
||||||
self.next_slot = 1
|
self.next_slot = 1
|
||||||
self.live.clear()
|
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():
|
def test_mamba_leak_diagnostic_does_not_report_reserved_slots():
|
||||||
pool, _, _, allocator = make_pool_and_req(capacity=69)
|
pool, _, _, allocator = make_pool_and_req(capacity=69)
|
||||||
allocator.free_pages = torch.arange(6, 70, dtype=torch.int64)
|
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_pool = SimpleNamespace(size=1)
|
||||||
pool.mamba_allocator = SimpleNamespace(
|
pool.mamba_allocator = SimpleNamespace(
|
||||||
size=1,
|
size=1,
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ def _make_checker(page_size=_PAGE_SIZE, row_width=4096, num_reqs=8, free_pages=N
|
|||||||
alloc = SimpleNamespace(
|
alloc = SimpleNamespace(
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
free_pages=free_pages,
|
free_pages=free_pages,
|
||||||
release_pages=torch.empty(0, dtype=torch.int64),
|
get_all_free_pages=lambda: free_pages,
|
||||||
)
|
)
|
||||||
tc = SimpleNamespace(slots={})
|
tc = SimpleNamespace(slots={})
|
||||||
_ps, _rtp, _alloc, _tc = page_size, rtp, alloc, tc
|
_ps, _rtp, _alloc, _tc = page_size, rtp, alloc, tc
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -67,11 +67,12 @@ class TestFreeSegment(unittest.TestCase):
|
|||||||
alloc.free_segment(row[:0], start_pos=0)
|
alloc.free_segment(row[:0], start_pos=0)
|
||||||
self.assertEqual(len(alloc.free_pages), before)
|
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)
|
alloc = _make_allocator(need_sort=True)
|
||||||
row = _make_kv_row(alloc, 2 * PAGE_SIZE)
|
row = _make_kv_row(alloc, 2 * PAGE_SIZE)
|
||||||
alloc.free_segment(row, start_pos=0)
|
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):
|
def test_group_defers_until_group_end(self):
|
||||||
alloc = _make_allocator()
|
alloc = _make_allocator()
|
||||||
@@ -108,8 +109,8 @@ class TestFreeSegment(unittest.TestCase):
|
|||||||
with self.assertRaises(AssertionError):
|
with self.assertRaises(AssertionError):
|
||||||
alloc.free_group_end()
|
alloc.free_group_end()
|
||||||
|
|
||||||
def test_group_end_debug_assert_covers_release_pages(self):
|
def test_group_end_debug_assert_covers_staged_releases(self):
|
||||||
# need_sort routes frees into release_pages; the duplicate check must
|
# need_sort stages frees in chunks; the duplicate check must
|
||||||
# not go vacuous there (PD disaggregation runs with need_sort=True).
|
# not go vacuous there (PD disaggregation runs with need_sort=True).
|
||||||
alloc = _make_allocator(need_sort=True)
|
alloc = _make_allocator(need_sort=True)
|
||||||
alloc.debug_mode = True
|
alloc.debug_mode = True
|
||||||
|
|||||||
Reference in New Issue
Block a user