[misc] Fold the allocator free-group flag into free_group (#36739)
This commit is contained in:
@@ -141,7 +141,7 @@ class NPUPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator):
|
||||
if free_index.numel() == 0:
|
||||
return
|
||||
|
||||
if self.is_not_in_free_group:
|
||||
if self.free_group is None:
|
||||
device = free_index.device
|
||||
free_page_indices = torch.unique(free_index.cpu() // self.page_size)
|
||||
free_page_indices = free_page_indices.to(device)
|
||||
|
||||
@@ -44,8 +44,8 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
||||
|
||||
self.free_pages = None
|
||||
self.release_pages = None
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
# None: free right away. A list: hold frees until free_group_end().
|
||||
self.free_group: list[torch.Tensor] | None = None
|
||||
|
||||
@property
|
||||
def size_full(self):
|
||||
@@ -61,13 +61,12 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
||||
return self._kvcache
|
||||
|
||||
def free_group_begin(self):
|
||||
self.is_not_in_free_group = False
|
||||
self.free_group = []
|
||||
|
||||
def free_group_end(self):
|
||||
self.is_not_in_free_group = True
|
||||
if self.free_group:
|
||||
self.free(torch.cat(self.free_group))
|
||||
pending, self.free_group = self.free_group, None
|
||||
if pending:
|
||||
self.free(torch.cat(pending))
|
||||
|
||||
@staticmethod
|
||||
def _copy_for_free_group(free_index: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
@@ -61,8 +61,7 @@ class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
|
||||
self.free_pages = None
|
||||
self.release_pages = None
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
self.clear()
|
||||
self._kvcache.register_mapping(
|
||||
weakref.proxy(self.full_to_hisparse_device_index_mapping)
|
||||
@@ -166,8 +165,6 @@ class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
return buffer_indices
|
||||
|
||||
def free_hisparse_indices(self, buffer_indices: torch.Tensor):
|
||||
# disable free group mechanism for device buffer free
|
||||
self.hisparse_attn_allocator.is_not_in_free_group = True
|
||||
self.hisparse_attn_allocator.free(buffer_indices[buffer_indices > 0])
|
||||
|
||||
def get_last_loc_compressed(self, last_locs: torch.Tensor):
|
||||
@@ -246,9 +243,9 @@ class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.hisparse_attn_allocator.clear()
|
||||
# Note: the last item is -1, we don't clear it, see the comment in __init__
|
||||
self.full_to_hisparse_device_index_mapping[:-1].fill_(0)
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
|
||||
# No deferred frees: free() must clear the full-to-hisparse mapping at once.
|
||||
def free_group_begin(self):
|
||||
return
|
||||
|
||||
@@ -258,11 +255,8 @@ class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
def free(self, free_index: torch.Tensor):
|
||||
if free_index.numel() == 0:
|
||||
return
|
||||
if self.is_not_in_free_group:
|
||||
self.logical_attn_allocator.free(free_index)
|
||||
self.free_hisparse(free_index)
|
||||
else:
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
self.logical_attn_allocator.free(free_index)
|
||||
self.free_hisparse(free_index)
|
||||
assert (
|
||||
self.logical_attn_allocator.available_size()
|
||||
<= self.logical_attn_allocator.size
|
||||
@@ -321,8 +315,7 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.need_sort = logical_attn_allocator.need_sort
|
||||
self.free_pages = None
|
||||
self.release_pages = None
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
self.clear()
|
||||
|
||||
self.hisparse_kvcache.register_mapping(
|
||||
@@ -444,7 +437,6 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
surplus_pages = torch.unique(surplus // self.hisparse_page_size)
|
||||
pure_surplus = surplus_pages[~torch.isin(surplus_pages, buffer_pages)]
|
||||
if pure_surplus.numel() > 0:
|
||||
self.hisparse_attn_allocator.is_not_in_free_group = True
|
||||
self.hisparse_attn_allocator.free(
|
||||
pure_surplus * self.hisparse_page_size
|
||||
)
|
||||
@@ -474,7 +466,6 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
return buffer_indices
|
||||
|
||||
def free_hisparse_indices(self, buffer_indices: torch.Tensor):
|
||||
self.hisparse_attn_allocator.is_not_in_free_group = True
|
||||
self.hisparse_attn_allocator.free(buffer_indices[buffer_indices > 0])
|
||||
|
||||
def get_last_loc_compressed(self, last_locs: torch.Tensor):
|
||||
@@ -575,14 +566,13 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.hisparse_attn_allocator.clear()
|
||||
|
||||
self.full_to_hisparse_device_index_mapping[:-1].fill_(0)
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
|
||||
def free(self, free_index: torch.Tensor):
|
||||
if free_index.numel() == 0:
|
||||
return
|
||||
|
||||
if self.is_not_in_free_group:
|
||||
if self.free_group is None:
|
||||
self.logical_attn_allocator.free(free_index)
|
||||
else:
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
|
||||
@@ -262,7 +262,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
if free_index.numel() == 0:
|
||||
return
|
||||
|
||||
if self.is_not_in_free_group:
|
||||
if self.free_group is None:
|
||||
self._release_page_ids(torch.unique(free_index // self.page_size))
|
||||
else:
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
@@ -293,7 +293,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
torch.unique(free_index.cpu() // ps),
|
||||
)
|
||||
|
||||
if self.is_not_in_free_group:
|
||||
if self.free_group is None:
|
||||
self._release_page_ids(*(p // ps for p in pieces))
|
||||
if self.debug_mode:
|
||||
self._debug_check_no_duplicate_pages()
|
||||
@@ -333,8 +333,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.free_pages = torch.arange(
|
||||
1, self.num_pages + 1, dtype=torch.int64, device=self.device
|
||||
)
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
self.free_page_reps_group = []
|
||||
self.release_pages = torch.empty((0,), dtype=torch.int64, device=self.device)
|
||||
|
||||
|
||||
@@ -93,8 +93,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.need_sort = need_sort
|
||||
self.free_pages = None
|
||||
self.release_pages = None
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
self.swa_free_group = []
|
||||
|
||||
self._kvcache = kvcache
|
||||
@@ -319,7 +318,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
return
|
||||
|
||||
# NOTE: the API is not idempotent.
|
||||
if self.is_not_in_free_group:
|
||||
if self.free_group is None:
|
||||
self.full_attn_allocator.free(free_index)
|
||||
self.free_swa(free_index)
|
||||
else:
|
||||
@@ -363,7 +362,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
swa_indices = swa_indices[swa_indices > 0]
|
||||
self.clear_full_to_swa_mapping(mapping_indices)
|
||||
|
||||
if not self.is_not_in_free_group:
|
||||
if self.free_group is not None:
|
||||
# Resolve ownership now. A cache action later in this group may
|
||||
# install a new mapping for the same full index.
|
||||
self.swa_free_group.append(swa_indices)
|
||||
@@ -408,8 +407,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.full_attn_allocator.clear()
|
||||
# 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 = []
|
||||
self.free_group = None
|
||||
self.swa_free_group = []
|
||||
|
||||
def get_cpu_copy(self, indices, mamba_indices=None):
|
||||
@@ -460,8 +458,7 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
|
||||
self.free_pages = None
|
||||
self.release_pages = None
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
|
||||
self._kvcache = kvcache
|
||||
self.swa_attn_allocator.clear()
|
||||
@@ -505,7 +502,7 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
def free(self, free_index: torch.Tensor):
|
||||
if free_index.numel() == 0:
|
||||
return
|
||||
if self.is_not_in_free_group:
|
||||
if self.free_group is None:
|
||||
self.swa_attn_allocator.free(free_index[free_index > 0])
|
||||
else:
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
@@ -514,22 +511,21 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
def free_swa(self, free_index: torch.Tensor):
|
||||
if free_index.numel() == 0:
|
||||
return
|
||||
if self.is_not_in_free_group:
|
||||
if self.free_group is None:
|
||||
self.swa_attn_allocator.free(free_index[free_index > 0])
|
||||
else:
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
|
||||
# Not inherited: the SWA parent's hooks drive swa_free_group,
|
||||
# which this pure-SWA variant does not have.
|
||||
def free_group_begin(self):
|
||||
self.is_not_in_free_group = False
|
||||
self.free_group = []
|
||||
|
||||
def free_group_end(self):
|
||||
self.is_not_in_free_group = True
|
||||
if self.free_group:
|
||||
self.free(torch.cat(self.free_group))
|
||||
self.free_group = []
|
||||
pending, self.free_group = self.free_group, None
|
||||
if pending:
|
||||
self.free(torch.cat(pending))
|
||||
|
||||
def clear(self):
|
||||
self.swa_attn_allocator.clear()
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
|
||||
@@ -44,8 +44,7 @@ class TokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self.free_pages = torch.arange(
|
||||
1, self.size + 1, dtype=torch.int64, device=self.device
|
||||
)
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
self.release_pages = torch.empty((0,), dtype=torch.int64, device=self.device)
|
||||
|
||||
def available_size(self):
|
||||
@@ -67,7 +66,7 @@ class TokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
if free_index.numel() == 0:
|
||||
return
|
||||
|
||||
if self.is_not_in_free_group:
|
||||
if self.free_group is None:
|
||||
if self.need_sort:
|
||||
self.release_pages = torch.cat((self.release_pages, free_index))
|
||||
else:
|
||||
|
||||
@@ -288,8 +288,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||
)
|
||||
else:
|
||||
self.free_virtual_ids = None
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group: List[torch.Tensor] = []
|
||||
self.free_group = None
|
||||
self._inverse_history.clear()
|
||||
self._free_phys_pages = torch.empty(0, dtype=torch.int64, device=self.device)
|
||||
self._pending_reuse.clear()
|
||||
@@ -980,7 +979,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||
with record_function("MultiEndedAlloc.free"):
|
||||
if free_index is None or free_index.numel() == 0:
|
||||
return
|
||||
if not self.is_not_in_free_group:
|
||||
if self.free_group is not None:
|
||||
self.free_group.append(self._copy_for_free_group(free_index))
|
||||
return
|
||||
if self.lazy_compaction:
|
||||
@@ -1703,19 +1702,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||
f"Caller: {callers}."
|
||||
)
|
||||
|
||||
# -- free-group --
|
||||
|
||||
def free_group_begin(self) -> None:
|
||||
self.is_not_in_free_group = False
|
||||
self.free_group = []
|
||||
|
||||
def free_group_end(self) -> None:
|
||||
self.is_not_in_free_group = True
|
||||
if self.free_group:
|
||||
merged = torch.cat(self.free_group)
|
||||
self.free_group = []
|
||||
self.free(merged)
|
||||
|
||||
|
||||
class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
"""Composite allocator for the MHA (full-attn) + Mamba hybrid pair.
|
||||
@@ -1786,8 +1772,7 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
# pure PHYSICAL store. The full-attn KV pool needs no allocator either —
|
||||
# write locations are resolved in the attention metadata.
|
||||
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group: List[torch.Tensor] = []
|
||||
self.free_group = None
|
||||
# Base init left these None; we use watermark math, not free-lists.
|
||||
self.free_pages = torch.empty(0, dtype=torch.int64, device=device)
|
||||
self.release_pages = torch.empty(0, dtype=torch.int64, device=device)
|
||||
@@ -1970,31 +1955,17 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
with record_function("UnifiedMambaAlloc.free"):
|
||||
if free_index is None or free_index.numel() == 0:
|
||||
return
|
||||
if not self.is_not_in_free_group:
|
||||
if self.free_group is not None:
|
||||
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()
|
||||
self.mamba_allocator.clear_inverse_history()
|
||||
|
||||
def free_group_begin(self) -> None:
|
||||
self.is_not_in_free_group = False
|
||||
self.free_group = []
|
||||
|
||||
def free_group_end(self) -> None:
|
||||
self.is_not_in_free_group = True
|
||||
if self.free_group:
|
||||
merged = torch.cat(self.free_group)
|
||||
self.free_group = []
|
||||
self.full_attn_allocator.free(merged)
|
||||
self.full_attn_allocator.clear_inverse_history()
|
||||
self.mamba_allocator.clear_inverse_history()
|
||||
|
||||
def clear(self) -> None:
|
||||
self.full_attn_allocator.clear()
|
||||
self.mamba_allocator.clear()
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
|
||||
# -- Lazy compaction hooks --
|
||||
|
||||
@@ -2131,8 +2102,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
swa_allocator=self.swa_attn_allocator,
|
||||
)
|
||||
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group: List[torch.Tensor] = []
|
||||
self.free_group = None
|
||||
# Empty (not None) for the leak checker.
|
||||
self.free_pages = torch.empty(0, dtype=torch.int64, device=device)
|
||||
self.release_pages = torch.empty(0, dtype=torch.int64, device=device)
|
||||
@@ -2439,7 +2409,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
with record_function("UnifiedSWAAlloc.free"):
|
||||
if free_index is None or free_index.numel() == 0:
|
||||
return
|
||||
if not self.is_not_in_free_group:
|
||||
if self.free_group is not None:
|
||||
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
|
||||
@@ -2491,24 +2461,10 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
||||
# Paired with set_full_to_swa_mapping: shared mode has no mapping tensor.
|
||||
return
|
||||
|
||||
# -- free-group --
|
||||
|
||||
def free_group_begin(self) -> None:
|
||||
self.is_not_in_free_group = False
|
||||
self.free_group = []
|
||||
|
||||
def free_group_end(self) -> None:
|
||||
self.is_not_in_free_group = True
|
||||
if self.free_group:
|
||||
merged = torch.cat(self.free_group)
|
||||
self.free_group = []
|
||||
self.free(merged)
|
||||
|
||||
def clear(self) -> None:
|
||||
self.full_attn_allocator.clear()
|
||||
self.swa_attn_allocator.clear()
|
||||
self.is_not_in_free_group = True
|
||||
self.free_group = []
|
||||
self.free_group = None
|
||||
|
||||
# -- Lazy compaction hooks --
|
||||
|
||||
|
||||
Reference in New Issue
Block a user