[SWA] Fix missing mamba_indices parameter in cpu copy interface (#24026)

This commit is contained in:
Lianmin Zheng
2026-04-29 17:33:38 -07:00
committed by GitHub
parent 2bbd30a27a
commit e60c60eff0
4 changed files with 34 additions and 28 deletions
+14 -10
View File
@@ -87,11 +87,11 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
(0,), dtype=self.release_pages.dtype, device=self.device
)
def get_cpu_copy(self, *args, **kwargs):
def get_cpu_copy(self, indices, mamba_indices=None):
# FIXME: reuse the get_cpu_copy after paged allocator is implemented
raise NotImplementedError()
def load_cpu_copy(self, *args, **kwargs):
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
# FIXME: reuse the load_cpu_copy after paged allocator is implemented
raise NotImplementedError()
@@ -164,11 +164,13 @@ class TokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
else:
self.free_group.append(free_index)
def get_cpu_copy(self, indices, **kwargs):
return self._kvcache.get_cpu_copy(indices, **kwargs)
def get_cpu_copy(self, indices, mamba_indices=None):
return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices)
def load_cpu_copy(self, kv_cache_cpu, indices, **kwargs):
return self._kvcache.load_cpu_copy(kv_cache_cpu, indices, **kwargs)
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
return self._kvcache.load_cpu_copy(
kv_cache_cpu, indices, mamba_indices=mamba_indices
)
def alloc_extend_naive(
@@ -512,8 +514,10 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.free_group = []
self.release_pages = torch.empty((0,), dtype=torch.int64, device=self.device)
def get_cpu_copy(self, indices, **kwargs):
return self._kvcache.get_cpu_copy(indices, **kwargs)
def get_cpu_copy(self, indices, mamba_indices=None):
return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices)
def load_cpu_copy(self, kv_cache_cpu, indices, **kwargs):
return self._kvcache.load_cpu_copy(kv_cache_cpu, indices, **kwargs)
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
return self._kvcache.load_cpu_copy(
kv_cache_cpu, indices, mamba_indices=mamba_indices
)
@@ -114,10 +114,10 @@ class HiSparseNSATokenToKVPool(NSATokenToKVPool):
num_layers=self.layer_num,
)
def get_cpu_copy(self, indices):
def get_cpu_copy(self, indices, mamba_indices=None):
raise NotImplementedError("HiSparseDevicePool does not support get_cpu_copy")
def load_cpu_copy(self, kv_cache_cpu, indices):
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
raise NotImplementedError("HiSparseDevicePool does not support load_cpu_copy")
+10 -10
View File
@@ -398,7 +398,7 @@ class MambaPool:
self.copy_from(src_index, dst_index)
return dst_index
def get_cpu_copy(self, indices, **kwargs):
def get_cpu_copy(self, indices):
torch.cuda.synchronize()
conv_cpu = [
conv[:, indices].to("cpu", non_blocking=True)
@@ -410,7 +410,7 @@ class MambaPool:
torch.cuda.synchronize()
return conv_cpu, temporal_cpu
def load_cpu_copy(self, mamba_cache_cpu, indices, **kwargs):
def load_cpu_copy(self, mamba_cache_cpu, indices):
conv_cpu, temporal_cpu = mamba_cache_cpu
torch.cuda.synchronize()
for i, conv in enumerate(self.mamba_cache.conv):
@@ -760,10 +760,10 @@ class KVCache(abc.ABC):
def register_layer_transfer_counter(self, layer_transfer_counter: LayerDoneCounter):
self.layer_transfer_counter = layer_transfer_counter
def get_cpu_copy(self, indices, **kwargs):
def get_cpu_copy(self, indices, mamba_indices=None):
raise NotImplementedError()
def load_cpu_copy(self, kv_cache_cpu, indices, **kwargs):
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
raise NotImplementedError()
def maybe_get_custom_mem_pool(self):
@@ -965,7 +965,7 @@ class MHATokenToKVPool(KVCache):
]
return kv_data_ptrs, kv_data_lens, kv_item_lens
def get_cpu_copy(self, indices, **kwargs):
def get_cpu_copy(self, indices, mamba_indices=None):
torch.cuda.synchronize()
kv_cache_cpu = []
chunk_size = self.cpu_offloading_chunk_size
@@ -983,7 +983,7 @@ class MHATokenToKVPool(KVCache):
torch.cuda.synchronize()
return kv_cache_cpu
def load_cpu_copy(self, kv_cache_cpu, indices, **kwargs):
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
torch.cuda.synchronize()
chunk_size = self.cpu_offloading_chunk_size
for layer_id in range(self.layer_num):
@@ -1449,7 +1449,7 @@ class HybridLinearKVPool(KVCache):
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
self.full_kv_pool.move_kv_cache(tgt_loc, src_loc)
def get_cpu_copy(self, indices, mamba_indices=None, **kwargs):
def get_cpu_copy(self, indices, mamba_indices=None):
kv_cpu = self.full_kv_pool.get_cpu_copy(indices)
mamba_cpu = (
self.mamba_pool.get_cpu_copy(mamba_indices)
@@ -1458,7 +1458,7 @@ class HybridLinearKVPool(KVCache):
)
return kv_cpu, mamba_cpu
def load_cpu_copy(self, cache_cpu, indices, mamba_indices=None, **kwargs):
def load_cpu_copy(self, cache_cpu, indices, mamba_indices=None):
kv_cpu, mamba_cpu = cache_cpu
self.full_kv_pool.load_cpu_copy(kv_cpu, indices)
if mamba_cpu is not None and mamba_indices is not None:
@@ -1695,7 +1695,7 @@ class MLATokenToKVPool(KVCache):
get_mla_kv_buffer_triton(kv_buffer, loc, cache_k_nope, cache_k_rope)
return cache_k_nope, cache_k_rope
def get_cpu_copy(self, indices, **kwargs):
def get_cpu_copy(self, indices, mamba_indices=None):
torch.cuda.synchronize()
kv_cache_cpu = []
chunk_size = self.cpu_offloading_chunk_size
@@ -1710,7 +1710,7 @@ class MLATokenToKVPool(KVCache):
torch.cuda.synchronize()
return kv_cache_cpu
def load_cpu_copy(self, kv_cache_cpu, indices, **kwargs):
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
torch.cuda.synchronize()
chunk_size = self.cpu_offloading_chunk_size
for layer_id in range(self.layer_num):
@@ -198,7 +198,7 @@ class SWAKVPool(KVCache):
src_loc_swa = self.translate_loc_from_full_to_swa(src_loc)
self.swa_kv_pool.move_kv_cache(tgt_loc_swa, src_loc_swa)
def get_cpu_copy(self, indices):
def get_cpu_copy(self, indices, mamba_indices=None):
# For SWA, we need to copy KV cache from both full and SWA pools
# The indices are for the full pool, and we use mapping to get SWA indices
full_kv_cpu = self.full_kv_pool.get_cpu_copy(indices)
@@ -213,7 +213,7 @@ class SWAKVPool(KVCache):
return {"full": full_kv_cpu, "swa": swa_kv_cpu}
def load_cpu_copy(self, kv_cache_cpu, indices):
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
# Load KV cache back from CPU to both full and SWA pools
# Note: indices here are NEW indices (newly allocated), different from get_cpu_copy indices
full_kv_cpu = kv_cache_cpu["full"]
@@ -484,8 +484,10 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.is_not_in_free_group = True
self.free_group = []
def get_cpu_copy(self, indices):
return self._kvcache.get_cpu_copy(indices)
def get_cpu_copy(self, indices, mamba_indices=None):
return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices)
def load_cpu_copy(self, kv_cache_cpu, indices):
return self._kvcache.load_cpu_copy(kv_cache_cpu, indices)
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
return self._kvcache.load_cpu_copy(
kv_cache_cpu, indices, mamba_indices=mamba_indices
)