[Refactor] Remove dead allocator backup_state / restore_state (#32709)
This commit is contained in:
@@ -776,39 +776,6 @@ class DSV4NPUTokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
if slots.numel() > 0:
|
if slots.numel() > 0:
|
||||||
allocator.free(slots.to(torch.int64))
|
allocator.free(slots.to(torch.int64))
|
||||||
|
|
||||||
def backup_state(self):
|
|
||||||
# EAGLE/NEXTN draft preprocess allocates speculative c{4,128} KV via
|
|
||||||
# alloc_extend(backup_state=True) and rolls it back with restore_state.
|
|
||||||
# The base SWATokenToKVPoolAllocator only snapshots the full + SWA pools,
|
|
||||||
# so without this override the draft's c{4,128} (+ state) slots are never
|
|
||||||
# rolled back -> they leak every draft step until the c4 pool exhausts.
|
|
||||||
# Snapshot the sub-allocators alongside the base pools.
|
|
||||||
return (
|
|
||||||
super().backup_state(),
|
|
||||||
self.c4_attn_allocator.backup_state(),
|
|
||||||
self.c128_attn_allocator.backup_state(),
|
|
||||||
(
|
|
||||||
self.c4_state_attn_allocator.backup_state()
|
|
||||||
if self.c4_state_attn_allocator is not None
|
|
||||||
else None
|
|
||||||
),
|
|
||||||
(
|
|
||||||
self.c128_state_attn_allocator.backup_state()
|
|
||||||
if self.c128_state_attn_allocator is not None
|
|
||||||
else None
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
def restore_state(self, state):
|
|
||||||
base, c4, c128, c4_state, c128_state = state
|
|
||||||
super().restore_state(base)
|
|
||||||
self.c4_attn_allocator.restore_state(c4)
|
|
||||||
self.c128_attn_allocator.restore_state(c128)
|
|
||||||
if self.c4_state_attn_allocator is not None and c4_state is not None:
|
|
||||||
self.c4_state_attn_allocator.restore_state(c4_state)
|
|
||||||
if self.c128_state_attn_allocator is not None and c128_state is not None:
|
|
||||||
self.c128_state_attn_allocator.restore_state(c128_state)
|
|
||||||
|
|
||||||
def clear(self):
|
def clear(self):
|
||||||
super().clear()
|
super().clear()
|
||||||
# super().__init__ calls clear() before our sub-allocators exist;
|
# super().__init__ calls clear() before our sub-allocators exist;
|
||||||
|
|||||||
@@ -146,15 +146,10 @@ def get_last_loc_torch(
|
|||||||
def alloc_token_slots(
|
def alloc_token_slots(
|
||||||
tree_cache: BasePrefixCache,
|
tree_cache: BasePrefixCache,
|
||||||
num_tokens: int,
|
num_tokens: int,
|
||||||
backup_state: bool = False,
|
|
||||||
):
|
):
|
||||||
allocator = tree_cache.token_to_kv_pool_allocator
|
allocator = tree_cache.token_to_kv_pool_allocator
|
||||||
evict_from_tree_cache(tree_cache, num_tokens)
|
evict_from_tree_cache(tree_cache, num_tokens)
|
||||||
|
|
||||||
state = None
|
|
||||||
if backup_state:
|
|
||||||
state = allocator.backup_state()
|
|
||||||
|
|
||||||
out_cache_loc = allocator.alloc(num_tokens)
|
out_cache_loc = allocator.alloc(num_tokens)
|
||||||
|
|
||||||
if out_cache_loc is None:
|
if out_cache_loc is None:
|
||||||
@@ -168,7 +163,7 @@ def alloc_token_slots(
|
|||||||
tree_cache.pretty_print()
|
tree_cache.pretty_print()
|
||||||
raise RuntimeError(error_msg)
|
raise RuntimeError(error_msg)
|
||||||
|
|
||||||
return (out_cache_loc, state) if backup_state else out_cache_loc
|
return out_cache_loc
|
||||||
|
|
||||||
|
|
||||||
def _compute_dsv4_state_lens(batch, *, is_decode: bool):
|
def _compute_dsv4_state_lens(batch, *, is_decode: bool):
|
||||||
@@ -203,7 +198,6 @@ def alloc_paged_token_slots_extend(
|
|||||||
seq_lens_cpu: torch.Tensor,
|
seq_lens_cpu: torch.Tensor,
|
||||||
last_loc: torch.Tensor,
|
last_loc: torch.Tensor,
|
||||||
extend_num_tokens: int,
|
extend_num_tokens: int,
|
||||||
backup_state: bool = False,
|
|
||||||
req_pool_indices: Optional[torch.Tensor] = None,
|
req_pool_indices: Optional[torch.Tensor] = None,
|
||||||
dsv4_state_lens: Optional[DSV4StateLens] = None,
|
dsv4_state_lens: Optional[DSV4StateLens] = None,
|
||||||
batch=None,
|
batch=None,
|
||||||
@@ -213,10 +207,6 @@ def alloc_paged_token_slots_extend(
|
|||||||
num_tokens = extend_num_tokens + len(seq_lens_cpu) * allocator.page_size
|
num_tokens = extend_num_tokens + len(seq_lens_cpu) * allocator.page_size
|
||||||
evict_from_tree_cache(tree_cache, num_tokens)
|
evict_from_tree_cache(tree_cache, num_tokens)
|
||||||
|
|
||||||
state = None
|
|
||||||
if backup_state:
|
|
||||||
state = allocator.backup_state()
|
|
||||||
|
|
||||||
is_dsv4 = req_pool_indices is not None and hasattr(allocator, "c4_attn_allocator")
|
is_dsv4 = req_pool_indices is not None and hasattr(allocator, "c4_attn_allocator")
|
||||||
extra_alloc_kwargs = {}
|
extra_alloc_kwargs = {}
|
||||||
if is_dsv4:
|
if is_dsv4:
|
||||||
@@ -256,7 +246,7 @@ def alloc_paged_token_slots_extend(
|
|||||||
tree_cache.pretty_print()
|
tree_cache.pretty_print()
|
||||||
raise RuntimeError(error_msg)
|
raise RuntimeError(error_msg)
|
||||||
|
|
||||||
return (out_cache_loc, state) if backup_state else out_cache_loc
|
return out_cache_loc
|
||||||
|
|
||||||
|
|
||||||
def alloc_req_slots(
|
def alloc_req_slots(
|
||||||
|
|||||||
@@ -60,12 +60,6 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
|||||||
def get_kvcache(self):
|
def get_kvcache(self):
|
||||||
return self._kvcache
|
return self._kvcache
|
||||||
|
|
||||||
def restore_state(self, state):
|
|
||||||
self.free_pages, self.release_pages = state
|
|
||||||
|
|
||||||
def backup_state(self):
|
|
||||||
return (self.free_pages, self.release_pages)
|
|
||||||
|
|
||||||
def free_group_begin(self):
|
def free_group_begin(self):
|
||||||
self.is_not_in_free_group = False
|
self.is_not_in_free_group = False
|
||||||
self.free_group = []
|
self.free_group = []
|
||||||
|
|||||||
@@ -365,17 +365,6 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
)
|
)
|
||||||
return (pages[:, None] * self.page_size + page_offsets[None, :]).reshape(-1)
|
return (pages[:, None] * self.page_size + page_offsets[None, :]).reshape(-1)
|
||||||
|
|
||||||
def backup_state(self):
|
|
||||||
return [
|
|
||||||
self.full_attn_allocator.backup_state(),
|
|
||||||
self.swa_attn_allocator.backup_state(),
|
|
||||||
]
|
|
||||||
|
|
||||||
def restore_state(self, state):
|
|
||||||
assert len(state) == 2
|
|
||||||
self.full_attn_allocator.restore_state(state[0])
|
|
||||||
self.swa_attn_allocator.restore_state(state[1])
|
|
||||||
|
|
||||||
def resize(self, config) -> None:
|
def resize(self, config) -> None:
|
||||||
size_full = int(config.full_max_total_num_tokens)
|
size_full = int(config.full_max_total_num_tokens)
|
||||||
size_swa = int(config.swa_max_total_num_tokens)
|
size_swa = int(config.swa_max_total_num_tokens)
|
||||||
@@ -512,12 +501,6 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
self.free(torch.cat(self.free_group))
|
self.free(torch.cat(self.free_group))
|
||||||
self.free_group = []
|
self.free_group = []
|
||||||
|
|
||||||
def backup_state(self):
|
|
||||||
return self.swa_attn_allocator.backup_state()
|
|
||||||
|
|
||||||
def restore_state(self, state):
|
|
||||||
self.swa_attn_allocator.restore_state(state)
|
|
||||||
|
|
||||||
def clear(self):
|
def clear(self):
|
||||||
self.swa_attn_allocator.clear()
|
self.swa_attn_allocator.clear()
|
||||||
self.is_not_in_free_group = True
|
self.is_not_in_free_group = True
|
||||||
|
|||||||
@@ -285,33 +285,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
self._inflight_forward = None
|
self._inflight_forward = None
|
||||||
self._latest_forward_done_event = None
|
self._latest_forward_done_event = None
|
||||||
|
|
||||||
def backup_state(self):
|
|
||||||
# Spec-decode allocates only inside a backup window (no free), so
|
|
||||||
# `_inverse_history` doesn't grow under correct usage.
|
|
||||||
return (
|
|
||||||
self.watermark_physical,
|
|
||||||
(len(self.free_virtual_ids) if self.is_id_owner else None),
|
|
||||||
len(self._inverse_history),
|
|
||||||
)
|
|
||||||
|
|
||||||
def restore_state(self, state):
|
|
||||||
watermark, n_free_virtual, n_inverse = state
|
|
||||||
self.watermark_physical = watermark
|
|
||||||
if self.is_id_owner and n_free_virtual is not None:
|
|
||||||
pass # spec asserted off; no free-list rollback.
|
|
||||||
new_entries = self._inverse_history[n_inverse:]
|
|
||||||
if new_entries:
|
|
||||||
logger.warning(
|
|
||||||
"MultiEndedAllocator.restore_state: %d relocation(s) recorded inside "
|
|
||||||
"a backup window (sub_pool=%s). Eager compaction is not fully "
|
|
||||||
"reversible; SGLang's spec path should not produce a free() inside a "
|
|
||||||
"backup window.",
|
|
||||||
len(new_entries),
|
|
||||||
self.sub_pool_name,
|
|
||||||
)
|
|
||||||
del self._inverse_history[n_inverse:]
|
|
||||||
return new_entries
|
|
||||||
|
|
||||||
def clear_inverse_history(self) -> None:
|
def clear_inverse_history(self) -> None:
|
||||||
self._inverse_history.clear()
|
self._inverse_history.clear()
|
||||||
|
|
||||||
@@ -1890,18 +1863,6 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
self.full_attn_allocator.clear_inverse_history()
|
self.full_attn_allocator.clear_inverse_history()
|
||||||
self.mamba_allocator.clear_inverse_history()
|
self.mamba_allocator.clear_inverse_history()
|
||||||
|
|
||||||
def backup_state(self):
|
|
||||||
return [
|
|
||||||
self.full_attn_allocator.backup_state(),
|
|
||||||
self.mamba_allocator.backup_state(),
|
|
||||||
]
|
|
||||||
|
|
||||||
def restore_state(self, state):
|
|
||||||
assert len(state) == 2
|
|
||||||
full_rollback = self.full_attn_allocator.restore_state(state[0])
|
|
||||||
mamba_rollback = self.mamba_allocator.restore_state(state[1])
|
|
||||||
return full_rollback + mamba_rollback
|
|
||||||
|
|
||||||
def clear(self) -> None:
|
def clear(self) -> None:
|
||||||
self.full_attn_allocator.clear()
|
self.full_attn_allocator.clear()
|
||||||
self.mamba_allocator.clear()
|
self.mamba_allocator.clear()
|
||||||
@@ -2412,20 +2373,6 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
self.free_group = []
|
self.free_group = []
|
||||||
self.free(merged)
|
self.free(merged)
|
||||||
|
|
||||||
# -- spec-decode hooks (asserted off; preserved for future use) --
|
|
||||||
|
|
||||||
def backup_state(self):
|
|
||||||
return [
|
|
||||||
self.full_attn_allocator.backup_state(),
|
|
||||||
self.swa_attn_allocator.backup_state(),
|
|
||||||
]
|
|
||||||
|
|
||||||
def restore_state(self, state):
|
|
||||||
assert len(state) == 2
|
|
||||||
full_rollback = self.full_attn_allocator.restore_state(state[0])
|
|
||||||
swa_rollback = self.swa_attn_allocator.restore_state(state[1])
|
|
||||||
return full_rollback + swa_rollback
|
|
||||||
|
|
||||||
def clear(self) -> None:
|
def clear(self) -> None:
|
||||||
self.full_attn_allocator.clear()
|
self.full_attn_allocator.clear()
|
||||||
self.swa_attn_allocator.clear()
|
self.swa_attn_allocator.clear()
|
||||||
|
|||||||
Reference in New Issue
Block a user