[PD] Optimize SWA allocation (#28085)

Co-authored-by: cctry <cctry@fb.com>
Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
This commit is contained in:
cctry
2026-06-15 11:01:55 -07:00
committed by GitHub
co-authored by cctry Lianmin Zheng
parent 19e85868f6
commit 33719cfb31
3 changed files with 34 additions and 22 deletions
@@ -1379,6 +1379,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
extend_num_tokens=fill_len,
swa_tail_len=self._swa_tail_len(fill_len),
)
req.swa_evicted_seqlen = fill_len - self._swa_tail_len(fill_len)
else:
kv_loc = self.token_to_kv_pool_allocator.alloc_extend(
prefix_lens=torch.tensor(
+17 -21
View File
@@ -151,14 +151,17 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
assert alloc_full_indices is not None
assert alloc_swa_indices is not None
if _is_npu:
self.full_to_swa_index_mapping[alloc_full_indices.to(torch.int64)] = (
alloc_swa_indices.to(torch.int64)
)
else:
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
self.set_full_to_swa_mapping(alloc_full_indices, alloc_swa_indices)
return alloc_full_indices
def new_pages_available(self, num_full_pages: int, num_swa_pages: int) -> bool:
return (
num_full_pages
<= self.full_attn_allocator.available_size() // self.page_size
and num_swa_pages
<= self.swa_attn_allocator.available_size() // self.page_size
)
def alloc_extend(
self,
prefix_lens: torch.Tensor,
@@ -173,9 +176,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
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 > self.full_attn_allocator.available_size() // self.page_size:
return None
if num_new_pages > self.swa_attn_allocator.available_size() // self.page_size:
if not self.new_pages_available(num_new_pages, num_new_pages):
return None
swa_last_loc = self.translate_loc_from_full_to_swa(last_loc)
@@ -201,12 +202,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
assert alloc_full_indices is not None
assert alloc_swa_indices is not None
if _is_npu:
self.full_to_swa_index_mapping[alloc_full_indices.to(torch.int64)] = (
alloc_swa_indices.to(torch.int64)
)
else:
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
self.set_full_to_swa_mapping(alloc_full_indices, alloc_swa_indices)
return alloc_full_indices
@@ -235,9 +231,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
seq_lens=seq_lens_cpu, page_size=self.page_size, prefix_lens=prefix_lens_cpu
)
num_swa_pages = (swa_tail_len + self.page_size - 1) // self.page_size
if num_full_pages > self.full_attn_allocator.available_size() // self.page_size:
return None
if num_swa_pages > self.swa_attn_allocator.available_size() // self.page_size:
if not self.new_pages_available(num_full_pages, num_swa_pages):
return None
alloc_full_indices = self.full_attn_allocator.alloc_extend(
@@ -247,6 +241,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
seq_lens_cpu,
last_loc,
extend_num_tokens,
num_new_pages=num_full_pages,
)
assert alloc_full_indices is not None
@@ -267,12 +262,13 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
swa_seq_lens_cpu,
swa_last_loc,
swa_tail_len,
num_new_pages=num_swa_pages,
)
assert alloc_swa_indices is not None
self.full_to_swa_index_mapping[
alloc_full_indices[-swa_tail_len:].to(torch.int64)
] = alloc_swa_indices.to(torch.int64)
self.set_full_to_swa_mapping(
alloc_full_indices[-swa_tail_len:], alloc_swa_indices
)
if swa_tail_len < extend_num_tokens:
self.full_to_swa_index_mapping[
alloc_full_indices[:-swa_tail_len].to(torch.int64)