feat(unified-memory): byte-budget sizing, feasibility floor, and a conservation verifier (#35158)
Co-authored-by: Caihua Li <caihua.li@bytedance.com> Co-authored-by: Claude Fable 5 <noreply@anthropic.com> Co-authored-by: Cheng Wan <cheng.wan@radixark.ai>
This commit is contained in:
co-authored by
Caihua Li
Claude Fable 5
Cheng Wan
parent
961beee9e5
commit
98cb3535b7
@@ -2979,11 +2979,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
return total
|
return total
|
||||||
|
|
||||||
def check_decode_mem(self, selected_indices: Optional[List[int]] = None):
|
def check_decode_mem(self, selected_indices: Optional[List[int]] = None):
|
||||||
"""Reclaim evictable tree-cache entries (shortfall only), then report
|
"""Whether the next decode step fits in the KV pool. The ALLOCATOR owns
|
||||||
whether the next decode step fits in the KV pool."""
|
the capacity gate (eviction + any per-step reservations of its own) —
|
||||||
|
the retract loop converges on this same check, so allocator-side
|
||||||
|
shortfalls retract gracefully instead of tripping fail-loud alloc
|
||||||
|
errors."""
|
||||||
num_tokens = self.new_tokens_required_next_decode(selected_indices)
|
num_tokens = self.new_tokens_required_next_decode(selected_indices)
|
||||||
evict_from_tree_cache(self.tree_cache, num_tokens)
|
return self.token_to_kv_pool_allocator.check_decode_capacity(
|
||||||
return self.token_to_kv_pool_allocator.available_size() >= num_tokens
|
num_tokens=num_tokens, tree_cache=self.tree_cache
|
||||||
|
)
|
||||||
|
|
||||||
def retract_decode(self) -> Tuple[List[Req], float, List[Req]]:
|
def retract_decode(self) -> Tuple[List[Req], float, List[Req]]:
|
||||||
"""Retract the decoding requests when there is not enough memory."""
|
"""Retract the decoding requests when there is not enough memory."""
|
||||||
|
|||||||
@@ -4383,6 +4383,13 @@ class Scheduler(
|
|||||||
if has_leak:
|
if has_leak:
|
||||||
self.invariant_checker._report_leak("pool", "\n".join(messages))
|
self.invariant_checker._report_leak("pool", "\n".join(messages))
|
||||||
self.invariant_checker._check_req_pool()
|
self.invariant_checker._check_req_pool()
|
||||||
|
# Byte-conservation diagnostic (allocator-owned; static pools
|
||||||
|
# return [] — the token identity above can't see byte leaks).
|
||||||
|
byte_violations = self.token_to_kv_pool_allocator.verify_byte_accounting()
|
||||||
|
if byte_violations:
|
||||||
|
self.invariant_checker._report_leak(
|
||||||
|
"pool-bytes", "\n".join(byte_violations)
|
||||||
|
)
|
||||||
|
|
||||||
# tree cache sanity check
|
# tree cache sanity check
|
||||||
self.invariant_checker._check_tree_cache()
|
self.invariant_checker._check_tree_cache()
|
||||||
|
|||||||
@@ -51,6 +51,39 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
|||||||
def size_full(self):
|
def size_full(self):
|
||||||
return self.size
|
return self.size
|
||||||
|
|
||||||
|
# -- scheduler-facing capacity hooks --
|
||||||
|
# The scheduler calls these UNCONDITIONALLY (zero feature branches on its
|
||||||
|
# side); the defaults reproduce the historical token behavior exactly, and
|
||||||
|
# unified composites override them with byte-denominated logic.
|
||||||
|
|
||||||
|
def evict_to_free_tokens(self, tree_cache, num_tokens: int) -> None:
|
||||||
|
"""Ask the prefix cache to evict unlocked entries until this allocator
|
||||||
|
can serve ``num_tokens`` (or nothing evictable remains). Default = the
|
||||||
|
shared token-count eviction; joint-byte composites override (evicting
|
||||||
|
one multi-lifetime tree node frees bytes on several sides at once).
|
||||||
|
"""
|
||||||
|
from sglang.srt.mem_cache.common import evict_from_tree_cache
|
||||||
|
|
||||||
|
evict_from_tree_cache(tree_cache, num_tokens)
|
||||||
|
|
||||||
|
def check_decode_capacity(self, *, num_tokens: int, tree_cache) -> bool:
|
||||||
|
"""Whether the NEXT decode step's ``num_tokens`` allocation fits,
|
||||||
|
evicting reclaimable cache first. The retract loop converges on this
|
||||||
|
same check, so allocator-side shortfalls retract gracefully instead of
|
||||||
|
tripping fail-loud alloc errors. Default reproduces the historical
|
||||||
|
``ScheduleBatch.check_decode_mem`` body; unified composites override
|
||||||
|
with byte gates + per-step reservations of their own.
|
||||||
|
"""
|
||||||
|
self.evict_to_free_tokens(tree_cache, num_tokens)
|
||||||
|
return self.available_size() >= num_tokens
|
||||||
|
|
||||||
|
def verify_byte_accounting(self) -> list:
|
||||||
|
"""Idle-time conservation diagnostic: recompute this allocator's
|
||||||
|
byte/slot accounting and return human-readable violation strings
|
||||||
|
(empty == healthy). Default: static pools have no byte model.
|
||||||
|
"""
|
||||||
|
return []
|
||||||
|
|
||||||
def debug_print(self) -> str:
|
def debug_print(self) -> str:
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
|
|||||||
@@ -230,6 +230,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_total_bytes: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True, kw_only=True)
|
@dataclass(slots=True, kw_only=True)
|
||||||
@@ -390,6 +391,7 @@ class KVCacheConfigurator:
|
|||||||
c128_state_pool_size=c128_state_pool_size,
|
c128_state_pool_size=c128_state_pool_size,
|
||||||
c4_state_dtype=c4_state_dtype,
|
c4_state_dtype=c4_state_dtype,
|
||||||
c128_state_dtype=c128_state_dtype,
|
c128_state_dtype=c128_state_dtype,
|
||||||
|
unified_total_bytes=config.unified_total_bytes,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_pools(
|
def _init_pools(
|
||||||
@@ -419,6 +421,7 @@ class KVCacheConfigurator:
|
|||||||
bundle = self._init_unified_mamba_pools(
|
bundle = self._init_unified_mamba_pools(
|
||||||
max_num_reqs=sizes.max_running_requests,
|
max_num_reqs=sizes.max_running_requests,
|
||||||
max_total_num_tokens=sizes.max_total_num_tokens,
|
max_total_num_tokens=sizes.max_total_num_tokens,
|
||||||
|
unified_total_bytes=sizes.unified_total_bytes,
|
||||||
)
|
)
|
||||||
elif self.is_hybrid_swa and not is_deepseek_v4(self.model_config.hf_config):
|
elif self.is_hybrid_swa and not is_deepseek_v4(self.model_config.hf_config):
|
||||||
if pd_enabled:
|
if pd_enabled:
|
||||||
@@ -432,6 +435,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_total_bytes=sizes.unified_total_bytes,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Fail loud, not silently fall through to the normal pools (which would
|
# Fail loud, not silently fall through to the normal pools (which would
|
||||||
@@ -568,7 +572,11 @@ class KVCacheConfigurator:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _init_unified_mamba_pools(
|
def _init_unified_mamba_pools(
|
||||||
self, *, max_num_reqs: int, max_total_num_tokens: int
|
self,
|
||||||
|
*,
|
||||||
|
max_num_reqs: int,
|
||||||
|
max_total_num_tokens: int,
|
||||||
|
unified_total_bytes: Optional[int] = None,
|
||||||
) -> UnifiedPoolBundle:
|
) -> UnifiedPoolBundle:
|
||||||
"""Build the shared-KV-pool stack for a hybrid-Mamba model:
|
"""Build the shared-KV-pool stack for a hybrid-Mamba model:
|
||||||
one byte buffer split between the full-attn MHA KV pool and the
|
one byte buffer split between the full-attn MHA KV pool and the
|
||||||
@@ -640,6 +648,9 @@ class KVCacheConfigurator:
|
|||||||
forward_stream=self.forward_stream,
|
forward_stream=self.forward_stream,
|
||||||
# Lazy compaction: default ON, env-var escape hatch for rollback / A/B.
|
# Lazy compaction: default ON, 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),
|
||||||
)
|
)
|
||||||
return bundle
|
return bundle
|
||||||
|
|
||||||
@@ -649,6 +660,7 @@ class KVCacheConfigurator:
|
|||||||
max_num_reqs: int,
|
max_num_reqs: int,
|
||||||
full_max_total_num_tokens: Optional[int],
|
full_max_total_num_tokens: Optional[int],
|
||||||
swa_max_total_num_tokens: Optional[int],
|
swa_max_total_num_tokens: Optional[int],
|
||||||
|
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
|
||||||
buffer split between the full-attention and SWA KV pools."""
|
buffer split between the full-attention and SWA KV pools."""
|
||||||
@@ -731,6 +743,14 @@ 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
|
||||||
|
# sliding window term only -- the full-attention side is not
|
||||||
|
# charged, see `_check_bs1_feasibility_floor`.
|
||||||
|
model_context_len=self.model_config.context_len,
|
||||||
|
sliding_window_size=self.model_config.sliding_window_size,
|
||||||
)
|
)
|
||||||
return UnifiedPoolBundle(
|
return UnifiedPoolBundle(
|
||||||
unified_memory_pool=bundle.unified_memory_pool,
|
unified_memory_pool=bundle.unified_memory_pool,
|
||||||
@@ -2034,10 +2054,18 @@ class KVCacheConfigurator:
|
|||||||
config = configurator.calculate_pool_sizes(
|
config = configurator.calculate_pool_sizes(
|
||||||
budget_bytes, get_schedule().page_size
|
budget_bytes, get_schedule().page_size
|
||||||
)
|
)
|
||||||
|
if get_memory().enable_unified_memory:
|
||||||
|
# Floor-align to 4096 B: the factories `.view()` the whole uint8
|
||||||
|
# buffer as the KV/state dtype, so the total must be a dtype-size
|
||||||
|
# multiple and a profiled budget is not. Flooring never overcommits.
|
||||||
|
config.unified_total_bytes = budget_bytes - (budget_bytes % 4096)
|
||||||
max_tokens = self._apply_token_constraints(config.max_total_num_tokens)
|
max_tokens = self._apply_token_constraints(config.max_total_num_tokens)
|
||||||
if cap_tokens is not None:
|
if cap_tokens is not None:
|
||||||
max_tokens = min(max_tokens, cap_tokens)
|
max_tokens = min(max_tokens, cap_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
|
||||||
|
# applies; the recalced config's unified_total_bytes stays None
|
||||||
|
# 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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -345,6 +345,31 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
return self.watermark_physical * self.entry_bytes_per_page
|
return self.watermark_physical * self.entry_bytes_per_page
|
||||||
return self.num_pages * self.entry_bytes_per_page
|
return self.num_pages * self.entry_bytes_per_page
|
||||||
|
|
||||||
|
def _byte_accounting_violations(self) -> List[str]:
|
||||||
|
"""Per-sub-pool conservation strings (empty == healthy): the watermark
|
||||||
|
span must equal live + holes + pending pages, and frontiers must lie
|
||||||
|
inside the buffer. Idle-time diagnostic — pure host arithmetic."""
|
||||||
|
out: List[str] = []
|
||||||
|
total = self.unified_buffer.total_bytes
|
||||||
|
lo_b, hi_b = self._byte_low_frontier(), self._byte_high_frontier()
|
||||||
|
if not (0 <= lo_b <= hi_b <= total):
|
||||||
|
out.append(
|
||||||
|
f"[{self.sub_pool_name}] frontier out of bounds: "
|
||||||
|
f"low={lo_b}, high={hi_b}, total={total}"
|
||||||
|
)
|
||||||
|
if self.lazy_compaction:
|
||||||
|
# Lazy end: the watermark span contains live + holes + pending
|
||||||
|
# (eager has no holes/pending — span == live by construction).
|
||||||
|
holes = int(self._free_phys_pages.numel())
|
||||||
|
pending = len(self._pending_reuse_pages_cpu)
|
||||||
|
wm_span = self._allocated_pages()
|
||||||
|
if wm_span != self.live_page_count + holes + pending:
|
||||||
|
out.append(
|
||||||
|
f"[{self.sub_pool_name}] span {wm_span} != live "
|
||||||
|
f"{self.live_page_count} + holes {holes} + pending {pending}"
|
||||||
|
)
|
||||||
|
return out
|
||||||
|
|
||||||
def _byte_low_frontier(self) -> int:
|
def _byte_low_frontier(self) -> int:
|
||||||
"""Byte starting this side's allocatable range (grow-up) / just below its lowest live page (grow-down)."""
|
"""Byte starting this side's allocatable range (grow-up) / just below its lowest live page (grow-down)."""
|
||||||
if self.grow_direction == "up":
|
if self.grow_direction == "up":
|
||||||
@@ -1750,6 +1775,40 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
self.free(reps, _pages=reps // self.page_size)
|
self.free(reps, _pages=reps // self.page_size)
|
||||||
|
|
||||||
|
|
||||||
|
def _chain_byte_accounting_violations(
|
||||||
|
chain: List[MultiEndedAllocator],
|
||||||
|
) -> List[str]:
|
||||||
|
"""Conservation for an ordered low→high chain of band allocators: each
|
||||||
|
member's own accounting, plus the frontier total order — a member's low
|
||||||
|
frontier must clear the previous member's high frontier, or the bands
|
||||||
|
overlap in the shared byte buffer.
|
||||||
|
|
||||||
|
Today's chains are the 2-pool end pairs; the N-pool track inserts float
|
||||||
|
middles here (and teaches the walk to skip empty/parked ones).
|
||||||
|
"""
|
||||||
|
out: List[str] = []
|
||||||
|
for a in chain:
|
||||||
|
out.extend(a._byte_accounting_violations())
|
||||||
|
frontier = 0
|
||||||
|
for a in chain:
|
||||||
|
lo_b, hi_b = a._byte_low_frontier(), a._byte_high_frontier()
|
||||||
|
if lo_b < frontier:
|
||||||
|
out.append(
|
||||||
|
f"[chain] {a.sub_pool_name} low frontier {lo_b} overlaps the "
|
||||||
|
f"previous pool's high frontier {frontier}"
|
||||||
|
)
|
||||||
|
frontier = max(frontier, hi_b)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def _end_pair_chain(
|
||||||
|
a: MultiEndedAllocator, b: MultiEndedAllocator
|
||||||
|
) -> List[MultiEndedAllocator]:
|
||||||
|
"""Order an end pair low→high by grow direction (the factories and the
|
||||||
|
unit fixtures orient the pair differently; the chain check must not care)."""
|
||||||
|
return sorted((a, b), key=lambda x: x.grow_direction != "up")
|
||||||
|
|
||||||
|
|
||||||
class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||||
"""Composite allocator for the MHA (full-attn) + Mamba hybrid pair.
|
"""Composite allocator for the MHA (full-attn) + Mamba hybrid pair.
|
||||||
|
|
||||||
@@ -2042,6 +2101,11 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
self.full_attn_allocator.clear_inverse_history()
|
self.full_attn_allocator.clear_inverse_history()
|
||||||
self.mamba_allocator.clear_inverse_history()
|
self.mamba_allocator.clear_inverse_history()
|
||||||
|
|
||||||
|
def verify_byte_accounting(self) -> List[str]:
|
||||||
|
return _chain_byte_accounting_violations(
|
||||||
|
_end_pair_chain(self.mamba_allocator, self.full_attn_allocator)
|
||||||
|
)
|
||||||
|
|
||||||
def free_group_begin(self) -> None:
|
def free_group_begin(self) -> None:
|
||||||
super().free_group_begin()
|
super().free_group_begin()
|
||||||
self.free_page_reps_group = []
|
self.free_page_reps_group = []
|
||||||
@@ -2602,6 +2666,11 @@ 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)
|
||||||
|
)
|
||||||
|
|
||||||
def clear(self) -> None:
|
def clear(self) -> None:
|
||||||
self.full_attn_allocator.clear()
|
self.full_attn_allocator.clear()
|
||||||
self.swa_attn_allocator.clear()
|
self.swa_attn_allocator.clear()
|
||||||
|
|||||||
@@ -253,6 +253,30 @@ def _assert_kernel_id_bound(*, sub_pool_name: str, n_rows: int) -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _reserved_floor_bytes(sub_pool_specs: List[SubPoolSpec], page_size: int) -> int:
|
||||||
|
"""Bytes at the bottom of the buffer reserved as the slot-0 padding sink.
|
||||||
|
|
||||||
|
Slot-0 dummy writes for every sub-pool land here; each sub-pool's first
|
||||||
|
allocatable slot is chosen so real data starts past it. For a PAGE-AWARE
|
||||||
|
sub-pool the slot-0 write touches layer blocks spread across the whole
|
||||||
|
page-0 envelope (page_size * entry_bytes), not just one slot envelope --
|
||||||
|
but a mamba sub-pool is page_size=1, so its entry is charged ONCE. Charging
|
||||||
|
a mamba entry per page would reserve page_size * ~100 MB of buffer that the
|
||||||
|
sink never touches.
|
||||||
|
|
||||||
|
Single source of truth: `UnifiedKVPool` reserves exactly this, and the
|
||||||
|
factories' bs=1 feasibility floors charge exactly this.
|
||||||
|
"""
|
||||||
|
return max(
|
||||||
|
[max(s.entry_bytes() for s in sub_pool_specs)]
|
||||||
|
+ [
|
||||||
|
page_size * s.entry_bytes()
|
||||||
|
for s in sub_pool_specs
|
||||||
|
if not isinstance(s, MambaSubPoolSpec) # mamba is page_size=1
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class UnifiedKVPool:
|
class UnifiedKVPool:
|
||||||
"""One physical `uint8` byte buffer shared by 2 sub-pools, each exposing
|
"""One physical `uint8` byte buffer shared by 2 sub-pools, each exposing
|
||||||
per-layer views over its own byte range (contiguous per layer for KV,
|
per-layer views over its own byte range (contiguous per layer for KV,
|
||||||
@@ -324,15 +348,7 @@ class UnifiedKVPool:
|
|||||||
# For a page-aware sub-pool the slot-0 write touches layer blocks spread
|
# For a page-aware sub-pool the slot-0 write touches layer blocks spread
|
||||||
# across the WHOLE page-0 envelope (up to page_size * entry_bytes), not
|
# across the WHOLE page-0 envelope (up to page_size * entry_bytes), not
|
||||||
# just one slot envelope — reserve the max of both.
|
# just one slot envelope — reserve the max of both.
|
||||||
entry_max = max(s.entry_bytes() for s in sub_pool_specs)
|
reserved_floor = _reserved_floor_bytes(sub_pool_specs, page_size)
|
||||||
reserved_floor = max(
|
|
||||||
[entry_max]
|
|
||||||
+ [
|
|
||||||
page_size * s.entry_bytes()
|
|
||||||
for s in sub_pool_specs
|
|
||||||
if not isinstance(s, MambaSubPoolSpec) # mamba is page_size=1
|
|
||||||
]
|
|
||||||
)
|
|
||||||
|
|
||||||
for spec in sub_pool_specs:
|
for spec in sub_pool_specs:
|
||||||
entry_bytes = spec.entry_bytes()
|
entry_bytes = spec.entry_bytes()
|
||||||
@@ -1086,6 +1102,31 @@ class UnifiedPoolBundle(NamedTuple):
|
|||||||
req_to_token_pool: object # UnifiedHybridReqToTokenPool
|
req_to_token_pool: object # UnifiedHybridReqToTokenPool
|
||||||
|
|
||||||
|
|
||||||
|
def _check_bs1_feasibility_floor(
|
||||||
|
*,
|
||||||
|
total_bytes: int,
|
||||||
|
floor_terms: List[Tuple[str, int]],
|
||||||
|
factory: str,
|
||||||
|
) -> None:
|
||||||
|
"""bs=1 feasibility FLOOR — the retract loop's terminal guarantee.
|
||||||
|
|
||||||
|
The scheduler retracts requests until the LAST one fits; if one worst-case
|
||||||
|
request running ALONE does not fit in the buffer, under-sizing is a retract
|
||||||
|
LIVELOCK at runtime, not a perf bug. Fail loud at boot, before any pool
|
||||||
|
construction, with the itemized requirement.
|
||||||
|
"""
|
||||||
|
floor = sum(b for _, b in floor_terms)
|
||||||
|
if total_bytes >= floor:
|
||||||
|
return
|
||||||
|
detail = " + ".join(f"{name}={b}" for name, b in floor_terms)
|
||||||
|
raise RuntimeError(
|
||||||
|
f"[unified-memory-pool] {factory}: byte budget {total_bytes} cannot fit "
|
||||||
|
f"ONE worst-case request (bs=1 floor {floor} = {detail}). A pool this "
|
||||||
|
f"size retract-livelocks at runtime. Raise --mem-fraction-static, lower "
|
||||||
|
f"the model context length, or reduce reserved memory."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def init_unified_mamba_pools(
|
def init_unified_mamba_pools(
|
||||||
*,
|
*,
|
||||||
device: str,
|
device: str,
|
||||||
@@ -1116,6 +1157,7 @@ def init_unified_mamba_pools(
|
|||||||
forward_stream: Optional[torch.cuda.Stream] = None,
|
forward_stream: Optional[torch.cuda.Stream] = None,
|
||||||
lazy_compaction: bool = False,
|
lazy_compaction: bool = False,
|
||||||
decode_pre_alloc_size: int = 0,
|
decode_pre_alloc_size: int = 0,
|
||||||
|
unified_total_bytes: Optional[int] = None,
|
||||||
) -> UnifiedPoolBundle:
|
) -> UnifiedPoolBundle:
|
||||||
"""Build the Mamba-hybrid unified-memory-pool stack."""
|
"""Build the Mamba-hybrid unified-memory-pool stack."""
|
||||||
from sglang.srt.mem_cache.multi_ended_allocator import (
|
from sglang.srt.mem_cache.multi_ended_allocator import (
|
||||||
@@ -1164,10 +1206,30 @@ def init_unified_mamba_pools(
|
|||||||
conv_slice_axis=getattr(cp.shape, "conv_slice_axis", 0),
|
conv_slice_axis=getattr(cp.shape, "conv_slice_axis", 0),
|
||||||
grow_direction="up",
|
grow_direction="up",
|
||||||
)
|
)
|
||||||
|
if unified_total_bytes is not None:
|
||||||
|
# PROFILED byte budget for the token side (captured pre-ratio-floor);
|
||||||
|
# the state pool's bytes ride on top. The token counts stay boot
|
||||||
|
# labels / conserve caps -- the runtime split floats.
|
||||||
|
total_bytes = (
|
||||||
|
unified_total_bytes + max_mamba_cache_size * mamba_spec.entry_bytes()
|
||||||
|
)
|
||||||
|
else:
|
||||||
total_bytes = (
|
total_bytes = (
|
||||||
max_total_num_tokens * full_spec.entry_bytes()
|
max_total_num_tokens * full_spec.entry_bytes()
|
||||||
+ max_mamba_cache_size * mamba_spec.entry_bytes()
|
+ max_mamba_cache_size * mamba_spec.entry_bytes()
|
||||||
)
|
)
|
||||||
|
# bs=1 floor: the state slots one running request locks (1 active + 2 radix
|
||||||
|
# checkpoints, a FLOOR not headroom) + the slot-0 sink. The token side is
|
||||||
|
# not charged -- `TpModelWorker.get_worker_info` already clamps max_req_len
|
||||||
|
# to the pool, so a too-long request is refused at admission, not livelocked.
|
||||||
|
_check_bs1_feasibility_floor(
|
||||||
|
total_bytes=total_bytes,
|
||||||
|
floor_terms=[
|
||||||
|
("bs1_state_slots", 3 * mamba_spec.entry_bytes()),
|
||||||
|
("sink", _reserved_floor_bytes([full_spec, mamba_spec], page_size)),
|
||||||
|
],
|
||||||
|
factory="init_unified_mamba_pools",
|
||||||
|
)
|
||||||
shared_pool = UnifiedKVPool(
|
shared_pool = UnifiedKVPool(
|
||||||
total_bytes=total_bytes,
|
total_bytes=total_bytes,
|
||||||
sub_pool_specs=[full_spec, mamba_spec],
|
sub_pool_specs=[full_spec, mamba_spec],
|
||||||
@@ -1576,6 +1638,9 @@ def init_unified_swa_pools(
|
|||||||
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,
|
||||||
|
sliding_window_size: Optional[int] = None,
|
||||||
) -> UnifiedSWAPoolBundle:
|
) -> UnifiedSWAPoolBundle:
|
||||||
"""Build the SWA-hybrid unified-memory-pool stack."""
|
"""Build the SWA-hybrid unified-memory-pool stack."""
|
||||||
from sglang.srt.mem_cache.multi_ended_allocator import (
|
from sglang.srt.mem_cache.multi_ended_allocator import (
|
||||||
@@ -1612,10 +1677,33 @@ 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:
|
||||||
|
# PROFILED byte budget, sized from directly: the re-sum's floor losses
|
||||||
|
# stay out of the buffer, and the token counts remain boot labels.
|
||||||
|
total_bytes = unified_total_bytes
|
||||||
|
else:
|
||||||
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()
|
||||||
)
|
)
|
||||||
|
if model_context_len is not None:
|
||||||
|
# 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
|
||||||
|
# (max_req_len clamps it); the swa sub-pool is sized independently of
|
||||||
|
# that clamp, which is why the window term stays.
|
||||||
|
swa_bs1_tokens = (
|
||||||
|
min(model_context_len, sliding_window_size + page_size)
|
||||||
|
if sliding_window_size is not None
|
||||||
|
else model_context_len
|
||||||
|
)
|
||||||
|
_check_bs1_feasibility_floor(
|
||||||
|
total_bytes=total_bytes,
|
||||||
|
floor_terms=[
|
||||||
|
("swa_window_kv", swa_bs1_tokens * swa_spec.entry_bytes()),
|
||||||
|
("sink", _reserved_floor_bytes([full_spec, swa_spec], page_size)),
|
||||||
|
],
|
||||||
|
factory="init_unified_swa_pools",
|
||||||
|
)
|
||||||
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],
|
||||||
|
|||||||
@@ -71,6 +71,13 @@ class MemoryPoolConfig:
|
|||||||
|
|
||||||
mem_fraction_static: Optional[float] = None
|
mem_fraction_static: Optional[float] = None
|
||||||
|
|
||||||
|
# Unified pool only: the PROFILED byte budget for the token-granular
|
||||||
|
# sub-pools. Set, the factories size the buffer from it directly instead of
|
||||||
|
# re-summing ratio-derived token counts, which keeps the re-sum's floor
|
||||||
|
# losses out of the buffer; the token counts stay boot labels / conserve
|
||||||
|
# caps. None on the token-capped path -- a user token cap IS the budget.
|
||||||
|
unified_total_bytes: Optional[int] = None
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
if self.max_total_num_tokens <= 0:
|
if self.max_total_num_tokens <= 0:
|
||||||
msg = "Not enough memory. Please try to increase --mem-fraction-static."
|
msg = "Not enough memory. Please try to increase --mem-fraction-static."
|
||||||
|
|||||||
@@ -0,0 +1,157 @@
|
|||||||
|
# Copyright 2023-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.
|
||||||
|
# ==============================================================================
|
||||||
|
"""Byte-conservation verifier for the unified 2-pool composites.
|
||||||
|
|
||||||
|
`verify_byte_accounting` is the idle-time tripwire the token-identity leak
|
||||||
|
check cannot provide: the unified pool's correctness rests on BYTE bookkeeping
|
||||||
|
(watermark spans, holes, pending compaction, frontier ordering inside one
|
||||||
|
shared buffer), and a drifted counter admits requests into memory that is not
|
||||||
|
actually free — silent corruption territory, not a crash.
|
||||||
|
|
||||||
|
Derived properties pinned here:
|
||||||
|
|
||||||
|
* Conservation: on a lazy end pool the watermark span must equal
|
||||||
|
live + holes + pending pages at EVERY point of a healthy lifecycle
|
||||||
|
(alloc, partial free, group free, flush) — not just at rest.
|
||||||
|
* The check is not vacuous: drifting any single term (live count, watermark,
|
||||||
|
a leaked hole) reports loudly, naming the sub-pool.
|
||||||
|
* Chain order: one member's low frontier clearing the other's high frontier
|
||||||
|
is what "two pools share one buffer without overlap" MEANS; the pair check
|
||||||
|
must hold regardless of which member grows up.
|
||||||
|
* The strict escalation env defaults OFF: promoting the diagnostic to a
|
||||||
|
RuntimeError is a validation posture, not the production one.
|
||||||
|
|
||||||
|
python -m pytest test/registered/unit/mem_cache/test_unified_byte_accounting.py -v
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from test_multi_ended_allocator import TestPagedMultiEndedAllocator as _PagedFixture
|
||||||
|
from test_multi_ended_allocator import (
|
||||||
|
TestUnifiedSWATokenToKVPoolAllocator as _SwaFixture,
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache import multi_ended_allocator as mea
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _swa_composite():
|
||||||
|
inst = _SwaFixture([m for m in dir(_SwaFixture) if m.startswith("test_")][0])
|
||||||
|
pool, allocator, kvcache = inst._build()
|
||||||
|
return inst, allocator, kvcache
|
||||||
|
|
||||||
|
|
||||||
|
def _paged_pair(lazy: bool):
|
||||||
|
inst = _PagedFixture([m for m in dir(_PagedFixture) if m.startswith("test_")][0])
|
||||||
|
_pool, full, swa, _fkv, _skv = inst._build()
|
||||||
|
full.lazy_compaction = lazy
|
||||||
|
return full, swa
|
||||||
|
|
||||||
|
|
||||||
|
class TestHealthyLifecycleReportsClean(unittest.TestCase):
|
||||||
|
def test_swa_composite_clean_at_every_step(self):
|
||||||
|
inst, allocator, kvcache = _swa_composite()
|
||||||
|
self.assertEqual(allocator.verify_byte_accounting(), [])
|
||||||
|
v = inst._alloc(allocator, kvcache, 8)
|
||||||
|
self.assertEqual(allocator.verify_byte_accounting(), [])
|
||||||
|
allocator.free_swa(v[:4]) # tombstone half the swa side
|
||||||
|
self.assertEqual(allocator.verify_byte_accounting(), [])
|
||||||
|
inst._free(allocator, kvcache, v)
|
||||||
|
self.assertEqual(allocator.verify_byte_accounting(), [])
|
||||||
|
allocator.clear()
|
||||||
|
self.assertEqual(allocator.verify_byte_accounting(), [])
|
||||||
|
|
||||||
|
def test_lazy_end_pool_clean_through_free_and_flush(self):
|
||||||
|
full, _swa = _paged_pair(lazy=True)
|
||||||
|
self.assertEqual(full._byte_accounting_violations(), [])
|
||||||
|
v = full.alloc(full.page_size * 4)
|
||||||
|
self.assertEqual(full._byte_accounting_violations(), [])
|
||||||
|
full.free(v[: full.page_size * 2]) # lazy: holes, no compaction yet
|
||||||
|
self.assertEqual(full._byte_accounting_violations(), [])
|
||||||
|
full._flush(urgent=True)
|
||||||
|
self.assertEqual(full._byte_accounting_violations(), [])
|
||||||
|
|
||||||
|
|
||||||
|
class TestDriftReportsLoudly(unittest.TestCase):
|
||||||
|
"""Each mutation below models a distinct bookkeeping bug; the verifier
|
||||||
|
must name the drifted sub-pool. Without these, a regression in any single
|
||||||
|
counter passes every other test (the pool still 'works' — it just lies
|
||||||
|
about capacity)."""
|
||||||
|
|
||||||
|
def _lazy_full(self):
|
||||||
|
full, _swa = _paged_pair(lazy=True)
|
||||||
|
v = full.alloc(full.page_size * 4)
|
||||||
|
full.free(v[: full.page_size]) # one hole so all three terms are live
|
||||||
|
self.assertEqual(full._byte_accounting_violations(), [])
|
||||||
|
return full
|
||||||
|
|
||||||
|
def test_drifted_live_count(self):
|
||||||
|
full = self._lazy_full()
|
||||||
|
full.live_page_count += 1
|
||||||
|
self.assertTrue(any("span" in s for s in full._byte_accounting_violations()))
|
||||||
|
|
||||||
|
def test_leaked_hole(self):
|
||||||
|
full = self._lazy_full()
|
||||||
|
full._free_phys_pages = full._free_phys_pages[:-1] # hole vanished
|
||||||
|
self.assertTrue(any("span" in s for s in full._byte_accounting_violations()))
|
||||||
|
|
||||||
|
def test_drifted_watermark(self):
|
||||||
|
full = self._lazy_full()
|
||||||
|
full.watermark_physical += 1
|
||||||
|
self.assertTrue(any("span" in s for s in full._byte_accounting_violations()))
|
||||||
|
|
||||||
|
def test_composite_report_names_the_sub_pool(self):
|
||||||
|
"""Frontier-bounds drift (checked in BOTH lazy and eager modes): push
|
||||||
|
the swa band's watermark outside the buffer."""
|
||||||
|
inst, allocator, kvcache = _swa_composite()
|
||||||
|
inst._alloc(allocator, kvcache, 8)
|
||||||
|
swa = allocator.swa_attn_allocator
|
||||||
|
# grow-down member: low frontier = (wm+1)*bytes; wm == num_pages puts
|
||||||
|
# it past the buffer top.
|
||||||
|
self.assertEqual(swa.grow_direction, "down")
|
||||||
|
swa.watermark_physical = swa.num_pages
|
||||||
|
out = allocator.verify_byte_accounting()
|
||||||
|
self.assertTrue(out and any("[swa]" in s for s in out), out)
|
||||||
|
|
||||||
|
|
||||||
|
class TestChainFrontierOrder(unittest.TestCase):
|
||||||
|
def test_overlapping_frontiers_report(self):
|
||||||
|
"""Both bands hold pages, then the up member's watermark is pushed past
|
||||||
|
the down member's LIVE low frontier: the two bands now claim the same
|
||||||
|
bytes of one buffer. (An empty down band cannot overlap — its low
|
||||||
|
frontier IS the buffer top — so both sides must be populated for the
|
||||||
|
scenario to be a real corruption.)"""
|
||||||
|
full, swa = _paged_pair(lazy=False)
|
||||||
|
chain = mea._end_pair_chain(full, swa)
|
||||||
|
up, down = chain
|
||||||
|
self.assertEqual(up.grow_direction, "up")
|
||||||
|
self.assertIsNotNone(down.alloc(down.page_size * 2)) # down side live
|
||||||
|
self.assertLess(down._byte_low_frontier(), up.unified_buffer.total_bytes)
|
||||||
|
up.watermark_physical = up.num_pages # up band swallows the buffer
|
||||||
|
out = mea._chain_byte_accounting_violations(chain)
|
||||||
|
self.assertTrue(any("overlap" in s for s in out), out)
|
||||||
|
|
||||||
|
def test_pair_order_is_direction_agnostic(self):
|
||||||
|
"""The factories and the unit fixtures orient the pair differently;
|
||||||
|
the check must order by grow direction, not by argument position."""
|
||||||
|
full, swa = _paged_pair(lazy=False)
|
||||||
|
a = mea._end_pair_chain(full, swa)
|
||||||
|
b = mea._end_pair_chain(swa, full)
|
||||||
|
self.assertEqual([x.sub_pool_name for x in a], [x.sub_pool_name for x in b])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,244 @@
|
|||||||
|
# Copyright 2023-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.
|
||||||
|
# ==============================================================================
|
||||||
|
"""Byte-budget buffer sizing for the unified 2-pool factories.
|
||||||
|
|
||||||
|
Derived properties pinned here:
|
||||||
|
|
||||||
|
* Budget honored EXACTLY: with ``unified_total_bytes`` set, the swa pair's
|
||||||
|
buffer is that many bytes (the mamba pair adds the state pool's bytes on
|
||||||
|
top — the budget is captured AFTER the state carve-out). Sizing from the
|
||||||
|
ratio-derived token counts instead re-introduces the configurator's
|
||||||
|
rounding: the swa split floors the budget by the cell size and then
|
||||||
|
page-aligns EACH side's token count, so the re-sum reconstructs less
|
||||||
|
than the profiled budget by up to about one page of tokens per side.
|
||||||
|
* Fallback: without the budget, sizing is the historical token-count re-sum,
|
||||||
|
bit-for-bit.
|
||||||
|
* bs=1 feasibility floor: a budget that cannot fit ONE worst-case request
|
||||||
|
(full KV at max context, plus one SWA window / the state slots a single
|
||||||
|
running request locks) raises at BOOT, before any pool construction —
|
||||||
|
under-sizing is a retract LIVELOCK at runtime, not a perf bug.
|
||||||
|
* The 4096-byte alignment exists because the factories ``.view()`` the whole
|
||||||
|
uint8 buffer as the KV dtype; an unaligned budget must be floored, never
|
||||||
|
rounded up (rounding up overcommits profiled memory).
|
||||||
|
|
||||||
|
python -m pytest test/registered/unit/mem_cache/test_unified_byte_budget_sizing.py -v
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.unified_memory_pool import (
|
||||||
|
MambaSubPoolSpec,
|
||||||
|
MHASubPoolSpec,
|
||||||
|
UnifiedKVPool,
|
||||||
|
_check_bs1_feasibility_floor,
|
||||||
|
_reserved_floor_bytes,
|
||||||
|
init_unified_swa_pools,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
_DEV = "cpu"
|
||||||
|
|
||||||
|
|
||||||
|
def _swa_factory(**over):
|
||||||
|
kw = dict(
|
||||||
|
device=_DEV,
|
||||||
|
kv_cache_dtype=torch.float16,
|
||||||
|
head_num=2,
|
||||||
|
head_dim=8,
|
||||||
|
v_head_dim=8,
|
||||||
|
swa_head_num=2,
|
||||||
|
swa_head_dim=8,
|
||||||
|
swa_v_head_dim=8,
|
||||||
|
page_size=1,
|
||||||
|
start_layer=0,
|
||||||
|
end_layer=4,
|
||||||
|
swa_attention_layer_ids=[1, 3],
|
||||||
|
full_attention_layer_ids=[0, 2],
|
||||||
|
full_max_total_num_tokens=64,
|
||||||
|
swa_max_total_num_tokens=32,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
need_sort=False,
|
||||||
|
)
|
||||||
|
kw.update(over)
|
||||||
|
return init_unified_swa_pools(**kw)
|
||||||
|
|
||||||
|
|
||||||
|
def _entry_bytes():
|
||||||
|
full = MHASubPoolSpec(
|
||||||
|
name="full",
|
||||||
|
layer_num=2,
|
||||||
|
head_num=2,
|
||||||
|
head_dim=8,
|
||||||
|
store_dtype=torch.float16,
|
||||||
|
grow_direction="up",
|
||||||
|
)
|
||||||
|
return full.entry_bytes()
|
||||||
|
|
||||||
|
|
||||||
|
class TestBudgetSizing(unittest.TestCase):
|
||||||
|
def test_swa_factory_honors_the_budget_exactly(self):
|
||||||
|
e = _entry_bytes()
|
||||||
|
budget = 96 * e + 512 # deliberately NOT a token-count multiple
|
||||||
|
bundle = _swa_factory(unified_total_bytes=budget)
|
||||||
|
self.assertEqual(bundle.unified_memory_pool.total_bytes, budget)
|
||||||
|
|
||||||
|
def test_fallback_is_the_token_count_resum(self):
|
||||||
|
e = _entry_bytes()
|
||||||
|
bundle = _swa_factory()
|
||||||
|
self.assertEqual(bundle.unified_memory_pool.total_bytes, (64 + 32) * e)
|
||||||
|
|
||||||
|
def test_budget_beats_resum_on_rounding(self):
|
||||||
|
"""The property that motivates the whole phase: the re-sum cannot
|
||||||
|
represent a budget that is not a whole-token multiple per side, so it
|
||||||
|
strands bytes the buffer could have held."""
|
||||||
|
e = _entry_bytes()
|
||||||
|
budget = (64 + 32) * e + (e - 2) # almost one more entry
|
||||||
|
bundle = _swa_factory(unified_total_bytes=budget)
|
||||||
|
self.assertEqual(bundle.unified_memory_pool.total_bytes, budget)
|
||||||
|
self.assertGreater(budget, (64 + 32) * e)
|
||||||
|
|
||||||
|
|
||||||
|
class TestReservedFloorIsOneSourceOfTruth(unittest.TestCase):
|
||||||
|
"""The bs=1 floor charges the slot-0 sink, and MUST charge exactly what
|
||||||
|
`UnifiedKVPool` actually reserves.
|
||||||
|
|
||||||
|
Regression (GPU eval_434/436, Falcon-H1 boot): the floor hand-copied the
|
||||||
|
formula as `page_size * max(entry_bytes)`, applying the page multiplier to
|
||||||
|
the MAMBA spec. The pool deliberately excludes mamba (it is page_size=1),
|
||||||
|
so with page_size=256 and a ~139 MB state entry the floor over-charged the
|
||||||
|
sink by 256x — ~33 GiB of phantom requirement — and a healthy config
|
||||||
|
failed to boot with 25 GiB of real headroom.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _specs(self, page_size):
|
||||||
|
full = MHASubPoolSpec(
|
||||||
|
name="full",
|
||||||
|
layer_num=2,
|
||||||
|
head_num=2,
|
||||||
|
head_dim=8,
|
||||||
|
store_dtype=torch.float16,
|
||||||
|
grow_direction="down",
|
||||||
|
)
|
||||||
|
# A state entry vastly larger than a KV token entry — the real ratio
|
||||||
|
# (~139 MB vs ~45 KB) is what made the over-charge fatal.
|
||||||
|
mamba = MambaSubPoolSpec(
|
||||||
|
name="mamba",
|
||||||
|
layer_num=2,
|
||||||
|
conv_state_shapes=((4, 256),),
|
||||||
|
conv_dtype=torch.float16,
|
||||||
|
temporal_state_shape=(4, 256, 64),
|
||||||
|
temporal_dtype=torch.float16,
|
||||||
|
grow_direction="up",
|
||||||
|
)
|
||||||
|
return full, mamba
|
||||||
|
|
||||||
|
def test_mamba_entry_is_not_multiplied_by_page_size(self):
|
||||||
|
full, mamba = self._specs(page_size=256)
|
||||||
|
got = _reserved_floor_bytes([full, mamba], 256)
|
||||||
|
self.assertEqual(got, max(mamba.entry_bytes(), 256 * full.entry_bytes()))
|
||||||
|
self.assertLess(got, 256 * mamba.entry_bytes()) # the bug's value
|
||||||
|
|
||||||
|
def test_floor_sink_equals_what_the_pool_reserves(self):
|
||||||
|
"""Pin the two against each other so the formula cannot drift again."""
|
||||||
|
for page_size in (1, 4, 256):
|
||||||
|
with self.subTest(page_size=page_size):
|
||||||
|
full, mamba = self._specs(page_size)
|
||||||
|
floor = _reserved_floor_bytes([full, mamba], page_size)
|
||||||
|
pool = UnifiedKVPool(
|
||||||
|
total_bytes=floor + 64 * mamba.entry_bytes(),
|
||||||
|
sub_pool_specs=[full, mamba],
|
||||||
|
device=_DEV,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
page_size=page_size,
|
||||||
|
)
|
||||||
|
# min_slot_index is ceil(reserved_floor / entry_bytes) per side.
|
||||||
|
for spec in (full, mamba):
|
||||||
|
self.assertEqual(
|
||||||
|
pool.min_slot_index(spec.name),
|
||||||
|
-(-floor // spec.entry_bytes()),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestBs1FeasibilityFloor(unittest.TestCase):
|
||||||
|
def test_infeasible_budget_raises_before_construction(self):
|
||||||
|
"""The buffer cannot hold one sliding window plus the sink, so boot
|
||||||
|
must fail loud instead of livelocking later."""
|
||||||
|
with self.assertRaises(RuntimeError) as ctx:
|
||||||
|
_swa_factory(
|
||||||
|
unified_total_bytes=8 * _entry_bytes(),
|
||||||
|
model_context_len=4096,
|
||||||
|
sliding_window_size=4096,
|
||||||
|
)
|
||||||
|
self.assertIn("bs=1 floor", str(ctx.exception))
|
||||||
|
self.assertIn("swa_window_kv", str(ctx.exception))
|
||||||
|
|
||||||
|
def test_context_longer_than_the_pool_is_not_rejected(self):
|
||||||
|
"""REGRESSION: the floor must NOT charge the full-attention token side.
|
||||||
|
`TpModelWorker.get_worker_info` clamps max_req_len to the pool, so a
|
||||||
|
context far larger than the buffer is refused at admission, not a
|
||||||
|
livelock -- and it is an ordinary way to serve a long-context model on
|
||||||
|
one GPU. Charging it here made such configs fail at boot."""
|
||||||
|
e = _entry_bytes()
|
||||||
|
bundle = _swa_factory(
|
||||||
|
unified_total_bytes=200 * e,
|
||||||
|
model_context_len=1_000_000, # far beyond what the buffer holds
|
||||||
|
sliding_window_size=16,
|
||||||
|
)
|
||||||
|
self.assertEqual(bundle.unified_memory_pool.total_bytes, 200 * e)
|
||||||
|
|
||||||
|
def test_feasible_config_boots_with_floor_inputs_present(self):
|
||||||
|
e = _entry_bytes()
|
||||||
|
bundle = _swa_factory(
|
||||||
|
unified_total_bytes=200 * e,
|
||||||
|
model_context_len=64,
|
||||||
|
sliding_window_size=16,
|
||||||
|
)
|
||||||
|
self.assertEqual(bundle.unified_memory_pool.total_bytes, 200 * e)
|
||||||
|
|
||||||
|
def test_window_term_is_clamped_to_context(self):
|
||||||
|
"""A window larger than the context must charge at most the context —
|
||||||
|
otherwise short-context models over-raise."""
|
||||||
|
e = _entry_bytes()
|
||||||
|
bundle = _swa_factory(
|
||||||
|
unified_total_bytes=200 * e,
|
||||||
|
model_context_len=64,
|
||||||
|
sliding_window_size=10_000, # window >> context
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(bundle)
|
||||||
|
|
||||||
|
def test_floor_message_itemizes_terms(self):
|
||||||
|
with self.assertRaises(RuntimeError) as ctx:
|
||||||
|
_check_bs1_feasibility_floor(
|
||||||
|
total_bytes=10,
|
||||||
|
floor_terms=[("a", 8), ("b", 8)],
|
||||||
|
factory="test",
|
||||||
|
)
|
||||||
|
msg = str(ctx.exception)
|
||||||
|
self.assertIn("a=8", msg)
|
||||||
|
self.assertIn("b=8", msg)
|
||||||
|
self.assertIn("16", msg)
|
||||||
|
|
||||||
|
def test_exact_floor_passes(self):
|
||||||
|
"""Boundary: total == floor must NOT raise (>= is the contract)."""
|
||||||
|
_check_bs1_feasibility_floor(
|
||||||
|
total_bytes=16, floor_terms=[("a", 8), ("b", 8)], factory="test"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user