Support swa allocator page size > 1 (#16296)

This commit is contained in:
Ke Bao
2026-01-03 08:49:47 +08:00
committed by GitHub
parent 31ed68e7c1
commit 7c1b4b1c4c
4 changed files with 205 additions and 60 deletions

View File

@@ -20,6 +20,7 @@ Page-aligned memory pool.
"""
import abc
import weakref
from typing import TYPE_CHECKING
import torch
@@ -179,39 +180,77 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self,
size: int,
size_swa: int,
page_size: int,
dtype: torch.dtype,
device: str,
kvcache: SWAKVPool,
need_sort: bool,
):
super().__init__(size, 1, dtype, device, kvcache, need_sort)
assert isinstance(kvcache, SWAKVPool)
self._size_full = size
self._size_swa = size_swa
self.full_attn_allocator = TokenToKVPoolAllocator(
size,
dtype,
device,
kvcache.full_kv_pool,
need_sort,
)
self.swa_attn_allocator = TokenToKVPoolAllocator(
size_swa,
dtype,
device,
kvcache.swa_kv_pool,
need_sort,
)
self.full_to_swa_index_mapping = torch.empty(
size + size_swa + 1,
dtype=torch.int64,
device=device,
)
self.clear()
self.dtype = dtype
self.device = device
self.page_size = page_size
self._kvcache.full_to_swa_index_mapping = self.full_to_swa_index_mapping
if page_size == 1:
self.full_attn_allocator = TokenToKVPoolAllocator(
size,
dtype,
device,
kvcache.full_kv_pool,
need_sort,
)
self.swa_attn_allocator = TokenToKVPoolAllocator(
size_swa,
dtype,
device,
kvcache.swa_kv_pool,
need_sort,
)
else:
self.full_attn_allocator = PagedTokenToKVPoolAllocator(
size,
page_size,
dtype,
device,
kvcache.full_kv_pool,
need_sort,
)
self.swa_attn_allocator = PagedTokenToKVPoolAllocator(
size_swa,
page_size,
dtype,
device,
kvcache.swa_kv_pool,
need_sort,
)
# Note: append one more item of value -1 in the end so -1 maps to -1.
# It is needed for the last_loc in alloc_extend, where the first full_last_loc
# is -1, and we need to map it to swa_last_loc -1 as well.
self.full_to_swa_index_mapping = torch.cat(
[
torch.zeros(
size + self.page_size,
dtype=torch.int64,
device=device,
),
torch.tensor([-1], dtype=torch.int64, device=device),
]
)
self.need_sort = need_sort
self.free_pages = None
self.release_pages = None
self.is_not_in_free_group = True
self.free_group = []
self.clear()
self._kvcache = kvcache
self._kvcache.register_mapping(weakref.proxy(self.full_to_swa_index_mapping))
def available_size(self):
# Note: use full_available_size() and swa_available_size() instead.
raise NotImplementedError()
def full_available_size(self):
@@ -221,13 +260,17 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
return self.swa_attn_allocator.available_size()
@property
def size_full(self):
return self._size_full
def size(self):
return min(self._size_full, self._size_swa)
@property
def size_swa(self):
return self._size_swa
@property
def size_full(self):
return self._size_full
def debug_print(self) -> str:
msg = ""
msg += f"#swa-available-size: {self.swa_attn_allocator.available_size()}, "
@@ -240,10 +283,11 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
return self._kvcache
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
assert self.full_to_swa_index_mapping is not None
return self.full_to_swa_index_mapping[kv_indices].to(torch.int32)
assert self._kvcache.full_to_swa_index_mapping is not None
return self._kvcache.translate_loc_from_full_to_swa(kv_indices)
def alloc(self, need_size: int):
assert self.page_size == 1
if need_size > self.full_attn_allocator.available_size():
return None
if need_size > self.swa_attn_allocator.available_size():
@@ -251,12 +295,81 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
alloc_full_indices = self.full_attn_allocator.alloc(need_size)
alloc_swa_indices = self.swa_attn_allocator.alloc(need_size)
assert alloc_full_indices is not None
assert alloc_swa_indices is not None
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
return alloc_full_indices
def alloc_extend(
self,
prefix_lens: torch.Tensor,
prefix_lens_cpu: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_cpu: torch.Tensor,
last_loc: torch.Tensor, # last_loc for full layers
extend_num_tokens: int,
):
assert self.page_size > 1
num_tokens = extend_num_tokens + len(seq_lens) * self.page_size
if num_tokens > self.full_attn_allocator.available_size():
return None
if num_tokens > self.swa_attn_allocator.available_size():
return None
swa_last_loc = self.translate_loc_from_full_to_swa(last_loc)
alloc_full_indices = self.full_attn_allocator.alloc_extend(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
)
alloc_swa_indices = self.swa_attn_allocator.alloc_extend(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
swa_last_loc,
extend_num_tokens,
)
assert alloc_full_indices is not None
assert alloc_swa_indices is not None
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
return alloc_full_indices
def alloc_decode(
self,
seq_lens: torch.Tensor,
seq_lens_cpu: torch.Tensor,
last_loc: torch.Tensor, # last_loc for full layers
):
assert self.page_size > 1
swa_last_loc = self.translate_loc_from_full_to_swa(last_loc)
alloc_full_indices = self.full_attn_allocator.alloc_decode(
seq_lens, seq_lens_cpu, last_loc
)
alloc_swa_indices = self.swa_attn_allocator.alloc_decode(
seq_lens, seq_lens_cpu, swa_last_loc
)
if alloc_full_indices is None or alloc_swa_indices is None:
return None
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
return alloc_full_indices
def free(self, free_index: torch.Tensor):
if free_index.numel() == 0:
return
# NOTE: the API is not idempotent.
if self.is_not_in_free_group:
self.full_attn_allocator.free(free_index)
self.free_swa(free_index)
@@ -287,7 +400,8 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
def clear(self):
self.swa_attn_allocator.clear()
self.full_attn_allocator.clear()
self.full_to_swa_index_mapping.fill_(0)
# Note: the last item is -1, we don't clear it, see the comment in __init__
self.full_to_swa_index_mapping[:-1].fill_(0)
self.is_not_in_free_group = True
self.free_group = []