Use a shared byte budget for unified hybrid-SWA memory (#36729)
Co-authored-by: yhzhuang <yhzhuang@fb.com> Co-authored-by: Cheng Wan <cheng.wan@radixark.ai>
This commit is contained in:
co-authored by
yhzhuang
Cheng Wan
parent
4da5599e93
commit
2929a39927
@@ -1808,6 +1808,9 @@ def post_capture_kv_sizing_planned(server_args: Any) -> bool:
|
|||||||
mla_enabled = use_mla_backend(server_args)
|
mla_enabled = use_mla_backend(server_args)
|
||||||
if not envs.SGLANG_ENABLE_POST_CAPTURE_KV_SIZING.get():
|
if not envs.SGLANG_ENABLE_POST_CAPTURE_KV_SIZING.get():
|
||||||
return False
|
return False
|
||||||
|
# Unified arenas are fully backed before capture and cannot resize afterward.
|
||||||
|
if cfg.enable_unified_memory:
|
||||||
|
return False
|
||||||
if cfg.device != "cuda":
|
if cfg.device != "cuda":
|
||||||
return False
|
return False
|
||||||
if cfg.dcp_size != 1:
|
if cfg.dcp_size != 1:
|
||||||
|
|||||||
@@ -50,9 +50,7 @@ from sglang.srt.mem_cache.allocator.hisparse import (
|
|||||||
DeepSeekV4HiSparseTokenToKVPoolAllocator,
|
DeepSeekV4HiSparseTokenToKVPoolAllocator,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator.swa import (
|
from sglang.srt.mem_cache.allocator.swa import (
|
||||||
PureSWATokenToKVPoolAllocator,
|
|
||||||
SWATokenToKVPoolAllocator,
|
SWATokenToKVPoolAllocator,
|
||||||
is_swa_req_ring,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
||||||
UnifiedMambaSWATokenToKVPoolAllocator,
|
UnifiedMambaSWATokenToKVPoolAllocator,
|
||||||
@@ -581,8 +579,6 @@ class PrefillAdder:
|
|||||||
self.prefill_tile_block_m = prefill_tile_block_m
|
self.prefill_tile_block_m = prefill_tile_block_m
|
||||||
self.tree_cache = tree_cache
|
self.tree_cache = tree_cache
|
||||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
||||||
# Per-request SWA ring: one fixed slot per request, not a token budget.
|
|
||||||
self._swa_req_ring = is_swa_req_ring(token_to_kv_pool_allocator)
|
|
||||||
self.running_batch = running_batch
|
self.running_batch = running_batch
|
||||||
self.new_token_ratio = new_token_ratio
|
self.new_token_ratio = new_token_ratio
|
||||||
self.rem_input_tokens = rem_input_tokens - num_mixed_decode_tokens
|
self.rem_input_tokens = rem_input_tokens - num_mixed_decode_tokens
|
||||||
@@ -595,8 +591,9 @@ class PrefillAdder:
|
|||||||
|
|
||||||
if self.rem_chunk_tokens is not None:
|
if self.rem_chunk_tokens is not None:
|
||||||
self.rem_chunk_tokens -= num_mixed_decode_tokens
|
self.rem_chunk_tokens -= num_mixed_decode_tokens
|
||||||
self.rem_total_token_offset = num_mixed_decode_tokens
|
self.memory_budget = token_to_kv_pool_allocator.create_prefill_budget(
|
||||||
self.cur_rem_token_offset = num_mixed_decode_tokens
|
tree_cache, num_mixed_decode_tokens=num_mixed_decode_tokens
|
||||||
|
)
|
||||||
|
|
||||||
self.req_states = None
|
self.req_states = None
|
||||||
self.can_run_list = []
|
self.can_run_list = []
|
||||||
@@ -612,7 +609,7 @@ class PrefillAdder:
|
|||||||
|
|
||||||
if running_batch is not None:
|
if running_batch is not None:
|
||||||
# Estimate the offset in the remaining token space
|
# Estimate the offset in the remaining token space
|
||||||
self.rem_total_token_offset += sum(
|
self.memory_budget.total_offset += sum(
|
||||||
[
|
[
|
||||||
self._get_running_request_total_token_offset(r)
|
self._get_running_request_total_token_offset(r)
|
||||||
for r in running_batch.reqs
|
for r in running_batch.reqs
|
||||||
@@ -625,13 +622,7 @@ class PrefillAdder:
|
|||||||
self.token_to_kv_pool_allocator,
|
self.token_to_kv_pool_allocator,
|
||||||
(SWATokenToKVPoolAllocator, DeepSeekV4HiSparseTokenToKVPoolAllocator),
|
(SWATokenToKVPoolAllocator, DeepSeekV4HiSparseTokenToKVPoolAllocator),
|
||||||
)
|
)
|
||||||
self.is_all_swa = isinstance(
|
|
||||||
self.token_to_kv_pool_allocator, PureSWATokenToKVPoolAllocator
|
|
||||||
)
|
|
||||||
self.is_hybrid_ssm_cache = self.tree_cache.supports_mamba()
|
self.is_hybrid_ssm_cache = self.tree_cache.supports_mamba()
|
||||||
|
|
||||||
self.rem_swa_token_offset = 0
|
|
||||||
|
|
||||||
# A new state slot eats shared-gap bytes that `rem_total_tokens` counts
|
# A new state slot eats shared-gap bytes that `rem_total_tokens` counts
|
||||||
# as free, so reserve per slot or admission over-commits. Gate on the
|
# as free, so reserve per slot or admission over-commits. Gate on the
|
||||||
# ALLOCATOR, not `is_hybrid_ssm_cache`: that is False for `ChunkCache`,
|
# ALLOCATOR, not `is_hybrid_ssm_cache`: that is False for `ChunkCache`,
|
||||||
@@ -718,114 +709,11 @@ class PrefillAdder:
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def rem_total_tokens(self):
|
def rem_total_tokens(self):
|
||||||
if self.is_all_swa:
|
return self.memory_budget.remaining_total
|
||||||
available_and_evictable = (
|
|
||||||
self.token_to_kv_pool_allocator.swa_available_size()
|
|
||||||
+ self.tree_cache.swa_evictable_size()
|
|
||||||
)
|
|
||||||
elif self.is_hybrid_swa:
|
|
||||||
available_and_evictable = (
|
|
||||||
self.token_to_kv_pool_allocator.full_available_size()
|
|
||||||
+ self.tree_cache.full_evictable_size()
|
|
||||||
)
|
|
||||||
elif self.is_hybrid_ssm_cache:
|
|
||||||
available_and_evictable = (
|
|
||||||
self.token_to_kv_pool_allocator.available_size()
|
|
||||||
+ self.tree_cache.full_evictable_size()
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
available_and_evictable = (
|
|
||||||
self.token_to_kv_pool_allocator.available_size()
|
|
||||||
+ self.tree_cache.evictable_size()
|
|
||||||
)
|
|
||||||
return available_and_evictable - self.rem_total_token_offset
|
|
||||||
|
|
||||||
@property
|
|
||||||
def rem_swa_tokens(self):
|
|
||||||
allocator = self.token_to_kv_pool_allocator
|
|
||||||
if self._swa_req_ring:
|
|
||||||
# swa_available_size() already reports ring capacity; tree
|
|
||||||
# swa_evictable is in linear token units and frees no ring space.
|
|
||||||
return allocator.swa_available_size() - self.rem_swa_token_offset
|
|
||||||
return (
|
|
||||||
allocator.swa_available_size()
|
|
||||||
+ self.tree_cache.swa_evictable_size()
|
|
||||||
- self.rem_swa_token_offset
|
|
||||||
)
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def cur_rem_tokens(self):
|
def cur_rem_tokens(self):
|
||||||
if self.is_all_swa:
|
return self.memory_budget.remaining_current
|
||||||
available_and_evictable = (
|
|
||||||
self.token_to_kv_pool_allocator.swa_available_size()
|
|
||||||
+ self.tree_cache.swa_evictable_size()
|
|
||||||
)
|
|
||||||
elif self.is_hybrid_swa:
|
|
||||||
available_and_evictable = (
|
|
||||||
self.token_to_kv_pool_allocator.full_available_size()
|
|
||||||
+ self.tree_cache.full_evictable_size()
|
|
||||||
)
|
|
||||||
elif self.is_hybrid_ssm_cache:
|
|
||||||
available_and_evictable = (
|
|
||||||
self.token_to_kv_pool_allocator.available_size()
|
|
||||||
+ self.tree_cache.full_evictable_size()
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
available_and_evictable = (
|
|
||||||
self.token_to_kv_pool_allocator.available_size()
|
|
||||||
+ self.tree_cache.evictable_size()
|
|
||||||
)
|
|
||||||
|
|
||||||
return available_and_evictable - self.cur_rem_token_offset
|
|
||||||
|
|
||||||
def _swa_budget_for_req(
|
|
||||||
self, extend_input_len: int, max_new_tokens: int, swa_host_hit_length: int = 0
|
|
||||||
) -> int:
|
|
||||||
"""SWA pool budget per request. Only valid when is_hybrid_swa is True.
|
|
||||||
|
|
||||||
With chunked prefill + overlap scheduler, the peak SWA occupancy is:
|
|
||||||
chunk N (running, not yet in tree) + sliding window (locked in tree)
|
|
||||||
+ chunk N+1 (new allocation)
|
|
||||||
Since chunk N and locked tokens are already excluded from
|
|
||||||
swa_available + swa_evictable, the budget only needs to cover the
|
|
||||||
chunk N+1 allocation plus decode headroom:
|
|
||||||
|
|
||||||
budget = max(alloc - window, 0) + min(extend + max_new_tokens, window) + page
|
|
||||||
|
|
||||||
where alloc = min(extend, rem_chunk); the min() cap keeps the two terms
|
|
||||||
from double-counting extend, so budget <= extend + max_new_tokens + page.
|
|
||||||
"""
|
|
||||||
allocator = self.token_to_kv_pool_allocator
|
|
||||||
if self._swa_req_ring:
|
|
||||||
# One ring slot per request, in the same unit as swa_available_size.
|
|
||||||
return allocator.swa_ring_cost_tokens
|
|
||||||
if self.rem_chunk_tokens is not None:
|
|
||||||
alloc = min(extend_input_len, self.rem_chunk_tokens)
|
|
||||||
else:
|
|
||||||
alloc = extend_input_len
|
|
||||||
window = self.tree_cache.sliding_window_size
|
|
||||||
return max(alloc - window, 0) + self._swa_reserved_tokens(
|
|
||||||
extend_input_len, max_new_tokens, swa_host_hit_length
|
|
||||||
)
|
|
||||||
|
|
||||||
def _swa_reserved_tokens(
|
|
||||||
self, extend_input_len: int, max_new_tokens: int, swa_host_hit_length: int = 0
|
|
||||||
) -> int:
|
|
||||||
"""SWA slots a request adds to its own sliding window + page slack + the
|
|
||||||
load-back charge. Shared floor of _swa_budget_for_req and _swa_chunk_cap.
|
|
||||||
|
|
||||||
The headroom is min(extend + decode, window), not a constant window: a
|
|
||||||
request contributes only extend + decode fresh tokens to its window and
|
|
||||||
a cached SWA prefix funds the rest. Charging a full window double-counted
|
|
||||||
a short cached-prefix resume and livelocked admission at a ~2-window
|
|
||||||
pool; keeping extend in the min() holds the reservation >= the prefill
|
|
||||||
allocation so an admitted request cannot OOM."""
|
|
||||||
window = self.tree_cache.sliding_window_size
|
|
||||||
headroom = min(extend_input_len + max_new_tokens, window)
|
|
||||||
reserved = headroom + self.page_size
|
|
||||||
if swa_host_hit_length > 0:
|
|
||||||
reserved += self.ceil_paged_tokens(swa_host_hit_length)
|
|
||||||
return reserved
|
|
||||||
|
|
||||||
def _swa_new_tokens(self, req: Req) -> int:
|
def _swa_new_tokens(self, req: Req) -> int:
|
||||||
"""Tokens a request may still decode, for SWA headroom sizing. Mirrors
|
"""Tokens a request may still decode, for SWA headroom sizing. Mirrors
|
||||||
@@ -837,80 +725,22 @@ class PrefillAdder:
|
|||||||
CLIP_MAX_NEW_TOKENS,
|
CLIP_MAX_NEW_TOKENS,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _swa_chunk_cap(self, max_new_tokens: int, swa_host_hit_length: int = 0) -> int:
|
def _check_prefill_budget(
|
||||||
"""Largest page-aligned extend chunk the SWA pool can admit right now,
|
|
||||||
keeping a sliding window of headroom below rem_swa_tokens; 0 if not
|
|
||||||
even one page fits. Only valid when is_hybrid_swa is True.
|
|
||||||
|
|
||||||
Escape hatch for a request whose budget can never pass the
|
|
||||||
_swa_budget_for_req gate (extend near/above the pool size, or a large
|
|
||||||
load-back charge): without shrinking its chunk it would be rejected
|
|
||||||
forever (head-of-line livelock). Shrinking is sound because past a
|
|
||||||
chunk boundary only the sliding window stays locked — the rest turns
|
|
||||||
evictable — so each pass's transient footprint fits the pool."""
|
|
||||||
# extend_input_len=0: this solves for the extend chunk itself, so the
|
|
||||||
# reserved headroom is the post-chunk decode window only.
|
|
||||||
cap = int(self.rem_swa_tokens) - self._swa_reserved_tokens(
|
|
||||||
0, max_new_tokens, swa_host_hit_length
|
|
||||||
)
|
|
||||||
if cap <= 0:
|
|
||||||
return 0
|
|
||||||
return cap // self.page_size * self.page_size
|
|
||||||
|
|
||||||
def _swa_req_never_fits(
|
|
||||||
self, extend_input_len: int, max_new_tokens: int, swa_host_hit_length: int = 0
|
|
||||||
) -> bool:
|
|
||||||
"""True when a request's SWA budget exceeds the *entire* SWA pool, so it
|
|
||||||
can never be admitted whole no matter how far the pool drains.
|
|
||||||
|
|
||||||
This is the head-of-line livelock the _swa_chunk_cap escape hatch exists
|
|
||||||
for; the hatch must fire only in this case. A request that merely
|
|
||||||
exceeds *current* rem_swa (transient pressure) would fit once running
|
|
||||||
decodes free their windows, so it must wait — admitting it into the
|
|
||||||
decode headroom collapses the SWA evictable cushion and forces running
|
|
||||||
requests to retract (observed as a severe retraction/re-prefill storm on
|
|
||||||
hybrid-SWA models at high concurrency)."""
|
|
||||||
capacity = self.token_to_kv_pool_allocator.size_swa
|
|
||||||
return (
|
|
||||||
self._swa_budget_for_req(
|
|
||||||
extend_input_len, max_new_tokens, swa_host_hit_length
|
|
||||||
)
|
|
||||||
>= capacity
|
|
||||||
)
|
|
||||||
|
|
||||||
def _swa_admission_gate(
|
|
||||||
self,
|
self,
|
||||||
req: Req,
|
req: Req,
|
||||||
|
*,
|
||||||
extend_input_len: int,
|
extend_input_len: int,
|
||||||
|
total_tokens: int,
|
||||||
swa_host_hit_length: int,
|
swa_host_hit_length: int,
|
||||||
chunk_tokens_limit: Optional[int],
|
) -> tuple[bool, Optional[int]]:
|
||||||
) -> tuple[Optional[AddReqResult], Optional[int]]:
|
return self.memory_budget.check_prefill(
|
||||||
"""SWA-pool gate: a non-None verdict rejects; otherwise the returned chunk
|
extend_input_len=extend_input_len,
|
||||||
limit stands, tightened to the pool cap when never-fits fires."""
|
total_tokens=total_tokens,
|
||||||
max_new_tokens = self._swa_new_tokens(req)
|
max_new_tokens=self._swa_new_tokens(req),
|
||||||
swa_needed = self._swa_budget_for_req(
|
input_tokens=len(req.full_untruncated_fill_ids),
|
||||||
extend_input_len, max_new_tokens, swa_host_hit_length=swa_host_hit_length
|
swa_host_hit_length=swa_host_hit_length,
|
||||||
|
chunk_limit=self.rem_chunk_tokens,
|
||||||
)
|
)
|
||||||
# Ring-slot capacity is exact, so needing exactly what is left still
|
|
||||||
# fits; the legacy SWA-token path keeps its conservative `>=`.
|
|
||||||
fits = (
|
|
||||||
swa_needed <= self.rem_swa_tokens
|
|
||||||
if self._swa_req_ring
|
|
||||||
else swa_needed < self.rem_swa_tokens
|
|
||||||
)
|
|
||||||
if fits:
|
|
||||||
return None, chunk_tokens_limit
|
|
||||||
if not self._swa_req_never_fits(
|
|
||||||
extend_input_len, max_new_tokens, swa_host_hit_length
|
|
||||||
):
|
|
||||||
return AddReqResult.NO_TOKEN, chunk_tokens_limit
|
|
||||||
swa_cap = self._swa_chunk_cap(max_new_tokens, swa_host_hit_length)
|
|
||||||
if self.rem_chunk_tokens is None or swa_cap <= 0:
|
|
||||||
return AddReqResult.NO_TOKEN, chunk_tokens_limit
|
|
||||||
current = (
|
|
||||||
self.rem_chunk_tokens if chunk_tokens_limit is None else chunk_tokens_limit
|
|
||||||
)
|
|
||||||
return None, min(current, swa_cap)
|
|
||||||
|
|
||||||
def _mamba_gap_budget_for_req(self, req: Req) -> int:
|
def _mamba_gap_budget_for_req(self, req: Req) -> int:
|
||||||
"""Shared-gap reservation (full-token-equivalents) for a request's new
|
"""Shared-gap reservation (full-token-equivalents) for a request's new
|
||||||
@@ -931,9 +761,7 @@ class PrefillAdder:
|
|||||||
return -(-tokens // self.page_size) * self.page_size
|
return -(-tokens // self.page_size) * self.page_size
|
||||||
|
|
||||||
def budget_state(self):
|
def budget_state(self):
|
||||||
no_token = self.rem_total_tokens <= 0 or self.cur_rem_tokens <= 0
|
no_token = not self.memory_budget.has_capacity()
|
||||||
if not no_token and self.is_hybrid_swa:
|
|
||||||
no_token = self.rem_swa_tokens <= 0
|
|
||||||
# Gate new mamba slots separately: rem_total_tokens' full_evictable can't
|
# Gate new mamba slots separately: rem_total_tokens' full_evictable can't
|
||||||
# cover a mamba slot, which needs mamba-recoverable bytes (see __init__).
|
# cover a mamba slot, which needs mamba-recoverable bytes (see __init__).
|
||||||
if not no_token and self.rem_mamba_slots is not None:
|
if not no_token and self.rem_mamba_slots is not None:
|
||||||
@@ -980,17 +808,12 @@ class PrefillAdder:
|
|||||||
if compute_charge is None:
|
if compute_charge is None:
|
||||||
compute_charge = extend_input_len
|
compute_charge = extend_input_len
|
||||||
|
|
||||||
# alloc_extend reserves an extra page_size per request to make sure the budget doesn't over-commit
|
self.memory_budget.reserve(
|
||||||
page_overhead = self.page_size
|
extend_input_len,
|
||||||
# `mamba_gap_reserve` (shared Mamba pool only; 0 otherwise) charges the new
|
max_new_tokens,
|
||||||
# mamba state's shared-gap cost to BOTH full budgets: the slot is allocated
|
extra_tokens=mamba_gap_reserve,
|
||||||
# immediately (counts against `cur_rem`) and held for the request lifetime
|
chunk_limit=self.rem_chunk_tokens,
|
||||||
# (counts against `rem_total`). See `_mamba_gap_budget_for_req`.
|
is_chunked_continuation=is_chunked_continuation,
|
||||||
self.rem_total_token_offset += (
|
|
||||||
extend_input_len + max_new_tokens + page_overhead + mamba_gap_reserve
|
|
||||||
)
|
|
||||||
self.cur_rem_token_offset += (
|
|
||||||
extend_input_len + page_overhead + mamba_gap_reserve
|
|
||||||
)
|
)
|
||||||
# The new mamba slot also consumes one mamba-recoverable slot (gated
|
# The new mamba slot also consumes one mamba-recoverable slot (gated
|
||||||
# separately so full_evictable can't cover it — see __init__).
|
# separately so full_evictable can't cover it — see __init__).
|
||||||
@@ -998,14 +821,6 @@ class PrefillAdder:
|
|||||||
self.rem_mamba_slots -= 1
|
self.rem_mamba_slots -= 1
|
||||||
self.rem_input_tokens -= compute_charge
|
self.rem_input_tokens -= compute_charge
|
||||||
|
|
||||||
if self.is_hybrid_swa:
|
|
||||||
# The ring slot is reserved once at first admission; charging it
|
|
||||||
# again on a continuation would double-count and over-throttle.
|
|
||||||
if not (self._swa_req_ring and is_chunked_continuation):
|
|
||||||
self.rem_swa_token_offset += self._swa_budget_for_req(
|
|
||||||
extend_input_len, max_new_tokens
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.dllm_config is not None:
|
if self.dllm_config is not None:
|
||||||
self.rem_dllm_tokens -= compute_charge
|
self.rem_dllm_tokens -= compute_charge
|
||||||
elif self.rem_chunk_tokens is not None:
|
elif self.rem_chunk_tokens is not None:
|
||||||
@@ -1136,20 +951,11 @@ class PrefillAdder:
|
|||||||
if self.dllm_config is not None:
|
if self.dllm_config is not None:
|
||||||
_rem_tokens = self._get_dllm_remain_tokens()
|
_rem_tokens = self._get_dllm_remain_tokens()
|
||||||
else:
|
else:
|
||||||
_rem_tokens = min(self.rem_chunk_tokens, int(self.rem_total_tokens))
|
_rem_tokens = self.memory_budget.available_chunk_tokens(
|
||||||
if self.is_hybrid_swa and not self._swa_req_ring:
|
self.rem_chunk_tokens
|
||||||
# alloc_extend needs extend_num_tokens + page_size per request,
|
)
|
||||||
# so reserve one page here to avoid OOM.
|
if _rem_tokens is None:
|
||||||
# Ring mode skips it: rem_swa_tokens counts slots, not chunk tokens.
|
return req
|
||||||
_rem_tokens = min(
|
|
||||||
_rem_tokens, int(self.rem_swa_tokens) - self.page_size
|
|
||||||
)
|
|
||||||
# The chunked_req must be added to the list; otherwise, it will cause a memory leak.
|
|
||||||
# Therefore, in certain cases where _rem_tokens <= 0, it should be replaced with rem_chunk_tokens.
|
|
||||||
if _rem_tokens <= 0:
|
|
||||||
if self.is_hybrid_swa:
|
|
||||||
return req
|
|
||||||
_rem_tokens = self.rem_chunk_tokens
|
|
||||||
|
|
||||||
# A mid-chunk rank prefills this pass regardless of the delayer
|
# A mid-chunk rank prefills this pass regardless of the delayer
|
||||||
# verdict, so report prefillable=True and ignore the result.
|
# verdict, so report prefillable=True and ignore the result.
|
||||||
@@ -1165,6 +971,13 @@ class PrefillAdder:
|
|||||||
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
||||||
req.prefix_indices
|
req.prefix_indices
|
||||||
)
|
)
|
||||||
|
_rem_tokens = self.memory_budget.fit_chunk(
|
||||||
|
extend_input_len=cand_extend_input_len,
|
||||||
|
max_new_tokens=self._swa_new_tokens(req),
|
||||||
|
chunk_limit=_rem_tokens,
|
||||||
|
)
|
||||||
|
if _rem_tokens is None:
|
||||||
|
return req
|
||||||
truncated = cand_extend_input_len > _rem_tokens
|
truncated = cand_extend_input_len > _rem_tokens
|
||||||
new_len = min(cand_extend_input_len, _rem_tokens)
|
new_len = min(cand_extend_input_len, _rem_tokens)
|
||||||
req.set_extend_range(len(req.prefix_indices), len(req.prefix_indices) + new_len)
|
req.set_extend_range(len(req.prefix_indices), len(req.prefix_indices) + new_len)
|
||||||
@@ -1210,16 +1023,14 @@ class PrefillAdder:
|
|||||||
# Shared Mamba pool: fold the new mamba state's shared-gap cost into the
|
# Shared Mamba pool: fold the new mamba state's shared-gap cost into the
|
||||||
# budget gate so admission can't over-commit (0 for baseline / non-Mamba).
|
# budget gate so admission can't over-commit (0 for baseline / non-Mamba).
|
||||||
paged_input += self._mamba_gap_budget_for_req(req)
|
paged_input += self._mamba_gap_budget_for_req(req)
|
||||||
if paged_input > min(self.cur_rem_tokens, self.rem_total_tokens):
|
fits = self.memory_budget.can_allocate_prefill(
|
||||||
|
paged_input=paged_input,
|
||||||
|
extend_input_len=cand_extend_input_len,
|
||||||
|
max_new_tokens=self._swa_new_tokens(req),
|
||||||
|
chunk_limit=self.rem_chunk_tokens,
|
||||||
|
)
|
||||||
|
if not fits:
|
||||||
return AddReqResult.NO_TOKEN
|
return AddReqResult.NO_TOKEN
|
||||||
if self.is_hybrid_swa:
|
|
||||||
if (
|
|
||||||
self._swa_budget_for_req(
|
|
||||||
cand_extend_input_len, self._swa_new_tokens(req)
|
|
||||||
)
|
|
||||||
> self.rem_swa_tokens
|
|
||||||
):
|
|
||||||
return AddReqResult.NO_TOKEN
|
|
||||||
|
|
||||||
def add_req_state(r, insert_sort=False):
|
def add_req_state(r, insert_sort=False):
|
||||||
new_token_ratio = (
|
new_token_ratio = (
|
||||||
@@ -1369,9 +1180,6 @@ class PrefillAdder:
|
|||||||
mamba_gap_reserve = self._mamba_gap_budget_for_req(req)
|
mamba_gap_reserve = self._mamba_gap_budget_for_req(req)
|
||||||
total_tokens += mamba_gap_reserve
|
total_tokens += mamba_gap_reserve
|
||||||
|
|
||||||
if total_tokens >= self.rem_total_tokens:
|
|
||||||
return AddReqResult.NO_TOKEN
|
|
||||||
|
|
||||||
# The temporary pin excludes this prefix from the evictable budget.
|
# The temporary pin excludes this prefix from the evictable budget.
|
||||||
# Selection itself neither allocates slots nor materializes host hits.
|
# Selection itself neither allocates slots nor materializes host hits.
|
||||||
with self._lock_node(req.last_node):
|
with self._lock_node(req.last_node):
|
||||||
@@ -1466,9 +1274,6 @@ class PrefillAdder:
|
|||||||
truncation_align_size: Optional[int],
|
truncation_align_size: Optional[int],
|
||||||
) -> _PrefillAdmission | AddReqResult:
|
) -> _PrefillAdmission | AddReqResult:
|
||||||
"""Select a prefill shape without allocating or publishing cached KV."""
|
"""Select a prefill shape without allocating or publishing cached KV."""
|
||||||
if total_tokens >= self.rem_total_tokens:
|
|
||||||
return AddReqResult.NO_TOKEN
|
|
||||||
|
|
||||||
prefix_len = len(req.prefix_indices) + host_hit_length
|
prefix_len = len(req.prefix_indices) + host_hit_length
|
||||||
extend_len = len(req.full_untruncated_fill_ids) - prefix_len
|
extend_len = len(req.full_untruncated_fill_ids) - prefix_len
|
||||||
input_tokens = self.ceil_paged_tokens(extend_len)
|
input_tokens = self.ceil_paged_tokens(extend_len)
|
||||||
@@ -1476,13 +1281,14 @@ class PrefillAdder:
|
|||||||
# exact-chunk-fill, so a request whose ceiled length would spill is
|
# exact-chunk-fill, so a request whose ceiled length would spill is
|
||||||
# not needlessly split into a second chunk.
|
# not needlessly split into a second chunk.
|
||||||
chunk_fit_tokens = extend_len if self.exact_chunk_fill else input_tokens
|
chunk_fit_tokens = extend_len if self.exact_chunk_fill else input_tokens
|
||||||
chunk_tokens_limit = self.rem_chunk_tokens
|
can_admit, chunk_tokens_limit = self._check_prefill_budget(
|
||||||
if self.is_hybrid_swa:
|
req,
|
||||||
verdict, chunk_tokens_limit = self._swa_admission_gate(
|
extend_input_len=extend_len,
|
||||||
req, input_tokens, swa_host_hit_length, chunk_tokens_limit
|
total_tokens=total_tokens,
|
||||||
)
|
swa_host_hit_length=swa_host_hit_length,
|
||||||
if verdict is not None:
|
)
|
||||||
return verdict
|
if not can_admit:
|
||||||
|
return AddReqResult.NO_TOKEN
|
||||||
|
|
||||||
# Without chunking, allow the first request even above the input cap.
|
# Without chunking, allow the first request even above the input cap.
|
||||||
if (
|
if (
|
||||||
@@ -1619,7 +1425,7 @@ class PrefillAdder:
|
|||||||
release_counter = 0
|
release_counter = 0
|
||||||
for i, running_req in enumerate(self.running_batch.reqs):
|
for i, running_req in enumerate(self.running_batch.reqs):
|
||||||
if running_req in preemptible_reqs:
|
if running_req in preemptible_reqs:
|
||||||
self.rem_total_token_offset -= (
|
self.memory_budget.total_offset -= (
|
||||||
self._get_running_request_total_token_offset(running_req)
|
self._get_running_request_total_token_offset(running_req)
|
||||||
)
|
)
|
||||||
release_counter += 1
|
release_counter += 1
|
||||||
|
|||||||
@@ -2523,23 +2523,28 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
max_new_tokens = min(max_new_tokens, self.max_new_tokens_limit)
|
max_new_tokens = min(max_new_tokens, self.max_new_tokens_limit)
|
||||||
|
|
||||||
# Keep this bound consistent with PrefillAdder's admission budget:
|
# Keep this bound consistent with PrefillAdder's admission budget.
|
||||||
# ceil_page(input_len) + max_new_tokens + page_size must be strictly
|
max_new_tokens = max(
|
||||||
# smaller than max_total_num_tokens. Otherwise a request can be accepted
|
|
||||||
# into the waiting queue but can never be scheduled, blocking the queue
|
|
||||||
# and eventually making health checks fail.
|
|
||||||
paged_input_len = -(-input_len // self.page_size) * self.page_size
|
|
||||||
req.sampling_params.max_new_tokens = max(
|
|
||||||
0,
|
0,
|
||||||
min(
|
min(
|
||||||
max_new_tokens,
|
max_new_tokens,
|
||||||
self.max_req_len - input_len - 1,
|
self.max_req_len - input_len - 1,
|
||||||
self.max_total_num_tokens * get_parallel().attn_dcp_size
|
|
||||||
- paged_input_len
|
|
||||||
- self.page_size
|
|
||||||
- 1,
|
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
max_new_tokens = self.token_to_kv_pool_allocator.max_new_tokens_for_memory(
|
||||||
|
input_len,
|
||||||
|
max_new_tokens,
|
||||||
|
token_capacity=self.max_total_num_tokens * get_parallel().attn_dcp_size,
|
||||||
|
sliding_window_size=self.sliding_window_size,
|
||||||
|
chunk_size=self.chunked_prefill_size,
|
||||||
|
)
|
||||||
|
if max_new_tokens is None:
|
||||||
|
req.set_finish_with_abort(
|
||||||
|
f"Request prompt exceeds the KV memory budget: input_len={input_len}."
|
||||||
|
)
|
||||||
|
max_new_tokens = 0
|
||||||
|
|
||||||
|
req.sampling_params.max_new_tokens = max(0, max_new_tokens)
|
||||||
# Clipping above can push max_new_tokens below min_new_tokens, which
|
# Clipping above can push max_new_tokens below min_new_tokens, which
|
||||||
# would suppress EOS for the whole generation. Restore the invariant.
|
# would suppress EOS for the whole generation. Restore the invariant.
|
||||||
if req.sampling_params.min_new_tokens > req.sampling_params.max_new_tokens:
|
if req.sampling_params.min_new_tokens > req.sampling_params.max_new_tokens:
|
||||||
|
|||||||
@@ -22,9 +22,6 @@ from sglang.srt.managers.scheduler_components.pool_stats_observer import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring
|
from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring
|
||||||
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
|
||||||
UnifiedMambaSWATokenToKVPoolAllocator,
|
|
||||||
)
|
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
from sglang.srt.observability.scheduler_stage_metrics import (
|
from sglang.srt.observability.scheduler_stage_metrics import (
|
||||||
@@ -94,12 +91,13 @@ class SchedulerInvariantChecker:
|
|||||||
return leak, msg
|
return leak, msg
|
||||||
|
|
||||||
def _check_full_pool(self, ps: PoolStats, uncached: int = 0) -> Tuple[bool, str]:
|
def _check_full_pool(self, ps: PoolStats, uncached: int = 0) -> Tuple[bool, str]:
|
||||||
if self.is_hybrid_swa and not self.full_tokens_per_layer:
|
allocator = self.token_to_kv_pool_allocator
|
||||||
|
if self.is_hybrid_swa and not ps.full_capacity:
|
||||||
return False, ""
|
return False, ""
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
protected = self.tree_cache.full_protected_size()
|
protected = self.tree_cache.full_protected_size()
|
||||||
session_held = self.pool_stats_observer.session_held_full_tokens()
|
session_held = self.pool_stats_observer.session_held_full_tokens()
|
||||||
total = self.full_tokens_per_layer
|
total = ps.full_capacity
|
||||||
elif self.is_hybrid_ssm:
|
elif self.is_hybrid_ssm:
|
||||||
# Branch on cache type for the protected accessor (MambaRadixCache
|
# Branch on cache type for the protected accessor (MambaRadixCache
|
||||||
# splits full/mamba; ChunkCache only has the single protected_size).
|
# splits full/mamba; ChunkCache only has the single protected_size).
|
||||||
@@ -119,7 +117,6 @@ class SchedulerInvariantChecker:
|
|||||||
session_held = self.pool_stats_observer.session_held_tokens()
|
session_held = self.pool_stats_observer.session_held_tokens()
|
||||||
total = self.max_total_num_tokens
|
total = self.max_total_num_tokens
|
||||||
full_evictable_size = ps.full_evictable_size
|
full_evictable_size = ps.full_evictable_size
|
||||||
allocator = self.token_to_kv_pool_allocator
|
|
||||||
if get_parallel().dcp_enabled and allocator.page_size > 1:
|
if get_parallel().dcp_enabled and allocator.page_size > 1:
|
||||||
# DCP stores logical tokens in widened physical pages. Prefix cache
|
# DCP stores logical tokens in widened physical pages. Prefix cache
|
||||||
# counters are logical-token based, while the allocator frees whole
|
# counters are logical-token based, while the allocator frees whole
|
||||||
@@ -129,14 +126,9 @@ class SchedulerInvariantChecker:
|
|||||||
// allocator.page_size
|
// allocator.page_size
|
||||||
* allocator.page_size
|
* allocator.page_size
|
||||||
)
|
)
|
||||||
full_available = ps.full_available_size
|
|
||||||
if isinstance(allocator, UnifiedMambaSWATokenToKVPoolAllocator):
|
|
||||||
# Pair the static per-layer total with the conserve view, never the
|
|
||||||
# byte-coordinated one -- see `conserve_full_available_size`.
|
|
||||||
full_available = allocator.conserve_full_available_size()
|
|
||||||
leak, msg = self._check_pool_invariant(
|
leak, msg = self._check_pool_invariant(
|
||||||
"full",
|
"full",
|
||||||
full_available,
|
ps.full_available_size,
|
||||||
full_evictable_size,
|
full_evictable_size,
|
||||||
protected,
|
protected,
|
||||||
session_held,
|
session_held,
|
||||||
@@ -162,18 +154,13 @@ class SchedulerInvariantChecker:
|
|||||||
f"evictable={ps.swa_evictable_size}, "
|
f"evictable={ps.swa_evictable_size}, "
|
||||||
f"total={self.swa_tokens_per_layer}"
|
f"total={self.swa_tokens_per_layer}"
|
||||||
)
|
)
|
||||||
swa_available = ps.swa_available_size
|
|
||||||
if isinstance(allocator, UnifiedMambaSWATokenToKVPoolAllocator):
|
|
||||||
# Tri-pool: same floating-boundary phantom as the full pool -- use the
|
|
||||||
# slot-conservation view, not the byte-coordinated min (see _check_full_pool).
|
|
||||||
swa_available = allocator.conserve_swa_available_size()
|
|
||||||
return self._check_pool_invariant(
|
return self._check_pool_invariant(
|
||||||
"swa",
|
"swa",
|
||||||
swa_available,
|
ps.swa_available_size,
|
||||||
ps.swa_evictable_size,
|
ps.swa_evictable_size,
|
||||||
self.tree_cache.swa_protected_size(),
|
self.tree_cache.swa_protected_size(),
|
||||||
self.pool_stats_observer.session_held_swa_tokens(),
|
self.pool_stats_observer.session_held_swa_tokens(),
|
||||||
self.swa_tokens_per_layer,
|
ps.swa_capacity,
|
||||||
uncached,
|
uncached,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -12,9 +12,6 @@ from typing import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring
|
from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring
|
||||||
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
|
||||||
UnifiedMambaSWATokenToKVPoolAllocator,
|
|
||||||
)
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
@@ -38,6 +35,8 @@ class PoolStats:
|
|||||||
is_hisparse: bool = False
|
is_hisparse: bool = False
|
||||||
|
|
||||||
# For hybrid-swa pools
|
# For hybrid-swa pools
|
||||||
|
full_capacity: Optional[int] = None
|
||||||
|
swa_capacity: Optional[int] = None
|
||||||
swa_num_used: Optional[int] = None
|
swa_num_used: Optional[int] = None
|
||||||
swa_token_usage: Optional[float] = None
|
swa_token_usage: Optional[float] = None
|
||||||
swa_available_size: Optional[int] = None
|
swa_available_size: Optional[int] = None
|
||||||
@@ -289,29 +288,20 @@ class SchedulerPoolStatsObserver:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _get_swa_token_info(self) -> PoolStats:
|
def _get_swa_token_info(self) -> PoolStats:
|
||||||
# `*_num_used` is `static_cap - (available + evictable)`, so the
|
(full_capacity, full_available_size), (swa_capacity, swa_available_size) = (
|
||||||
# available term must match the static cap's denomination: the conserve
|
self.token_to_kv_pool_allocator.swa_capacity_and_available(
|
||||||
# view, never the byte-coordinated one (see
|
full_capacity=self.full_tokens_per_layer,
|
||||||
# `conserve_full_available_size`). Measured ~25-90x inflated otherwise.
|
swa_capacity=self.swa_tokens_per_layer,
|
||||||
allocator = self.token_to_kv_pool_allocator
|
)
|
||||||
if isinstance(allocator, UnifiedMambaSWATokenToKVPoolAllocator):
|
)
|
||||||
full_available_size = allocator.conserve_full_available_size()
|
|
||||||
swa_available_size = allocator.conserve_swa_available_size()
|
|
||||||
else:
|
|
||||||
full_available_size = allocator.full_available_size()
|
|
||||||
swa_available_size = allocator.swa_available_size()
|
|
||||||
full_evictable_size = self.tree_cache.full_evictable_size()
|
full_evictable_size = self.tree_cache.full_evictable_size()
|
||||||
swa_evictable_size = self.tree_cache.swa_evictable_size()
|
swa_evictable_size = self.tree_cache.swa_evictable_size()
|
||||||
# Per-request SWA ring: released with the req slot, yet cached radix
|
# Per-request SWA ring: released with the req slot, yet cached radix
|
||||||
# prefixes still report swa_evictable; counting it drives usage negative.
|
# prefixes still report swa_evictable; counting it drives usage negative.
|
||||||
if is_swa_req_ring(self.token_to_kv_pool_allocator):
|
if is_swa_req_ring(self.token_to_kv_pool_allocator):
|
||||||
swa_evictable_size = 0
|
swa_evictable_size = 0
|
||||||
full_num_used = self.full_tokens_per_layer - (
|
full_num_used = full_capacity - (full_available_size + full_evictable_size)
|
||||||
full_available_size + full_evictable_size
|
swa_num_used = swa_capacity - (swa_available_size + swa_evictable_size)
|
||||||
)
|
|
||||||
swa_num_used = self.swa_tokens_per_layer - (
|
|
||||||
swa_available_size + swa_evictable_size
|
|
||||||
)
|
|
||||||
# FIXME(hisparse): host-backup transiently over-releases the device pool
|
# FIXME(hisparse): host-backup transiently over-releases the device pool
|
||||||
# counter, producing negative full_num_used / swa_num_used. We clamp to 0
|
# counter, producing negative full_num_used / swa_num_used. We clamp to 0
|
||||||
# to keep token_usage / leak checks sane, but the underlying accounting
|
# to keep token_usage / leak checks sane, but the underlying accounting
|
||||||
@@ -319,16 +309,23 @@ class SchedulerPoolStatsObserver:
|
|||||||
if self.enable_hisparse:
|
if self.enable_hisparse:
|
||||||
full_num_used = max(0, full_num_used)
|
full_num_used = max(0, full_num_used)
|
||||||
swa_num_used = max(0, swa_num_used)
|
swa_num_used = max(0, swa_num_used)
|
||||||
if not self.full_tokens_per_layer:
|
if not full_capacity:
|
||||||
full_num_used = 0
|
full_num_used = 0
|
||||||
full_available_size = 0
|
full_available_size = 0
|
||||||
full_token_usage = 0.0
|
full_token_usage = 0.0
|
||||||
else:
|
else:
|
||||||
full_token_usage = full_num_used / self.full_tokens_per_layer
|
full_token_usage = full_num_used / full_capacity
|
||||||
swa_token_usage = swa_num_used / self.swa_tokens_per_layer
|
if not swa_capacity:
|
||||||
|
swa_num_used = 0
|
||||||
|
swa_available_size = 0
|
||||||
|
swa_token_usage = 0.0
|
||||||
|
else:
|
||||||
|
swa_token_usage = swa_num_used / swa_capacity
|
||||||
|
|
||||||
return PoolStats(
|
return PoolStats(
|
||||||
is_hybrid_swa=True,
|
is_hybrid_swa=True,
|
||||||
|
full_capacity=full_capacity,
|
||||||
|
swa_capacity=swa_capacity,
|
||||||
full_num_used=full_num_used,
|
full_num_used=full_num_used,
|
||||||
full_token_usage=full_token_usage,
|
full_token_usage=full_token_usage,
|
||||||
full_available_size=full_available_size,
|
full_available_size=full_available_size,
|
||||||
|
|||||||
@@ -69,12 +69,47 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
|||||||
# The scheduler calls these unconditionally, with no allocator-type branches
|
# The scheduler calls these unconditionally, with no allocator-type branches
|
||||||
# on its side; byte-accounted composites override the token-count defaults.
|
# on its side; byte-accounted composites override the token-count defaults.
|
||||||
|
|
||||||
def evict_to_free_tokens(self, tree_cache, num_tokens: int) -> None:
|
def create_prefill_budget(self, tree_cache, *, num_mixed_decode_tokens=0):
|
||||||
"""Evict unlocked prefix-cache entries until this allocator can serve
|
from sglang.srt.mem_cache.prefill_budget import PrefillBudget
|
||||||
``num_tokens`` or nothing evictable remains."""
|
|
||||||
from sglang.srt.mem_cache.common import evict_from_tree_cache
|
|
||||||
|
|
||||||
evict_from_tree_cache(tree_cache, num_tokens)
|
return PrefillBudget(
|
||||||
|
self, tree_cache, num_mixed_decode_tokens=num_mixed_decode_tokens
|
||||||
|
)
|
||||||
|
|
||||||
|
def max_new_tokens_for_memory(
|
||||||
|
self,
|
||||||
|
input_tokens: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
*,
|
||||||
|
token_capacity: int,
|
||||||
|
sliding_window_size: int | None,
|
||||||
|
chunk_size: int | None,
|
||||||
|
) -> int | None:
|
||||||
|
"""Clip generation to the empty-pool budget; None means prompt cannot fit.
|
||||||
|
|
||||||
|
token_capacity is the scheduler's configured capacity, including its
|
||||||
|
distributed token scaling. Shared pools use their physical byte layout.
|
||||||
|
"""
|
||||||
|
paged_input = -(-input_tokens // self.page_size) * self.page_size
|
||||||
|
return max(
|
||||||
|
0, min(max_new_tokens, token_capacity - paged_input - self.page_size - 1)
|
||||||
|
)
|
||||||
|
|
||||||
|
def evict_to_free_tokens(self, tree_cache, num_tokens: int) -> bool | None:
|
||||||
|
"""Evict unlocked prefix-cache entries until this allocator can serve
|
||||||
|
``num_tokens`` or nothing evictable remains.
|
||||||
|
|
||||||
|
Return whether capacity was realized, or None if it still needs checking.
|
||||||
|
"""
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||||
|
from sglang.srt.mem_cache.common import _evict_until_allocatable
|
||||||
|
|
||||||
|
if tree_cache is None or tree_cache.is_chunk_cache():
|
||||||
|
return
|
||||||
|
shortfall = num_tokens - self.available_size()
|
||||||
|
if shortfall > 0:
|
||||||
|
tree_cache.evict_for_alloc(EvictParams(num_tokens=shortfall))
|
||||||
|
_evict_until_allocatable(tree_cache, self, num_tokens)
|
||||||
|
|
||||||
def check_decode_capacity(
|
def check_decode_capacity(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -357,6 +357,20 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
kv_indices
|
kv_indices
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def create_prefill_budget(self, tree_cache, *, num_mixed_decode_tokens=0):
|
||||||
|
# Use this wrapper's capacity, which also includes the compressed pool.
|
||||||
|
from sglang.srt.mem_cache.prefill_budget import SWAPrefillBudget
|
||||||
|
|
||||||
|
return SWAPrefillBudget(
|
||||||
|
self, tree_cache, num_mixed_decode_tokens=num_mixed_decode_tokens
|
||||||
|
)
|
||||||
|
|
||||||
|
def swa_capacity_and_available(self, *, full_capacity, swa_capacity):
|
||||||
|
return (
|
||||||
|
(full_capacity, self.full_available_size()),
|
||||||
|
(swa_capacity, self.swa_available_size()),
|
||||||
|
)
|
||||||
|
|
||||||
def full_available_size(self):
|
def full_available_size(self):
|
||||||
return min(
|
return min(
|
||||||
self.logical_attn_allocator.full_available_size(),
|
self.logical_attn_allocator.full_available_size(),
|
||||||
|
|||||||
@@ -158,6 +158,31 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
self.swa_attn_allocator.available_size(),
|
self.swa_attn_allocator.available_size(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def create_prefill_budget(self, tree_cache, *, num_mixed_decode_tokens=0):
|
||||||
|
from sglang.srt.mem_cache.prefill_budget import SWAPrefillBudget
|
||||||
|
|
||||||
|
return SWAPrefillBudget(
|
||||||
|
self, tree_cache, num_mixed_decode_tokens=num_mixed_decode_tokens
|
||||||
|
)
|
||||||
|
|
||||||
|
def swa_capacity_and_available(self, *, full_capacity, swa_capacity):
|
||||||
|
return (
|
||||||
|
(full_capacity, self.full_available_size()),
|
||||||
|
(swa_capacity, self.swa_available_size()),
|
||||||
|
)
|
||||||
|
|
||||||
|
def evict_to_free_tokens(self, tree_cache, num_tokens: int) -> None:
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||||
|
|
||||||
|
if tree_cache is None or tree_cache.is_chunk_cache():
|
||||||
|
return
|
||||||
|
full_shortfall = max(0, num_tokens - self.full_available_size())
|
||||||
|
swa_shortfall = max(0, num_tokens - self.swa_available_size())
|
||||||
|
if full_shortfall or swa_shortfall:
|
||||||
|
tree_cache.evict_for_alloc(
|
||||||
|
EvictParams(num_tokens=full_shortfall, swa_num_tokens=swa_shortfall)
|
||||||
|
)
|
||||||
|
|
||||||
def full_available_size(self):
|
def full_available_size(self):
|
||||||
return self.full_attn_allocator.available_size()
|
return self.full_attn_allocator.available_size()
|
||||||
|
|
||||||
@@ -687,6 +712,16 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
def available_size(self):
|
def available_size(self):
|
||||||
return self.swa_attn_allocator.available_size()
|
return self.swa_attn_allocator.available_size()
|
||||||
|
|
||||||
|
def create_prefill_budget(self, tree_cache, *, num_mixed_decode_tokens=0):
|
||||||
|
from sglang.srt.mem_cache.prefill_budget import SWAPrefillBudget
|
||||||
|
|
||||||
|
return SWAPrefillBudget(
|
||||||
|
self,
|
||||||
|
tree_cache,
|
||||||
|
num_mixed_decode_tokens=num_mixed_decode_tokens,
|
||||||
|
all_swa=True,
|
||||||
|
)
|
||||||
|
|
||||||
def full_available_size(self):
|
def full_available_size(self):
|
||||||
return self.swa_attn_allocator.available_size()
|
return self.swa_attn_allocator.available_size()
|
||||||
|
|
||||||
|
|||||||
@@ -17,6 +17,8 @@ sub-pools of one `UnifiedKVPool`, and the tri-pool variant that adds mamba state
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import math
|
||||||
|
from abc import abstractmethod
|
||||||
from typing import Callable, List, Optional, Sequence, Tuple
|
from typing import Callable, List, Optional, Sequence, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -42,11 +44,10 @@ from sglang.srt.utils.common import get_num_new_pages
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
class UnifiedSWAAllocatorBase(SWATokenToKVPoolAllocator):
|
||||||
"""Composite allocator for the hybrid SWA pair (full + swa MHA sub-pools).
|
"""Shared FULL/SWA virtual IDs, allocation lifecycle, and index translation.
|
||||||
|
|
||||||
One alloc(N) binds N pages on BOTH sides under the same virtual id, so
|
Concrete allocators define the two-ended or Mamba/SWA/FULL capacity policy.
|
||||||
`available_size()` (joint bytes, in TOKENS) is the only safe alloc pre-check.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Parent's `size` property has no setter but base init does `self.size = size`;
|
# Parent's `size` property has no setter but base init does `self.size = size`;
|
||||||
@@ -65,26 +66,36 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
unified_buffer: UnifiedKVPool,
|
unified_buffer: UnifiedKVPool,
|
||||||
kvcache, # UnifiedSWAKVPool
|
kvcache, # UnifiedSWAKVPool
|
||||||
device: str,
|
device: str,
|
||||||
full_max_total_num_tokens: int,
|
full_max_total_num_tokens: Optional[int] = None,
|
||||||
swa_max_total_num_tokens: int,
|
swa_max_total_num_tokens: Optional[int] = None,
|
||||||
page_size: int = 1,
|
page_size: int = 1,
|
||||||
need_sort: bool = False,
|
need_sort: bool = False,
|
||||||
forward_stream: Optional[torch.cuda.Stream] = None,
|
forward_stream: Optional[torch.cuda.Stream] = None,
|
||||||
lazy_compaction: bool = False,
|
lazy_compaction: bool = False,
|
||||||
):
|
):
|
||||||
# Set _size_full / _size_swa BEFORE base init (read during it). STATIC
|
if (full_max_total_num_tokens is None) != (swa_max_total_num_tokens is None):
|
||||||
# partition caps -- the slot-conservation value the leak invariant expects.
|
raise ValueError(
|
||||||
self._size_full = full_max_total_num_tokens
|
"full_max_total_num_tokens and swa_max_total_num_tokens must "
|
||||||
self._size_swa = swa_max_total_num_tokens
|
"either both be set or both be omitted"
|
||||||
self._full_max_total_num_tokens = full_max_total_num_tokens
|
)
|
||||||
self._swa_max_total_num_tokens = swa_max_total_num_tokens
|
legacy_capacities = full_max_total_num_tokens is not None
|
||||||
|
self._size_full = (
|
||||||
|
int(full_max_total_num_tokens)
|
||||||
|
if legacy_capacities
|
||||||
|
else unified_buffer.max_slots("full") - 1
|
||||||
|
)
|
||||||
|
self._size_swa = (
|
||||||
|
int(swa_max_total_num_tokens)
|
||||||
|
if legacy_capacities
|
||||||
|
else unified_buffer.max_slots("swa") - 1
|
||||||
|
)
|
||||||
self.page_size = page_size
|
self.page_size = page_size
|
||||||
|
|
||||||
# The parent is inherited only for the isinstance contract: skip its
|
# The parent is inherited only for the isinstance contract: skip its
|
||||||
# static-partition sub-pool allocation, which the unified pool replaces.
|
# static-partition sub-pool allocation, which the unified pool replaces.
|
||||||
BaseTokenToKVPoolAllocator.__init__(
|
BaseTokenToKVPoolAllocator.__init__(
|
||||||
self,
|
self,
|
||||||
size=full_max_total_num_tokens,
|
size=self._size_full,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
dtype=unified_buffer.mha_spec("full").store_dtype,
|
dtype=unified_buffer.mha_spec("full").store_dtype,
|
||||||
device=device,
|
device=device,
|
||||||
@@ -120,6 +131,16 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
)
|
)
|
||||||
self._wire_peers()
|
self._wire_peers()
|
||||||
|
|
||||||
|
self._empty_shared_gap_bytes = self.full_attn_allocator._current_gap_bytes()
|
||||||
|
if not legacy_capacities:
|
||||||
|
self._size_full = self.full_attn_allocator.available_size()
|
||||||
|
self._size_swa = min(
|
||||||
|
self.swa_attn_allocator.available_size(),
|
||||||
|
len(self.full_attn_allocator.free_virtual_ids) * page_size,
|
||||||
|
)
|
||||||
|
self._full_max_total_num_tokens = self._size_full
|
||||||
|
self._swa_max_total_num_tokens = self._size_swa
|
||||||
|
|
||||||
# Epoch-keyed memo for the joint capacity view (any chain member's
|
# Epoch-keyed memo for the joint capacity view (any chain member's
|
||||||
# mutation invalidates -- see `MultiEndedAllocator._chain_capacity_epoch`).
|
# mutation invalidates -- see `MultiEndedAllocator._chain_capacity_epoch`).
|
||||||
self._joint_avail_memo_epoch: Optional[int] = None
|
self._joint_avail_memo_epoch: Optional[int] = None
|
||||||
@@ -143,7 +164,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
"[unified-memory-pool] UnifiedSWATokenToKVPoolAllocator ready: "
|
"[unified-memory-pool] UnifiedSWATokenToKVPoolAllocator ready: "
|
||||||
"full max_slots=%d (min_slot_index=%d, entry_bytes=%d), "
|
"full max_slots=%d (min_slot_index=%d, entry_bytes=%d), "
|
||||||
"swa max_slots=%d (min_slot_index=%d, entry_bytes=%d), "
|
"swa max_slots=%d (min_slot_index=%d, entry_bytes=%d), "
|
||||||
"static caps full=%d swa=%d, joint available=%d",
|
"max capacity full=%d swa=%d, joint available=%d",
|
||||||
self.full_attn_allocator.max_slots,
|
self.full_attn_allocator.max_slots,
|
||||||
self.full_attn_allocator.min_slot_index,
|
self.full_attn_allocator.min_slot_index,
|
||||||
self.full_attn_allocator.entry_bytes,
|
self.full_attn_allocator.entry_bytes,
|
||||||
@@ -157,18 +178,6 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
|
|
||||||
# -- construction hooks (the tri-pool subclass overrides both) --
|
# -- construction hooks (the tri-pool subclass overrides both) --
|
||||||
|
|
||||||
def _build_swa_attn_allocator(self, **kwargs) -> MultiEndedAllocator:
|
|
||||||
"""The swa sub-allocator: an END pool in the 2-pool pair."""
|
|
||||||
return MultiEndedAllocator(
|
|
||||||
sub_pool_name="swa",
|
|
||||||
is_id_owner=False, # non-owner; consumes virtuals minted by full
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _wire_peers(self) -> None:
|
|
||||||
self.full_attn_allocator.bind_peer(self.swa_attn_allocator)
|
|
||||||
self.swa_attn_allocator.bind_peer(self.full_attn_allocator)
|
|
||||||
|
|
||||||
# -- capacity reporting (three-way split) --
|
# -- capacity reporting (three-way split) --
|
||||||
|
|
||||||
def available_size(self) -> int:
|
def available_size(self) -> int:
|
||||||
@@ -179,46 +188,6 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
self._joint_avail_memo_epoch = epoch
|
self._joint_avail_memo_epoch = epoch
|
||||||
return self._joint_avail_memo_tokens
|
return self._joint_avail_memo_tokens
|
||||||
|
|
||||||
def _compute_available_size(self) -> int:
|
|
||||||
"""Joint byte budget in TOKENS: each composite alloc(1) consumes one
|
|
||||||
full-side AND one swa-side page under the same virtual id."""
|
|
||||||
fa, sa = self.full_attn_allocator, self.swa_attn_allocator
|
|
||||||
e_f = fa.entry_bytes_per_page
|
|
||||||
e_s = sa.entry_bytes_per_page
|
|
||||||
# Direction-agnostic shared gap: the free byte band between the two pools.
|
|
||||||
if fa.grow_direction == "up":
|
|
||||||
gap_bytes = max(0, sa._byte_low_frontier() - fa._byte_high_frontier())
|
|
||||||
else:
|
|
||||||
gap_bytes = max(0, fa._byte_low_frontier() - sa._byte_high_frontier())
|
|
||||||
R_f = fa.num_pages - fa.min_page_index - fa._allocated_pages()
|
|
||||||
R_s = sa.num_pages - sa.min_page_index - sa._allocated_pages()
|
|
||||||
|
|
||||||
if not self.lazy_compaction:
|
|
||||||
pages_by_bytes = gap_bytes // (e_f + e_s)
|
|
||||||
return min(pages_by_bytes, R_f, R_s) * self.page_size
|
|
||||||
|
|
||||||
H_f = len(fa._free_phys_pages)
|
|
||||||
H_s = len(sa._free_phys_pages)
|
|
||||||
|
|
||||||
K1 = min(H_f, H_s) # Phase 1: both drain
|
|
||||||
|
|
||||||
# Phase 2: fewer-holes side extends; more-holes side keeps draining.
|
|
||||||
if H_f <= H_s:
|
|
||||||
e_phase2 = e_f
|
|
||||||
K_phase2_max = H_s
|
|
||||||
else:
|
|
||||||
e_phase2 = e_s
|
|
||||||
K_phase2_max = H_f
|
|
||||||
K2_room = K_phase2_max - K1
|
|
||||||
K2 = min(K2_room, gap_bytes // e_phase2) if e_phase2 > 0 else K2_room
|
|
||||||
gap_bytes -= K2 * e_phase2
|
|
||||||
|
|
||||||
K3 = gap_bytes // (e_f + e_s) # Phase 3: both extend
|
|
||||||
|
|
||||||
K_total = K1 + K2 + K3
|
|
||||||
K_total = min(K_total, H_f + R_f, H_s + R_s) # index-space caps
|
|
||||||
return K_total * self.page_size
|
|
||||||
|
|
||||||
# Slot-conservation views for the leak invariant only; the byte-coordinated
|
# Slot-conservation views for the leak invariant only; the byte-coordinated
|
||||||
# value would flag spurious leaks. `allocated_count()` is in TOKENS.
|
# value would flag spurious leaks. `allocated_count()` is in TOKENS.
|
||||||
def _conserve_full_available_size(self) -> int:
|
def _conserve_full_available_size(self) -> int:
|
||||||
@@ -261,17 +230,16 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
def schedulable_swa_available_size(self) -> int:
|
def schedulable_swa_available_size(self) -> int:
|
||||||
return self.swa_attn_allocator.schedulable_available_size()
|
return self.swa_attn_allocator.schedulable_available_size()
|
||||||
|
|
||||||
def _flush_targets(self):
|
# `size_full` / `size_swa` bound each side independently; current capacities
|
||||||
"""Flush ALL members, including ones that are not short themselves: a
|
# also account for the peer's live byte usage.
|
||||||
one-sided hole is unusable, and compacting it yields SHARED gap."""
|
|
||||||
return (self.full_attn_allocator, self.swa_attn_allocator)
|
|
||||||
|
|
||||||
def _ask_float_for_room(self, need_tokens: int) -> None:
|
@property
|
||||||
"""No float in a two-END chain -- nothing can slide."""
|
def current_full_capacity(self) -> int:
|
||||||
return None
|
return self.full_available_size() + self.full_attn_allocator.allocated_count()
|
||||||
|
|
||||||
# `size_full` / `size_swa` are inherited and read the static caps; reporting
|
@property
|
||||||
# `max_slots - 1` here would be ~= full_max + swa_max and over-promise.
|
def current_swa_capacity(self) -> int:
|
||||||
|
return self.swa_available_size() + self.swa_attn_allocator.allocated_count()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def draft_virtual_id_space(self) -> int:
|
def draft_virtual_id_space(self) -> int:
|
||||||
@@ -396,11 +364,8 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
|
|
||||||
def alloc(self, need_size: int) -> Optional[torch.Tensor]:
|
def alloc(self, need_size: int) -> Optional[torch.Tensor]:
|
||||||
with record_function("UnifiedSWAAlloc.alloc"):
|
with record_function("UnifiedSWAAlloc.alloc"):
|
||||||
# Joint pre-check. Both sides are mutual peers (each side's compaction
|
if not self.ensure_capacity(need_size, need_size):
|
||||||
# opens gap for the other), so flush BOTH on shortfall.
|
return None
|
||||||
if need_size > self.available_size():
|
|
||||||
if not _relieve_for_alloc(self, need_size):
|
|
||||||
return None
|
|
||||||
# Snapshot the virtual PAGES full will consume, to bind them on swa too.
|
# Snapshot the virtual PAGES full will consume, to bind them on swa too.
|
||||||
num_pages = need_size // self.page_size
|
num_pages = need_size // self.page_size
|
||||||
fa = self.full_attn_allocator
|
fa = self.full_attn_allocator
|
||||||
@@ -437,9 +402,8 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
prefix_lens=prefix_lens_cpu,
|
prefix_lens=prefix_lens_cpu,
|
||||||
)
|
)
|
||||||
need_tokens = num_new_pages * self.page_size
|
need_tokens = num_new_pages * self.page_size
|
||||||
if need_tokens > self.available_size():
|
if not self.ensure_capacity(need_tokens, need_tokens):
|
||||||
if not _relieve_for_alloc(self, need_tokens):
|
return None
|
||||||
return None
|
|
||||||
|
|
||||||
# Snapshot the virtual PAGES the kernel will consume; clone so swa keeps
|
# Snapshot the virtual PAGES the kernel will consume; clone so swa keeps
|
||||||
# its view after the slice is consumed.
|
# its view after the slice is consumed.
|
||||||
@@ -555,9 +519,8 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
seq_lens=seq_lens_cpu, page_size=self.page_size, decode=True
|
seq_lens=seq_lens_cpu, page_size=self.page_size, decode=True
|
||||||
)
|
)
|
||||||
need_tokens = num_new_pages * self.page_size
|
need_tokens = num_new_pages * self.page_size
|
||||||
if need_tokens > self.available_size():
|
if not self.ensure_capacity(need_tokens, need_tokens):
|
||||||
if not _relieve_for_alloc(self, need_tokens):
|
return None
|
||||||
return None
|
|
||||||
|
|
||||||
fa = self.full_attn_allocator
|
fa = self.full_attn_allocator
|
||||||
new_virtual_pages = fa.free_virtual_ids[:num_new_pages].clone()
|
new_virtual_pages = fa.free_virtual_ids[:num_new_pages].clone()
|
||||||
@@ -727,14 +690,6 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
self.full_attn_allocator.clear_inverse_history()
|
self.full_attn_allocator.clear_inverse_history()
|
||||||
self.swa_attn_allocator.clear_inverse_history()
|
self.swa_attn_allocator.clear_inverse_history()
|
||||||
|
|
||||||
def verify_byte_accounting(self) -> List[str]:
|
|
||||||
return (
|
|
||||||
_chain_byte_accounting_violations(
|
|
||||||
_end_pair_chain(self.full_attn_allocator, self.swa_attn_allocator)
|
|
||||||
)
|
|
||||||
+ self._joint_capacity_memo_violations()
|
|
||||||
)
|
|
||||||
|
|
||||||
def _joint_capacity_memo_violations(self) -> List[str]:
|
def _joint_capacity_memo_violations(self) -> List[str]:
|
||||||
"""Idle-time twin of `MultiEndedAllocator._capacity_memo_violations`
|
"""Idle-time twin of `MultiEndedAllocator._capacity_memo_violations`
|
||||||
for the composite joint view. Empty == healthy."""
|
for the composite joint view. Empty == healthy."""
|
||||||
@@ -781,6 +736,385 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
forward_done, out_cache_loc_virtual
|
forward_done, out_cache_loc_virtual
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def _build_swa_attn_allocator(self, **kwargs) -> MultiEndedAllocator: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def _wire_peers(self) -> None: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def _compute_available_size(self) -> int: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def ensure_capacity(self, full_tokens: int, swa_tokens: int) -> bool: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def _flush_targets(self): ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def _ask_float_for_room(self, need_tokens: int) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
|
class UnifiedSWATokenToKVPoolAllocator(UnifiedSWAAllocatorBase):
|
||||||
|
"""Two-ended FULL/SWA allocator with asymmetric shared-byte reservations."""
|
||||||
|
|
||||||
|
def _build_swa_attn_allocator(self, **kwargs) -> MultiEndedAllocator:
|
||||||
|
"""The swa sub-allocator: an END pool in the 2-pool pair."""
|
||||||
|
return MultiEndedAllocator(
|
||||||
|
sub_pool_name="swa",
|
||||||
|
is_id_owner=False, # non-owner; consumes virtuals minted by full
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _wire_peers(self) -> None:
|
||||||
|
self.full_attn_allocator.bind_peer(self.swa_attn_allocator)
|
||||||
|
self.swa_attn_allocator.bind_peer(self.full_attn_allocator)
|
||||||
|
|
||||||
|
def _compute_available_size(self) -> int:
|
||||||
|
"""Joint byte budget in TOKENS: each composite alloc(1) consumes one
|
||||||
|
full-side AND one swa-side page under the same virtual id."""
|
||||||
|
fa, sa = self.full_attn_allocator, self.swa_attn_allocator
|
||||||
|
e_f = fa.entry_bytes_per_page
|
||||||
|
e_s = sa.entry_bytes_per_page
|
||||||
|
# Direction-agnostic shared gap: the free byte band between the two pools.
|
||||||
|
if fa.grow_direction == "up":
|
||||||
|
gap_bytes = max(0, sa._byte_low_frontier() - fa._byte_high_frontier())
|
||||||
|
else:
|
||||||
|
gap_bytes = max(0, fa._byte_low_frontier() - sa._byte_high_frontier())
|
||||||
|
R_f = fa.num_pages - fa.min_page_index - fa._allocated_pages()
|
||||||
|
R_s = sa.num_pages - sa.min_page_index - sa._allocated_pages()
|
||||||
|
|
||||||
|
if not self.lazy_compaction:
|
||||||
|
pages_by_bytes = gap_bytes // (e_f + e_s)
|
||||||
|
return min(pages_by_bytes, R_f, R_s) * self.page_size
|
||||||
|
|
||||||
|
H_f = len(fa._free_phys_pages)
|
||||||
|
H_s = len(sa._free_phys_pages)
|
||||||
|
|
||||||
|
K1 = min(H_f, H_s) # Phase 1: both drain
|
||||||
|
|
||||||
|
# Phase 2: fewer-holes side extends; more-holes side keeps draining.
|
||||||
|
if H_f <= H_s:
|
||||||
|
e_phase2 = e_f
|
||||||
|
K_phase2_max = H_s
|
||||||
|
else:
|
||||||
|
e_phase2 = e_s
|
||||||
|
K_phase2_max = H_f
|
||||||
|
K2_room = K_phase2_max - K1
|
||||||
|
K2 = min(K2_room, gap_bytes // e_phase2) if e_phase2 > 0 else K2_room
|
||||||
|
gap_bytes -= K2 * e_phase2
|
||||||
|
|
||||||
|
K3 = gap_bytes // (e_f + e_s) # Phase 3: both extend
|
||||||
|
|
||||||
|
K_total = K1 + K2 + K3
|
||||||
|
K_total = min(K_total, H_f + R_f, H_s + R_s) # index-space caps
|
||||||
|
return K_total * self.page_size
|
||||||
|
|
||||||
|
def _flush_targets(self):
|
||||||
|
"""Flush ALL members, including ones that are not short themselves: a
|
||||||
|
one-sided hole is unusable, and compacting it yields SHARED gap."""
|
||||||
|
return (self.full_attn_allocator, self.swa_attn_allocator)
|
||||||
|
|
||||||
|
def _ask_float_for_room(self, need_tokens: int) -> None:
|
||||||
|
"""No float in a two-END chain -- nothing can slide."""
|
||||||
|
return None
|
||||||
|
|
||||||
|
def reclaim_plan(
|
||||||
|
self,
|
||||||
|
full_tokens: int | float,
|
||||||
|
swa_tokens: int | float,
|
||||||
|
*,
|
||||||
|
full_evictable_tokens: int = 0,
|
||||||
|
swa_evictable_tokens: int = 0,
|
||||||
|
empty_pool: bool = False,
|
||||||
|
) -> Optional[Tuple[int, int]]:
|
||||||
|
"""Return cumulative FULL/SWA eviction targets, or None if impossible.
|
||||||
|
|
||||||
|
The tree evicts FULL before SWA. Minimize required SWA reclaim with all
|
||||||
|
evictable FULL available, then trim excess FULL reclaim. FULL eviction's
|
||||||
|
actual SWA cascade is counted by the tree's shared eviction tracker.
|
||||||
|
"""
|
||||||
|
if full_tokens < 0 or swa_tokens < 0:
|
||||||
|
return None
|
||||||
|
|
||||||
|
page_size = self.page_size
|
||||||
|
full_pages = (math.ceil(full_tokens) + page_size - 1) // page_size
|
||||||
|
swa_pages = (math.ceil(swa_tokens) + page_size - 1) // page_size
|
||||||
|
# Restoring host SWA for device-resident FULL can require more new
|
||||||
|
# SWA pages than new FULL pages; only the shared budget constrains it.
|
||||||
|
|
||||||
|
compacted = empty_pool or not self.lazy_compaction or self._compaction_allowed()
|
||||||
|
|
||||||
|
def fits(full_reclaim_pages: int, swa_reclaim_pages: int) -> bool:
|
||||||
|
return self._fits_page_demand(
|
||||||
|
full_pages,
|
||||||
|
swa_pages,
|
||||||
|
full_reclaim_pages=full_reclaim_pages,
|
||||||
|
swa_reclaim_pages=swa_reclaim_pages,
|
||||||
|
compacted=compacted,
|
||||||
|
empty_pool=empty_pool,
|
||||||
|
)
|
||||||
|
|
||||||
|
if empty_pool:
|
||||||
|
return (0, 0) if fits(0, 0) else None
|
||||||
|
if fits(0, 0):
|
||||||
|
return (0, 0)
|
||||||
|
|
||||||
|
max_full_pages = min(
|
||||||
|
self.full_attn_allocator.allocated_count() // page_size,
|
||||||
|
max(0, int(full_evictable_tokens)) // page_size,
|
||||||
|
)
|
||||||
|
max_swa_pages = min(
|
||||||
|
self.swa_attn_allocator.allocated_count() // page_size,
|
||||||
|
max(0, int(swa_evictable_tokens)) // page_size,
|
||||||
|
)
|
||||||
|
if not fits(max_full_pages, max_swa_pages):
|
||||||
|
return None
|
||||||
|
|
||||||
|
def first_fit(high: int, predicate: Callable[[int], bool]) -> int:
|
||||||
|
low = 0
|
||||||
|
while low < high:
|
||||||
|
mid = (low + high) // 2
|
||||||
|
if predicate(mid):
|
||||||
|
high = mid
|
||||||
|
else:
|
||||||
|
low = mid + 1
|
||||||
|
return low
|
||||||
|
|
||||||
|
swa_reclaim_pages = first_fit(
|
||||||
|
max_swa_pages, lambda value: fits(max_full_pages, value)
|
||||||
|
)
|
||||||
|
full_reclaim_pages = first_fit(
|
||||||
|
max_full_pages, lambda value: fits(value, swa_reclaim_pages)
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
full_reclaim_pages * page_size,
|
||||||
|
swa_reclaim_pages * page_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
def can_reserve(
|
||||||
|
self,
|
||||||
|
full_tokens: int | float,
|
||||||
|
swa_tokens: int | float,
|
||||||
|
*,
|
||||||
|
full_evictable_tokens: int = 0,
|
||||||
|
swa_evictable_tokens: int = 0,
|
||||||
|
empty_pool: bool = False,
|
||||||
|
require_token_slack: bool = False,
|
||||||
|
) -> bool:
|
||||||
|
"""Check pending FULL/SWA demand against the shared byte envelope.
|
||||||
|
|
||||||
|
Scheduler admission keeps the historical one-token strict slack at an
|
||||||
|
empty-pool boundary. Live admission checks the state reachable after
|
||||||
|
reclaiming the currently evictable pages from both sides.
|
||||||
|
"""
|
||||||
|
if full_tokens < 0 or swa_tokens < 0:
|
||||||
|
return False
|
||||||
|
if require_token_slack and (
|
||||||
|
full_tokens >= self.size_full or swa_tokens >= self.size_swa
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
|
||||||
|
return (
|
||||||
|
self.reclaim_plan(
|
||||||
|
full_tokens,
|
||||||
|
swa_tokens,
|
||||||
|
full_evictable_tokens=full_evictable_tokens,
|
||||||
|
swa_evictable_tokens=swa_evictable_tokens,
|
||||||
|
empty_pool=empty_pool,
|
||||||
|
)
|
||||||
|
is not None
|
||||||
|
)
|
||||||
|
|
||||||
|
def _compaction_allowed(self) -> bool:
|
||||||
|
return all(
|
||||||
|
allocator.disagg_move_gate is None or allocator.disagg_move_gate()
|
||||||
|
for allocator in (self.full_attn_allocator, self.swa_attn_allocator)
|
||||||
|
)
|
||||||
|
|
||||||
|
def _fits_page_demand(
|
||||||
|
self,
|
||||||
|
num_full_pages: int,
|
||||||
|
num_swa_pages: int,
|
||||||
|
*,
|
||||||
|
full_reclaim_pages: int = 0,
|
||||||
|
swa_reclaim_pages: int = 0,
|
||||||
|
compacted: bool,
|
||||||
|
empty_pool: bool = False,
|
||||||
|
) -> bool:
|
||||||
|
"""Check one FULL/SWA page demand against a single allocator snapshot."""
|
||||||
|
if min(num_full_pages, num_swa_pages) < 0:
|
||||||
|
return False
|
||||||
|
|
||||||
|
fa, sa = self.full_attn_allocator, self.swa_attn_allocator
|
||||||
|
if empty_pool:
|
||||||
|
full_live_pages = swa_live_pages = 0
|
||||||
|
full_reclaim_pages = swa_reclaim_pages = 0
|
||||||
|
else:
|
||||||
|
full_live_pages = fa.allocated_count() // self.page_size
|
||||||
|
swa_live_pages = sa.allocated_count() // self.page_size
|
||||||
|
full_reclaim_pages = min(full_live_pages, max(0, int(full_reclaim_pages)))
|
||||||
|
swa_reclaim_pages = min(swa_live_pages, max(0, int(swa_reclaim_pages)))
|
||||||
|
|
||||||
|
full_live_pages -= full_reclaim_pages
|
||||||
|
swa_live_pages -= swa_reclaim_pages
|
||||||
|
full_total_pages = full_live_pages + num_full_pages
|
||||||
|
swa_total_pages = swa_live_pages + num_swa_pages
|
||||||
|
virtual_page_capacity = fa.num_virtual_ids - fa.min_page_index
|
||||||
|
full_page_capacity = min(
|
||||||
|
virtual_page_capacity,
|
||||||
|
fa.num_pages - fa.min_page_index,
|
||||||
|
)
|
||||||
|
swa_page_capacity = min(
|
||||||
|
virtual_page_capacity,
|
||||||
|
sa.num_pages - sa.min_page_index,
|
||||||
|
)
|
||||||
|
if full_total_pages > full_page_capacity or swa_total_pages > swa_page_capacity:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if compacted:
|
||||||
|
full_virtual_room = virtual_page_capacity - full_live_pages
|
||||||
|
full_holes = swa_holes = 0
|
||||||
|
full_index_room = full_page_capacity - full_live_pages
|
||||||
|
swa_index_room = swa_page_capacity - swa_live_pages
|
||||||
|
gap_bytes = (
|
||||||
|
self._empty_shared_gap_bytes
|
||||||
|
- full_live_pages * fa.entry_bytes_per_page
|
||||||
|
- swa_live_pages * sa.entry_bytes_per_page
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
full_virtual_room = len(fa.free_virtual_ids) + full_reclaim_pages
|
||||||
|
full_holes = len(fa._free_phys_pages) + full_reclaim_pages
|
||||||
|
swa_holes = len(sa._free_phys_pages) + swa_reclaim_pages
|
||||||
|
full_index_room = fa.num_pages - fa.min_page_index - fa._allocated_pages()
|
||||||
|
swa_index_room = sa.num_pages - sa.min_page_index - sa._allocated_pages()
|
||||||
|
gap_bytes = fa._current_gap_bytes()
|
||||||
|
|
||||||
|
if num_full_pages > full_virtual_room:
|
||||||
|
return False
|
||||||
|
if num_full_pages > full_holes + full_index_room:
|
||||||
|
return False
|
||||||
|
if num_swa_pages > swa_holes + swa_index_room:
|
||||||
|
return False
|
||||||
|
full_extensions = max(0, num_full_pages - full_holes)
|
||||||
|
swa_extensions = max(0, num_swa_pages - swa_holes)
|
||||||
|
return (
|
||||||
|
full_extensions * fa.entry_bytes_per_page
|
||||||
|
+ swa_extensions * sa.entry_bytes_per_page
|
||||||
|
<= max(0, gap_bytes)
|
||||||
|
)
|
||||||
|
|
||||||
|
def ensure_capacity(self, full_tokens: int, swa_tokens: int) -> bool:
|
||||||
|
"""Gate one allocation and compact both sides on shortfall."""
|
||||||
|
if full_tokens < 0 or swa_tokens < 0:
|
||||||
|
return False
|
||||||
|
if full_tokens == 0 and swa_tokens == 0:
|
||||||
|
return True
|
||||||
|
page_size = self.page_size
|
||||||
|
num_full_pages = (int(full_tokens) + page_size - 1) // page_size
|
||||||
|
num_swa_pages = (int(swa_tokens) + page_size - 1) // page_size
|
||||||
|
if self._fits_page_demand(
|
||||||
|
num_full_pages,
|
||||||
|
num_swa_pages,
|
||||||
|
compacted=False,
|
||||||
|
):
|
||||||
|
return True
|
||||||
|
if not self.lazy_compaction or not self._compaction_allowed():
|
||||||
|
return False
|
||||||
|
if not self._fits_page_demand(
|
||||||
|
num_full_pages,
|
||||||
|
num_swa_pages,
|
||||||
|
compacted=True,
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
self.full_attn_allocator.flush_for_allocation()
|
||||||
|
self.swa_attn_allocator.flush_for_allocation()
|
||||||
|
return self._fits_page_demand(
|
||||||
|
num_full_pages,
|
||||||
|
num_swa_pages,
|
||||||
|
compacted=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
def create_prefill_budget(self, tree_cache, *, num_mixed_decode_tokens=0):
|
||||||
|
from sglang.srt.mem_cache.prefill_budget import SharedSWAPrefillBudget
|
||||||
|
|
||||||
|
return SharedSWAPrefillBudget(
|
||||||
|
self, tree_cache, num_mixed_decode_tokens=num_mixed_decode_tokens
|
||||||
|
)
|
||||||
|
|
||||||
|
def swa_capacity_and_available(self, *, full_capacity, swa_capacity):
|
||||||
|
return (
|
||||||
|
(self.current_full_capacity, self.full_available_size()),
|
||||||
|
(self.current_swa_capacity, self.swa_available_size()),
|
||||||
|
)
|
||||||
|
|
||||||
|
def max_new_tokens_for_memory(
|
||||||
|
self,
|
||||||
|
input_tokens: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
*,
|
||||||
|
token_capacity: int,
|
||||||
|
sliding_window_size: int | None,
|
||||||
|
chunk_size: int | None,
|
||||||
|
) -> int | None:
|
||||||
|
from sglang.srt.mem_cache.prefill_budget import estimate_swa_kv_tokens
|
||||||
|
|
||||||
|
def fits(candidate):
|
||||||
|
return self.can_reserve(
|
||||||
|
input_tokens + candidate + self.page_size,
|
||||||
|
estimate_swa_kv_tokens(
|
||||||
|
input_tokens,
|
||||||
|
candidate,
|
||||||
|
sliding_window_size=sliding_window_size,
|
||||||
|
page_size=self.page_size,
|
||||||
|
allocation_limit=chunk_size,
|
||||||
|
),
|
||||||
|
empty_pool=True,
|
||||||
|
require_token_slack=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not fits(0):
|
||||||
|
return None
|
||||||
|
if fits(max_new_tokens):
|
||||||
|
return max_new_tokens
|
||||||
|
lo, hi = 0, max_new_tokens
|
||||||
|
while lo < hi:
|
||||||
|
mid = (lo + hi + 1) // 2
|
||||||
|
if fits(mid):
|
||||||
|
lo = mid
|
||||||
|
else:
|
||||||
|
hi = mid - 1
|
||||||
|
return lo
|
||||||
|
|
||||||
|
def evict_to_free_tokens(self, tree_cache, num_tokens: int) -> bool | None:
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||||
|
|
||||||
|
if tree_cache is None or tree_cache.is_chunk_cache():
|
||||||
|
return
|
||||||
|
reclaim_plan = self.reclaim_plan(
|
||||||
|
num_tokens,
|
||||||
|
num_tokens,
|
||||||
|
full_evictable_tokens=tree_cache.full_evictable_size(),
|
||||||
|
swa_evictable_tokens=tree_cache.swa_evictable_size(),
|
||||||
|
)
|
||||||
|
if reclaim_plan is None:
|
||||||
|
return
|
||||||
|
full_reclaim, swa_reclaim = reclaim_plan
|
||||||
|
if full_reclaim or swa_reclaim:
|
||||||
|
tree_cache.evict_for_alloc(
|
||||||
|
EvictParams(num_tokens=full_reclaim, swa_num_tokens=swa_reclaim)
|
||||||
|
)
|
||||||
|
# A zero-reclaim plan can still depend on compaction before allocation.
|
||||||
|
return self.ensure_capacity(num_tokens, num_tokens)
|
||||||
|
|
||||||
|
def verify_byte_accounting(self) -> List[str]:
|
||||||
|
return (
|
||||||
|
_chain_byte_accounting_violations(
|
||||||
|
_end_pair_chain(self.full_attn_allocator, self.swa_attn_allocator)
|
||||||
|
)
|
||||||
|
+ self._joint_capacity_memo_violations()
|
||||||
|
)
|
||||||
|
|
||||||
def flush_opportunistic(self) -> int:
|
def flush_opportunistic(self) -> int:
|
||||||
"""Non-urgent flush of BOTH sub-allocators; sync-free."""
|
"""Non-urgent flush of BOTH sub-allocators; sync-free."""
|
||||||
with record_function("UnifiedSWAAlloc.flush_opportunistic"):
|
with record_function("UnifiedSWAAlloc.flush_opportunistic"):
|
||||||
@@ -796,7 +1130,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
|
|||||||
return fa.flush_opportunistic() + sa.flush_opportunistic()
|
return fa.flush_opportunistic() + sa.flush_opportunistic()
|
||||||
|
|
||||||
|
|
||||||
class UnifiedMambaSWATokenToKVPoolAllocator(UnifiedSWATokenToKVPoolAllocator):
|
class UnifiedMambaSWATokenToKVPoolAllocator(UnifiedSWAAllocatorBase):
|
||||||
"""Tri-pool composite for models with full KV + SWA KV + mamba/conv state
|
"""Tri-pool composite for models with full KV + SWA KV + mamba/conv state
|
||||||
(both `mambaish_config` and `is_hybrid_swa`).
|
(both `mambaish_config` and `is_hybrid_swa`).
|
||||||
|
|
||||||
@@ -887,6 +1221,37 @@ class UnifiedMambaSWATokenToKVPoolAllocator(UnifiedSWATokenToKVPoolAllocator):
|
|||||||
|
|
||||||
# -- capacity --
|
# -- capacity --
|
||||||
|
|
||||||
|
def can_reserve(
|
||||||
|
self,
|
||||||
|
full_tokens: int | float,
|
||||||
|
swa_tokens: int | float,
|
||||||
|
*,
|
||||||
|
full_evictable_tokens: int = 0,
|
||||||
|
swa_evictable_tokens: int = 0,
|
||||||
|
empty_pool: bool = False,
|
||||||
|
require_token_slack: bool = False,
|
||||||
|
) -> bool:
|
||||||
|
if (
|
||||||
|
full_tokens < 0
|
||||||
|
or swa_tokens < 0
|
||||||
|
or full_tokens != swa_tokens
|
||||||
|
or full_evictable_tokens
|
||||||
|
or swa_evictable_tokens
|
||||||
|
or empty_pool
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
return full_tokens <= self.available_size()
|
||||||
|
|
||||||
|
def ensure_capacity(self, full_tokens: int, swa_tokens: int) -> bool:
|
||||||
|
if full_tokens < 0 or swa_tokens < 0 or full_tokens != swa_tokens:
|
||||||
|
return False
|
||||||
|
if full_tokens == 0:
|
||||||
|
return True
|
||||||
|
need_tokens = int(full_tokens)
|
||||||
|
if need_tokens <= self.available_size():
|
||||||
|
return True
|
||||||
|
return _relieve_for_alloc(self, need_tokens)
|
||||||
|
|
||||||
def _compute_available_size(self) -> int:
|
def _compute_available_size(self) -> int:
|
||||||
"""Joint TOKENS for `alloc(N)`: N costs N full pages AND N swa pages, drawn
|
"""Joint TOKENS for `alloc(N)`: N costs N full pages AND N swa pages, drawn
|
||||||
from DIFFERENT bands -- full extends only into the high band, the float into
|
from DIFFERENT bands -- full extends only into the high band, the float into
|
||||||
@@ -1036,17 +1401,33 @@ class UnifiedMambaSWATokenToKVPoolAllocator(UnifiedSWATokenToKVPoolAllocator):
|
|||||||
super().set_inflight_forward(forward_done, out_cache_loc_virtual)
|
super().set_inflight_forward(forward_done, out_cache_loc_virtual)
|
||||||
self.mamba_allocator.set_inflight_forward(forward_done, None)
|
self.mamba_allocator.set_inflight_forward(forward_done, None)
|
||||||
|
|
||||||
|
def create_prefill_budget(self, tree_cache, *, num_mixed_decode_tokens=0):
|
||||||
|
# Mamba competes for the shared gap too; retain the tri-pool's existing
|
||||||
|
# token and state-slot admission until it has a three-way reservation.
|
||||||
|
return SWATokenToKVPoolAllocator.create_prefill_budget(
|
||||||
|
self, tree_cache, num_mixed_decode_tokens=num_mixed_decode_tokens
|
||||||
|
)
|
||||||
|
|
||||||
|
def max_new_tokens_for_memory(self, *args, **kwargs):
|
||||||
|
return BaseTokenToKVPoolAllocator.max_new_tokens_for_memory(
|
||||||
|
self, *args, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
def swa_capacity_and_available(self, *, full_capacity, swa_capacity):
|
||||||
|
return (
|
||||||
|
(full_capacity, self.conserve_full_available_size()),
|
||||||
|
(swa_capacity, self.conserve_swa_available_size()),
|
||||||
|
)
|
||||||
|
|
||||||
def evict_to_free_tokens(self, tree_cache, num_tokens: int) -> None:
|
def evict_to_free_tokens(self, tree_cache, num_tokens: int) -> None:
|
||||||
"""Joint-aware eviction: one tri-lifetime node frees bytes on several sides
|
"""Joint-aware eviction: one tri-lifetime node frees bytes on several sides
|
||||||
at once, so re-check the JOINT gate instead of the per-side shortfall."""
|
at once, so re-check the JOINT gate instead of the per-side shortfall."""
|
||||||
from sglang.srt.mem_cache.common import evict_from_tree_cache
|
|
||||||
|
|
||||||
# Arbitrary retry bound; a round that frees nothing ends the loop anyway.
|
# Arbitrary retry bound; a round that frees nothing ends the loop anyway.
|
||||||
for _ in range(4):
|
for _ in range(4):
|
||||||
before = self.available_size()
|
before = self.available_size()
|
||||||
if before >= num_tokens:
|
if before >= num_tokens:
|
||||||
return
|
return
|
||||||
evict_from_tree_cache(tree_cache, num_tokens)
|
SWATokenToKVPoolAllocator.evict_to_free_tokens(self, tree_cache, num_tokens)
|
||||||
if self.available_size() <= before:
|
if self.available_size() <= before:
|
||||||
return # no progress
|
return # no progress
|
||||||
|
|
||||||
|
|||||||
@@ -345,9 +345,13 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
|
|
||||||
# v2p is indexed by VIRTUAL page id, p2v by PHYSICAL page id. A non-owner
|
# v2p is indexed by VIRTUAL page id, p2v by PHYSICAL page id. A non-owner
|
||||||
# consumes the owner's ids, so the two counts are unrelated.
|
# consumes the owner's ids, so the two counts are unrelated.
|
||||||
|
assert virtual_num_pages is None or not is_id_owner, (
|
||||||
|
"only a non-owner allocator may use another pool's virtual-id space"
|
||||||
|
)
|
||||||
self.num_virtual_ids = (
|
self.num_virtual_ids = (
|
||||||
self.num_pages if virtual_num_pages is None else virtual_num_pages
|
self.num_pages if virtual_num_pages is None else virtual_num_pages
|
||||||
)
|
)
|
||||||
|
assert self.num_virtual_ids > 0, "virtual page count must be positive"
|
||||||
# Page 0 is the padding anchor; the trailing row is the -1 sentinel.
|
# Page 0 is the padding anchor; the trailing row is the -1 sentinel.
|
||||||
self.virtual_to_physical = torch.full(
|
self.virtual_to_physical = torch.full(
|
||||||
(self.num_virtual_ids + 1,),
|
(self.num_virtual_ids + 1,),
|
||||||
@@ -554,6 +558,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
f"is_id_owner={self.is_id_owner}, page_size={self.page_size}, "
|
f"is_id_owner={self.is_id_owner}, page_size={self.page_size}, "
|
||||||
f"min_page_index={self.min_page_index}, "
|
f"min_page_index={self.min_page_index}, "
|
||||||
f"num_pages={self.num_pages}, "
|
f"num_pages={self.num_pages}, "
|
||||||
|
f"num_virtual_ids={self.num_virtual_ids}, "
|
||||||
f"watermark_physical={self.watermark_physical}, "
|
f"watermark_physical={self.watermark_physical}, "
|
||||||
f"allocated_pages={self._allocated_pages()}"
|
f"allocated_pages={self._allocated_pages()}"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ from sglang.kernels.ops.memory.common import (
|
|||||||
)
|
)
|
||||||
from sglang.kernels.ops.memory.common import get_last_loc_kernel as get_last_loc_kernel
|
from sglang.kernels.ops.memory.common import get_last_loc_kernel as get_last_loc_kernel
|
||||||
from sglang.srt.mem_cache.allocator.page_interleave import page_interleave_shard_size
|
from sglang.srt.mem_cache.allocator.page_interleave import page_interleave_shard_size
|
||||||
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
|
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
||||||
from sglang.srt.mem_cache.hicache_storage import PoolTransfer
|
from sglang.srt.mem_cache.hicache_storage import PoolTransfer
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||||
@@ -161,34 +160,13 @@ def maybe_cache_unfinished_req(req: Req, tree_cache: BasePrefixCache, **kwargs):
|
|||||||
tree_cache.cache_unfinished_req(req, **kwargs)
|
tree_cache.cache_unfinished_req(req, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
def evict_from_tree_cache(tree_cache: BasePrefixCache | None, num_tokens: int):
|
def evict_from_tree_cache(
|
||||||
if tree_cache is None:
|
tree_cache: BasePrefixCache | None, num_tokens: int
|
||||||
return
|
) -> bool | None:
|
||||||
|
if tree_cache is not None and not tree_cache.is_chunk_cache():
|
||||||
if tree_cache.is_chunk_cache():
|
return tree_cache.token_to_kv_pool_allocator.evict_to_free_tokens(
|
||||||
return
|
tree_cache, num_tokens
|
||||||
|
)
|
||||||
allocator = tree_cache.token_to_kv_pool_allocator
|
|
||||||
|
|
||||||
if isinstance(allocator, SWATokenToKVPoolAllocator):
|
|
||||||
# Hybrid allocator
|
|
||||||
full_available_size = allocator.full_available_size()
|
|
||||||
swa_available_size = allocator.swa_available_size()
|
|
||||||
|
|
||||||
if full_available_size < num_tokens or swa_available_size < num_tokens:
|
|
||||||
full_num_tokens = max(0, num_tokens - full_available_size)
|
|
||||||
swa_num_tokens = max(0, num_tokens - swa_available_size)
|
|
||||||
tree_cache.evict_for_alloc(
|
|
||||||
EvictParams(num_tokens=full_num_tokens, swa_num_tokens=swa_num_tokens)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# Standard allocator: evict only the shortfall (mirrors the SWA arm)
|
|
||||||
available_size = allocator.available_size()
|
|
||||||
if available_size < num_tokens:
|
|
||||||
tree_cache.evict_for_alloc(
|
|
||||||
EvictParams(num_tokens=num_tokens - available_size)
|
|
||||||
)
|
|
||||||
_evict_until_allocatable(tree_cache, allocator, num_tokens)
|
|
||||||
|
|
||||||
|
|
||||||
def _evict_until_allocatable(
|
def _evict_until_allocatable(
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ from sglang.srt.mem_cache.allocator.swa import (
|
|||||||
is_swa_req_ring,
|
is_swa_req_ring,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
||||||
UnifiedSWATokenToKVPoolAllocator,
|
UnifiedSWAAllocatorBase,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator.unified_mamba import (
|
from sglang.srt.mem_cache.allocator.unified_mamba import (
|
||||||
UnifiedMambaTokenToKVPoolAllocator,
|
UnifiedMambaTokenToKVPoolAllocator,
|
||||||
@@ -240,6 +240,7 @@ class _PoolSizes(msgspec.Struct, frozen=True, kw_only=True):
|
|||||||
c128_state_pool_size: int
|
c128_state_pool_size: int
|
||||||
c4_state_dtype: Optional[torch.dtype]
|
c4_state_dtype: Optional[torch.dtype]
|
||||||
c128_state_dtype: Optional[torch.dtype]
|
c128_state_dtype: Optional[torch.dtype]
|
||||||
|
unified_memory_pool_bytes: Optional[int] = None
|
||||||
unified_total_bytes: Optional[int] = None
|
unified_total_bytes: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
@@ -289,6 +290,23 @@ class KVCacheConfigurator:
|
|||||||
self.draft_model_idx in self.model_config.swa_attention_layer_ids
|
self.draft_model_idx in self.model_config.swa_attention_layer_ids
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def hybrid_swa_token_capacity(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
allocator: BaseTokenToKVPoolAllocator,
|
||||||
|
full_capacity: Optional[int],
|
||||||
|
swa_capacity: Optional[int],
|
||||||
|
) -> int:
|
||||||
|
if get_memory().enable_unified_memory:
|
||||||
|
capacity = allocator.size_full
|
||||||
|
max_total_tokens = get_schedule().max_total_tokens
|
||||||
|
return (
|
||||||
|
min(capacity, max_total_tokens)
|
||||||
|
if max_total_tokens is not None
|
||||||
|
else capacity
|
||||||
|
)
|
||||||
|
return full_capacity or swa_capacity
|
||||||
|
|
||||||
def _build_fp4_quant_method(self, *, num_layers: int):
|
def _build_fp4_quant_method(self, *, num_layers: int):
|
||||||
if not is_float4_e2m1fn_x2(self.kv_cache_dtype):
|
if not is_float4_e2m1fn_x2(self.kv_cache_dtype):
|
||||||
return None
|
return None
|
||||||
@@ -423,6 +441,10 @@ class KVCacheConfigurator:
|
|||||||
max_running_requests=max_running_requests,
|
max_running_requests=max_running_requests,
|
||||||
full_max_total_num_tokens=full_max_total_num_tokens,
|
full_max_total_num_tokens=full_max_total_num_tokens,
|
||||||
swa_max_total_num_tokens=swa_max_total_num_tokens,
|
swa_max_total_num_tokens=swa_max_total_num_tokens,
|
||||||
|
# The target's byte envelope excludes the separate draft allocation.
|
||||||
|
unified_memory_pool_bytes=(
|
||||||
|
None if self.is_draft_worker else config.unified_memory_pool_bytes
|
||||||
|
),
|
||||||
c4_max_total_num_tokens=c4_max_total_num_tokens,
|
c4_max_total_num_tokens=c4_max_total_num_tokens,
|
||||||
c128_max_total_num_tokens=c128_max_total_num_tokens,
|
c128_max_total_num_tokens=c128_max_total_num_tokens,
|
||||||
c4_state_pool_size=c4_state_pool_size,
|
c4_state_pool_size=c4_state_pool_size,
|
||||||
@@ -469,6 +491,7 @@ class KVCacheConfigurator:
|
|||||||
max_num_reqs=sizes.max_running_requests,
|
max_num_reqs=sizes.max_running_requests,
|
||||||
full_max_total_num_tokens=sizes.full_max_total_num_tokens,
|
full_max_total_num_tokens=sizes.full_max_total_num_tokens,
|
||||||
swa_max_total_num_tokens=sizes.swa_max_total_num_tokens,
|
swa_max_total_num_tokens=sizes.swa_max_total_num_tokens,
|
||||||
|
unified_memory_pool_bytes=sizes.unified_memory_pool_bytes,
|
||||||
unified_total_bytes=sizes.unified_total_bytes,
|
unified_total_bytes=sizes.unified_total_bytes,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -497,7 +520,7 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator,
|
||||||
(
|
(
|
||||||
UnifiedMambaTokenToKVPoolAllocator,
|
UnifiedMambaTokenToKVPoolAllocator,
|
||||||
UnifiedSWATokenToKVPoolAllocator,
|
UnifiedSWAAllocatorBase,
|
||||||
),
|
),
|
||||||
):
|
):
|
||||||
draft_virtual_id_space = (
|
draft_virtual_id_space = (
|
||||||
@@ -520,7 +543,7 @@ class KVCacheConfigurator:
|
|||||||
if (
|
if (
|
||||||
isinstance(
|
isinstance(
|
||||||
token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator,
|
||||||
UnifiedSWATokenToKVPoolAllocator,
|
UnifiedSWAAllocatorBase,
|
||||||
)
|
)
|
||||||
and self.is_hybrid_swa
|
and self.is_hybrid_swa
|
||||||
):
|
):
|
||||||
@@ -821,8 +844,9 @@ class KVCacheConfigurator:
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
max_num_reqs: int,
|
max_num_reqs: int,
|
||||||
full_max_total_num_tokens: Optional[int],
|
full_max_total_num_tokens: Optional[int] = None,
|
||||||
swa_max_total_num_tokens: Optional[int],
|
swa_max_total_num_tokens: Optional[int] = None,
|
||||||
|
unified_memory_pool_bytes: Optional[int] = None,
|
||||||
unified_total_bytes: Optional[int] = None,
|
unified_total_bytes: Optional[int] = None,
|
||||||
) -> UnifiedPoolBundle:
|
) -> UnifiedPoolBundle:
|
||||||
"""Build the unified-pool stack for a hybrid-SWA model (Triton): one byte
|
"""Build the unified-pool stack for a hybrid-SWA model (Triton): one byte
|
||||||
@@ -899,6 +923,12 @@ class KVCacheConfigurator:
|
|||||||
if self.layer_info.start_layer <= i < self.layer_info.end_layer
|
if self.layer_info.start_layer <= i < self.layer_info.end_layer
|
||||||
]
|
]
|
||||||
|
|
||||||
|
total_bytes = unified_memory_pool_bytes
|
||||||
|
# An uncapped, draft-free pool owns the profiled budget, including bytes
|
||||||
|
# left over after rounding the FULL/SWA boot capacities to pages.
|
||||||
|
if unified_total_bytes is not None and self.spec_algorithm.is_none():
|
||||||
|
total_bytes = unified_total_bytes
|
||||||
|
|
||||||
bundle = init_unified_swa_pools(
|
bundle = init_unified_swa_pools(
|
||||||
device=self.device,
|
device=self.device,
|
||||||
kv_cache_dtype=self.kv_cache_dtype,
|
kv_cache_dtype=self.kv_cache_dtype,
|
||||||
@@ -915,6 +945,7 @@ class KVCacheConfigurator:
|
|||||||
full_attention_layer_ids=full_attention_layer_ids,
|
full_attention_layer_ids=full_attention_layer_ids,
|
||||||
full_max_total_num_tokens=full_max_total_num_tokens,
|
full_max_total_num_tokens=full_max_total_num_tokens,
|
||||||
swa_max_total_num_tokens=swa_max_total_num_tokens,
|
swa_max_total_num_tokens=swa_max_total_num_tokens,
|
||||||
|
total_bytes=total_bytes,
|
||||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||||
need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"),
|
need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"),
|
||||||
# Overlap mode: same wait_stream(forward_stream) rationale as
|
# Overlap mode: same wait_stream(forward_stream) rationale as
|
||||||
@@ -922,9 +953,6 @@ class KVCacheConfigurator:
|
|||||||
forward_stream=self.forward_stream,
|
forward_stream=self.forward_stream,
|
||||||
# Lazy compaction: default ON, with env var escape hatch for rollback / A/B.
|
# Lazy compaction: default ON, with env var escape hatch for rollback / A/B.
|
||||||
lazy_compaction=_should_enable_lazy_compaction(),
|
lazy_compaction=_should_enable_lazy_compaction(),
|
||||||
# Draft workers keep the token-count byte sum (spec is asserted
|
|
||||||
# off under unified; belt only).
|
|
||||||
unified_total_bytes=(None if self.is_draft_worker else unified_total_bytes),
|
|
||||||
# bs=1 feasibility floor inputs. `model_context_len` bounds the
|
# bs=1 feasibility floor inputs. `model_context_len` bounds the
|
||||||
# sliding window term only -- the full-attention side is not
|
# sliding window term only -- the full-attention side is not
|
||||||
# charged, see `_check_bs1_feasibility_floor`.
|
# charged, see `_check_bs1_feasibility_floor`.
|
||||||
@@ -2120,7 +2148,7 @@ class KVCacheConfigurator:
|
|||||||
else:
|
else:
|
||||||
swa_allocator = token_to_kv_pool_allocator
|
swa_allocator = token_to_kv_pool_allocator
|
||||||
uses_unified_virtual_ids = isinstance(
|
uses_unified_virtual_ids = isinstance(
|
||||||
swa_allocator, UnifiedSWATokenToKVPoolAllocator
|
swa_allocator, UnifiedSWAAllocatorBase
|
||||||
)
|
)
|
||||||
has_draft_swa_layers = (
|
has_draft_swa_layers = (
|
||||||
not self.is_hybrid_swa_mtp_draft or self.draft_swa_full_capacity
|
not self.is_hybrid_swa_mtp_draft or self.draft_swa_full_capacity
|
||||||
@@ -2363,9 +2391,8 @@ class KVCacheConfigurator:
|
|||||||
f"{config.max_total_num_tokens}"
|
f"{config.max_total_num_tokens}"
|
||||||
)
|
)
|
||||||
if max_tokens != config.max_total_num_tokens:
|
if max_tokens != config.max_total_num_tokens:
|
||||||
# Token-capped re-derivation: the profiled budget no longer
|
# Re-derive the capped budget: SWA carries unified_memory_pool_bytes;
|
||||||
# applies; the recalced config's unified_total_bytes stays None
|
# Mamba factories fall back to token-count sizing without unified_total_bytes.
|
||||||
# and the factories fall back to the token-count byte sum.
|
|
||||||
config = configurator.calculate_pool_sizes_from_max_tokens(
|
config = configurator.calculate_pool_sizes_from_max_tokens(
|
||||||
max_tokens, get_schedule().page_size
|
max_tokens, get_schedule().page_size
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -65,7 +65,7 @@ from sglang.kernels.ops.kvcache.kv_read_table import (
|
|||||||
build_kv_read_table_packed,
|
build_kv_read_table_packed,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
||||||
UnifiedSWATokenToKVPoolAllocator,
|
UnifiedSWAAllocatorBase,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator.unified_mamba import (
|
from sglang.srt.mem_cache.allocator.unified_mamba import (
|
||||||
UnifiedMambaTokenToKVPoolAllocator,
|
UnifiedMambaTokenToKVPoolAllocator,
|
||||||
@@ -121,12 +121,13 @@ class KVIndexTranslator:
|
|||||||
self.is_translating = (
|
self.is_translating = (
|
||||||
isinstance(
|
isinstance(
|
||||||
token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator,
|
||||||
(UnifiedMambaTokenToKVPoolAllocator, UnifiedSWATokenToKVPoolAllocator),
|
(UnifiedMambaTokenToKVPoolAllocator, UnifiedSWAAllocatorBase),
|
||||||
)
|
)
|
||||||
and token_to_kv_pool_allocator.get_kvcache() is token_to_kv_pool
|
and token_to_kv_pool_allocator.get_kvcache() is token_to_kv_pool
|
||||||
)
|
)
|
||||||
if self.is_translating:
|
if self.is_translating:
|
||||||
alloc = token_to_kv_pool_allocator
|
alloc = token_to_kv_pool_allocator
|
||||||
|
self._capture_page_size = alloc.page_size
|
||||||
self._full_v2p_table = alloc.full_v2p_page_table
|
self._full_v2p_table = alloc.full_v2p_page_table
|
||||||
self._full_p2v_table = alloc.full_p2v_page_table
|
self._full_p2v_table = alloc.full_p2v_page_table
|
||||||
self._full_page_multiplier = alloc.kernel_page_multiplier
|
self._full_page_multiplier = alloc.kernel_page_multiplier
|
||||||
@@ -139,7 +140,7 @@ class KVIndexTranslator:
|
|||||||
# DCP read ids stay WIDENED to the consumer: selecting this rank's
|
# DCP read ids stay WIDENED to the consumer: selecting this rank's
|
||||||
# share changes the length, so only the production site can do it.
|
# share changes the length, so only the production site can do it.
|
||||||
self.defer_read_translate = get_parallel().attn_dcp_size > 1
|
self.defer_read_translate = get_parallel().attn_dcp_size > 1
|
||||||
if isinstance(alloc, UnifiedSWATokenToKVPoolAllocator):
|
if isinstance(alloc, UnifiedSWAAllocatorBase):
|
||||||
self._swa_v2p_table = alloc.swa_v2p_page_table
|
self._swa_v2p_table = alloc.swa_v2p_page_table
|
||||||
self._swa_page_multiplier = alloc.swa_kernel_page_multiplier
|
self._swa_page_multiplier = alloc.swa_kernel_page_multiplier
|
||||||
self._swa_write_loc_from_full = self._swa_write_loc_unified
|
self._swa_write_loc_from_full = self._swa_write_loc_unified
|
||||||
@@ -171,6 +172,16 @@ class KVIndexTranslator:
|
|||||||
)
|
)
|
||||||
self._index_table_memo: Optional[Tuple[weakref.ref, KVIndexTable]] = None
|
self._index_table_memo: Optional[Tuple[weakref.ref, KVIndexTable]] = None
|
||||||
|
|
||||||
|
def capture_token_capacity(self, max_token_pool_size: int) -> int:
|
||||||
|
"""Host capture rows are indexed by request-token IDs, not kernel IDs.
|
||||||
|
|
||||||
|
Unified IDs span the whole virtual table even when admission is capped.
|
||||||
|
DCP widens allocator pages; the runner's page size stays physical.
|
||||||
|
"""
|
||||||
|
if self.is_translating:
|
||||||
|
return self._full_v2p_table.numel() * self._capture_page_size
|
||||||
|
return max_token_pool_size + self.page_size
|
||||||
|
|
||||||
# -- per-batch view --------------------------------------------------------
|
# -- per-batch view --------------------------------------------------------
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|||||||
@@ -0,0 +1,413 @@
|
|||||||
|
# Copyright 2026 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
"""Memory reservations for one prefill pass.
|
||||||
|
|
||||||
|
The scheduler supplies token demand and its chunk/decode limits. These objects
|
||||||
|
account for admitted but not yet allocated work and query live cache capacity:
|
||||||
|
locking a prefix or preempting a request must affect the next admission check.
|
||||||
|
They neither select requests nor mutate the prefix cache or allocator.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.allocator.swa import is_swa_req_ring
|
||||||
|
|
||||||
|
|
||||||
|
def estimate_swa_kv_tokens(
|
||||||
|
extend_input_len: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
*,
|
||||||
|
sliding_window_size: Optional[int],
|
||||||
|
page_size: int,
|
||||||
|
allocation_limit: Optional[int] = None,
|
||||||
|
host_hit_length: int = 0,
|
||||||
|
) -> int:
|
||||||
|
"""Peak SWA reservation for one prefill/decode request."""
|
||||||
|
if sliding_window_size is None or sliding_window_size <= 0:
|
||||||
|
reserved = extend_input_len + max_new_tokens + page_size
|
||||||
|
else:
|
||||||
|
allocated = (
|
||||||
|
extend_input_len
|
||||||
|
if allocation_limit is None
|
||||||
|
else min(extend_input_len, allocation_limit)
|
||||||
|
)
|
||||||
|
allocated_tail = max(allocated - sliding_window_size, 0)
|
||||||
|
# With a roughly two-window SWA pool, a cached prefix can already lock
|
||||||
|
# one window. Charging another whole window for a short resume can then
|
||||||
|
# block admission forever on an idle pool. Reserve only its uncached
|
||||||
|
# tail plus decode headroom, capped at the window.
|
||||||
|
# Including the extension also keeps the reservation large enough for
|
||||||
|
# this pass's prefill allocation.
|
||||||
|
reserved = (
|
||||||
|
allocated_tail
|
||||||
|
+ min(extend_input_len + max_new_tokens, sliding_window_size)
|
||||||
|
+ page_size
|
||||||
|
)
|
||||||
|
if host_hit_length > 0:
|
||||||
|
reserved += -(-host_hit_length // page_size) * page_size
|
||||||
|
return reserved
|
||||||
|
|
||||||
|
|
||||||
|
class PrefillBudget:
|
||||||
|
"""Fixed token pool. Offsets include pending allocations and decode headroom."""
|
||||||
|
|
||||||
|
def __init__(self, allocator, tree_cache, *, num_mixed_decode_tokens: int = 0):
|
||||||
|
self.allocator = allocator
|
||||||
|
self.tree_cache = tree_cache
|
||||||
|
self.page_size = allocator.page_size
|
||||||
|
self.total_offset = num_mixed_decode_tokens
|
||||||
|
self.current_offset = num_mixed_decode_tokens
|
||||||
|
self.swa_offset = 0
|
||||||
|
|
||||||
|
def ceil_paged_tokens(self, tokens: int) -> int:
|
||||||
|
return -(-tokens // self.page_size) * self.page_size
|
||||||
|
|
||||||
|
def _available_and_evictable(self):
|
||||||
|
evictable = (
|
||||||
|
self.tree_cache.full_evictable_size()
|
||||||
|
if self.tree_cache.supports_mamba()
|
||||||
|
else self.tree_cache.evictable_size()
|
||||||
|
)
|
||||||
|
return self.allocator.available_size() + evictable
|
||||||
|
|
||||||
|
@property
|
||||||
|
def remaining_total(self):
|
||||||
|
return self._available_and_evictable() - self.total_offset
|
||||||
|
|
||||||
|
@property
|
||||||
|
def remaining_current(self):
|
||||||
|
return self._available_and_evictable() - self.current_offset
|
||||||
|
|
||||||
|
@property
|
||||||
|
def remaining_swa(self):
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def has_capacity(self) -> bool:
|
||||||
|
return self.remaining_total > 0 and self.remaining_current > 0
|
||||||
|
|
||||||
|
def check_prefill(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
extend_input_len: int,
|
||||||
|
total_tokens: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
input_tokens: int,
|
||||||
|
swa_host_hit_length: int,
|
||||||
|
chunk_limit: int | None,
|
||||||
|
) -> tuple[bool, int | None]:
|
||||||
|
"""Return admission feasibility and a memory bound on the requested chunk."""
|
||||||
|
return (
|
||||||
|
(True, chunk_limit)
|
||||||
|
if total_tokens < self.remaining_total
|
||||||
|
else (False, None)
|
||||||
|
)
|
||||||
|
|
||||||
|
def can_allocate_prefill(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
paged_input: int,
|
||||||
|
extend_input_len: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
chunk_limit: int | None,
|
||||||
|
) -> bool:
|
||||||
|
return paged_input <= min(self.remaining_current, self.remaining_total)
|
||||||
|
|
||||||
|
def available_chunk_tokens(self, chunk_limit: int) -> int | None:
|
||||||
|
available = min(chunk_limit, int(self.remaining_total))
|
||||||
|
# Single-pool continuation must make progress to release its KV.
|
||||||
|
return available if available > 0 else chunk_limit
|
||||||
|
|
||||||
|
def fit_chunk(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
extend_input_len: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
chunk_limit: int,
|
||||||
|
) -> int | None:
|
||||||
|
return chunk_limit
|
||||||
|
|
||||||
|
def reserve(
|
||||||
|
self,
|
||||||
|
extend_input_len: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
*,
|
||||||
|
extra_tokens: int = 0,
|
||||||
|
chunk_limit: int | None = None,
|
||||||
|
is_chunked_continuation: bool = False,
|
||||||
|
) -> None:
|
||||||
|
extend_input_len = self.ceil_paged_tokens(extend_input_len)
|
||||||
|
immediate = extend_input_len + self.page_size + extra_tokens
|
||||||
|
self.total_offset += immediate + max_new_tokens
|
||||||
|
self.current_offset += immediate
|
||||||
|
|
||||||
|
|
||||||
|
class SWAPrefillBudget(PrefillBudget):
|
||||||
|
"""Separate FULL/SWA partitions, including per-request SWA rings."""
|
||||||
|
|
||||||
|
def __init__(self, *args, all_swa=False, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.all_swa = all_swa
|
||||||
|
self.req_ring = is_swa_req_ring(self.allocator)
|
||||||
|
|
||||||
|
def _available_and_evictable(self):
|
||||||
|
if self.all_swa:
|
||||||
|
return (
|
||||||
|
self.allocator.swa_available_size()
|
||||||
|
+ self.tree_cache.swa_evictable_size()
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
self.allocator.full_available_size() + self.tree_cache.full_evictable_size()
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def remaining_swa(self):
|
||||||
|
evictable = 0 if self.req_ring else self.tree_cache.swa_evictable_size()
|
||||||
|
return self.allocator.swa_available_size() + evictable - self.swa_offset
|
||||||
|
|
||||||
|
def swa_tokens(
|
||||||
|
self,
|
||||||
|
extend_input_len,
|
||||||
|
max_new_tokens,
|
||||||
|
*,
|
||||||
|
chunk_limit=None,
|
||||||
|
swa_host_hit_length=0,
|
||||||
|
):
|
||||||
|
if self.req_ring:
|
||||||
|
return self.allocator.swa_ring_cost_tokens
|
||||||
|
return estimate_swa_kv_tokens(
|
||||||
|
extend_input_len,
|
||||||
|
max_new_tokens,
|
||||||
|
sliding_window_size=self.tree_cache.sliding_window_size,
|
||||||
|
page_size=self.page_size,
|
||||||
|
allocation_limit=chunk_limit,
|
||||||
|
host_hit_length=swa_host_hit_length,
|
||||||
|
)
|
||||||
|
|
||||||
|
def swa_never_fits(self, extend_input_len, max_new_tokens, **kwargs):
|
||||||
|
needed = self.swa_tokens(extend_input_len, max_new_tokens, **kwargs)
|
||||||
|
return (
|
||||||
|
needed > self.allocator.size_swa
|
||||||
|
if self.req_ring
|
||||||
|
else needed >= self.allocator.size_swa
|
||||||
|
)
|
||||||
|
|
||||||
|
def _chunk_cap(self, max_new_tokens, swa_host_hit_length=0):
|
||||||
|
# Only the sliding window stays locked between chunks, so a smaller
|
||||||
|
# chunk can bound the transient SWA footprint of a longer prompt.
|
||||||
|
headroom = self.swa_tokens(
|
||||||
|
0, max_new_tokens, swa_host_hit_length=swa_host_hit_length
|
||||||
|
)
|
||||||
|
cap = int(self.remaining_swa) - headroom
|
||||||
|
return max(0, cap // self.page_size * self.page_size)
|
||||||
|
|
||||||
|
def check_prefill(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
extend_input_len: int,
|
||||||
|
total_tokens: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
input_tokens: int,
|
||||||
|
swa_host_hit_length: int,
|
||||||
|
chunk_limit: int | None,
|
||||||
|
) -> tuple[bool, int | None]:
|
||||||
|
if total_tokens >= self.remaining_total:
|
||||||
|
return False, None
|
||||||
|
extend_input_len = self.ceil_paged_tokens(extend_input_len)
|
||||||
|
needed = self.swa_tokens(
|
||||||
|
extend_input_len,
|
||||||
|
max_new_tokens,
|
||||||
|
chunk_limit=chunk_limit,
|
||||||
|
swa_host_hit_length=swa_host_hit_length,
|
||||||
|
)
|
||||||
|
fits = (
|
||||||
|
needed <= self.remaining_swa
|
||||||
|
if self.req_ring
|
||||||
|
else needed < self.remaining_swa
|
||||||
|
)
|
||||||
|
if fits:
|
||||||
|
return True, chunk_limit
|
||||||
|
# Only permanent shortfalls may shrink a chunk. Transient pressure waits
|
||||||
|
# so a new prefill does not consume running decodes' window headroom.
|
||||||
|
cap = 0
|
||||||
|
if self.swa_never_fits(
|
||||||
|
extend_input_len,
|
||||||
|
max_new_tokens,
|
||||||
|
chunk_limit=chunk_limit,
|
||||||
|
swa_host_hit_length=swa_host_hit_length,
|
||||||
|
):
|
||||||
|
cap = self._chunk_cap(max_new_tokens, swa_host_hit_length)
|
||||||
|
if chunk_limit is None or cap <= 0:
|
||||||
|
return False, None
|
||||||
|
return True, min(chunk_limit, cap)
|
||||||
|
|
||||||
|
def has_capacity(self) -> bool:
|
||||||
|
return super().has_capacity() and self.remaining_swa > 0
|
||||||
|
|
||||||
|
def available_chunk_tokens(self, chunk_limit: int) -> int | None:
|
||||||
|
available = min(chunk_limit, int(self.remaining_total))
|
||||||
|
if not self.req_ring:
|
||||||
|
available = min(available, int(self.remaining_swa) - self.page_size)
|
||||||
|
return available if available > 0 else None
|
||||||
|
|
||||||
|
def can_allocate_prefill(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
paged_input: int,
|
||||||
|
extend_input_len: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
chunk_limit: int | None,
|
||||||
|
) -> bool:
|
||||||
|
return (
|
||||||
|
super().can_allocate_prefill(
|
||||||
|
paged_input=paged_input,
|
||||||
|
extend_input_len=extend_input_len,
|
||||||
|
max_new_tokens=max_new_tokens,
|
||||||
|
chunk_limit=chunk_limit,
|
||||||
|
)
|
||||||
|
and self.swa_tokens(
|
||||||
|
extend_input_len, max_new_tokens, chunk_limit=chunk_limit
|
||||||
|
)
|
||||||
|
<= self.remaining_swa
|
||||||
|
)
|
||||||
|
|
||||||
|
def reserve(
|
||||||
|
self,
|
||||||
|
extend_input_len: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
*,
|
||||||
|
extra_tokens: int = 0,
|
||||||
|
chunk_limit: int | None = None,
|
||||||
|
is_chunked_continuation: bool = False,
|
||||||
|
) -> None:
|
||||||
|
super().reserve(extend_input_len, max_new_tokens, extra_tokens=extra_tokens)
|
||||||
|
# A continuation already owns its ring slot.
|
||||||
|
if not (self.req_ring and is_chunked_continuation):
|
||||||
|
self.swa_offset += self.swa_tokens(
|
||||||
|
self.ceil_paged_tokens(extend_input_len),
|
||||||
|
max_new_tokens,
|
||||||
|
chunk_limit=chunk_limit,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class SharedSWAPrefillBudget(SWAPrefillBudget):
|
||||||
|
"""FULL and SWA reservations compete for the same physical byte budget."""
|
||||||
|
|
||||||
|
def __init__(self, *args, num_mixed_decode_tokens=0, **kwargs):
|
||||||
|
super().__init__(
|
||||||
|
*args, num_mixed_decode_tokens=num_mixed_decode_tokens, **kwargs
|
||||||
|
)
|
||||||
|
self.swa_offset = num_mixed_decode_tokens
|
||||||
|
|
||||||
|
def _fits(self, full_tokens, swa_tokens, *, empty_pool=False):
|
||||||
|
return self.allocator.can_reserve(
|
||||||
|
full_tokens + (0 if empty_pool else self.total_offset),
|
||||||
|
swa_tokens + (0 if empty_pool else self.swa_offset),
|
||||||
|
full_evictable_tokens=0
|
||||||
|
if empty_pool
|
||||||
|
else self.tree_cache.full_evictable_size(),
|
||||||
|
swa_evictable_tokens=0
|
||||||
|
if empty_pool
|
||||||
|
else self.tree_cache.swa_evictable_size(),
|
||||||
|
empty_pool=empty_pool,
|
||||||
|
require_token_slack=empty_pool,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _joint_chunk_cap(self, *, max_chunk_tokens, chunk_limit, swa_host_hit_length=0):
|
||||||
|
lo, hi = 0, max(0, max_chunk_tokens) // self.page_size
|
||||||
|
while lo < hi:
|
||||||
|
mid = (lo + hi + 1) // 2
|
||||||
|
tokens = mid * self.page_size
|
||||||
|
if self._fits(
|
||||||
|
tokens + self.page_size,
|
||||||
|
self.swa_tokens(
|
||||||
|
tokens,
|
||||||
|
0,
|
||||||
|
chunk_limit=chunk_limit,
|
||||||
|
swa_host_hit_length=swa_host_hit_length,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
lo = mid
|
||||||
|
else:
|
||||||
|
hi = mid - 1
|
||||||
|
return lo * self.page_size
|
||||||
|
|
||||||
|
def check_prefill(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
extend_input_len: int,
|
||||||
|
total_tokens: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
input_tokens: int,
|
||||||
|
swa_host_hit_length: int,
|
||||||
|
chunk_limit: int | None,
|
||||||
|
) -> tuple[bool, int | None]:
|
||||||
|
needed = self.swa_tokens(
|
||||||
|
extend_input_len,
|
||||||
|
max_new_tokens,
|
||||||
|
chunk_limit=chunk_limit,
|
||||||
|
swa_host_hit_length=swa_host_hit_length,
|
||||||
|
)
|
||||||
|
if self._fits(total_tokens, needed):
|
||||||
|
return True, chunk_limit
|
||||||
|
if chunk_limit is None or self._fits(
|
||||||
|
input_tokens + max_new_tokens + self.page_size,
|
||||||
|
self.swa_tokens(input_tokens, max_new_tokens, chunk_limit=chunk_limit),
|
||||||
|
empty_pool=True,
|
||||||
|
):
|
||||||
|
return False, None
|
||||||
|
cap = self._joint_chunk_cap(
|
||||||
|
max_chunk_tokens=min(chunk_limit, max(0, extend_input_len - 1)),
|
||||||
|
chunk_limit=chunk_limit,
|
||||||
|
swa_host_hit_length=swa_host_hit_length,
|
||||||
|
)
|
||||||
|
return (True, min(chunk_limit, cap)) if cap > 0 else (False, None)
|
||||||
|
|
||||||
|
def has_capacity(self) -> bool:
|
||||||
|
return self._fits(0, 0)
|
||||||
|
|
||||||
|
def available_chunk_tokens(self, chunk_limit: int) -> int | None:
|
||||||
|
return chunk_limit if chunk_limit > 0 else None
|
||||||
|
|
||||||
|
def fit_chunk(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
extend_input_len: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
chunk_limit: int,
|
||||||
|
) -> int | None:
|
||||||
|
candidate = min(extend_input_len, chunk_limit)
|
||||||
|
finishes = candidate >= extend_input_len
|
||||||
|
headroom = max_new_tokens if finishes else 0
|
||||||
|
if self._fits(
|
||||||
|
candidate + headroom + self.page_size,
|
||||||
|
self.swa_tokens(candidate, headroom, chunk_limit=chunk_limit),
|
||||||
|
):
|
||||||
|
return chunk_limit
|
||||||
|
cap = self._joint_chunk_cap(
|
||||||
|
max_chunk_tokens=max(0, candidate - 1) if finishes else candidate,
|
||||||
|
chunk_limit=chunk_limit,
|
||||||
|
)
|
||||||
|
return min(chunk_limit, cap) if cap > 0 else None
|
||||||
|
|
||||||
|
def can_allocate_prefill(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
paged_input: int,
|
||||||
|
extend_input_len: int,
|
||||||
|
max_new_tokens: int,
|
||||||
|
chunk_limit: int | None,
|
||||||
|
) -> bool:
|
||||||
|
return self._fits(
|
||||||
|
extend_input_len + max_new_tokens + self.page_size,
|
||||||
|
self.swa_tokens(extend_input_len, max_new_tokens, chunk_limit=chunk_limit),
|
||||||
|
)
|
||||||
@@ -226,11 +226,11 @@ class SWAComponent(TreeComponent):
|
|||||||
def _unified_allocator(self):
|
def _unified_allocator(self):
|
||||||
"""The unified SWA composite, or None when running on the static pool."""
|
"""The unified SWA composite, or None when running on the static pool."""
|
||||||
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
||||||
UnifiedSWATokenToKVPoolAllocator,
|
UnifiedSWAAllocatorBase,
|
||||||
)
|
)
|
||||||
|
|
||||||
allocator = self.cache.token_to_kv_pool_allocator
|
allocator = self.cache.token_to_kv_pool_allocator
|
||||||
if isinstance(allocator, UnifiedSWATokenToKVPoolAllocator):
|
if isinstance(allocator, UnifiedSWAAllocatorBase):
|
||||||
return allocator
|
return allocator
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|||||||
@@ -1713,13 +1713,13 @@ def init_unified_swa_pools(
|
|||||||
end_layer: int,
|
end_layer: int,
|
||||||
swa_attention_layer_ids: List[int],
|
swa_attention_layer_ids: List[int],
|
||||||
full_attention_layer_ids: List[int],
|
full_attention_layer_ids: List[int],
|
||||||
full_max_total_num_tokens: int,
|
full_max_total_num_tokens: Optional[int] = None,
|
||||||
swa_max_total_num_tokens: int,
|
swa_max_total_num_tokens: Optional[int] = None,
|
||||||
|
total_bytes: Optional[int] = None,
|
||||||
enable_memory_saver: bool,
|
enable_memory_saver: bool,
|
||||||
need_sort: bool,
|
need_sort: bool,
|
||||||
forward_stream: Optional[torch.cuda.Stream] = None,
|
forward_stream: Optional[torch.cuda.Stream] = None,
|
||||||
lazy_compaction: bool = False,
|
lazy_compaction: bool = False,
|
||||||
unified_total_bytes: Optional[int] = None,
|
|
||||||
model_context_len: Optional[int] = None,
|
model_context_len: Optional[int] = None,
|
||||||
sliding_window_size: Optional[int] = None,
|
sliding_window_size: Optional[int] = None,
|
||||||
) -> UnifiedSWAPoolBundle:
|
) -> UnifiedSWAPoolBundle:
|
||||||
@@ -1758,15 +1758,20 @@ def init_unified_swa_pools(
|
|||||||
store_dtype=store_dtype,
|
store_dtype=store_dtype,
|
||||||
grow_direction="up",
|
grow_direction="up",
|
||||||
)
|
)
|
||||||
if unified_total_bytes is not None:
|
legacy_allocator_capacities = {}
|
||||||
# PROFILED byte budget, sized from directly: the re-sum's floor losses
|
if total_bytes is None:
|
||||||
# stay out of the buffer, and the token counts remain boot labels.
|
if full_max_total_num_tokens is None or swa_max_total_num_tokens is None:
|
||||||
total_bytes = unified_total_bytes
|
raise ValueError(
|
||||||
else:
|
"total_bytes or both legacy full/SWA capacities must be provided"
|
||||||
|
)
|
||||||
total_bytes = (
|
total_bytes = (
|
||||||
full_max_total_num_tokens * full_spec.entry_bytes()
|
full_max_total_num_tokens * full_spec.entry_bytes()
|
||||||
+ swa_max_total_num_tokens * swa_spec.entry_bytes()
|
+ swa_max_total_num_tokens * swa_spec.entry_bytes()
|
||||||
)
|
)
|
||||||
|
legacy_allocator_capacities = {
|
||||||
|
"full_max_total_num_tokens": full_max_total_num_tokens,
|
||||||
|
"swa_max_total_num_tokens": swa_max_total_num_tokens,
|
||||||
|
}
|
||||||
if model_context_len is not None:
|
if model_context_len is not None:
|
||||||
# bs=1 floor: ONE sliding window of swa KV (+ a page of slack for the
|
# bs=1 floor: ONE sliding window of swa KV (+ a page of slack for the
|
||||||
# page-granular walk) + the slot-0 sink. The full side is not charged
|
# page-granular walk) + the slot-0 sink. The full side is not charged
|
||||||
@@ -1785,6 +1790,8 @@ def init_unified_swa_pools(
|
|||||||
],
|
],
|
||||||
factory="init_unified_swa_pools",
|
factory="init_unified_swa_pools",
|
||||||
)
|
)
|
||||||
|
if total_bytes <= 0:
|
||||||
|
raise ValueError(f"total_bytes must be positive, got {total_bytes}")
|
||||||
shared_pool = UnifiedKVPool(
|
shared_pool = UnifiedKVPool(
|
||||||
total_bytes=total_bytes,
|
total_bytes=total_bytes,
|
||||||
sub_pool_specs=[full_spec, swa_spec],
|
sub_pool_specs=[full_spec, swa_spec],
|
||||||
@@ -1805,12 +1812,11 @@ def init_unified_swa_pools(
|
|||||||
unified_buffer=shared_pool,
|
unified_buffer=shared_pool,
|
||||||
kvcache=token_to_kv_pool,
|
kvcache=token_to_kv_pool,
|
||||||
device=device,
|
device=device,
|
||||||
full_max_total_num_tokens=full_max_total_num_tokens,
|
|
||||||
swa_max_total_num_tokens=swa_max_total_num_tokens,
|
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
need_sort=need_sort,
|
need_sort=need_sort,
|
||||||
forward_stream=forward_stream,
|
forward_stream=forward_stream,
|
||||||
lazy_compaction=lazy_compaction,
|
lazy_compaction=lazy_compaction,
|
||||||
|
**legacy_allocator_capacities,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -1841,12 +1847,12 @@ def init_unified_swa_pools(
|
|||||||
page_size,
|
page_size,
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"[unified-memory-pool] total_bytes=%d (=%.2f GB), full_max_total_num_tokens=%d, "
|
"[unified-memory-pool] total_bytes=%d (=%.2f GB), "
|
||||||
"swa_max_total_num_tokens=%d, joint_available=%d slots",
|
"full_capacity=%d, swa_capacity=%d, joint_available=%d slots",
|
||||||
total_bytes,
|
total_bytes,
|
||||||
total_bytes / GB,
|
total_bytes / GB,
|
||||||
full_max_total_num_tokens,
|
allocator.size_full,
|
||||||
swa_max_total_num_tokens,
|
allocator.size_swa,
|
||||||
allocator.available_size(),
|
allocator.available_size(),
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
|
|||||||
@@ -1117,7 +1117,9 @@ class ModelRunner:
|
|||||||
RoutedExpertsCapturer.create(
|
RoutedExpertsCapturer.create(
|
||||||
model=self.model,
|
model=self.model,
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
num_tokens=self.max_token_pool_size + self.page_size,
|
num_tokens=self.kv_index_translator.capture_token_capacity(
|
||||||
|
self.max_token_pool_size
|
||||||
|
),
|
||||||
max_running_requests=self.max_running_requests,
|
max_running_requests=self.max_running_requests,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
@@ -1127,7 +1129,9 @@ class ModelRunner:
|
|||||||
set_global_indexer_capturer(
|
set_global_indexer_capturer(
|
||||||
create_indexer_capturer(
|
create_indexer_capturer(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
num_tokens=self.max_token_pool_size + self.page_size,
|
num_tokens=self.kv_index_translator.capture_token_capacity(
|
||||||
|
self.max_token_pool_size
|
||||||
|
),
|
||||||
max_running_requests=self.max_running_requests,
|
max_running_requests=self.max_running_requests,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
@@ -1373,7 +1377,11 @@ class ModelRunner:
|
|||||||
def effective_max_total_num_tokens(self):
|
def effective_max_total_num_tokens(self):
|
||||||
"""Return the max token pool size considering hybrid swa settings."""
|
"""Return the max token pool size considering hybrid swa settings."""
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
capacity = self.full_max_total_num_tokens or self.swa_max_total_num_tokens
|
capacity = self.kv_cache_configurator.hybrid_swa_token_capacity(
|
||||||
|
allocator=self.token_to_kv_pool_allocator,
|
||||||
|
full_capacity=self.full_max_total_num_tokens,
|
||||||
|
swa_capacity=self.swa_max_total_num_tokens,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
capacity = self.max_total_num_tokens
|
capacity = self.max_total_num_tokens
|
||||||
if (req_to_token_pool := getattr(self, "req_to_token_pool", None)) is not None:
|
if (req_to_token_pool := getattr(self, "req_to_token_pool", None)) is not None:
|
||||||
|
|||||||
@@ -71,6 +71,7 @@ class MemoryPoolConfig:
|
|||||||
max_running_requests: Optional[int] = None
|
max_running_requests: Optional[int] = None
|
||||||
full_max_total_num_tokens: Optional[int] = None
|
full_max_total_num_tokens: Optional[int] = None
|
||||||
swa_max_total_num_tokens: Optional[int] = None
|
swa_max_total_num_tokens: Optional[int] = None
|
||||||
|
unified_memory_pool_bytes: Optional[int] = None
|
||||||
|
|
||||||
# DSV4 compressed-attention pool sizes (target only; draft workers leave at 0).
|
# DSV4 compressed-attention pool sizes (target only; draft workers leave at 0).
|
||||||
c4_max_total_num_tokens: int = 0
|
c4_max_total_num_tokens: int = 0
|
||||||
@@ -698,6 +699,12 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
+ self._draft_cell_size
|
+ self._draft_cell_size
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _unified_pool_bytes(self, full_tokens: int, swa_tokens: int) -> int:
|
||||||
|
return (
|
||||||
|
full_tokens * self._full_per_token * self._full_layers_num
|
||||||
|
+ swa_tokens * self._swa_per_token * self._swa_layers_num
|
||||||
|
)
|
||||||
|
|
||||||
def _max_unified_full_tokens(
|
def _max_unified_full_tokens(
|
||||||
self,
|
self,
|
||||||
available_bytes: int,
|
available_bytes: int,
|
||||||
@@ -707,7 +714,6 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
"""Find the largest page-aligned full capacity whose allocations fit."""
|
"""Find the largest page-aligned full capacity whose allocations fit."""
|
||||||
draft_bytes_per_token = self._draft_pool_bytes_per_token()
|
draft_bytes_per_token = self._draft_pool_bytes_per_token()
|
||||||
target_full_bytes_per_token = self._full_per_token * self._full_layers_num
|
target_full_bytes_per_token = self._full_per_token * self._full_layers_num
|
||||||
target_swa_bytes_per_token = self._swa_per_token * self._swa_layers_num
|
|
||||||
assert target_full_bytes_per_token > 0
|
assert target_full_bytes_per_token > 0
|
||||||
|
|
||||||
def allocation_bytes(full_pages: int) -> int:
|
def allocation_bytes(full_pages: int) -> int:
|
||||||
@@ -719,10 +725,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
// page_size
|
// page_size
|
||||||
* page_size
|
* page_size
|
||||||
)
|
)
|
||||||
target_bytes = (
|
target_bytes = self._unified_pool_bytes(full_tokens, swa_tokens)
|
||||||
full_tokens * target_full_bytes_per_token
|
|
||||||
+ swa_tokens * target_swa_bytes_per_token
|
|
||||||
)
|
|
||||||
virtual_span = max(target_bytes // target_full_bytes_per_token - 1, 0)
|
virtual_span = max(target_bytes // target_full_bytes_per_token - 1, 0)
|
||||||
draft_tokens = ceil_align(virtual_span, page_size) + page_size
|
draft_tokens = ceil_align(virtual_span, page_size) + page_size
|
||||||
return target_bytes + draft_tokens * draft_bytes_per_token
|
return target_bytes + draft_tokens * draft_bytes_per_token
|
||||||
@@ -760,29 +763,34 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
full_tokens = align_page_size(max_total_num_tokens)
|
full_tokens = align_page_size(max_total_num_tokens)
|
||||||
swa_tokens = align_page_size(int(full_tokens * self._swa_full_tokens_ratio))
|
swa_tokens = align_page_size(int(full_tokens * self._swa_full_tokens_ratio))
|
||||||
|
|
||||||
self.validate_swa_pool_size(
|
if not self._enable_unified_memory:
|
||||||
swa_tokens, self._sliding_window_size, self._page_size
|
self.validate_swa_pool_size(
|
||||||
)
|
swa_tokens, self._sliding_window_size, self._page_size
|
||||||
|
)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Use sliding window memory pool. "
|
f"Use sliding window memory pool. "
|
||||||
f"full_layer_tokens={full_tokens}, swa_layer_tokens={swa_tokens}"
|
f"full_layer_tokens={full_tokens}, swa_layer_tokens={swa_tokens}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
return self._make_pool_config(full_tokens, swa_tokens)
|
||||||
|
|
||||||
|
def _make_pool_config(self, full_tokens: int, swa_tokens: int) -> MemoryPoolConfig:
|
||||||
return MemoryPoolConfig(
|
return MemoryPoolConfig(
|
||||||
max_total_num_tokens=full_tokens,
|
max_total_num_tokens=full_tokens,
|
||||||
full_max_total_num_tokens=full_tokens,
|
full_max_total_num_tokens=full_tokens,
|
||||||
swa_max_total_num_tokens=swa_tokens,
|
swa_max_total_num_tokens=swa_tokens,
|
||||||
|
unified_memory_pool_bytes=(
|
||||||
|
self._unified_pool_bytes(full_tokens, swa_tokens)
|
||||||
|
if self._enable_unified_memory
|
||||||
|
else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def calculate_pool_sizes(
|
def calculate_pool_sizes(
|
||||||
self, available_bytes: int, page_size: int
|
self, available_bytes: int, page_size: int
|
||||||
) -> MemoryPoolConfig:
|
) -> MemoryPoolConfig:
|
||||||
if (
|
if self._enable_unified_memory and self._full_layers_num > 0:
|
||||||
self._enable_unified_memory
|
|
||||||
and self._full_layers_num > 0
|
|
||||||
and self._draft_pool_bytes_per_token() > 0
|
|
||||||
):
|
|
||||||
max_total_num_tokens = self._max_unified_full_tokens(
|
max_total_num_tokens = self._max_unified_full_tokens(
|
||||||
available_bytes, page_size
|
available_bytes, page_size
|
||||||
)
|
)
|
||||||
@@ -873,7 +881,7 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
|
|||||||
* self._swa_per_token
|
* self._swa_per_token
|
||||||
* (self._swa_layers_num + self._draft_swa_layers_num)
|
* (self._swa_layers_num + self._draft_swa_layers_num)
|
||||||
)
|
)
|
||||||
if self._enable_unified_memory and self._draft_pool_bytes_per_token() > 0:
|
if self._enable_unified_memory:
|
||||||
full_tokens = self._max_unified_full_tokens(
|
full_tokens = self._max_unified_full_tokens(
|
||||||
available_bytes, page_size, fixed_swa_tokens=swa_tokens
|
available_bytes, page_size, fixed_swa_tokens=swa_tokens
|
||||||
)
|
)
|
||||||
@@ -894,11 +902,7 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
|
|||||||
f"Reduce --max-running-requests, lower SGLANG_SWA_EVICTION_INTERVAL, "
|
f"Reduce --max-running-requests, lower SGLANG_SWA_EVICTION_INTERVAL, "
|
||||||
f"or increase --mem-fraction-static."
|
f"or increase --mem-fraction-static."
|
||||||
)
|
)
|
||||||
return MemoryPoolConfig(
|
return self._make_pool_config(full_tokens, swa_tokens)
|
||||||
max_total_num_tokens=full_tokens,
|
|
||||||
full_max_total_num_tokens=full_tokens,
|
|
||||||
swa_max_total_num_tokens=swa_tokens,
|
|
||||||
)
|
|
||||||
|
|
||||||
def calculate_pool_sizes_from_max_tokens(
|
def calculate_pool_sizes_from_max_tokens(
|
||||||
self, max_total_num_tokens: int, page_size: int
|
self, max_total_num_tokens: int, page_size: int
|
||||||
@@ -906,10 +910,8 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
|
|||||||
# Constrained max_total goes to the full pool; SWA stays at its cap.
|
# Constrained max_total goes to the full pool; SWA stays at its cap.
|
||||||
swa_tokens = ceil_align(self._swa_cap, page_size)
|
swa_tokens = ceil_align(self._swa_cap, page_size)
|
||||||
full_tokens = (max_total_num_tokens // page_size) * page_size
|
full_tokens = (max_total_num_tokens // page_size) * page_size
|
||||||
return MemoryPoolConfig(
|
return self._make_pool_config(
|
||||||
max_total_num_tokens=full_tokens,
|
full_tokens, min(swa_tokens, max_total_num_tokens)
|
||||||
full_max_total_num_tokens=full_tokens,
|
|
||||||
swa_max_total_num_tokens=min(swa_tokens, max_total_num_tokens),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -201,13 +201,21 @@ class BaseTopkCapturer:
|
|||||||
slice_gpu = self._get_local_slice(
|
slice_gpu = self._get_local_slice(
|
||||||
forward_batch, can_run_graph, cuda_graph_batch
|
forward_batch, can_run_graph, cuda_graph_batch
|
||||||
)
|
)
|
||||||
|
# get_topk reads req_to_token IDs; attention may have rebound the write
|
||||||
|
# loc to kernel-facing IDs, which are neither stable nor the same space.
|
||||||
|
out_cache_loc = forward_batch.out_cache_loc_virtual
|
||||||
|
if out_cache_loc is None:
|
||||||
|
out_cache_loc = forward_batch.out_cache_loc
|
||||||
|
else:
|
||||||
|
# Kernel batches may be padded; only real request tokens are stored.
|
||||||
|
slice_gpu = slice_gpu[: out_cache_loc.shape[0]]
|
||||||
if no_copy_to_cpu:
|
if no_copy_to_cpu:
|
||||||
# Clone before the next overlapping forward reuses these buffers.
|
# Clone before the next overlapping forward reuses these buffers.
|
||||||
return TopkCaptureOutput(
|
return TopkCaptureOutput(
|
||||||
out_cache_loc=forward_batch.out_cache_loc.clone(),
|
out_cache_loc=out_cache_loc.clone(),
|
||||||
topk=slice_gpu.clone(),
|
topk=slice_gpu.clone(),
|
||||||
host_cache=self.host_cache,
|
host_cache=self.host_cache,
|
||||||
)
|
)
|
||||||
out_cache_loc_cpu = forward_batch.out_cache_loc.cpu()
|
out_cache_loc_cpu = out_cache_loc.cpu()
|
||||||
self.host_cache.buffer[out_cache_loc_cpu] = slice_gpu.cpu()
|
self.host_cache.buffer[out_cache_loc_cpu] = slice_gpu.cpu()
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -11,13 +11,17 @@ corruption with no crash.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
from typing import List, Optional, Set
|
from typing import List, Optional, Set
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
from sglang.srt.disaggregation.utils import (
|
from sglang.srt.disaggregation.utils import (
|
||||||
DisaggregationMode,
|
DisaggregationMode,
|
||||||
unified_memory_disagg_move_gate,
|
unified_memory_disagg_move_gate,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator.unified_sub_pool import MultiEndedAllocator
|
from sglang.srt.mem_cache.allocator.unified_sub_pool import MultiEndedAllocator
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -187,58 +191,41 @@ class TestMoveGateRejectsNonPdNode(CustomTestCase):
|
|||||||
|
|
||||||
|
|
||||||
class TestUnifiedAllocatorsPublishTheTransferContract(CustomTestCase):
|
class TestUnifiedAllocatorsPublishTheTransferContract(CustomTestCase):
|
||||||
"""Every unified composite allocator must OVERRIDE the two PD hooks.
|
"""Unified composites must translate virtual IDs before PD transfer.
|
||||||
|
|
||||||
`BaseTokenToKVPoolAllocator.translate_kv_indices_for_transfer` is the
|
The implementation may be inherited from a shared unified allocator base,
|
||||||
IDENTITY, and `set_disagg_move_gate` exists only where a composite defines
|
but inheriting the static allocator's identity would put virtual IDs on the
|
||||||
it. Inheriting either is silent, not loud: identity puts VIRTUAL ids on the
|
wire and silently corrupt KV. Gate installation must reach every member.
|
||||||
wire (they address real bytes, so the peer gets plausible garbage), and a
|
|
||||||
missing gate lets lazy compaction relocate pages under in-flight RDMA.
|
|
||||||
An AST-level check because instantiating these composites needs a GPU.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Composites that own the full-side virtual ids and so must define the
|
|
||||||
# transfer translate themselves.
|
|
||||||
_COMPOSITES = (
|
|
||||||
"UnifiedMambaTokenToKVPoolAllocator",
|
|
||||||
"UnifiedSWATokenToKVPoolAllocator",
|
|
||||||
)
|
|
||||||
# Every composite must define the gate setter, including the tri-pool,
|
|
||||||
# which inherits the SWA translates (same full side) but has a THIRD
|
|
||||||
# member the 2-pool setter does not reach.
|
|
||||||
_GATE_COMPOSITES = _COMPOSITES + ("UnifiedMambaSWATokenToKVPoolAllocator",)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _own_methods(cls_name: str) -> Set[str]:
|
def _allocator_class(name):
|
||||||
"""Names this class defines ITSELF, inheritance excluded.
|
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
||||||
|
UnifiedMambaSWATokenToKVPoolAllocator,
|
||||||
Resolved off the class object rather than by parsing a named module:
|
UnifiedSWATokenToKVPoolAllocator,
|
||||||
these composites have already been moved once (out of
|
)
|
||||||
`multi_ended_allocator` into `allocator/unified_*`), and a hardcoded
|
from sglang.srt.mem_cache.allocator.unified_mamba import (
|
||||||
module path turns that kind of move into a test failure that says
|
UnifiedMambaTokenToKVPoolAllocator,
|
||||||
nothing about the contract. `__dict__` needs no GPU -- it is the class
|
|
||||||
body, not an instance.
|
|
||||||
"""
|
|
||||||
from sglang.srt.mem_cache.allocator import (
|
|
||||||
unified_hybrid_swa,
|
|
||||||
unified_mamba,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
for mod in (unified_mamba, unified_hybrid_swa):
|
classes = (
|
||||||
cls = getattr(mod, cls_name, None)
|
UnifiedMambaTokenToKVPoolAllocator,
|
||||||
if cls is not None:
|
UnifiedSWATokenToKVPoolAllocator,
|
||||||
return set(vars(cls))
|
UnifiedMambaSWATokenToKVPoolAllocator,
|
||||||
raise AssertionError(f"class {cls_name} not found in the unified allocators")
|
)
|
||||||
|
return {cls.__name__: cls for cls in classes}[name]
|
||||||
|
|
||||||
def test_transfer_translate_is_not_inherited_identity(self):
|
def test_transfer_translate_is_not_inherited_identity(self):
|
||||||
for name in self._COMPOSITES:
|
virtual = torch.tensor([1, 3], dtype=torch.int32)
|
||||||
with self.subTest(composite=name):
|
for name in self._EXPECTED_COVERAGE:
|
||||||
self.assertIn(
|
with self.subTest(composite=name), get_parallel().override(attn_dcp_size=1):
|
||||||
"translate_kv_indices_for_transfer",
|
alloc = object.__new__(self._allocator_class(name))
|
||||||
self._own_methods(name),
|
alloc.full_attn_allocator = SimpleNamespace(
|
||||||
f"{name} inherits the identity transfer translate; PD would "
|
translate_kv_loc=lambda ids: ids + 16
|
||||||
"ship VIRTUAL ids and corrupt KV without any error",
|
|
||||||
)
|
)
|
||||||
|
physical = alloc.translate_kv_indices_for_transfer(virtual)
|
||||||
|
self.assertEqual(physical.dtype, torch.int64)
|
||||||
|
self.assertEqual(physical.tolist(), [17, 19])
|
||||||
|
|
||||||
# Every sub-allocator attribute a composite can hold. The stub carries all
|
# Every sub-allocator attribute a composite can hold. The stub carries all
|
||||||
# of them regardless of composite, so the assertion is on what installation
|
# of them regardless of composite, so the assertion is on what installation
|
||||||
@@ -270,11 +257,7 @@ class TestUnifiedAllocatorsPublishTheTransferContract(CustomTestCase):
|
|||||||
`object.__new__` skips `__init__` (which needs a GPU); the setter reads
|
`object.__new__` skips `__init__` (which needs a GPU); the setter reads
|
||||||
only `lazy_compaction` and the member attributes.
|
only `lazy_compaction` and the member attributes.
|
||||||
"""
|
"""
|
||||||
from sglang.srt.mem_cache.allocator import unified_hybrid_swa, unified_mamba
|
cls = self._allocator_class(cls_name)
|
||||||
|
|
||||||
cls = getattr(unified_mamba, cls_name, None) or getattr(
|
|
||||||
unified_hybrid_swa, cls_name
|
|
||||||
)
|
|
||||||
alloc = object.__new__(cls)
|
alloc = object.__new__(cls)
|
||||||
alloc.lazy_compaction = True
|
alloc.lazy_compaction = True
|
||||||
for attr in self._MEMBER_ATTRS:
|
for attr in self._MEMBER_ATTRS:
|
||||||
@@ -313,14 +296,8 @@ class TestUnifiedAllocatorsPublishTheTransferContract(CustomTestCase):
|
|||||||
"""
|
"""
|
||||||
import inspect
|
import inspect
|
||||||
|
|
||||||
from sglang.srt.mem_cache.allocator import unified_hybrid_swa, unified_mamba
|
|
||||||
|
|
||||||
for name in self._EXPECTED_COVERAGE:
|
for name in self._EXPECTED_COVERAGE:
|
||||||
cls = getattr(unified_mamba, name, None) or getattr(
|
cls = self._allocator_class(name)
|
||||||
unified_hybrid_swa, name
|
|
||||||
)
|
|
||||||
if "set_disagg_move_gate" not in vars(cls):
|
|
||||||
continue # inherited, and the inherited one is checked above
|
|
||||||
with self.subTest(composite=name):
|
with self.subTest(composite=name):
|
||||||
body = inspect.getsource(cls.set_disagg_move_gate)
|
body = inspect.getsource(cls.set_disagg_move_gate)
|
||||||
self.assertIn("install_move_gate", body)
|
self.assertIn("install_move_gate", body)
|
||||||
@@ -331,10 +308,23 @@ class TestUnifiedAllocatorsPublishTheTransferContract(CustomTestCase):
|
|||||||
does not name the SWA page holding the same virtual token. The read-path
|
does not name the SWA page holding the same virtual token. The read-path
|
||||||
`translate_loc_from_full_to_swa` cannot stand in either: it returns
|
`translate_loc_from_full_to_swa` cannot stand in either: it returns
|
||||||
kernel-facing ids, and the transfer addresses raw page envelopes."""
|
kernel-facing ids, and the transfer addresses raw page envelopes."""
|
||||||
self.assertIn(
|
virtual = torch.tensor([1, 3], dtype=torch.int32)
|
||||||
"translate_swa_indices_for_transfer",
|
for name in (
|
||||||
self._own_methods("UnifiedSWATokenToKVPoolAllocator"),
|
"UnifiedSWATokenToKVPoolAllocator",
|
||||||
)
|
"UnifiedMambaSWATokenToKVPoolAllocator",
|
||||||
|
):
|
||||||
|
with self.subTest(composite=name), get_parallel().override(attn_dcp_size=1):
|
||||||
|
alloc = object.__new__(self._allocator_class(name))
|
||||||
|
alloc.full_attn_allocator = SimpleNamespace(
|
||||||
|
translate_kv_loc=lambda ids: ids + 16
|
||||||
|
)
|
||||||
|
alloc.swa_attn_allocator = SimpleNamespace(
|
||||||
|
translate_kv_loc=lambda ids: ids + 32,
|
||||||
|
translate_kv_loc_for_kernel=lambda ids: ids + 64,
|
||||||
|
)
|
||||||
|
physical = alloc.translate_swa_indices_for_transfer(virtual)
|
||||||
|
self.assertEqual(physical.dtype, torch.int64)
|
||||||
|
self.assertEqual(physical.tolist(), [33, 35])
|
||||||
|
|
||||||
|
|
||||||
class TestEverySwaAllocatorAnswersTheTransferTranslate(CustomTestCase):
|
class TestEverySwaAllocatorAnswersTheTransferTranslate(CustomTestCase):
|
||||||
|
|||||||
@@ -16,6 +16,12 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
|||||||
DecLockRefResult,
|
DecLockRefResult,
|
||||||
IncLockRefResult,
|
IncLockRefResult,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.prefill_budget import (
|
||||||
|
PrefillBudget,
|
||||||
|
SWAPrefillBudget,
|
||||||
|
estimate_swa_kv_tokens,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.unified_memory_pool import init_unified_swa_pools
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||||
from sglang.srt.utils.common import Range
|
from sglang.srt.utils.common import Range
|
||||||
@@ -75,6 +81,10 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
allocator.swa_available_size.return_value = swa_available_size
|
allocator.swa_available_size.return_value = swa_available_size
|
||||||
allocator.available_size.return_value = available_size
|
allocator.available_size.return_value = available_size
|
||||||
allocator.size_swa = size_swa
|
allocator.size_swa = size_swa
|
||||||
|
allocator.swa_req_ring = False
|
||||||
|
allocator.create_prefill_budget.side_effect = lambda tree_cache, **kwargs: (
|
||||||
|
PrefillBudget(allocator, tree_cache, **kwargs)
|
||||||
|
)
|
||||||
return allocator
|
return allocator
|
||||||
|
|
||||||
def create_running_batch(self, reqs=None) -> MagicMock:
|
def create_running_batch(self, reqs=None) -> MagicMock:
|
||||||
@@ -129,8 +139,115 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
priority_scheduling_preemption_threshold=0,
|
priority_scheduling_preemption_threshold=0,
|
||||||
)
|
)
|
||||||
defaults.update(kwargs)
|
defaults.update(kwargs)
|
||||||
|
defaults["token_to_kv_pool_allocator"].page_size = defaults["page_size"]
|
||||||
return PrefillAdder(**defaults)
|
return PrefillAdder(**defaults)
|
||||||
|
|
||||||
|
def create_shared_adder(self, *, num_mixed_decode_tokens=0):
|
||||||
|
self.mock_tree_cache.supports_mamba.return_value = False
|
||||||
|
self.mock_tree_cache.sliding_window_size = 8
|
||||||
|
self.mock_tree_cache.is_tree_cache.return_value = False
|
||||||
|
allocator = init_unified_swa_pools(
|
||||||
|
device="cpu",
|
||||||
|
kv_cache_dtype=torch.float16,
|
||||||
|
head_num=1,
|
||||||
|
head_dim=4,
|
||||||
|
v_head_dim=4,
|
||||||
|
swa_head_num=1,
|
||||||
|
swa_head_dim=4,
|
||||||
|
swa_v_head_dim=4,
|
||||||
|
page_size=4,
|
||||||
|
start_layer=0,
|
||||||
|
end_layer=2,
|
||||||
|
swa_attention_layer_ids=[1],
|
||||||
|
full_attention_layer_ids=[0],
|
||||||
|
total_bytes=1024,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
need_sort=False,
|
||||||
|
lazy_compaction=True,
|
||||||
|
).token_to_kv_pool_allocator
|
||||||
|
return self.create_adder(
|
||||||
|
self.create_running_batch(),
|
||||||
|
page_size=4,
|
||||||
|
rem_chunk_tokens=16,
|
||||||
|
num_mixed_decode_tokens=num_mixed_decode_tokens,
|
||||||
|
token_to_kv_pool_allocator=allocator,
|
||||||
|
)
|
||||||
|
|
||||||
|
def create_shared_req(self, rid, max_new_tokens=4):
|
||||||
|
req = self.create_mock_req(rid, priority=0, max_new_tokens=max_new_tokens)
|
||||||
|
req.sampling_params.ignore_eos = False
|
||||||
|
req.swa_host_hit_length = 0
|
||||||
|
req.last_node = MagicMock()
|
||||||
|
req.full_untruncated_fill_ids = list(range(12))
|
||||||
|
req.set_extend_range = MagicMock(
|
||||||
|
side_effect=lambda start, end: setattr(
|
||||||
|
req, "extend_range", Range(start, end)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return req
|
||||||
|
|
||||||
|
def test_shared_admission_reserves_all_pending_requests(self):
|
||||||
|
adder = self.create_shared_adder()
|
||||||
|
first, second = (
|
||||||
|
self.create_shared_req("first"),
|
||||||
|
self.create_shared_req("second"),
|
||||||
|
)
|
||||||
|
adder.add_one_req(first, has_chunked_req=False, truncation_align_size=None)
|
||||||
|
self.assertEqual(adder.can_run_list, [first])
|
||||||
|
self.assertEqual(
|
||||||
|
adder.add_one_req(
|
||||||
|
second, has_chunked_req=False, truncation_align_size=None
|
||||||
|
),
|
||||||
|
AddReqResult.NO_TOKEN,
|
||||||
|
)
|
||||||
|
self.assertEqual(adder.can_run_list, [first])
|
||||||
|
|
||||||
|
def test_shared_admission_rechecks_after_prefix_lock(self):
|
||||||
|
adder = self.create_shared_adder()
|
||||||
|
self.assertIsNotNone(adder.token_to_kv_pool_allocator.alloc(24))
|
||||||
|
self.mock_tree_cache.full_evictable_size.return_value = 24
|
||||||
|
self.mock_tree_cache.swa_evictable_size.return_value = 24
|
||||||
|
|
||||||
|
def lock_prefix(_):
|
||||||
|
self.mock_tree_cache.full_evictable_size.return_value = 0
|
||||||
|
self.mock_tree_cache.swa_evictable_size.return_value = 0
|
||||||
|
return IncLockRefResult()
|
||||||
|
|
||||||
|
self.mock_tree_cache.inc_lock_ref.side_effect = lock_prefix
|
||||||
|
req = self.create_shared_req("locked-prefix")
|
||||||
|
self.assertEqual(
|
||||||
|
adder.add_one_req(req, has_chunked_req=False, truncation_align_size=None),
|
||||||
|
AddReqResult.NO_TOKEN,
|
||||||
|
)
|
||||||
|
self.mock_tree_cache.inc_lock_ref.assert_called_once()
|
||||||
|
self.assertEqual(adder.can_run_list, [])
|
||||||
|
|
||||||
|
def test_shared_continuation_defers_when_decode_consumes_chunk_budget(self):
|
||||||
|
"""Mixed decode must not commit an empty or negatively sliced prompt."""
|
||||||
|
for decode_tokens in (16, 28):
|
||||||
|
with self.subTest(decode_tokens=decode_tokens):
|
||||||
|
adder = self.create_shared_adder(num_mixed_decode_tokens=decode_tokens)
|
||||||
|
req = self.create_shared_req("continuation")
|
||||||
|
before = (
|
||||||
|
adder.memory_budget.total_offset,
|
||||||
|
adder.memory_budget.swa_offset,
|
||||||
|
)
|
||||||
|
self.assertIs(adder.add_chunked_req(req), req)
|
||||||
|
self.assertEqual(adder.can_run_list, [])
|
||||||
|
req.set_extend_range.assert_not_called()
|
||||||
|
self.assertEqual(
|
||||||
|
(adder.memory_budget.total_offset, adder.memory_budget.swa_offset),
|
||||||
|
before,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_shared_continuation_uses_memory_chunk_limit(self):
|
||||||
|
adder = self.create_shared_adder()
|
||||||
|
req = self.create_shared_req("continuation", max_new_tokens=80)
|
||||||
|
self.assertIs(adder.add_chunked_req(req), req)
|
||||||
|
self.assertEqual(req.extend_range.length, 8)
|
||||||
|
self.assertEqual(adder.memory_budget.total_offset, 12)
|
||||||
|
self.assertEqual(adder.memory_budget.swa_offset, 12)
|
||||||
|
|
||||||
def test_storage_prefetch_fulfillment_resolves_at_admission(self):
|
def test_storage_prefetch_fulfillment_resolves_at_admission(self):
|
||||||
adder = self.create_adder(self.create_running_batch())
|
adder = self.create_adder(self.create_running_batch())
|
||||||
req = self.create_mock_req("storage-hit", priority=0, max_new_tokens=1)
|
req = self.create_mock_req("storage-hit", priority=0, max_new_tokens=1)
|
||||||
@@ -208,7 +325,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
running_batch = self.create_running_batch(running_reqs)
|
running_batch = self.create_running_batch(running_reqs)
|
||||||
adder = self.create_adder(running_batch)
|
adder = self.create_adder(running_batch)
|
||||||
|
|
||||||
self.assertEqual(adder.rem_total_token_offset, 225)
|
self.assertEqual(adder.memory_budget.total_offset, 225)
|
||||||
|
|
||||||
self.mock_token_allocator.full_available_size.return_value = (
|
self.mock_token_allocator.full_available_size.return_value = (
|
||||||
225 # full occupation of GRam
|
225 # full occupation of GRam
|
||||||
@@ -221,7 +338,9 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
|
|
||||||
self.assertTrue(success)
|
self.assertTrue(success)
|
||||||
self.assertIn(running_reqs[0], adder.preempt_list)
|
self.assertIn(running_reqs[0], adder.preempt_list)
|
||||||
self.assertEqual(adder.rem_total_token_offset, 175) # 50 + 75 + 100 - 50 = 175
|
self.assertEqual(
|
||||||
|
adder.memory_budget.total_offset, 175
|
||||||
|
) # 50 + 75 + 100 - 50 = 175
|
||||||
running_batch.release_req.assert_called_once()
|
running_batch.release_req.assert_called_once()
|
||||||
|
|
||||||
def test_preempt_success_low_priority_values_first(self):
|
def test_preempt_success_low_priority_values_first(self):
|
||||||
@@ -238,7 +357,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
running_batch = self.create_running_batch(running_reqs)
|
running_batch = self.create_running_batch(running_reqs)
|
||||||
adder = self.create_adder(running_batch)
|
adder = self.create_adder(running_batch)
|
||||||
|
|
||||||
self.assertEqual(adder.rem_total_token_offset, 225)
|
self.assertEqual(adder.memory_budget.total_offset, 225)
|
||||||
|
|
||||||
self.mock_token_allocator.full_available_size.return_value = (
|
self.mock_token_allocator.full_available_size.return_value = (
|
||||||
225 # full occupation of GRam
|
225 # full occupation of GRam
|
||||||
@@ -251,7 +370,9 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
|
|
||||||
self.assertTrue(success)
|
self.assertTrue(success)
|
||||||
self.assertIn(running_reqs[2], adder.preempt_list)
|
self.assertIn(running_reqs[2], adder.preempt_list)
|
||||||
self.assertEqual(adder.rem_total_token_offset, 125) # 50 + 75 + 100 - 100 = 125
|
self.assertEqual(
|
||||||
|
adder.memory_budget.total_offset, 125
|
||||||
|
) # 50 + 75 + 100 - 100 = 125
|
||||||
running_batch.release_req.assert_called_once()
|
running_batch.release_req.assert_called_once()
|
||||||
|
|
||||||
def test_preempt_fail_low_priority_values_first(self):
|
def test_preempt_fail_low_priority_values_first(self):
|
||||||
@@ -268,7 +389,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
running_batch = self.create_running_batch(running_reqs)
|
running_batch = self.create_running_batch(running_reqs)
|
||||||
adder = self.create_adder(running_batch)
|
adder = self.create_adder(running_batch)
|
||||||
|
|
||||||
self.assertEqual(adder.rem_total_token_offset, 225)
|
self.assertEqual(adder.memory_budget.total_offset, 225)
|
||||||
|
|
||||||
self.mock_token_allocator.full_available_size.return_value = (
|
self.mock_token_allocator.full_available_size.return_value = (
|
||||||
225 # full occupation of GRam
|
225 # full occupation of GRam
|
||||||
@@ -306,7 +427,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
running_batch = self.create_running_batch(running_reqs)
|
running_batch = self.create_running_batch(running_reqs)
|
||||||
adder = self.create_adder(running_batch)
|
adder = self.create_adder(running_batch)
|
||||||
|
|
||||||
self.assertEqual(adder.rem_total_token_offset, 225)
|
self.assertEqual(adder.memory_budget.total_offset, 225)
|
||||||
|
|
||||||
self.mock_token_allocator.full_available_size.return_value = (
|
self.mock_token_allocator.full_available_size.return_value = (
|
||||||
225 # full occupation of GRam
|
225 # full occupation of GRam
|
||||||
@@ -344,7 +465,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
running_batch = self.create_running_batch(running_reqs)
|
running_batch = self.create_running_batch(running_reqs)
|
||||||
adder = self.create_adder(running_batch)
|
adder = self.create_adder(running_batch)
|
||||||
|
|
||||||
self.assertEqual(adder.rem_total_token_offset, 225)
|
self.assertEqual(adder.memory_budget.total_offset, 225)
|
||||||
|
|
||||||
self.mock_token_allocator.full_available_size.return_value = 225
|
self.mock_token_allocator.full_available_size.return_value = 225
|
||||||
self.mock_token_allocator.available_size.return_value = 225
|
self.mock_token_allocator.available_size.return_value = 225
|
||||||
@@ -356,7 +477,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
first_success = adder.preempt_to_schedule(first_req)
|
first_success = adder.preempt_to_schedule(first_req)
|
||||||
self.assertTrue(first_success)
|
self.assertTrue(first_success)
|
||||||
self.assertIn(running_reqs[0], adder.preempt_list)
|
self.assertIn(running_reqs[0], adder.preempt_list)
|
||||||
self.assertEqual(adder.rem_total_token_offset, 175)
|
self.assertEqual(adder.memory_budget.total_offset, 175)
|
||||||
running_batch.release_req.assert_called_once()
|
running_batch.release_req.assert_called_once()
|
||||||
|
|
||||||
# Second call needs more tokens than currently free, so it would need to
|
# Second call needs more tokens than currently free, so it would need to
|
||||||
@@ -367,7 +488,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
second_success = adder.preempt_to_schedule(second_req)
|
second_success = adder.preempt_to_schedule(second_req)
|
||||||
|
|
||||||
self.assertFalse(second_success)
|
self.assertFalse(second_success)
|
||||||
self.assertEqual(adder.rem_total_token_offset, 175)
|
self.assertEqual(adder.memory_budget.total_offset, 175)
|
||||||
self.assertEqual(adder.preempt_list.count(running_reqs[0]), 1)
|
self.assertEqual(adder.preempt_list.count(running_reqs[0]), 1)
|
||||||
running_batch.release_req.assert_called_once()
|
running_batch.release_req.assert_called_once()
|
||||||
|
|
||||||
@@ -387,7 +508,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
running_batch = self.create_running_batch(running_reqs)
|
running_batch = self.create_running_batch(running_reqs)
|
||||||
adder = self.create_adder(running_batch)
|
adder = self.create_adder(running_batch)
|
||||||
|
|
||||||
self.assertEqual(adder.rem_total_token_offset, 475)
|
self.assertEqual(adder.memory_budget.total_offset, 475)
|
||||||
|
|
||||||
self.mock_token_allocator.full_available_size.return_value = (
|
self.mock_token_allocator.full_available_size.return_value = (
|
||||||
475 # full occupation of GRam
|
475 # full occupation of GRam
|
||||||
@@ -400,7 +521,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
self.assertTrue(success)
|
self.assertTrue(success)
|
||||||
self.assertIn(running_reqs[2], adder.preempt_list)
|
self.assertIn(running_reqs[2], adder.preempt_list)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
adder.rem_total_token_offset, 375
|
adder.memory_budget.total_offset, 375
|
||||||
) # 50 + 75 + 100 + 125 + 125 - 100 = 375
|
) # 50 + 75 + 100 + 125 + 125 - 100 = 375
|
||||||
running_batch.release_req.assert_called_once()
|
running_batch.release_req.assert_called_once()
|
||||||
|
|
||||||
@@ -420,7 +541,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
running_batch = self.create_running_batch(running_reqs)
|
running_batch = self.create_running_batch(running_reqs)
|
||||||
adder = self.create_adder(running_batch)
|
adder = self.create_adder(running_batch)
|
||||||
|
|
||||||
self.assertEqual(adder.rem_total_token_offset, 475)
|
self.assertEqual(adder.memory_budget.total_offset, 475)
|
||||||
|
|
||||||
self.mock_token_allocator.full_available_size.return_value = (
|
self.mock_token_allocator.full_available_size.return_value = (
|
||||||
475 # full occupation of GRam
|
475 # full occupation of GRam
|
||||||
@@ -434,7 +555,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
self.assertIn(running_reqs[2], adder.preempt_list)
|
self.assertIn(running_reqs[2], adder.preempt_list)
|
||||||
self.assertIn(running_reqs[3], adder.preempt_list)
|
self.assertIn(running_reqs[3], adder.preempt_list)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
adder.rem_total_token_offset, 250
|
adder.memory_budget.total_offset, 250
|
||||||
) # 50 + 75 + 100 + 125 + 125 - 100 - 125 = 250
|
) # 50 + 75 + 100 + 125 + 125 - 100 - 125 = 250
|
||||||
self.assertEqual(running_batch.release_req.call_count, 2)
|
self.assertEqual(running_batch.release_req.call_count, 2)
|
||||||
|
|
||||||
@@ -456,8 +577,8 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
|
|
||||||
self.assertEqual(adder.rem_input_tokens, 192) # 200 - 8
|
self.assertEqual(adder.rem_input_tokens, 192) # 200 - 8
|
||||||
self.assertEqual(adder.rem_chunk_tokens, 56) # 64 - 8
|
self.assertEqual(adder.rem_chunk_tokens, 56) # 64 - 8
|
||||||
self.assertEqual(adder.rem_total_token_offset, 408) # 8 + 8 * 50
|
self.assertEqual(adder.memory_budget.total_offset, 408) # 8 + 8 * 50
|
||||||
self.assertEqual(adder.cur_rem_token_offset, 8)
|
self.assertEqual(adder.memory_budget.current_offset, 8)
|
||||||
self.assertEqual(adder.budget_state(), AddReqResult.CONTINUE)
|
self.assertEqual(adder.budget_state(), AddReqResult.CONTINUE)
|
||||||
|
|
||||||
# Add a prefill that exactly consumes the chunk budget
|
# Add a prefill that exactly consumes the chunk budget
|
||||||
@@ -497,7 +618,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
|
|
||||||
self.assertEqual(adder2.rem_input_tokens, 195) # 200 - 5
|
self.assertEqual(adder2.rem_input_tokens, 195) # 200 - 5
|
||||||
self.assertEqual(adder2.rem_chunk_tokens, 59) # 64 - 5
|
self.assertEqual(adder2.rem_chunk_tokens, 59) # 64 - 5
|
||||||
self.assertEqual(adder2.rem_total_token_offset, 255) # 5 + 5 * 50
|
self.assertEqual(adder2.memory_budget.total_offset, 255) # 5 + 5 * 50
|
||||||
self.assertEqual(adder2.budget_state(), AddReqResult.CONTINUE)
|
self.assertEqual(adder2.budget_state(), AddReqResult.CONTINUE)
|
||||||
|
|
||||||
# Same prefill no longer exhausts the chunk budget
|
# Same prefill no longer exhausts the chunk budget
|
||||||
@@ -562,6 +683,10 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
rem_chunk_tokens=rem_chunk,
|
rem_chunk_tokens=rem_chunk,
|
||||||
)
|
)
|
||||||
adder.is_hybrid_swa = is_hybrid_swa
|
adder.is_hybrid_swa = is_hybrid_swa
|
||||||
|
if is_hybrid_swa:
|
||||||
|
adder.memory_budget = SWAPrefillBudget(
|
||||||
|
self.mock_token_allocator, self.mock_tree_cache
|
||||||
|
)
|
||||||
|
|
||||||
req = self.create_mock_req("chunked", priority=0, max_new_tokens=128)
|
req = self.create_mock_req("chunked", priority=0, max_new_tokens=128)
|
||||||
req.prefix_indices = []
|
req.prefix_indices = []
|
||||||
@@ -640,7 +765,16 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
page_size=page,
|
page_size=page,
|
||||||
rem_chunk_tokens=rem_chunk,
|
rem_chunk_tokens=rem_chunk,
|
||||||
)
|
)
|
||||||
self.assertEqual(adder._swa_budget_for_req(extend, max_new), expected)
|
self.assertEqual(
|
||||||
|
estimate_swa_kv_tokens(
|
||||||
|
extend,
|
||||||
|
max_new,
|
||||||
|
sliding_window_size=window,
|
||||||
|
page_size=page,
|
||||||
|
allocation_limit=rem_chunk,
|
||||||
|
),
|
||||||
|
expected,
|
||||||
|
)
|
||||||
|
|
||||||
def test_swa_admission_admits_short_cached_resume_at_two_window_pool(self):
|
def test_swa_admission_admits_short_cached_resume_at_two_window_pool(self):
|
||||||
# Livelock regression (real incident). At an SWA pool ~= 2 sliding
|
# Livelock regression (real incident). At an SWA pool ~= 2 sliding
|
||||||
@@ -660,6 +794,9 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
self.mock_tree_cache.is_tree_cache.return_value = False
|
self.mock_tree_cache.is_tree_cache.return_value = False
|
||||||
adder = self.create_adder(self.create_running_batch(), page_size=PAGE)
|
adder = self.create_adder(self.create_running_batch(), page_size=PAGE)
|
||||||
adder.is_hybrid_swa = True
|
adder.is_hybrid_swa = True
|
||||||
|
adder.memory_budget = SWAPrefillBudget(
|
||||||
|
self.mock_token_allocator, self.mock_tree_cache
|
||||||
|
)
|
||||||
|
|
||||||
req = self.create_mock_req(
|
req = self.create_mock_req(
|
||||||
"resume", priority=0, max_new_tokens=40, output_len=10
|
"resume", priority=0, max_new_tokens=40, output_len=10
|
||||||
@@ -677,7 +814,9 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
req.sampling_params = SimpleNamespace(max_new_tokens=40, ignore_eos=False)
|
req.sampling_params = SimpleNamespace(max_new_tokens=40, ignore_eos=False)
|
||||||
|
|
||||||
# Pre-fix: a constant sliding-window reservation rejects the resume.
|
# Pre-fix: a constant sliding-window reservation rejects the resume.
|
||||||
with patch.object(adder, "_swa_reserved_tokens", return_value=WINDOW + PAGE):
|
with patch.object(
|
||||||
|
adder.memory_budget, "swa_tokens", return_value=WINDOW + PAGE
|
||||||
|
):
|
||||||
self.assertIs(
|
self.assertIs(
|
||||||
adder.add_one_req(
|
adder.add_one_req(
|
||||||
req, has_chunked_req=False, truncation_align_size=None
|
req, has_chunked_req=False, truncation_align_size=None
|
||||||
@@ -708,6 +847,9 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
self.mock_token_allocator.swa_available_size.return_value = 400
|
self.mock_token_allocator.swa_available_size.return_value = 400
|
||||||
adder = self.create_adder(self.create_running_batch(), page_size=PAGE)
|
adder = self.create_adder(self.create_running_batch(), page_size=PAGE)
|
||||||
adder.is_hybrid_swa = True
|
adder.is_hybrid_swa = True
|
||||||
|
adder.memory_budget = SWAPrefillBudget(
|
||||||
|
self.mock_token_allocator, self.mock_tree_cache
|
||||||
|
)
|
||||||
req = self.create_mock_req("dropped-fetch", priority=0, max_new_tokens=8)
|
req = self.create_mock_req("dropped-fetch", priority=0, max_new_tokens=8)
|
||||||
req.prefix_indices = torch.empty(0, dtype=torch.int64)
|
req.prefix_indices = torch.empty(0, dtype=torch.int64)
|
||||||
req.full_untruncated_fill_ids = list(range(SPAN))
|
req.full_untruncated_fill_ids = list(range(SPAN))
|
||||||
@@ -844,7 +986,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
extend if req.retracted_stain else 0,
|
extend if req.retracted_stain else 0,
|
||||||
)
|
)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
adder.rem_total_token_offset,
|
adder.memory_budget.total_offset,
|
||||||
adder.ceil_paged_tokens(extend) + decode + 2,
|
adder.ceil_paged_tokens(extend) + decode + 2,
|
||||||
)
|
)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
@@ -1092,6 +1234,11 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
) -> PrefillAdder:
|
) -> PrefillAdder:
|
||||||
self.mock_tree_cache.sliding_window_size = sliding_window
|
self.mock_tree_cache.sliding_window_size = sliding_window
|
||||||
self.mock_token_allocator = self.create_token_allocator(size_swa=size_swa)
|
self.mock_token_allocator = self.create_token_allocator(size_swa=size_swa)
|
||||||
|
self.mock_token_allocator.create_prefill_budget.side_effect = (
|
||||||
|
lambda tree_cache, **kwargs: SWAPrefillBudget(
|
||||||
|
self.mock_token_allocator, tree_cache, **kwargs
|
||||||
|
)
|
||||||
|
)
|
||||||
return self.create_adder(
|
return self.create_adder(
|
||||||
self.create_running_batch(),
|
self.create_running_batch(),
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
@@ -1103,7 +1250,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
# once running decodes drain -> must wait, not take the hatch.
|
# once running decodes drain -> must wait, not take the hatch.
|
||||||
adder = self.create_swa_adder(size_swa=1024, sliding_window=128)
|
adder = self.create_swa_adder(size_swa=1024, sliding_window=128)
|
||||||
self.assertFalse(
|
self.assertFalse(
|
||||||
adder._swa_req_never_fits(extend_input_len=256, max_new_tokens=64)
|
adder.memory_budget.swa_never_fits(extend_input_len=256, max_new_tokens=64)
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_swa_never_fits_true_when_budget_exceeds_whole_pool(self):
|
def test_swa_never_fits_true_when_budget_exceeds_whole_pool(self):
|
||||||
@@ -1111,7 +1258,7 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
# pool: it can never fit however far the pool drains -> hatch.
|
# pool: it can never fit however far the pool drains -> hatch.
|
||||||
adder = self.create_swa_adder(size_swa=1024, sliding_window=128)
|
adder = self.create_swa_adder(size_swa=1024, sliding_window=128)
|
||||||
self.assertTrue(
|
self.assertTrue(
|
||||||
adder._swa_req_never_fits(
|
adder.memory_budget.swa_never_fits(
|
||||||
extend_input_len=256, max_new_tokens=64, swa_host_hit_length=4096
|
extend_input_len=256, max_new_tokens=64, swa_host_hit_length=4096
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -1121,14 +1268,14 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
# the budget against size_swa (guards against a wrong-accessor bug).
|
# the budget against size_swa (guards against a wrong-accessor bug).
|
||||||
req = dict(extend_input_len=256, max_new_tokens=64, swa_host_hit_length=600)
|
req = dict(extend_input_len=256, max_new_tokens=64, swa_host_hit_length=600)
|
||||||
self.assertTrue(
|
self.assertTrue(
|
||||||
self.create_swa_adder(size_swa=512, sliding_window=128)._swa_req_never_fits(
|
self.create_swa_adder(
|
||||||
**req
|
size_swa=512, sliding_window=128
|
||||||
)
|
).memory_budget.swa_never_fits(**req)
|
||||||
)
|
)
|
||||||
self.assertFalse(
|
self.assertFalse(
|
||||||
self.create_swa_adder(
|
self.create_swa_adder(
|
||||||
size_swa=4096, sliding_window=128
|
size_swa=4096, sliding_window=128
|
||||||
)._swa_req_never_fits(**req)
|
).memory_budget.swa_never_fits(**req)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2,8 +2,12 @@ import logging
|
|||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.scheduler import Scheduler
|
from sglang.srt.managers.scheduler import Scheduler
|
||||||
|
from sglang.srt.mem_cache.allocator.token import TokenToKVPoolAllocator
|
||||||
|
from sglang.srt.mem_cache.unified_memory_pool import init_unified_swa_pools
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
@@ -56,6 +60,16 @@ class TestSchedulerInitReqMaxNewTokens(unittest.TestCase):
|
|||||||
scheduler.max_total_num_tokens = max_total_num_tokens
|
scheduler.max_total_num_tokens = max_total_num_tokens
|
||||||
scheduler.page_size = page_size
|
scheduler.page_size = page_size
|
||||||
scheduler.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get()
|
scheduler.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get()
|
||||||
|
scheduler.sliding_window_size = None
|
||||||
|
scheduler.chunked_prefill_size = None
|
||||||
|
scheduler.token_to_kv_pool_allocator = TokenToKVPoolAllocator(
|
||||||
|
size=max_total_num_tokens,
|
||||||
|
dtype=torch.int64,
|
||||||
|
device="cpu",
|
||||||
|
kvcache=None,
|
||||||
|
need_sort=False,
|
||||||
|
)
|
||||||
|
scheduler.token_to_kv_pool_allocator.page_size = page_size
|
||||||
return scheduler
|
return scheduler
|
||||||
|
|
||||||
def _new_req(self, max_new_tokens, input_len: int = 8, min_new_tokens: int = 0):
|
def _new_req(self, max_new_tokens, input_len: int = 8, min_new_tokens: int = 0):
|
||||||
@@ -180,6 +194,37 @@ class TestSchedulerInitReqMaxNewTokens(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
self._init_and_check(scheduler, req)
|
self._init_and_check(scheduler, req)
|
||||||
|
|
||||||
|
def test_unified_budget_rounds_prompt_and_decode_together(self):
|
||||||
|
bundle = init_unified_swa_pools(
|
||||||
|
device="cpu",
|
||||||
|
kv_cache_dtype=torch.float16,
|
||||||
|
head_num=1,
|
||||||
|
head_dim=4,
|
||||||
|
v_head_dim=4,
|
||||||
|
swa_head_num=1,
|
||||||
|
swa_head_dim=4,
|
||||||
|
swa_v_head_dim=4,
|
||||||
|
page_size=4,
|
||||||
|
start_layer=0,
|
||||||
|
end_layer=2,
|
||||||
|
swa_attention_layer_ids=[1],
|
||||||
|
full_attention_layer_ids=[0],
|
||||||
|
total_bytes=384,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
need_sort=False,
|
||||||
|
lazy_compaction=True,
|
||||||
|
)
|
||||||
|
scheduler = self._new_scheduler(page_size=4)
|
||||||
|
scheduler.token_to_kv_pool_allocator = bundle.token_to_kv_pool_allocator
|
||||||
|
scheduler.sliding_window_size = 4
|
||||||
|
scheduler.chunked_prefill_size = 4
|
||||||
|
scheduler.max_new_tokens_limit = None
|
||||||
|
for prompt_len in (4, 5, 6, 7):
|
||||||
|
with self.subTest(prompt_len=prompt_len):
|
||||||
|
req = self._new_req(max_new_tokens=1, input_len=prompt_len)
|
||||||
|
scheduler.init_req_max_new_tokens(req)
|
||||||
|
self.assertEqual(req.sampling_params.max_new_tokens, 1)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -14,7 +14,9 @@ from types import SimpleNamespace
|
|||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
from sglang.srt.disaggregation.decode import DecodePreallocQueue
|
from sglang.srt.disaggregation.decode import DecodePreallocQueue
|
||||||
|
from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -28,6 +30,7 @@ def _make_model_runner(**attrs):
|
|||||||
raise AttributeError on the internal `self.effective_max_total_num_tokens`
|
raise AttributeError on the internal `self.effective_max_total_num_tokens`
|
||||||
read inside `max_token_pool_size`."""
|
read inside `max_token_pool_size`."""
|
||||||
instance = object.__new__(ModelRunner)
|
instance = object.__new__(ModelRunner)
|
||||||
|
instance.kv_cache_configurator = object.__new__(KVCacheConfigurator)
|
||||||
for name, value in attrs.items():
|
for name, value in attrs.items():
|
||||||
setattr(instance, name, value)
|
setattr(instance, name, value)
|
||||||
return instance
|
return instance
|
||||||
@@ -85,8 +88,9 @@ class TestMaxTokenPoolSize(CustomTestCase):
|
|||||||
full_max_total_num_tokens=3000,
|
full_max_total_num_tokens=3000,
|
||||||
swa_max_total_num_tokens=500,
|
swa_max_total_num_tokens=500,
|
||||||
)
|
)
|
||||||
self.assertEqual(instance.max_token_pool_size, 3000)
|
with get_context().override_server_args(enable_unified_memory=False):
|
||||||
self.assertEqual(instance.effective_max_total_num_tokens, 3000)
|
self.assertEqual(instance.max_token_pool_size, 3000)
|
||||||
|
self.assertEqual(instance.effective_max_total_num_tokens, 3000)
|
||||||
|
|
||||||
|
|
||||||
def _make_prealloc_queue(
|
def _make_prealloc_queue(
|
||||||
|
|||||||
@@ -35,6 +35,8 @@ from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
|||||||
from sglang.srt.mem_cache.kv_index_translator import KVIndexTranslator, KVReadTables
|
from sglang.srt.mem_cache.kv_index_translator import KVIndexTranslator, KVReadTables
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.mem_cache.unified_memory_pool import MHASubPoolSpec, UnifiedKVPool
|
from sglang.srt.mem_cache.unified_memory_pool import MHASubPoolSpec, UnifiedKVPool
|
||||||
|
from sglang.srt.state_capturer.base import BaseTopkCapturer
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
_DEV = "cpu"
|
_DEV = "cpu"
|
||||||
_FULL_L = 2
|
_FULL_L = 2
|
||||||
@@ -135,6 +137,7 @@ class TestPassthrough(unittest.TestCase):
|
|||||||
device=_DEV,
|
device=_DEV,
|
||||||
)
|
)
|
||||||
self.assertFalse(src.is_translating)
|
self.assertFalse(src.is_translating)
|
||||||
|
self.assertEqual(src.capture_token_capacity(17), 18)
|
||||||
rows = torch.tensor([2, 0])
|
rows = torch.tensor([2, 0])
|
||||||
view = src.build_index_table(
|
view = src.build_index_table(
|
||||||
req_pool_indices=rows, seq_lens=torch.tensor([5, 3])
|
req_pool_indices=rows, seq_lens=torch.tensor([5, 3])
|
||||||
@@ -594,7 +597,7 @@ class TestViewMemo(unittest.TestCase):
|
|||||||
self.assertEqual(v2.ids.shape[0], 1)
|
self.assertEqual(v2.ids.shape[0], 1)
|
||||||
|
|
||||||
|
|
||||||
class TestWriteLoc(unittest.TestCase):
|
class TestWriteLoc(CustomTestCase):
|
||||||
"""The two-phase write contract: `rebind_write_loc` rebinds the full side
|
"""The two-phase write contract: `rebind_write_loc` rebinds the full side
|
||||||
once at ForwardBatch construction, and the sliding-window write loc derives
|
once at ForwardBatch construction, and the sliding-window write loc derives
|
||||||
POINTWISE from the full-side values -- pads, slices, and fresh copies
|
POINTWISE from the full-side values -- pads, slices, and fresh copies
|
||||||
@@ -630,6 +633,121 @@ class TestWriteLoc(unittest.TestCase):
|
|||||||
self.assertTrue(torch.equal(fb.out_cache_loc, want_full))
|
self.assertTrue(torch.equal(fb.out_cache_loc, want_full))
|
||||||
self.assertTrue(torch.equal(virt, keep))
|
self.assertTrue(torch.equal(virt, keep))
|
||||||
|
|
||||||
|
def test_capture_capacity_covers_dcp_widened_mamba_ids(self):
|
||||||
|
"""A capturer must store the highest virtual IDs issued under DCP."""
|
||||||
|
from sglang.srt.mem_cache.allocator.unified_mamba import (
|
||||||
|
UnifiedMambaTokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.unified_memory_pool import MambaSubPoolSpec
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
|
for dcp_size in (1, 2, 4):
|
||||||
|
with (
|
||||||
|
self.subTest(dcp_size=dcp_size),
|
||||||
|
get_parallel().override(attn_dcp_size=dcp_size),
|
||||||
|
):
|
||||||
|
pool = UnifiedKVPool(
|
||||||
|
total_bytes=2048,
|
||||||
|
sub_pool_specs=[
|
||||||
|
MHASubPoolSpec(
|
||||||
|
name="full",
|
||||||
|
layer_num=1,
|
||||||
|
head_num=1,
|
||||||
|
head_dim=4,
|
||||||
|
store_dtype=torch.float16,
|
||||||
|
grow_direction="up",
|
||||||
|
),
|
||||||
|
MambaSubPoolSpec(
|
||||||
|
name="mamba",
|
||||||
|
layer_num=1,
|
||||||
|
conv_state_shapes=((2, 2),),
|
||||||
|
conv_dtype=torch.float16,
|
||||||
|
temporal_state_shape=(2, 2),
|
||||||
|
temporal_dtype=torch.float16,
|
||||||
|
grow_direction="down",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
device="cpu",
|
||||||
|
enable_memory_saver=False,
|
||||||
|
page_size=4,
|
||||||
|
)
|
||||||
|
allocator = UnifiedMambaTokenToKVPoolAllocator(
|
||||||
|
unified_buffer=pool,
|
||||||
|
kvcache=SimpleNamespace(full_kv_pool=None, mamba_pool=None),
|
||||||
|
device="cpu",
|
||||||
|
page_size=4,
|
||||||
|
)
|
||||||
|
virt = allocator.alloc(allocator.available_size())
|
||||||
|
self.assertIsNotNone(virt)
|
||||||
|
src = _make_source(allocator, virt[None, :], 4)
|
||||||
|
cap = object.__new__(BaseTopkCapturer)
|
||||||
|
cap.topk_size = 1
|
||||||
|
expected = torch.arange(len(virt), dtype=torch.int32).reshape(-1, 1, 1)
|
||||||
|
cap.device_cache = SimpleNamespace(buffer=expected)
|
||||||
|
cap.host_cache = SimpleNamespace(
|
||||||
|
buffer=torch.zeros(
|
||||||
|
src.capture_token_capacity(1), 1, 1, dtype=torch.int32
|
||||||
|
)
|
||||||
|
)
|
||||||
|
fb = _FakeForwardBatch(out_cache_loc=virt)
|
||||||
|
fb.out_cache_loc_virtual = virt
|
||||||
|
cap.on_forward_end(fb, False, None, no_copy_to_cpu=False)
|
||||||
|
req_pool = SimpleNamespace(req_to_token=virt[None, :])
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(cap.get_topk(0, len(virt) + 1, req_pool), expected)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_topk_capture_round_trips_request_token_ids(self):
|
||||||
|
for ps in (1, 4, 64):
|
||||||
|
for translating in (False, True):
|
||||||
|
for overlap in (False, True):
|
||||||
|
with self.subTest(
|
||||||
|
page_size=ps, translating=translating, overlap=overlap
|
||||||
|
):
|
||||||
|
src, _, _, _, virt, _, _ = self._built(ps=ps, n=3 * ps)
|
||||||
|
req_pool = SimpleNamespace(req_to_token=virt.clone()[None, :])
|
||||||
|
fb = _FakeForwardBatch(out_cache_loc=virt.clone())
|
||||||
|
if translating:
|
||||||
|
src.rebind_write_loc(fb)
|
||||||
|
fb.out_cache_loc = torch.cat(
|
||||||
|
[fb.out_cache_loc, virt.new_zeros(2)]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
fb.out_cache_loc_virtual = None
|
||||||
|
|
||||||
|
# Admission may be capped below IDs issued after reuse.
|
||||||
|
capacity = src.capture_token_capacity(ps)
|
||||||
|
self.assertGreater(capacity, int(virt.max()))
|
||||||
|
expected = (
|
||||||
|
torch.arange(len(virt) * 4, dtype=torch.int32).reshape(
|
||||||
|
-1, 2, 2
|
||||||
|
)
|
||||||
|
+ 1
|
||||||
|
)
|
||||||
|
cap = object.__new__(BaseTopkCapturer)
|
||||||
|
cap.topk_size = 2
|
||||||
|
cap.device_cache = SimpleNamespace(
|
||||||
|
buffer=torch.cat([expected, expected.new_zeros(2, 2, 2)])
|
||||||
|
)
|
||||||
|
cap.host_cache = SimpleNamespace(
|
||||||
|
buffer=torch.zeros(capacity, 2, 2, dtype=torch.int32)
|
||||||
|
)
|
||||||
|
result = cap.on_forward_end(
|
||||||
|
fb, False, None, no_copy_to_cpu=overlap
|
||||||
|
)
|
||||||
|
if overlap:
|
||||||
|
fb.out_cache_loc.zero_()
|
||||||
|
if translating:
|
||||||
|
fb.out_cache_loc_virtual.zero_()
|
||||||
|
cap.device_cache.buffer.zero_()
|
||||||
|
result.map_device_tensors(lambda value: value.cpu())
|
||||||
|
result.finalize()
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(
|
||||||
|
cap.get_topk(0, len(virt) + 1, req_pool), expected
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def test_swa_write_loc_round_trips_from_full_side(self):
|
def test_swa_write_loc_round_trips_from_full_side(self):
|
||||||
"""Derived property: `field(full(t)) == swa(t)` for any virtual run t,
|
"""Derived property: `field(full(t)) == swa(t)` for any virtual run t,
|
||||||
across page sizes and multipliers."""
|
across page sizes and multipliers."""
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ from sglang.srt.mem_cache.allocator.unified_sub_pool import (
|
|||||||
MultiEndedAllocator,
|
MultiEndedAllocator,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||||
|
from sglang.srt.mem_cache.prefill_budget import estimate_swa_kv_tokens
|
||||||
from sglang.srt.mem_cache.unified_cache.components import ComponentType
|
from sglang.srt.mem_cache.unified_cache.components import ComponentType
|
||||||
from sglang.srt.mem_cache.unified_memory_pool import (
|
from sglang.srt.mem_cache.unified_memory_pool import (
|
||||||
MambaSubPoolSpec,
|
MambaSubPoolSpec,
|
||||||
@@ -558,6 +559,7 @@ class TestUnifiedSWATokenToKVPoolAllocator(unittest.TestCase):
|
|||||||
swa_layer_num=2,
|
swa_layer_num=2,
|
||||||
head_num=2,
|
head_num=2,
|
||||||
head_dim=4,
|
head_dim=4,
|
||||||
|
page_size=1,
|
||||||
):
|
):
|
||||||
full_spec = MHASubPoolSpec(
|
full_spec = MHASubPoolSpec(
|
||||||
name="full",
|
name="full",
|
||||||
@@ -584,6 +586,7 @@ class TestUnifiedSWATokenToKVPoolAllocator(unittest.TestCase):
|
|||||||
sub_pool_specs=[full_spec, swa_spec],
|
sub_pool_specs=[full_spec, swa_spec],
|
||||||
device=_DEV,
|
device=_DEV,
|
||||||
enable_memory_saver=False,
|
enable_memory_saver=False,
|
||||||
|
page_size=page_size,
|
||||||
)
|
)
|
||||||
kvcache = _FakeUnifiedSWAKVPool(pool)
|
kvcache = _FakeUnifiedSWAKVPool(pool)
|
||||||
allocator = UnifiedSWATokenToKVPoolAllocator(
|
allocator = UnifiedSWATokenToKVPoolAllocator(
|
||||||
@@ -592,11 +595,175 @@ class TestUnifiedSWATokenToKVPoolAllocator(unittest.TestCase):
|
|||||||
device=_DEV,
|
device=_DEV,
|
||||||
full_max_total_num_tokens=n_full_slots,
|
full_max_total_num_tokens=n_full_slots,
|
||||||
swa_max_total_num_tokens=n_swa_slots,
|
swa_max_total_num_tokens=n_swa_slots,
|
||||||
|
page_size=page_size,
|
||||||
need_sort=False,
|
need_sort=False,
|
||||||
forward_stream=None,
|
forward_stream=None,
|
||||||
)
|
)
|
||||||
return pool, allocator, kvcache
|
return pool, allocator, kvcache
|
||||||
|
|
||||||
|
def test_reclaim_plan_matches_exhaustive_page_targets(self):
|
||||||
|
page_size = 4
|
||||||
|
_, allocator, _ = self._build(
|
||||||
|
n_full_slots=40, n_swa_slots=24, page_size=page_size
|
||||||
|
)
|
||||||
|
allocator.lazy_compaction = True
|
||||||
|
for sub_pool in (allocator.full_attn_allocator, allocator.swa_attn_allocator):
|
||||||
|
sub_pool.lazy_compaction = True
|
||||||
|
sub_pool.disagg_move_gate = lambda: False
|
||||||
|
live = allocator.alloc(16)
|
||||||
|
self.assertIsNotNone(live)
|
||||||
|
allocator.free(live[4:8])
|
||||||
|
allocator.free_swa(live[8:12])
|
||||||
|
|
||||||
|
for compacted in (False, True):
|
||||||
|
for sub_pool in (
|
||||||
|
allocator.full_attn_allocator,
|
||||||
|
allocator.swa_attn_allocator,
|
||||||
|
):
|
||||||
|
sub_pool.disagg_move_gate = lambda: compacted
|
||||||
|
for full_evictable, swa_evictable in ((0, 0), (7, 5), (12, 8), (100, 100)):
|
||||||
|
max_full = min(12, full_evictable) // page_size
|
||||||
|
max_swa = min(8, swa_evictable) // page_size
|
||||||
|
for full_pages in range(9):
|
||||||
|
for swa_pages in range(9):
|
||||||
|
feasible = [
|
||||||
|
(full * page_size, swa * page_size)
|
||||||
|
for swa in range(max_swa + 1)
|
||||||
|
for full in range(max_full + 1)
|
||||||
|
if allocator._fits_page_demand(
|
||||||
|
full_pages,
|
||||||
|
swa_pages,
|
||||||
|
full_reclaim_pages=full,
|
||||||
|
swa_reclaim_pages=swa,
|
||||||
|
compacted=compacted,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
with self.subTest(
|
||||||
|
compacted=compacted,
|
||||||
|
evictable=(full_evictable, swa_evictable),
|
||||||
|
pages=(full_pages, swa_pages),
|
||||||
|
):
|
||||||
|
self.assertEqual(
|
||||||
|
allocator.reclaim_plan(
|
||||||
|
full_pages * page_size,
|
||||||
|
swa_pages * page_size,
|
||||||
|
full_evictable_tokens=full_evictable,
|
||||||
|
swa_evictable_tokens=swa_evictable,
|
||||||
|
),
|
||||||
|
feasible[0] if feasible else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_restore_swa_without_allocating_more_full(self):
|
||||||
|
_, allocator, _ = self._build(page_size=4)
|
||||||
|
indices = allocator.alloc(8)
|
||||||
|
full_before = allocator.translate_kv_indices_for_transfer(indices).clone()
|
||||||
|
allocator.free_swa(indices)
|
||||||
|
|
||||||
|
self.assertEqual(allocator.reclaim_plan(0, 8), (0, 0))
|
||||||
|
self.assertTrue(allocator.can_reserve(0, 8))
|
||||||
|
self.assertTrue(allocator.ensure_capacity(0, 8))
|
||||||
|
allocator.swa_attn_allocator.alloc_with_virtual((indices // 4).unique())
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(
|
||||||
|
allocator.translate_kv_indices_for_transfer(indices), full_before
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
bool((allocator.swa_attn_allocator.translate_kv_loc(indices) > 0).all())
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_empty_pool_reservation_matches_packed_byte_boundary(self):
|
||||||
|
page_size = 4
|
||||||
|
_, allocator, _ = self._build(
|
||||||
|
n_full_slots=40,
|
||||||
|
n_swa_slots=24,
|
||||||
|
full_layer_num=4,
|
||||||
|
swa_layer_num=2,
|
||||||
|
page_size=page_size,
|
||||||
|
)
|
||||||
|
full_allocator = allocator.full_attn_allocator
|
||||||
|
swa_allocator = allocator.swa_attn_allocator
|
||||||
|
swa_pages = 2
|
||||||
|
full_pages = (
|
||||||
|
allocator._empty_shared_gap_bytes
|
||||||
|
- swa_pages * swa_allocator.entry_bytes_per_page
|
||||||
|
) // full_allocator.entry_bytes_per_page
|
||||||
|
packed_bytes = (
|
||||||
|
full_pages * full_allocator.entry_bytes_per_page
|
||||||
|
+ swa_pages * swa_allocator.entry_bytes_per_page
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertLessEqual(
|
||||||
|
full_pages + 1, full_allocator.num_pages - full_allocator.min_page_index
|
||||||
|
)
|
||||||
|
self.assertLessEqual(
|
||||||
|
swa_pages, swa_allocator.num_pages - swa_allocator.min_page_index
|
||||||
|
)
|
||||||
|
self.assertLessEqual(packed_bytes, allocator._empty_shared_gap_bytes)
|
||||||
|
self.assertGreater(
|
||||||
|
packed_bytes + full_allocator.entry_bytes_per_page,
|
||||||
|
allocator._empty_shared_gap_bytes,
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
allocator.can_reserve(
|
||||||
|
full_pages * page_size,
|
||||||
|
swa_pages * page_size,
|
||||||
|
empty_pool=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertFalse(
|
||||||
|
allocator.can_reserve(
|
||||||
|
full_pages * page_size + 1,
|
||||||
|
swa_pages * page_size,
|
||||||
|
empty_pool=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
extend_tokens = 32
|
||||||
|
max_new_tokens = 0
|
||||||
|
reservation_full_tokens = extend_tokens + max_new_tokens + page_size
|
||||||
|
reservation_swa_tokens = estimate_swa_kv_tokens(
|
||||||
|
extend_tokens,
|
||||||
|
max_new_tokens,
|
||||||
|
sliding_window_size=16,
|
||||||
|
page_size=page_size,
|
||||||
|
allocation_limit=16,
|
||||||
|
)
|
||||||
|
reservation_swa_with_tail = estimate_swa_kv_tokens(
|
||||||
|
extend_tokens,
|
||||||
|
max_new_tokens,
|
||||||
|
sliding_window_size=16,
|
||||||
|
page_size=page_size,
|
||||||
|
)
|
||||||
|
reservation_bytes = (
|
||||||
|
reservation_full_tokens // page_size
|
||||||
|
) * full_allocator.entry_bytes_per_page + (
|
||||||
|
reservation_swa_tokens // page_size
|
||||||
|
) * swa_allocator.entry_bytes_per_page
|
||||||
|
reservation_bytes_with_tail = (
|
||||||
|
reservation_full_tokens // page_size
|
||||||
|
) * full_allocator.entry_bytes_per_page + (
|
||||||
|
reservation_swa_with_tail // page_size
|
||||||
|
) * swa_allocator.entry_bytes_per_page
|
||||||
|
self.assertLessEqual(reservation_bytes, allocator._empty_shared_gap_bytes)
|
||||||
|
self.assertGreater(
|
||||||
|
reservation_bytes_with_tail, allocator._empty_shared_gap_bytes
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
allocator.can_reserve(
|
||||||
|
reservation_full_tokens,
|
||||||
|
reservation_swa_tokens,
|
||||||
|
empty_pool=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertFalse(
|
||||||
|
allocator.can_reserve(
|
||||||
|
reservation_full_tokens,
|
||||||
|
reservation_swa_with_tail,
|
||||||
|
empty_pool=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def _alloc(self, allocator, kvcache, n):
|
def _alloc(self, allocator, kvcache, n):
|
||||||
"""Allocate N virtual ids; stamp the data marker on both sub-pools."""
|
"""Allocate N virtual ids; stamp the data marker on both sub-pools."""
|
||||||
v = allocator.alloc(n)
|
v = allocator.alloc(n)
|
||||||
|
|||||||
@@ -0,0 +1,300 @@
|
|||||||
|
"""CPU regressions for allocator-owned prefill admission and pending demand."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from array import array
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
|
from sglang.srt.managers.schedule_policy import PrefillAdder
|
||||||
|
from sglang.srt.managers.scheduler import Scheduler
|
||||||
|
from sglang.srt.mem_cache.allocator.hisparse import (
|
||||||
|
DeepSeekV4HiSparseTokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.allocator.swa import (
|
||||||
|
PureSWATokenToKVPoolAllocator,
|
||||||
|
SWATokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
||||||
|
UnifiedMambaSWATokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.common import evict_from_tree_cache
|
||||||
|
from sglang.srt.mem_cache.prefill_budget import SWAPrefillBudget
|
||||||
|
from sglang.srt.mem_cache.unified_memory_pool import init_unified_swa_pools
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _shared_allocator(*, page_size=4, total_bytes=1024):
|
||||||
|
return init_unified_swa_pools(
|
||||||
|
device="cpu",
|
||||||
|
kv_cache_dtype=torch.float16,
|
||||||
|
head_num=1,
|
||||||
|
head_dim=4,
|
||||||
|
v_head_dim=4,
|
||||||
|
swa_head_num=1,
|
||||||
|
swa_head_dim=4,
|
||||||
|
swa_v_head_dim=4,
|
||||||
|
page_size=page_size,
|
||||||
|
start_layer=0,
|
||||||
|
end_layer=2,
|
||||||
|
swa_attention_layer_ids=[1],
|
||||||
|
full_attention_layer_ids=[0],
|
||||||
|
total_bytes=total_bytes,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
need_sort=False,
|
||||||
|
lazy_compaction=True,
|
||||||
|
).token_to_kv_pool_allocator
|
||||||
|
|
||||||
|
|
||||||
|
def _cache():
|
||||||
|
return SimpleNamespace(
|
||||||
|
sliding_window_size=8,
|
||||||
|
full_evictable_size=lambda: 0,
|
||||||
|
swa_evictable_size=lambda: 0,
|
||||||
|
is_chunk_cache=lambda: False,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestSharedPrefillMemoryBudget(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.allocator = _shared_allocator()
|
||||||
|
self.cache = _cache()
|
||||||
|
self.budget = self.allocator.create_prefill_budget(self.cache)
|
||||||
|
self.request = dict(
|
||||||
|
extend_input_len=12,
|
||||||
|
total_tokens=20,
|
||||||
|
max_new_tokens=4,
|
||||||
|
input_tokens=12,
|
||||||
|
swa_host_hit_length=0,
|
||||||
|
chunk_limit=16,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_pending_batch_cannot_spend_shared_bytes_twice(self):
|
||||||
|
self.assertEqual(self.budget.check_prefill(**self.request), (True, 16))
|
||||||
|
self.budget.reserve(12, 4, chunk_limit=16)
|
||||||
|
# Each side separately has room, but their combined reservation does not.
|
||||||
|
self.assertGreater(self.budget.remaining_total, 20)
|
||||||
|
self.assertGreater(self.budget.remaining_swa, 16)
|
||||||
|
self.assertEqual(self.budget.check_prefill(**self.request), (False, None))
|
||||||
|
self.assertEqual(self.allocator.full_attn_allocator.allocated_count(), 0)
|
||||||
|
self.assertEqual(self.allocator.swa_attn_allocator.allocated_count(), 0)
|
||||||
|
|
||||||
|
def test_mixed_decode_reserves_both_sides(self):
|
||||||
|
budget = self.allocator.create_prefill_budget(
|
||||||
|
self.cache, num_mixed_decode_tokens=4
|
||||||
|
)
|
||||||
|
budget.reserve(12, 4, chunk_limit=16)
|
||||||
|
self.assertEqual(
|
||||||
|
(budget.total_offset, budget.current_offset, budget.swa_offset),
|
||||||
|
(24, 20, 20),
|
||||||
|
)
|
||||||
|
self.assertEqual(budget.check_prefill(**self.request), (False, None))
|
||||||
|
|
||||||
|
def test_prefix_lock_changes_admission_without_rebuilding_budget(self):
|
||||||
|
self.assertIsNotNone(self.allocator.alloc(24))
|
||||||
|
self.cache.full_evictable_size = lambda: 24
|
||||||
|
self.cache.swa_evictable_size = lambda: 24
|
||||||
|
self.assertEqual(self.budget.check_prefill(**self.request), (True, 16))
|
||||||
|
# Locking the cached prefix removes its eviction credit.
|
||||||
|
self.cache.full_evictable_size = lambda: 0
|
||||||
|
self.cache.swa_evictable_size = lambda: 0
|
||||||
|
self.assertEqual(self.budget.check_prefill(**self.request), (False, None))
|
||||||
|
|
||||||
|
def test_final_chunk_reserves_decode_headroom(self):
|
||||||
|
limit = self.budget.fit_chunk(
|
||||||
|
extend_input_len=12, max_new_tokens=80, chunk_limit=16
|
||||||
|
)
|
||||||
|
self.assertEqual(limit, 8)
|
||||||
|
self.budget.reserve(limit, 0, chunk_limit=16, is_chunked_continuation=True)
|
||||||
|
self.assertIsNotNone(self.allocator.alloc(limit))
|
||||||
|
|
||||||
|
def test_host_swa_load_is_part_of_joint_demand(self):
|
||||||
|
self.assertEqual(self.budget.check_prefill(**self.request), (True, 16))
|
||||||
|
request = {**self.request, "swa_host_hit_length": 32}
|
||||||
|
self.assertEqual(self.budget.check_prefill(**request), (False, None))
|
||||||
|
|
||||||
|
def test_prompt_clipping_uses_the_empty_pool(self):
|
||||||
|
kwargs = dict(token_capacity=1, sliding_window_size=8, chunk_size=16)
|
||||||
|
limit = self.allocator.max_new_tokens_for_memory(12, 80, **kwargs)
|
||||||
|
self.assertIsNotNone(limit)
|
||||||
|
self.assertGreater(limit, 0)
|
||||||
|
self.assertIsNotNone(self.allocator.alloc(24))
|
||||||
|
self.assertEqual(
|
||||||
|
self.allocator.max_new_tokens_for_memory(12, 80, **kwargs), limit
|
||||||
|
)
|
||||||
|
self.assertIsNone(self.allocator.max_new_tokens_for_memory(100, 0, **kwargs))
|
||||||
|
|
||||||
|
def test_shared_stats_pair_available_tokens_with_current_capacity(self):
|
||||||
|
self.assertIsNotNone(self.allocator.alloc(12))
|
||||||
|
(full_capacity, full_free), (swa_capacity, swa_free) = (
|
||||||
|
self.allocator.swa_capacity_and_available(full_capacity=1, swa_capacity=1)
|
||||||
|
)
|
||||||
|
self.assertEqual(full_capacity - full_free, 12)
|
||||||
|
self.assertEqual(swa_capacity - swa_free, 12)
|
||||||
|
|
||||||
|
def test_common_eviction_dispatches_joint_reclaim(self):
|
||||||
|
self.cache.token_to_kv_pool_allocator = self.allocator
|
||||||
|
self.allocator.evict_to_free_tokens = MagicMock()
|
||||||
|
evict_from_tree_cache(self.cache, 8)
|
||||||
|
self.allocator.evict_to_free_tokens.assert_called_once_with(self.cache, 8)
|
||||||
|
|
||||||
|
|
||||||
|
class TestSharedPrefillAdmission(unittest.TestCase):
|
||||||
|
def _new_admission(self, page_size, pool_pages, *, ignore_eos=False):
|
||||||
|
allocator = _shared_allocator(
|
||||||
|
page_size=page_size, total_bytes=pool_pages * page_size * 16
|
||||||
|
)
|
||||||
|
req = Req(
|
||||||
|
rid="unaligned-prompt",
|
||||||
|
origin_input_text=None,
|
||||||
|
origin_input_ids=array("q", [1] * (page_size + 1)),
|
||||||
|
sampling_params=SamplingParams(max_new_tokens=1, ignore_eos=ignore_eos),
|
||||||
|
)
|
||||||
|
scheduler = Scheduler.__new__(Scheduler)
|
||||||
|
scheduler.max_req_len = 16 * page_size
|
||||||
|
scheduler.max_total_num_tokens = allocator.size_full
|
||||||
|
scheduler.page_size = page_size
|
||||||
|
scheduler.max_new_tokens_limit = None
|
||||||
|
scheduler.sliding_window_size = page_size
|
||||||
|
scheduler.chunked_prefill_size = page_size
|
||||||
|
scheduler.token_to_kv_pool_allocator = allocator
|
||||||
|
with get_parallel().override(attn_dcp_size=1):
|
||||||
|
scheduler.init_req_max_new_tokens(req)
|
||||||
|
self.assertEqual(req.sampling_params.max_new_tokens, 1)
|
||||||
|
req._refresh_fill_ids()
|
||||||
|
|
||||||
|
cache = SimpleNamespace(
|
||||||
|
sliding_window_size=page_size,
|
||||||
|
disable=True,
|
||||||
|
full_evictable_size=lambda: 0,
|
||||||
|
swa_evictable_size=lambda: 0,
|
||||||
|
is_chunk_cache=lambda: True,
|
||||||
|
supports_mamba=lambda: False,
|
||||||
|
)
|
||||||
|
adder = PrefillAdder(
|
||||||
|
page_size=page_size,
|
||||||
|
tree_cache=cache,
|
||||||
|
token_to_kv_pool_allocator=allocator,
|
||||||
|
running_batch=None,
|
||||||
|
new_token_ratio=1.0,
|
||||||
|
rem_input_tokens=16 * page_size,
|
||||||
|
rem_chunk_tokens=page_size,
|
||||||
|
)
|
||||||
|
return allocator, req, adder
|
||||||
|
|
||||||
|
def test_unaligned_final_chunk_makes_progress(self):
|
||||||
|
for page_size in (4, 64):
|
||||||
|
with self.subTest(page_size=page_size):
|
||||||
|
allocator, req, adder = self._new_admission(page_size, pool_pages=7)
|
||||||
|
req.prefix_indices = allocator.alloc(page_size)
|
||||||
|
self.assertIsNotNone(req.prefix_indices)
|
||||||
|
self.assertTrue(allocator.can_reserve(page_size + 2, page_size + 2))
|
||||||
|
|
||||||
|
self.assertIsNone(adder.add_chunked_req(req))
|
||||||
|
self.assertEqual(adder.can_run_list, [req])
|
||||||
|
self.assertEqual(req.extend_range.length, 1)
|
||||||
|
|
||||||
|
def test_unaligned_ignore_eos_enters_empty_pool(self):
|
||||||
|
for page_size in (4, 64):
|
||||||
|
with self.subTest(page_size=page_size):
|
||||||
|
allocator, req, adder = self._new_admission(
|
||||||
|
page_size, pool_pages=6, ignore_eos=True
|
||||||
|
)
|
||||||
|
self.assertEqual(len(req.prefix_indices), 0)
|
||||||
|
self.assertTrue(allocator.can_reserve(2 * page_size + 2, 2 * page_size))
|
||||||
|
|
||||||
|
adder.add_one_req(
|
||||||
|
req, has_chunked_req=False, truncation_align_size=None
|
||||||
|
)
|
||||||
|
self.assertEqual(adder.can_run_list, [req])
|
||||||
|
|
||||||
|
|
||||||
|
class TestFixedPrefillMemoryBudget(unittest.TestCase):
|
||||||
|
def _allocator(self, cls=SWATokenToKVPoolAllocator):
|
||||||
|
allocator = object.__new__(cls)
|
||||||
|
allocator.page_size = 4
|
||||||
|
allocator._size_full = 128
|
||||||
|
allocator._size_swa = 64
|
||||||
|
allocator.full_available_size = lambda: 128
|
||||||
|
allocator.swa_available_size = lambda: 64
|
||||||
|
return allocator
|
||||||
|
|
||||||
|
def test_ring_slot_reserved_once_and_evictable_tokens_give_no_credit(self):
|
||||||
|
allocator = self._allocator()
|
||||||
|
allocator._swa_req_ring = True
|
||||||
|
allocator._swa_ring_cost = 32
|
||||||
|
cache = _cache()
|
||||||
|
cache.swa_evictable_size = lambda: 1000
|
||||||
|
budget = allocator.create_prefill_budget(cache)
|
||||||
|
budget.reserve(12, 4, chunk_limit=16)
|
||||||
|
self.assertEqual(budget.remaining_swa, 32)
|
||||||
|
budget.reserve(12, 4, chunk_limit=16, is_chunked_continuation=True)
|
||||||
|
self.assertEqual(budget.remaining_swa, 32)
|
||||||
|
# The last exact ring slot remains admissible.
|
||||||
|
self.assertEqual(
|
||||||
|
budget.check_prefill(
|
||||||
|
extend_input_len=4,
|
||||||
|
total_tokens=12,
|
||||||
|
max_new_tokens=4,
|
||||||
|
input_tokens=4,
|
||||||
|
swa_host_hit_length=0,
|
||||||
|
chunk_limit=16,
|
||||||
|
),
|
||||||
|
(True, 16),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_pure_swa_budget_reads_swa_capacity(self):
|
||||||
|
allocator = self._allocator(PureSWATokenToKVPoolAllocator)
|
||||||
|
allocator.full_available_size = lambda: 0
|
||||||
|
budget = allocator.create_prefill_budget(_cache())
|
||||||
|
self.assertEqual(budget.remaining_total, 64)
|
||||||
|
self.assertTrue(budget.has_capacity())
|
||||||
|
|
||||||
|
def test_hisparse_budget_reads_wrapper_capacity(self):
|
||||||
|
allocator = self._allocator(DeepSeekV4HiSparseTokenToKVPoolAllocator)
|
||||||
|
allocator.full_available_size = lambda: 12
|
||||||
|
budget = allocator.create_prefill_budget(_cache())
|
||||||
|
self.assertEqual(budget.remaining_total, 12)
|
||||||
|
self.assertEqual(budget.remaining_swa, 64)
|
||||||
|
|
||||||
|
def test_tri_pool_keeps_fixed_admission_and_clipping(self):
|
||||||
|
allocator = self._allocator(UnifiedMambaSWATokenToKVPoolAllocator)
|
||||||
|
allocator.can_reserve = MagicMock(
|
||||||
|
side_effect=AssertionError("two-pool reservation")
|
||||||
|
)
|
||||||
|
budget = allocator.create_prefill_budget(_cache())
|
||||||
|
self.assertIs(type(budget), SWAPrefillBudget)
|
||||||
|
self.assertTrue(budget.has_capacity())
|
||||||
|
self.assertEqual(
|
||||||
|
allocator.max_new_tokens_for_memory(
|
||||||
|
5,
|
||||||
|
100,
|
||||||
|
token_capacity=32,
|
||||||
|
sliding_window_size=8,
|
||||||
|
chunk_size=16,
|
||||||
|
),
|
||||||
|
19,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_tri_pool_eviction_does_not_reenter_common(self):
|
||||||
|
allocator = self._allocator(UnifiedMambaSWATokenToKVPoolAllocator)
|
||||||
|
allocator.available_size = lambda: 0
|
||||||
|
allocator.full_available_size = lambda: 0
|
||||||
|
allocator.swa_available_size = lambda: 0
|
||||||
|
cache = _cache()
|
||||||
|
cache.token_to_kv_pool_allocator = allocator
|
||||||
|
cache.evict_for_alloc = MagicMock()
|
||||||
|
evict_from_tree_cache(cache, 8)
|
||||||
|
cache.evict_for_alloc.assert_called_once()
|
||||||
|
params = cache.evict_for_alloc.call_args.args[0]
|
||||||
|
self.assertEqual((params.num_tokens, params.swa_num_tokens), (8, 8))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -13,9 +13,9 @@
|
|||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
"""Byte-budget buffer sizing for the unified 2-pool factories.
|
"""Byte-budget buffer sizing for the unified 2-pool factories.
|
||||||
|
|
||||||
With ``unified_total_bytes`` set the buffer is that many bytes exactly (the
|
The SWA factory's ``total_bytes`` sets the exact buffer size. The Mamba pair's
|
||||||
mamba pair adds the state pool's bytes on top -- the budget is captured AFTER
|
``unified_total_bytes`` adds state bytes because its budget excludes that carve-out.
|
||||||
the state carve-out); without it, sizing falls back to the token-count re-sum.
|
Without an explicit budget, sizing falls back to the token-count re-sum.
|
||||||
Sizing from the ratio-derived token counts instead would re-introduce the
|
Sizing from the ratio-derived token counts instead would re-introduce the
|
||||||
configurator's rounding, which floors by the cell size and then page-aligns
|
configurator's rounding, which floors by the cell size and then page-aligns
|
||||||
EACH side, losing up to about a page of tokens per side.
|
EACH side, losing up to about a page of tokens per side.
|
||||||
@@ -91,7 +91,7 @@ class TestBudgetSizing(unittest.TestCase):
|
|||||||
(64 + 32) * e + (e - 2), # almost one more entry
|
(64 + 32) * e + (e - 2), # almost one more entry
|
||||||
):
|
):
|
||||||
with self.subTest(budget=budget):
|
with self.subTest(budget=budget):
|
||||||
bundle = _swa_factory(unified_total_bytes=budget)
|
bundle = _swa_factory(total_bytes=budget)
|
||||||
self.assertEqual(bundle.unified_memory_pool.total_bytes, budget)
|
self.assertEqual(bundle.unified_memory_pool.total_bytes, budget)
|
||||||
|
|
||||||
def test_fallback_is_the_token_count_resum(self):
|
def test_fallback_is_the_token_count_resum(self):
|
||||||
@@ -166,7 +166,7 @@ class TestBs1FeasibilityFloor(unittest.TestCase):
|
|||||||
must fail loud instead of livelocking later."""
|
must fail loud instead of livelocking later."""
|
||||||
with self.assertRaises(RuntimeError) as ctx:
|
with self.assertRaises(RuntimeError) as ctx:
|
||||||
_swa_factory(
|
_swa_factory(
|
||||||
unified_total_bytes=8 * _entry_bytes(),
|
total_bytes=8 * _entry_bytes(),
|
||||||
model_context_len=4096,
|
model_context_len=4096,
|
||||||
sliding_window_size=4096,
|
sliding_window_size=4096,
|
||||||
)
|
)
|
||||||
@@ -184,7 +184,7 @@ class TestBs1FeasibilityFloor(unittest.TestCase):
|
|||||||
):
|
):
|
||||||
with self.subTest(case=case):
|
with self.subTest(case=case):
|
||||||
bundle = _swa_factory(
|
bundle = _swa_factory(
|
||||||
unified_total_bytes=200 * e,
|
total_bytes=200 * e,
|
||||||
model_context_len=model_context_len,
|
model_context_len=model_context_len,
|
||||||
sliding_window_size=sliding_window_size,
|
sliding_window_size=sliding_window_size,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -327,6 +327,7 @@ class TestEveryUnifiedAllocatorOverridesFreeSegment(unittest.TestCase):
|
|||||||
mea.MultiEndedAllocator,
|
mea.MultiEndedAllocator,
|
||||||
unified_mamba.UnifiedMambaTokenToKVPoolAllocator,
|
unified_mamba.UnifiedMambaTokenToKVPoolAllocator,
|
||||||
unified_hybrid_swa.UnifiedSWATokenToKVPoolAllocator,
|
unified_hybrid_swa.UnifiedSWATokenToKVPoolAllocator,
|
||||||
|
unified_hybrid_swa.UnifiedMambaSWATokenToKVPoolAllocator,
|
||||||
):
|
):
|
||||||
with self.subTest(cls=cls.__name__):
|
with self.subTest(cls=cls.__name__):
|
||||||
self.assertIsNot(
|
self.assertIsNot(
|
||||||
@@ -346,9 +347,13 @@ class TestEveryUnifiedAllocatorOverridesFreeSegment(unittest.TestCase):
|
|||||||
mea.MultiEndedAllocator,
|
mea.MultiEndedAllocator,
|
||||||
unified_mamba.UnifiedMambaTokenToKVPoolAllocator,
|
unified_mamba.UnifiedMambaTokenToKVPoolAllocator,
|
||||||
unified_hybrid_swa.UnifiedSWATokenToKVPoolAllocator,
|
unified_hybrid_swa.UnifiedSWATokenToKVPoolAllocator,
|
||||||
|
unified_hybrid_swa.UnifiedMambaSWATokenToKVPoolAllocator,
|
||||||
):
|
):
|
||||||
with self.subTest(cls=cls.__name__):
|
with self.subTest(cls=cls.__name__):
|
||||||
self.assertIn("free_page_reps_group", inspect.getsource(cls))
|
alloc = object.__new__(cls)
|
||||||
|
alloc.free_group = None
|
||||||
|
alloc.free_group_begin()
|
||||||
|
self.assertEqual(alloc.free_page_reps_group, [])
|
||||||
|
|
||||||
|
|
||||||
class TestUnifiedSwaFullSideGroup(unittest.TestCase):
|
class TestUnifiedSwaFullSideGroup(unittest.TestCase):
|
||||||
|
|||||||
@@ -554,7 +554,11 @@ class TestUnifiedRadixAllocationEviction(CustomTestCase):
|
|||||||
def test_common_helper_uses_allocation_aware_entry_point(self):
|
def test_common_helper_uses_allocation_aware_entry_point(self):
|
||||||
tree_cache = MagicMock()
|
tree_cache = MagicMock()
|
||||||
tree_cache.is_chunk_cache.return_value = False
|
tree_cache.is_chunk_cache.return_value = False
|
||||||
tree_cache.token_to_kv_pool_allocator.available_size.return_value = 30
|
from sglang.srt.mem_cache.allocator.token import TokenToKVPoolAllocator
|
||||||
|
|
||||||
|
allocator = object.__new__(TokenToKVPoolAllocator)
|
||||||
|
allocator.available_size = lambda: 30
|
||||||
|
tree_cache.token_to_kv_pool_allocator = allocator
|
||||||
|
|
||||||
evict_from_tree_cache(tree_cache, num_tokens=100)
|
evict_from_tree_cache(tree_cache, num_tokens=100)
|
||||||
|
|
||||||
|
|||||||
@@ -4319,7 +4319,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
|
|
||||||
def test_buffer_load_back_swa_window_charged_at_admission(self):
|
def test_buffer_load_back_swa_window_charged_at_admission(self):
|
||||||
"""Admission contract: a request the SWA budget gate accepts must be
|
"""Admission contract: a request the SWA budget gate accepts must be
|
||||||
allocatable at batch time (_swa_reserved_tokens: "an admitted request
|
allocatable at batch time (estimate_swa_kv_tokens: "an admitted request
|
||||||
cannot OOM"). Regression: buffer mode surfaced a staged prefetch as
|
cannot OOM"). Regression: buffer mode surfaced a staged prefetch as
|
||||||
host_hit_length only, so the gate never charged the SWA window that
|
host_hit_length only, so the gate never charged the SWA window that
|
||||||
consumption (init_load_back -> cc.load) allocates and the request
|
consumption (init_load_back -> cc.load) allocates and the request
|
||||||
|
|||||||
@@ -20,7 +20,7 @@ Pure CPU; fakes stand in for the KV pools (data markers verify moves).
|
|||||||
|
|
||||||
import inspect
|
import inspect
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -175,6 +175,33 @@ class TestUnifiedTriPool(unittest.TestCase):
|
|||||||
self.assertIs(kvcache._full_allocator, fa)
|
self.assertIs(kvcache._full_allocator, fa)
|
||||||
self.assertIs(kvcache._swa_allocator, sa)
|
self.assertIs(kvcache._swa_allocator, sa)
|
||||||
|
|
||||||
|
def test_pd_preallocation_binds_only_swa_tail_pages(self):
|
||||||
|
_, allocator, _, _ = self._build(page_size=4)
|
||||||
|
before = allocator.available_size()
|
||||||
|
prefix = torch.tensor([0], dtype=torch.int64)
|
||||||
|
seq = torch.tensor([12], dtype=torch.int64)
|
||||||
|
# With an empty prefix, ordinary allocation supplies the same virtual
|
||||||
|
# pages without launching the GPU extend kernel. All page binding and
|
||||||
|
# capacity accounting still run through the real sub-allocators.
|
||||||
|
full = allocator.full_attn_allocator
|
||||||
|
with patch.object(
|
||||||
|
full, "alloc_extend", side_effect=lambda *a, **kw: full.alloc(12)
|
||||||
|
):
|
||||||
|
virtual = allocator.alloc_extend_swa_tail(
|
||||||
|
prefix, prefix, seq, seq, torch.tensor([-1]), 12, swa_tail_len=5
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(virtual)
|
||||||
|
self.assertEqual(len(virtual), 12)
|
||||||
|
self.assertEqual(allocator.full_attn_allocator.allocated_count(), 12)
|
||||||
|
self.assertEqual(allocator.swa_attn_allocator.allocated_count(), 8)
|
||||||
|
swa_pages = allocator.swa_v2p_page_table[virtual[::4] // 4]
|
||||||
|
self.assertLessEqual(swa_pages[0].item(), 0)
|
||||||
|
self.assertTrue(torch.all(swa_pages[1:] > 0).item())
|
||||||
|
allocator.free(virtual)
|
||||||
|
self.assertEqual(allocator.full_attn_allocator.allocated_count(), 0)
|
||||||
|
self.assertEqual(allocator.swa_attn_allocator.allocated_count(), 0)
|
||||||
|
self.assertEqual(allocator.available_size(), before)
|
||||||
|
|
||||||
def test_empty_float_is_transparent_to_the_ends(self):
|
def test_empty_float_is_transparent_to_the_ends(self):
|
||||||
_, allocator, _, _ = self._build()
|
_, allocator, _, _ = self._build()
|
||||||
fa = allocator.full_attn_allocator
|
fa = allocator.full_attn_allocator
|
||||||
@@ -961,11 +988,19 @@ class TestTriFactorySizing(unittest.TestCase):
|
|||||||
return kw
|
return kw
|
||||||
|
|
||||||
def test_budget_sizing_and_boot_signature(self):
|
def test_budget_sizing_and_boot_signature(self):
|
||||||
|
from sglang.srt.mem_cache.allocator.unified_hybrid_swa import (
|
||||||
|
UnifiedSWAAllocatorBase,
|
||||||
|
UnifiedSWATokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
|
|
||||||
budget = 1 << 20
|
budget = 1 << 20
|
||||||
bundle = init_unified_mamba_swa_pools(
|
bundle = init_unified_mamba_swa_pools(
|
||||||
**self._factory_kwargs(unified_total_bytes=budget)
|
**self._factory_kwargs(unified_total_bytes=budget)
|
||||||
)
|
)
|
||||||
pool = bundle.unified_memory_pool
|
pool = bundle.unified_memory_pool
|
||||||
|
allocator = bundle.token_to_kv_pool_allocator
|
||||||
|
self.assertIsInstance(allocator, UnifiedSWAAllocatorBase)
|
||||||
|
self.assertNotIsInstance(allocator, UnifiedSWATokenToKVPoolAllocator)
|
||||||
# Buffer = budget + the state pool's bytes (budget captured AFTER the
|
# Buffer = budget + the state pool's bytes (budget captured AFTER the
|
||||||
# state carve-out), never the token-count re-sum.
|
# state carve-out), never the token-count re-sum.
|
||||||
state_bytes = 4 * pool.spec("mamba").entry_bytes()
|
state_bytes = 4 * pool.spec("mamba").entry_bytes()
|
||||||
|
|||||||
@@ -83,6 +83,7 @@ def _make_model_runner(
|
|||||||
disaggregation_mode="null",
|
disaggregation_mode="null",
|
||||||
max_running_requests=None,
|
max_running_requests=None,
|
||||||
disaggregation_decode_extra_slots=0,
|
disaggregation_decode_extra_slots=0,
|
||||||
|
enable_unified_memory=False,
|
||||||
kv_lora_rank=512,
|
kv_lora_rank=512,
|
||||||
qk_rope_head_dim=64,
|
qk_rope_head_dim=64,
|
||||||
swa_kv_lora_rank=128,
|
swa_kv_lora_rank=128,
|
||||||
@@ -150,6 +151,7 @@ def _make_model_runner(
|
|||||||
disaggregation_mode=disaggregation_mode,
|
disaggregation_mode=disaggregation_mode,
|
||||||
max_running_requests=max_running_requests,
|
max_running_requests=max_running_requests,
|
||||||
disaggregation_decode_extra_slots=disaggregation_decode_extra_slots,
|
disaggregation_decode_extra_slots=disaggregation_decode_extra_slots,
|
||||||
|
enable_unified_memory=enable_unified_memory,
|
||||||
enable_hisparse=False,
|
enable_hisparse=False,
|
||||||
enable_hierarchical_cache=False,
|
enable_hierarchical_cache=False,
|
||||||
enable_dsa_cache_layer_split=False,
|
enable_dsa_cache_layer_split=False,
|
||||||
@@ -299,7 +301,14 @@ class TestDefaultConfigurator(CustomTestCase):
|
|||||||
class TestHybridSWAConfigurator(CustomTestCase):
|
class TestHybridSWAConfigurator(CustomTestCase):
|
||||||
"""Hybrid SWA: full/swa split, ratio, memory invariant."""
|
"""Hybrid SWA: full/swa split, ratio, memory invariant."""
|
||||||
|
|
||||||
def _make_swa_runner(self, full_layers=16, swa_layers=16, ratio=0.5, page_size=1):
|
def _make_swa_runner(
|
||||||
|
self,
|
||||||
|
full_layers=16,
|
||||||
|
swa_layers=16,
|
||||||
|
ratio=0.5,
|
||||||
|
page_size=1,
|
||||||
|
enable_unified_memory=False,
|
||||||
|
):
|
||||||
return _make_model_runner(
|
return _make_model_runner(
|
||||||
self,
|
self,
|
||||||
is_hybrid_swa=True,
|
is_hybrid_swa=True,
|
||||||
@@ -308,6 +317,7 @@ class TestHybridSWAConfigurator(CustomTestCase):
|
|||||||
swa_num_kv_heads=4,
|
swa_num_kv_heads=4,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
swa_full_tokens_ratio=ratio,
|
swa_full_tokens_ratio=ratio,
|
||||||
|
enable_unified_memory=enable_unified_memory,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _run(self, available_bytes, **kwargs):
|
def _run(self, available_bytes, **kwargs):
|
||||||
@@ -329,6 +339,92 @@ class TestHybridSWAConfigurator(CustomTestCase):
|
|||||||
self.assertLessEqual(used, available)
|
self.assertLessEqual(used, available)
|
||||||
self.assertGreater(used, available * 0.99)
|
self.assertGreater(used, available * 0.99)
|
||||||
|
|
||||||
|
def test_draft_does_not_inherit_target_shared_byte_budget(self):
|
||||||
|
"""A separate draft pool must not allocate the target's byte envelope again."""
|
||||||
|
from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator
|
||||||
|
|
||||||
|
mr, _, config = self._run(1 << 20, enable_unified_memory=True)
|
||||||
|
self.assertIsNotNone(config.unified_memory_pool_bytes)
|
||||||
|
configurator = object.__new__(KVCacheConfigurator)
|
||||||
|
configurator.model_config = mr.model_config
|
||||||
|
configurator.is_hybrid_swa = True
|
||||||
|
configurator.is_draft_worker = False
|
||||||
|
target = configurator._derive_pool_sizes(config=config)
|
||||||
|
configurator.is_draft_worker = True
|
||||||
|
draft = configurator._derive_pool_sizes(config=config)
|
||||||
|
self.assertEqual(
|
||||||
|
target.unified_memory_pool_bytes, config.unified_memory_pool_bytes
|
||||||
|
)
|
||||||
|
self.assertIsNone(draft.unified_memory_pool_bytes)
|
||||||
|
self.assertEqual(
|
||||||
|
draft.full_max_total_num_tokens, config.full_max_total_num_tokens
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
draft.swa_max_total_num_tokens, config.swa_max_total_num_tokens
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_unified_capacity_is_maximal_with_draft_pool(self):
|
||||||
|
page_size = 8
|
||||||
|
full_layers = 2
|
||||||
|
swa_layers = 1
|
||||||
|
draft_layers = 2
|
||||||
|
draft_swa_layers = 1
|
||||||
|
ratio = 0.5
|
||||||
|
mr = _make_model_runner(
|
||||||
|
self,
|
||||||
|
is_hybrid_swa=True,
|
||||||
|
full_attention_layer_ids=list(range(full_layers)),
|
||||||
|
swa_attention_layer_ids=list(range(full_layers, full_layers + swa_layers)),
|
||||||
|
swa_num_kv_heads=4,
|
||||||
|
swa_full_tokens_ratio=ratio,
|
||||||
|
page_size=page_size,
|
||||||
|
enable_unified_memory=True,
|
||||||
|
speculative_algorithm="EAGLE",
|
||||||
|
)
|
||||||
|
mr.spec_algorithm.is_eagle.return_value = True
|
||||||
|
mr.spec_algorithm.is_none.return_value = False
|
||||||
|
mr.spec_aux_config.eagle_draft_num_layers = draft_layers
|
||||||
|
mr.spec_aux_config.eagle_draft_swa_num_layers = draft_swa_layers
|
||||||
|
|
||||||
|
full_bytes_per_token = _full_per_token(mr)
|
||||||
|
swa_bytes_per_token = _swa_per_token(mr)
|
||||||
|
target_full_bytes_per_token = full_bytes_per_token * full_layers
|
||||||
|
draft_bytes_per_token = (
|
||||||
|
full_bytes_per_token * (draft_layers - draft_swa_layers)
|
||||||
|
+ swa_bytes_per_token * draft_swa_layers
|
||||||
|
)
|
||||||
|
|
||||||
|
def allocation_bytes(full_tokens, *, include_reserved_draft_page=True):
|
||||||
|
swa_tokens = int(full_tokens * ratio) // page_size * page_size
|
||||||
|
target_bytes = (
|
||||||
|
full_tokens * target_full_bytes_per_token
|
||||||
|
+ swa_tokens * swa_bytes_per_token * swa_layers
|
||||||
|
)
|
||||||
|
virtual_span = max(target_bytes // target_full_bytes_per_token - 1, 0)
|
||||||
|
draft_tokens = (virtual_span + page_size - 1) // page_size * page_size
|
||||||
|
if include_reserved_draft_page:
|
||||||
|
draft_tokens += page_size
|
||||||
|
return target_bytes + draft_tokens * draft_bytes_per_token
|
||||||
|
|
||||||
|
expected_full_tokens = 10 * page_size
|
||||||
|
available = allocation_bytes(
|
||||||
|
expected_full_tokens + page_size,
|
||||||
|
include_reserved_draft_page=False,
|
||||||
|
)
|
||||||
|
with mock_cpu_env():
|
||||||
|
from sglang.srt.model_executor.pool_configurator import (
|
||||||
|
create_memory_pool_configurator,
|
||||||
|
)
|
||||||
|
|
||||||
|
cfg = create_memory_pool_configurator(mr)
|
||||||
|
config = cfg.calculate_pool_sizes(available, page_size)
|
||||||
|
|
||||||
|
full_tokens = config.full_max_total_num_tokens
|
||||||
|
self.assertEqual(full_tokens % page_size, 0)
|
||||||
|
self.assertEqual(full_tokens, expected_full_tokens)
|
||||||
|
self.assertLessEqual(allocation_bytes(full_tokens), available)
|
||||||
|
self.assertGreater(allocation_bytes(full_tokens + page_size), available)
|
||||||
|
|
||||||
@patch(
|
@patch(
|
||||||
"sglang.srt.mem_cache.kv_cache_configurator.calculate_mla_kv_cache_dim",
|
"sglang.srt.mem_cache.kv_cache_configurator.calculate_mla_kv_cache_dim",
|
||||||
return_value=576,
|
return_value=576,
|
||||||
|
|||||||
Reference in New Issue
Block a user