diff --git a/python/sglang/srt/hardware_backend/npu/allocator_npu.py b/python/sglang/srt/hardware_backend/npu/allocator_npu.py index e59dbf6a0..bbc3e9dec 100644 --- a/python/sglang/srt/hardware_backend/npu/allocator_npu.py +++ b/python/sglang/srt/hardware_backend/npu/allocator_npu.py @@ -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) diff --git a/python/sglang/srt/mem_cache/allocator/base.py b/python/sglang/srt/mem_cache/allocator/base.py index 103324dcd..ef1de7121 100644 --- a/python/sglang/srt/mem_cache/allocator/base.py +++ b/python/sglang/srt/mem_cache/allocator/base.py @@ -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: diff --git a/python/sglang/srt/mem_cache/allocator/hisparse.py b/python/sglang/srt/mem_cache/allocator/hisparse.py index 8ebfee249..dc283b617 100644 --- a/python/sglang/srt/mem_cache/allocator/hisparse.py +++ b/python/sglang/srt/mem_cache/allocator/hisparse.py @@ -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)) diff --git a/python/sglang/srt/mem_cache/allocator/paged.py b/python/sglang/srt/mem_cache/allocator/paged.py index b092911f2..f21c077f6 100755 --- a/python/sglang/srt/mem_cache/allocator/paged.py +++ b/python/sglang/srt/mem_cache/allocator/paged.py @@ -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) diff --git a/python/sglang/srt/mem_cache/allocator/swa.py b/python/sglang/srt/mem_cache/allocator/swa.py index f3540340f..2deffeeb1 100644 --- a/python/sglang/srt/mem_cache/allocator/swa.py +++ b/python/sglang/srt/mem_cache/allocator/swa.py @@ -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 diff --git a/python/sglang/srt/mem_cache/allocator/token.py b/python/sglang/srt/mem_cache/allocator/token.py index 06aceef12..fd0f32682 100644 --- a/python/sglang/srt/mem_cache/allocator/token.py +++ b/python/sglang/srt/mem_cache/allocator/token.py @@ -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: diff --git a/python/sglang/srt/mem_cache/multi_ended_allocator.py b/python/sglang/srt/mem_cache/multi_ended_allocator.py index 46eef0302..4e8680a8e 100644 --- a/python/sglang/srt/mem_cache/multi_ended_allocator.py +++ b/python/sglang/srt/mem_cache/multi_ended_allocator.py @@ -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 --