[misc] Fold the allocator free-group flag into free_group (#36739)

This commit is contained in:
Liangsheng Yin
2026-08-27 16:12:47 -07:00
committed by GitHub
parent 6ccfeb59bc
commit 5d52f02f22
7 changed files with 40 additions and 101 deletions
@@ -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)
+13 -17
View File
@@ -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 --