[Bugfix] Fix batched KV free aliasing (#34067)
This commit is contained in:
@@ -149,7 +149,7 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator):
|
||||
else:
|
||||
self.free_pages = torch.cat((free_page_indices, self.free_pages))
|
||||
else:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
|
||||
if self.debug_mode:
|
||||
assert len(torch.unique(self.free_pages)) == len(self.free_pages)
|
||||
|
||||
@@ -69,6 +69,11 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
||||
if self.free_group:
|
||||
self.free(torch.cat(self.free_group))
|
||||
|
||||
@staticmethod
|
||||
def _copy_for_free_group(free_index: torch.Tensor) -> torch.Tensor:
|
||||
"""Take ownership before a caller can mutate a deferred tensor view."""
|
||||
return free_index.clone()
|
||||
|
||||
def merge_and_sort_free(self):
|
||||
if len(self.release_pages) > 0:
|
||||
self.free_pages = torch.cat((self.free_pages, self.release_pages))
|
||||
|
||||
@@ -262,7 +262,7 @@ class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.logical_attn_allocator.free(free_index)
|
||||
self.free_hisparse(free_index)
|
||||
else:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
assert (
|
||||
self.logical_attn_allocator.available_size()
|
||||
<= self.logical_attn_allocator.size
|
||||
@@ -585,4 +585,4 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
if self.is_not_in_free_group:
|
||||
self.logical_attn_allocator.free(free_index)
|
||||
else:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
|
||||
@@ -265,7 +265,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
if self.is_not_in_free_group:
|
||||
self._release_page_ids(torch.unique(free_index // self.page_size))
|
||||
else:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
|
||||
if self.debug_mode:
|
||||
self._debug_check_no_duplicate_pages()
|
||||
@@ -298,7 +298,9 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
if self.debug_mode:
|
||||
self._debug_check_no_duplicate_pages()
|
||||
else:
|
||||
self.free_page_reps_group.extend(pieces)
|
||||
self.free_page_reps_group.extend(
|
||||
self._copy_for_free_group(piece) for piece in pieces
|
||||
)
|
||||
|
||||
def _debug_check_no_duplicate_pages(self):
|
||||
# span both containers: need_sort (PD disagg) routes frees into release_pages
|
||||
|
||||
@@ -325,7 +325,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.full_attn_allocator.free(free_index)
|
||||
self.free_swa(free_index)
|
||||
else:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
assert (
|
||||
self.full_attn_allocator.available_size() <= self.full_attn_allocator.size
|
||||
)
|
||||
@@ -350,7 +350,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
return
|
||||
|
||||
if not self.is_not_in_free_group:
|
||||
self.swa_free_group.append(free_index)
|
||||
self.swa_free_group.append(self._copy_for_free_group(free_index))
|
||||
return
|
||||
|
||||
if self.page_size == 1:
|
||||
@@ -500,7 +500,7 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
if self.is_not_in_free_group:
|
||||
self.swa_attn_allocator.free(free_index[free_index > 0])
|
||||
else:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
assert self.swa_attn_allocator.available_size() <= self.swa_attn_allocator.size
|
||||
|
||||
def free_swa(self, free_index: torch.Tensor):
|
||||
@@ -509,7 +509,7 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
if self.is_not_in_free_group:
|
||||
self.swa_attn_allocator.free(free_index[free_index > 0])
|
||||
else:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
|
||||
def free_group_begin(self):
|
||||
self.is_not_in_free_group = False
|
||||
|
||||
@@ -73,7 +73,7 @@ class TokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
else:
|
||||
self.free_pages = torch.cat((self.free_pages, free_index))
|
||||
else:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
|
||||
def get_cpu_copy(self, indices, mamba_indices=None):
|
||||
return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices)
|
||||
|
||||
@@ -954,7 +954,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||
if free_index is None or free_index.numel() == 0:
|
||||
return
|
||||
if not self.is_not_in_free_group:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
return
|
||||
if self.lazy_compaction:
|
||||
self._free_lazy(free_index)
|
||||
@@ -1928,7 +1928,7 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
if free_index is None or free_index.numel() == 0:
|
||||
return
|
||||
if not self.is_not_in_free_group:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
return
|
||||
self.full_attn_allocator.free(free_index)
|
||||
self.full_attn_allocator.clear_inverse_history()
|
||||
@@ -2397,7 +2397,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
if free_index is None or free_index.numel() == 0:
|
||||
return
|
||||
if not self.is_not_in_free_group:
|
||||
self.free_group.append(free_index)
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
return
|
||||
# Free both peers; the per-sub-pool v2p IS the mapping, so order isn't
|
||||
# load-bearing. Filter the swa side to skip already-tombstoned virtuals
|
||||
|
||||
@@ -511,10 +511,8 @@ class RadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||
)
|
||||
new_prefix_len = result.prefix_len
|
||||
|
||||
# Use the out-of-place values copy so the allocator can safely defer or group
|
||||
# this free after req_to_token is overwritten below.
|
||||
self.token_to_kv_pool_allocator.free_segment(
|
||||
values[req.cache_protected_len : new_prefix_len],
|
||||
kv_indices[req.cache_protected_len : new_prefix_len],
|
||||
start_pos=req.cache_protected_len,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user