[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:
wangwenmingaa
2026-09-02 16:51:37 -07:00
committed by GitHub
co-authored by wangwenming.41 hnyls2002 Liangsheng Yin
parent 4bc34117f1
commit c05f8ae830
8 changed files with 118 additions and 40 deletions
@@ -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))
@@ -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(
@@ -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 = []
+22 -9
View File
@@ -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)