fix: preserve HiCache load-back allocator headroom
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user