fix: preserve HiCache load-back allocator headroom

This commit is contained in:
2026-05-08 00:18:33 +08:00
parent 1f074f434e
commit 95bacb8862
10 changed files with 172 additions and 37 deletions

View File

@@ -69,6 +69,38 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
def available_size(self):
return (len(self.free_pages) + len(self.release_pages)) * self.page_size
def immediate_available_pages(self) -> int:
if self.free_pages is None:
return 0
return len(self.free_pages)
def deferred_available_pages(self) -> int:
if self.release_pages is None:
return 0
return len(self.release_pages)
def immediate_available_size(self) -> int:
if self.free_pages is None:
return self.available_size()
return self.immediate_available_pages() * self.page_size
def deferred_available_size(self) -> int:
if self.release_pages is None:
return 0
return self.deferred_available_pages() * self.page_size
def allocator_state_str(self) -> str:
if self.free_pages is None:
return f"allocator_available_size={self.available_size()}"
return (
"allocator_available_size="
f"{self.available_size()} "
f"(allocator_free_size={self.immediate_available_size()} "
f"[free_pages={self.immediate_available_pages()}] + "
f"allocator_release_size={self.deferred_available_size()} "
f"[release_pages={self.deferred_available_pages()}])"
)
def get_kvcache(self):
return self._kvcache
@@ -434,10 +466,15 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
)
bs = len(prefix_lens)
if self.need_sort and extend_num_tokens // self.page_size + bs + 1 > len(
self.free_pages
):
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
prefix_lens=prefix_lens_cpu,
)
if self.need_sort and num_new_pages > len(self.free_pages):
self.merge_and_sort_free()
if num_new_pages > len(self.free_pages):
return None
out_indices = torch.empty(
(extend_num_tokens,), dtype=torch.int64, device=self.device
@@ -456,14 +493,6 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
if self.debug_mode:
assert len(torch.unique(out_indices)) == len(out_indices)
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
prefix_lens=prefix_lens_cpu,
)
if num_new_pages > len(self.free_pages):
return None
self.free_pages = self.free_pages[num_new_pages:]
return out_indices
@@ -479,8 +508,15 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
)
bs = len(seq_lens)
if self.need_sort and bs > len(self.free_pages):
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
decode=True,
)
if self.need_sort and num_new_pages > len(self.free_pages):
self.merge_and_sort_free()
if num_new_pages > len(self.free_pages):
return None
out_indices = torch.empty((bs,), dtype=torch.int64, device=self.device)
alloc_decode_kernel[(bs,)](
@@ -495,14 +531,6 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
if self.debug_mode:
assert len(torch.unique(out_indices)) == len(out_indices)
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
decode=True,
)
if num_new_pages > len(self.free_pages):
return None
self.free_pages = self.free_pages[num_new_pages:]
return out_indices