Support unified memory page-envelope transfers in PD (#39477)
Co-authored-by: yhzhuang <yhzhuang@fb.com> Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com> Co-authored-by: Yonghao Zhuang <yhzhuang@users.noreply.github.com> Co-authored-by: Cheng Wan <cheng.wan@radixark.ai>
This commit is contained in:
co-authored by
yhzhuang
Lianmin Zheng
Yonghao Zhuang
Cheng Wan
parent
d0730a0e8b
commit
5931fd60ee
@@ -88,10 +88,8 @@ class Memory(msgspec.Struct):
|
||||
"Replace the statically-partitioned hybrid-model pools (full-attn KV + "
|
||||
"SWA/Mamba state) with one byte buffer split dynamically between "
|
||||
"sub-pools. Requires the Triton attention / linear-attn / Mamba "
|
||||
"backends. PD disaggregation is supported over mooncake at equal "
|
||||
"attention TP with pp=1; not yet compatible with hierarchical / "
|
||||
"host-tiered KV cache, prefill cuda-graph capture, or speculative "
|
||||
"decoding other than DSPARK.",
|
||||
"backends. Supported PD-disaggregation and speculative-decoding "
|
||||
"configurations are validated at startup.",
|
||||
] = False
|
||||
enable_session_radix_cache: A[
|
||||
bool,
|
||||
|
||||
@@ -437,11 +437,12 @@ def handle_unified_memory_pool(server_args: Any) -> None:
|
||||
if not cfg.enable_unified_memory:
|
||||
return
|
||||
if cfg.disaggregation_mode != "null":
|
||||
# Constraints of the whole-envelope transfer; see
|
||||
# UnifiedMLATokenToKVPool.get_contiguous_buf_infos.
|
||||
assert cfg.disaggregation_transfer_backend == "mooncake", (
|
||||
"--enable-unified-memory with PD disaggregation supports only "
|
||||
"the mooncake transfer backend; got "
|
||||
# Constraints of the whole-envelope transfer; see the unified MHA and
|
||||
# MLA pool get_contiguous_buf_infos implementations.
|
||||
supported_backends = server_args._unified_memory_pd_transfer_backends()
|
||||
assert cfg.disaggregation_transfer_backend in supported_backends, (
|
||||
"--enable-unified-memory with PD disaggregation supports only these "
|
||||
f"transfer backends: {', '.join(sorted(supported_backends))}; got "
|
||||
f"{cfg.disaggregation_transfer_backend!r}."
|
||||
)
|
||||
assert cfg.pp_size == 1, (
|
||||
@@ -449,6 +450,13 @@ def handle_unified_memory_pool(server_args: Any) -> None:
|
||||
"pipeline parallelism (whole-envelope transfer has no per-layer "
|
||||
"entries to subset)."
|
||||
)
|
||||
assert not (
|
||||
cfg.disaggregation_transfer_backend == "mooncake"
|
||||
and cfg.speculative_algorithm is not None
|
||||
), (
|
||||
"--enable-unified-memory with PD disaggregation does not support "
|
||||
"speculative decoding with the Mooncake transfer backend."
|
||||
)
|
||||
assert not envs.SGLANG_DISABLE_LAZY_COMPACTION.get(), (
|
||||
"--enable-unified-memory with PD disaggregation requires lazy "
|
||||
"compaction; unset SGLANG_DISABLE_LAZY_COMPACTION."
|
||||
@@ -459,6 +467,15 @@ def handle_unified_memory_pool(server_args: Any) -> None:
|
||||
"ships host/C4 rows straight from the allocator, bypassing the "
|
||||
"virtual->physical translation the unified pool needs."
|
||||
)
|
||||
assert cfg.disaggregation_decode_retraction_backup != "host_pool", (
|
||||
"--enable-unified-memory with PD disaggregation does not support "
|
||||
"--disaggregation-decode-retraction-backup=host_pool; use "
|
||||
"cpu_tensor (the automatic default for unified pools)."
|
||||
)
|
||||
assert not cfg.disaggregation_decode_enable_offload_kvcache, (
|
||||
"--enable-unified-memory with PD disaggregation does not yet support "
|
||||
"--disaggregation-decode-enable-offload-kvcache."
|
||||
)
|
||||
assert cfg.speculative_algorithm in (None, "DSPARK"), (
|
||||
"--enable-unified-memory only supports --speculative-algorithm "
|
||||
"DSPARK (chain draft); other speculative algorithms are not yet "
|
||||
|
||||
@@ -455,6 +455,28 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
and hasattr(self.token_to_kv_pool_allocator, "alloc_extend_swa_tail")
|
||||
)
|
||||
|
||||
def _uses_swa_reservation(self) -> bool:
|
||||
return (
|
||||
self._uses_swa_tail_prealloc()
|
||||
or self.token_to_kv_pool_allocator.prealloc_fits_assumes_reclaim()
|
||||
)
|
||||
|
||||
def _prealloc_reservation_fits(
|
||||
self,
|
||||
full_tokens: int,
|
||||
swa_tokens: int,
|
||||
*,
|
||||
full_allocatable_tokens: int,
|
||||
swa_allocatable_tokens: Optional[int],
|
||||
) -> bool:
|
||||
return self.token_to_kv_pool_allocator.prealloc_fits(
|
||||
self.tree_cache,
|
||||
full_tokens,
|
||||
swa_tokens,
|
||||
full_budget_tokens=full_allocatable_tokens,
|
||||
swa_budget_tokens=swa_allocatable_tokens,
|
||||
)
|
||||
|
||||
def _release_matched_prefix_lock(self, req: Req) -> None:
|
||||
if req.swa_prefix_lock_released:
|
||||
self.tree_cache.dec_lock_ref(req.last_node, req.lock_receipt, skip_swa=True)
|
||||
@@ -463,24 +485,16 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
self.tree_cache.dec_lock_ref(req.last_node, req.lock_receipt)
|
||||
|
||||
def _reclaim_swa_tail_capacity(
|
||||
self, swa_tail_len: int, req_id: str
|
||||
self, swa_tail_len: int, req_id: str, *, full_len: int = 0
|
||||
) -> Optional[str]:
|
||||
page_size = self.token_to_kv_pool_allocator.page_size
|
||||
required = ceil_align(swa_tail_len, page_size)
|
||||
available = self.token_to_kv_pool_allocator.swa_available_size()
|
||||
if available < required:
|
||||
self.tree_cache.evict_for_alloc(
|
||||
EvictParams(swa_num_tokens=required - available)
|
||||
)
|
||||
available = self.token_to_kv_pool_allocator.swa_available_size()
|
||||
|
||||
if available < required:
|
||||
return (
|
||||
f"SWA eviction insufficient: needed={required}, "
|
||||
f"available={available}, req={req_id}"
|
||||
)
|
||||
|
||||
return None
|
||||
allocator = self.token_to_kv_pool_allocator
|
||||
page_size = allocator.page_size
|
||||
shortfall = allocator.reclaim_for_prealloc(
|
||||
self.tree_cache,
|
||||
ceil_align(full_len, page_size),
|
||||
ceil_align(swa_tail_len, page_size),
|
||||
)
|
||||
return None if shortfall is None else f"{shortfall}, req={req_id}"
|
||||
|
||||
# SWA caches expose full-attention accounting through full_* accessors.
|
||||
def _radix_full_evictable(self) -> int:
|
||||
@@ -581,8 +595,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
kv_data_mem_kinds += ["VRAM"] * len(device_kv_data_ptrs[c4_layer_num:])
|
||||
num_draft_entries = 0
|
||||
if self.draft_token_to_kv_pool is not None:
|
||||
# We should also transfer draft model kv cache. The indices are
|
||||
# always shared with a target model.
|
||||
# Draft KV shares target virtual ids. Unified target KV is transferred
|
||||
# with physical ids, so it needs a separate draft index vector.
|
||||
draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = (
|
||||
self.draft_token_to_kv_pool.get_contiguous_buf_infos()
|
||||
)
|
||||
@@ -814,30 +828,40 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
return len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0)
|
||||
|
||||
def _check_if_req_exceed_kv_capacity(self, req: Req) -> bool:
|
||||
# HiSparse admits up to the host-backed logical capacity.
|
||||
if self.scheduler.enable_hisparse:
|
||||
capacity = self.scheduler.tp_worker.model_runner.max_token_pool_size
|
||||
else:
|
||||
capacity = self.max_total_num_tokens
|
||||
input_len = self._rebootstrap_prefill_len(req)
|
||||
if input_len > capacity:
|
||||
message = f"Request {req.rid} exceeds the maximum number of tokens: {input_len} > {capacity}"
|
||||
message = None
|
||||
allocator = self.token_to_kv_pool_allocator
|
||||
full_required, swa_required = self._prealloc_required_tokens(req)
|
||||
if not self._uses_swa_tail_prealloc():
|
||||
swa_required = full_required
|
||||
ceiling_fits = allocator.prealloc_ceiling_fits(full_required, swa_required)
|
||||
if ceiling_fits is False:
|
||||
message = (
|
||||
f"Request {req.rid} exceeds the unified FULL/SWA KV byte "
|
||||
f"budget: full={full_required}, swa={swa_required}"
|
||||
)
|
||||
elif ceiling_fits is None:
|
||||
# HiSparse admits up to the host-backed logical capacity.
|
||||
capacity = (
|
||||
self.scheduler.tp_worker.model_runner.max_token_pool_size
|
||||
if self.scheduler.enable_hisparse
|
||||
else self.max_total_num_tokens
|
||||
)
|
||||
input_len = self._rebootstrap_prefill_len(req)
|
||||
if input_len > capacity:
|
||||
message = f"Request {req.rid} exceeds the maximum number of tokens: {input_len} > {capacity}"
|
||||
elif self._uses_swa_tail_prealloc():
|
||||
_, swa_required = self._prealloc_required_tokens(req)
|
||||
swa_capacity = self.token_to_kv_pool_allocator.size_swa
|
||||
if swa_required > swa_capacity:
|
||||
message = (
|
||||
f"Request {req.rid} requires too many SWA KV tokens for "
|
||||
f"decode preallocation: {swa_required} > {swa_capacity}"
|
||||
)
|
||||
if message is not None:
|
||||
logger.error(message)
|
||||
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
||||
self.scheduler.output_streamer.stream_output([req], req.return_logprob)
|
||||
return True
|
||||
if self._uses_swa_tail_prealloc():
|
||||
_, swa_required = self._prealloc_required_tokens(req)
|
||||
swa_capacity = self.token_to_kv_pool_allocator.size_swa
|
||||
if swa_required > swa_capacity:
|
||||
message = (
|
||||
f"Request {req.rid} requires too many SWA KV tokens for "
|
||||
f"decode preallocation: {swa_required} > {swa_capacity}"
|
||||
)
|
||||
logger.error(message)
|
||||
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
||||
self.scheduler.output_streamer.stream_output([req], req.return_logprob)
|
||||
return True
|
||||
return False
|
||||
|
||||
def extend(self, reqs: List[Req], is_retracted: bool = False) -> None:
|
||||
@@ -870,8 +894,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
# allocate memory
|
||||
resumed_reqs = []
|
||||
indices_to_remove = set()
|
||||
uses_swa_tail_prealloc = self._uses_swa_tail_prealloc()
|
||||
if uses_swa_tail_prealloc:
|
||||
swa_allocatable_tokens = None
|
||||
if self._uses_swa_reservation():
|
||||
full_allocatable_tokens, swa_allocatable_tokens = (
|
||||
self._swa_aware_allocatable_token_budgets(count_retracted=False)
|
||||
)
|
||||
@@ -888,17 +912,31 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
break
|
||||
|
||||
full_required, swa_required = self._prealloc_required_tokens(req)
|
||||
if full_required > full_allocatable_tokens:
|
||||
break
|
||||
if uses_swa_tail_prealloc and swa_required > swa_allocatable_tokens:
|
||||
if not self._prealloc_reservation_fits(
|
||||
full_required,
|
||||
swa_required,
|
||||
full_allocatable_tokens=full_allocatable_tokens,
|
||||
swa_allocatable_tokens=swa_allocatable_tokens,
|
||||
):
|
||||
break
|
||||
|
||||
if self.token_to_kv_pool_allocator.prealloc_fits_assumes_reclaim():
|
||||
full_len, swa_len = self._prealloc_kv_lens(req)
|
||||
if (
|
||||
self._reclaim_swa_tail_capacity(swa_len, req.rid, full_len=full_len)
|
||||
is not None
|
||||
):
|
||||
break
|
||||
|
||||
resumed_reqs.append(req)
|
||||
indices_to_remove.add(i)
|
||||
req.is_retracted = False
|
||||
self._pre_alloc(req)
|
||||
full_allocatable_tokens -= full_required
|
||||
if uses_swa_tail_prealloc:
|
||||
full_allocatable_tokens = self._allocatable_token_budgets(
|
||||
count_retracted=False,
|
||||
extra_reserved_reqs=len(resumed_reqs),
|
||||
)
|
||||
if swa_allocatable_tokens is not None:
|
||||
swa_allocatable_tokens = self._swa_tail_allocatable_token_budget(
|
||||
count_retracted=False,
|
||||
extra_reserved_reqs=len(resumed_reqs),
|
||||
@@ -1159,8 +1197,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
)
|
||||
|
||||
uses_swa_tail_prealloc = self._uses_swa_tail_prealloc()
|
||||
swa_allocatable_tokens = 0
|
||||
if uses_swa_tail_prealloc:
|
||||
swa_allocatable_tokens = None
|
||||
if self._uses_swa_reservation():
|
||||
retractable_swa_tokens = sum(
|
||||
self._swa_retractable_len(r) for r in self.scheduler.running_batch.reqs
|
||||
)
|
||||
@@ -1352,27 +1390,18 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
required_alloc_tokens + self.num_reserved_decode_tokens
|
||||
)
|
||||
|
||||
if (
|
||||
max(
|
||||
required_tokens_for_request,
|
||||
origin_input_len
|
||||
- prefix_len
|
||||
+ min(
|
||||
decode_req.req.sampling_params.max_new_tokens,
|
||||
CLIP_MAX_NEW_TOKEN,
|
||||
)
|
||||
- retractable_tokens,
|
||||
full_required_for_admission = max(
|
||||
required_tokens_for_request,
|
||||
origin_input_len
|
||||
- prefix_len
|
||||
+ min(
|
||||
decode_req.req.sampling_params.max_new_tokens,
|
||||
CLIP_MAX_NEW_TOKEN,
|
||||
)
|
||||
> full_allocatable_tokens
|
||||
):
|
||||
if prefix_match is not None and prefix_match.l1_prefix_len > 0:
|
||||
self._release_matched_prefix_lock(decode_req.req)
|
||||
break
|
||||
if required_tokens_for_request > full_allocatable_tokens:
|
||||
if prefix_match is not None and prefix_match.l1_prefix_len > 0:
|
||||
self._release_matched_prefix_lock(decode_req.req)
|
||||
break
|
||||
|
||||
- retractable_tokens,
|
||||
)
|
||||
swa_required_for_admission = 0
|
||||
swa_len = required_alloc_tokens
|
||||
if uses_swa_tail_prealloc:
|
||||
_, swa_required = self._prealloc_required_tokens(decode_req.req)
|
||||
_, swa_len = self._prealloc_kv_lens(decode_req.req)
|
||||
@@ -1380,19 +1409,28 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
decode_req.req.sampling_params.max_new_tokens,
|
||||
CLIP_MAX_NEW_TOKEN,
|
||||
)
|
||||
if (
|
||||
max(
|
||||
swa_required,
|
||||
swa_len + max_new_tokens - retractable_swa_tokens,
|
||||
)
|
||||
> swa_allocatable_tokens
|
||||
):
|
||||
if prefix_match is not None and prefix_match.l1_prefix_len > 0:
|
||||
self._release_matched_prefix_lock(decode_req.req)
|
||||
break
|
||||
swa_required_for_admission = max(
|
||||
swa_required,
|
||||
swa_len + max_new_tokens - retractable_swa_tokens,
|
||||
)
|
||||
elif swa_allocatable_tokens is not None:
|
||||
swa_required_for_admission = full_required_for_admission
|
||||
|
||||
if not self._prealloc_reservation_fits(
|
||||
full_required_for_admission,
|
||||
swa_required_for_admission,
|
||||
full_allocatable_tokens=full_allocatable_tokens,
|
||||
swa_allocatable_tokens=swa_allocatable_tokens,
|
||||
):
|
||||
if prefix_match is not None and prefix_match.l1_prefix_len > 0:
|
||||
self._release_matched_prefix_lock(decode_req.req)
|
||||
break
|
||||
|
||||
if swa_allocatable_tokens is not None:
|
||||
reclaim_error = self._reclaim_swa_tail_capacity(
|
||||
swa_len, decode_req.req.rid
|
||||
swa_len,
|
||||
decode_req.req.rid,
|
||||
full_len=required_alloc_tokens,
|
||||
)
|
||||
if reclaim_error is not None:
|
||||
if prefix_match is not None and prefix_match.l1_prefix_len > 0:
|
||||
@@ -1431,7 +1469,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
extra_reserved_reqs=len(preallocated_reqs) + 1,
|
||||
hicache_reserved_tokens=reserved_restore_tokens,
|
||||
)
|
||||
if uses_swa_tail_prealloc:
|
||||
if swa_allocatable_tokens is not None:
|
||||
swa_allocatable_tokens = self._swa_tail_allocatable_token_budget(
|
||||
retractable_tokens=retractable_tokens,
|
||||
retractable_swa_tokens=retractable_swa_tokens,
|
||||
@@ -1442,6 +1480,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
|
||||
page_size = self.token_to_kv_pool_allocator.page_size
|
||||
kv_transfer_page_size = page_size
|
||||
raw_kv_indices = self.req_to_token_pool.req_to_token[
|
||||
decode_req.req.kv.req_pool_idx
|
||||
][total_prefix_len:origin_input_len]
|
||||
if self.scheduler.enable_hisparse:
|
||||
# Direct-to-host sends host/C4 rows; keep allocator.page_size
|
||||
# logical and use the compressed page size only for these indices.
|
||||
@@ -1453,12 +1494,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
kv_indices = dst_kv_indices[: origin_input_len - prefix_len]
|
||||
else:
|
||||
# Only send delta indices (beyond prefix) to prefill.
|
||||
kv_indices = self.req_to_token_pool.req_to_token[
|
||||
decode_req.req.kv.req_pool_idx
|
||||
][total_prefix_len:origin_input_len]
|
||||
kv_indices = (
|
||||
self.token_to_kv_pool_allocator.translate_kv_indices_for_transfer(
|
||||
kv_indices
|
||||
raw_kv_indices
|
||||
)
|
||||
)
|
||||
|
||||
@@ -1735,7 +1773,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
# HiSparse pre-alloc only allocates logical indices, so the
|
||||
# logical pool is the binding constraint for admission control.
|
||||
available_size = logical_allocator.available_size()
|
||||
elif self._uses_swa_tail_prealloc():
|
||||
elif self._uses_swa_reservation():
|
||||
available_size = self.token_to_kv_pool_allocator.full_available_size()
|
||||
if get_disagg().disaggregation_decode_enable_radix_cache:
|
||||
available_size += self._radix_full_evictable()
|
||||
@@ -1802,8 +1840,10 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
# pool) over-reserves SWA in steady state. Cap by the actual
|
||||
# remaining headroom up to per-req window cap.
|
||||
window_size = self.scheduler.sliding_window_size or 0
|
||||
swa_total = self.token_to_kv_pool_allocator.size_swa
|
||||
swa_available = self.token_to_kv_pool_allocator.swa_available_size()
|
||||
allocator = self.token_to_kv_pool_allocator
|
||||
_, (swa_total, swa_available) = allocator.swa_capacity_and_available(
|
||||
full_capacity=allocator.size_full, swa_capacity=allocator.size_swa
|
||||
)
|
||||
# Per-request SWA ring: cached prefixes still report swa_evictable, but
|
||||
# evicting them frees no ring space.
|
||||
swa_evictable = (
|
||||
|
||||
@@ -790,11 +790,18 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
# Published layer IDs give exact pairing; plain-MHA peers publish none
|
||||
# and keep positional slicing.
|
||||
has_layer_ids = bool(src_layer_ids or dst_layer_ids)
|
||||
# Unified SWA publishes one page-envelope region even on an MHA backend.
|
||||
is_single_region_swa = (
|
||||
state_type == StateType.SWA
|
||||
and len(src_data_ptrs) == 1
|
||||
and len(dst_data_ptrs) == 1
|
||||
)
|
||||
if (
|
||||
self.is_mla_backend
|
||||
or self.is_hybrid_mla_backend
|
||||
or force_flat
|
||||
or has_layer_ids
|
||||
or is_single_region_swa
|
||||
):
|
||||
# Layer IDs map PP-local buffers to global decode entries.
|
||||
# Registrations without them retain the existing PP mapping.
|
||||
@@ -1060,12 +1067,6 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
||||
prefill_data_indices=prefill_kv_indices,
|
||||
dst_data_indices=dst_kv_indices,
|
||||
executor=executor,
|
||||
# The unified pool registers ONE region holding every layer's K and
|
||||
# V inside each page envelope. The MHA branch would half-split that
|
||||
# single region into K and V halves and compute num_kv_layers = 0,
|
||||
# transferring nothing at all; the flat branch addresses the region
|
||||
# as-is. MLA-unified already reaches the flat branch via
|
||||
# is_mla_backend, so this only adds the MHA-unified peer.
|
||||
force_flat=get_memory().enable_unified_memory,
|
||||
src_layer_ids=self.kv_args.kv_layer_ids,
|
||||
dst_layer_ids=dst_layer_ids,
|
||||
|
||||
@@ -271,8 +271,8 @@ class PrefillBootstrapQueue:
|
||||
)
|
||||
num_draft_entries = 0
|
||||
if draft_kv_pool is not None:
|
||||
# We should also transfer draft model kv cache. The indices are
|
||||
# always shared with a target model.
|
||||
# Draft KV shares target virtual ids. Unified target KV is transferred
|
||||
# with physical ids, so it needs a separate draft index vector.
|
||||
draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = (
|
||||
draft_kv_pool.get_contiguous_buf_infos()
|
||||
)
|
||||
@@ -1463,14 +1463,14 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
|
||||
for seg_start, seg_end in segments:
|
||||
is_final_segment = seg_end == end_idx
|
||||
kv_indices = self.req_to_token_pool.req_to_token[
|
||||
raw_kv_indices = self.req_to_token_pool.req_to_token[
|
||||
req.kv.req_pool_idx, seg_start:seg_end
|
||||
]
|
||||
# Unified memory: req_to_token holds VIRTUAL ids; the transfer needs
|
||||
# physical ones. Per segment, since each is its own gather.
|
||||
kv_indices = (
|
||||
self.token_to_kv_pool_allocator.translate_kv_indices_for_transfer(
|
||||
kv_indices
|
||||
raw_kv_indices
|
||||
)
|
||||
)
|
||||
page_indices = kv_to_page_indices(kv_indices, page_size)
|
||||
@@ -1479,9 +1479,10 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
len(page_indices), segment_is_last
|
||||
):
|
||||
continue
|
||||
send_state_indices = state_indices if segment_is_last else None
|
||||
req.disagg_kv_sender.send(
|
||||
page_indices,
|
||||
state_indices if segment_is_last else None,
|
||||
send_state_indices,
|
||||
num_kv_tokens=seg_end - seg_start,
|
||||
)
|
||||
req.start_send_idx = end_idx
|
||||
|
||||
@@ -95,6 +95,43 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
|
||||
0, min(max_new_tokens, token_capacity - paged_input - self.page_size - 1)
|
||||
)
|
||||
|
||||
def prealloc_fits_assumes_reclaim(self) -> bool:
|
||||
"""Whether `prealloc_fits` answers about the state reachable AFTER
|
||||
reclaiming the evictable pages, so admitting on it still owes the
|
||||
reclaim. False when the answer describes the pool as it stands.
|
||||
"""
|
||||
return False
|
||||
|
||||
def prealloc_ceiling_fits(self, full_tokens: int, swa_tokens: int) -> bool | None:
|
||||
"""Whether a demand this size could EVER be preallocated, or None when
|
||||
this pool has no ceiling of its own and the caller's token capacity is
|
||||
the only bound.
|
||||
"""
|
||||
return None
|
||||
|
||||
def prealloc_fits(
|
||||
self,
|
||||
tree_cache,
|
||||
full_tokens: int,
|
||||
swa_tokens: int,
|
||||
*,
|
||||
full_budget_tokens: int,
|
||||
swa_budget_tokens: int | None = None,
|
||||
) -> bool:
|
||||
"""Whether a decode-node preallocation of this size fits.
|
||||
|
||||
The budgets are the scheduler's policy: what each side has left once
|
||||
decode headroom and retraction are reserved. Separate buffers make the
|
||||
two sides independent, so each is checked against its own budget and
|
||||
``tree_cache`` is never read -- what it could reclaim is already
|
||||
inside that budget. A pool that cuts both sides from one buffer
|
||||
overrides this to price them together, since a per-side token budget
|
||||
cannot express a shared byte envelope.
|
||||
"""
|
||||
return full_tokens <= full_budget_tokens and (
|
||||
swa_budget_tokens is None or swa_tokens <= swa_budget_tokens
|
||||
)
|
||||
|
||||
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.
|
||||
|
||||
@@ -349,10 +349,6 @@ class DeepSeekV4HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
def translate_swa_indices_for_transfer(
|
||||
self, kv_indices: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
# Delegated like the read-path translate above: this composite is not a
|
||||
# SWA allocator itself, so it inherits neither the default nor an
|
||||
# override, and the PD payload path calls this on whatever allocator
|
||||
# the scheduler holds.
|
||||
return self.logical_attn_allocator.translate_swa_indices_for_transfer(
|
||||
kv_indices
|
||||
)
|
||||
|
||||
@@ -166,6 +166,30 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
self, tree_cache, num_mixed_decode_tokens=num_mixed_decode_tokens
|
||||
)
|
||||
|
||||
def reclaim_for_prealloc(
|
||||
self, tree_cache, full_tokens: int, swa_tokens: int
|
||||
) -> str | None:
|
||||
"""Free room for a decode-node preallocation; None means it is ready.
|
||||
|
||||
Returns a description of the shortfall when it cannot be met, for the
|
||||
caller to attach to whichever request it was admitting. Separate
|
||||
buffers make the sliding-window side the only one that needs
|
||||
reclaiming here, since the full side is priced by the caller's budget.
|
||||
"""
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||
|
||||
available = self.swa_available_size()
|
||||
if available < swa_tokens:
|
||||
tree_cache.evict_for_alloc(
|
||||
EvictParams(swa_num_tokens=swa_tokens - available)
|
||||
)
|
||||
available = self.swa_available_size()
|
||||
if available < swa_tokens:
|
||||
return (
|
||||
f"SWA eviction insufficient: needed={swa_tokens}, available={available}"
|
||||
)
|
||||
return None
|
||||
|
||||
def swa_capacity_and_available(self, *, full_capacity, swa_capacity):
|
||||
return (
|
||||
(full_capacity, self.full_available_size()),
|
||||
@@ -232,14 +256,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
def translate_swa_indices_for_transfer(
|
||||
self, kv_indices: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""Sliding-window token ids as the PD transfer engine addresses them.
|
||||
|
||||
The sibling of `translate_kv_indices_for_transfer` for the SWA state
|
||||
component. On a static pool the sliding-window buffers are indexed by
|
||||
the same ids the kernels use, so the read-path translate IS the answer.
|
||||
A virtual-id pool must override: the transfer addresses raw bytes and
|
||||
needs PHYSICAL ids, not kernel-facing ones.
|
||||
"""
|
||||
"""Map full-pool token ids to SWA-buffer token ids for PD transfer."""
|
||||
return self.translate_loc_from_full_to_swa(kv_indices)
|
||||
|
||||
def alloc(self, need_size: int):
|
||||
|
||||
@@ -388,6 +388,8 @@ class UnifiedSWAAllocatorBase(SWATokenToKVPoolAllocator):
|
||||
seq_lens_cpu: torch.Tensor,
|
||||
last_loc: torch.Tensor,
|
||||
extend_num_tokens: int,
|
||||
*,
|
||||
num_swa_pages: Optional[int] = None,
|
||||
) -> Optional[Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""Run the full side's paged extend and report which virtual PAGES it
|
||||
newly took. Returns (virtual TOKEN ids, new virtual PAGE ids), or None
|
||||
@@ -402,7 +404,10 @@ class UnifiedSWAAllocatorBase(SWATokenToKVPoolAllocator):
|
||||
prefix_lens=prefix_lens_cpu,
|
||||
)
|
||||
need_tokens = num_new_pages * self.page_size
|
||||
if not self.ensure_capacity(need_tokens, need_tokens):
|
||||
swa_tokens = (
|
||||
need_tokens if num_swa_pages is None else num_swa_pages * self.page_size
|
||||
)
|
||||
if not self.ensure_capacity(need_tokens, swa_tokens):
|
||||
return None
|
||||
|
||||
# Snapshot the virtual PAGES the kernel will consume; clone so swa keeps
|
||||
@@ -478,14 +483,22 @@ class UnifiedSWAAllocatorBase(SWATokenToKVPoolAllocator):
|
||||
sink and is skipped by `free`'s `swa_v2p_page > 0` mask -- exactly the
|
||||
out-of-window state the ratchet produces via `free_swa`.
|
||||
|
||||
Admission is priced at the FULL side's page count, as plain
|
||||
`alloc_extend` is: pessimistic when the tail is short, but it reuses
|
||||
the composite's audited joint capacity path, and the bytes actually
|
||||
held still follow the tail.
|
||||
Admission prices FULL's new pages and only the new pages in the SWA
|
||||
tail. A partial prefix page is already bound and costs no new SWA page.
|
||||
"""
|
||||
assert len(prefix_lens_cpu) == 1
|
||||
assert 0 <= swa_tail_len <= extend_num_tokens
|
||||
with record_function("UnifiedSWAAlloc.alloc_extend_swa_tail"):
|
||||
prefix_len = int(prefix_lens_cpu[0])
|
||||
seq_len = int(seq_lens_cpu[0])
|
||||
first_new_page = (prefix_len + self.page_size - 1) // self.page_size
|
||||
first_tail_page = (seq_len - swa_tail_len) // self.page_size
|
||||
num_swa_pages = (
|
||||
(seq_len + self.page_size - 1) // self.page_size
|
||||
- max(first_new_page, first_tail_page)
|
||||
if swa_tail_len
|
||||
else 0
|
||||
)
|
||||
extended = self._extend_in_virtual_space(
|
||||
prefix_lens,
|
||||
prefix_lens_cpu,
|
||||
@@ -493,6 +506,7 @@ class UnifiedSWAAllocatorBase(SWATokenToKVPoolAllocator):
|
||||
seq_lens_cpu,
|
||||
last_loc,
|
||||
extend_num_tokens,
|
||||
num_swa_pages=num_swa_pages,
|
||||
)
|
||||
if extended is None:
|
||||
return None
|
||||
@@ -819,6 +833,63 @@ class UnifiedSWATokenToKVPoolAllocator(UnifiedSWAAllocatorBase):
|
||||
"""No float in a two-END chain -- nothing can slide."""
|
||||
return None
|
||||
|
||||
def prealloc_fits_assumes_reclaim(self) -> bool:
|
||||
return True
|
||||
|
||||
def prealloc_ceiling_fits(self, full_tokens: int, swa_tokens: int) -> bool | None:
|
||||
return self.can_reserve(full_tokens, swa_tokens, empty_pool=True)
|
||||
|
||||
def reclaim_for_prealloc(
|
||||
self, tree_cache, full_tokens: int, swa_tokens: int
|
||||
) -> str | None:
|
||||
"""Reclaim both sides together: freeing FULL pages can open SWA room
|
||||
and the reverse, so the shared envelope is the only gate worth
|
||||
re-checking."""
|
||||
ready = self.evict_to_free_tokens(
|
||||
tree_cache, full_tokens, swa_num_tokens=swa_tokens
|
||||
)
|
||||
if ready is None:
|
||||
ready = self.ensure_capacity(full_tokens, swa_tokens)
|
||||
if ready:
|
||||
return None
|
||||
return (
|
||||
"Unified FULL/SWA byte reclamation insufficient: "
|
||||
f"needed=({full_tokens}, {swa_tokens})"
|
||||
)
|
||||
|
||||
def prealloc_fits(
|
||||
self,
|
||||
tree_cache,
|
||||
full_tokens: int,
|
||||
swa_tokens: int,
|
||||
*,
|
||||
full_budget_tokens: int,
|
||||
swa_budget_tokens: int | None = None,
|
||||
) -> bool:
|
||||
"""Price both sides against the shared byte envelope.
|
||||
|
||||
There is no per-side capacity for the scheduler's budget to be
|
||||
compared against, so the gap between that budget and what this side
|
||||
can currently hand out is folded back into the demand; `can_reserve`
|
||||
then prices the whole ask in bytes. Reachable only for hybrid-SWA
|
||||
models, so the tree's `full_*` accounting is the full-attention one.
|
||||
"""
|
||||
full_evictable_tokens = tree_cache.full_evictable_size()
|
||||
swa_evictable_tokens = tree_cache.swa_evictable_size()
|
||||
full_tokens += (
|
||||
self.full_available_size() + full_evictable_tokens - full_budget_tokens
|
||||
)
|
||||
if swa_budget_tokens is not None:
|
||||
swa_tokens += (
|
||||
self.swa_available_size() + swa_evictable_tokens - swa_budget_tokens
|
||||
)
|
||||
return self.can_reserve(
|
||||
full_tokens,
|
||||
swa_tokens,
|
||||
full_evictable_tokens=full_evictable_tokens,
|
||||
swa_evictable_tokens=swa_evictable_tokens,
|
||||
)
|
||||
|
||||
def reclaim_plan(
|
||||
self,
|
||||
full_tokens: int | float,
|
||||
@@ -1086,14 +1157,17 @@ class UnifiedSWATokenToKVPoolAllocator(UnifiedSWAAllocatorBase):
|
||||
hi = mid - 1
|
||||
return lo
|
||||
|
||||
def evict_to_free_tokens(self, tree_cache, num_tokens: int) -> bool | None:
|
||||
def evict_to_free_tokens(
|
||||
self, tree_cache, num_tokens: int, *, swa_num_tokens: Optional[int] = None
|
||||
) -> bool | None:
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||
|
||||
if tree_cache is None or tree_cache.is_chunk_cache():
|
||||
return
|
||||
required_swa = num_tokens if swa_num_tokens is None else swa_num_tokens
|
||||
reclaim_plan = self.reclaim_plan(
|
||||
num_tokens,
|
||||
num_tokens,
|
||||
required_swa,
|
||||
full_evictable_tokens=tree_cache.full_evictable_size(),
|
||||
swa_evictable_tokens=tree_cache.swa_evictable_size(),
|
||||
)
|
||||
@@ -1105,7 +1179,7 @@ class UnifiedSWATokenToKVPoolAllocator(UnifiedSWAAllocatorBase):
|
||||
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)
|
||||
return self.ensure_capacity(num_tokens, required_swa)
|
||||
|
||||
def verify_byte_accounting(self) -> List[str]:
|
||||
return (
|
||||
@@ -1234,69 +1308,104 @@ class UnifiedMambaSWATokenToKVPoolAllocator(UnifiedSWAAllocatorBase):
|
||||
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()
|
||||
return self._fits_page_demand(
|
||||
math.ceil(full_tokens / self.page_size),
|
||||
math.ceil(swa_tokens / self.page_size),
|
||||
)
|
||||
|
||||
def prealloc_fits(
|
||||
self,
|
||||
tree_cache,
|
||||
full_tokens: int,
|
||||
swa_tokens: int,
|
||||
*,
|
||||
full_budget_tokens: int,
|
||||
swa_budget_tokens: int | None = None,
|
||||
) -> bool:
|
||||
"""Price the pair on the float chain's grid, then against the budgets.
|
||||
|
||||
Each side's `available_size` takes `schedulable_available_size()`,
|
||||
which credits the peer's drainable holes, so the two are backed by the
|
||||
same bytes and a pair that fits each side alone can fail together. The
|
||||
budgets still apply on top: they carry decode headroom this allocator
|
||||
cannot see.
|
||||
"""
|
||||
page_size = self.page_size
|
||||
if not self._fits_page_demand(
|
||||
-(-full_tokens // page_size), -(-swa_tokens // page_size)
|
||||
):
|
||||
return False
|
||||
return full_tokens <= full_budget_tokens and (
|
||||
swa_budget_tokens is None or swa_tokens <= swa_budget_tokens
|
||||
)
|
||||
|
||||
def ensure_capacity(self, full_tokens: int, swa_tokens: int) -> bool:
|
||||
if full_tokens < 0 or swa_tokens < 0 or full_tokens != swa_tokens:
|
||||
if full_tokens < 0 or swa_tokens < 0:
|
||||
return False
|
||||
if full_tokens == 0:
|
||||
if self.can_reserve(full_tokens, swa_tokens):
|
||||
return True
|
||||
need_tokens = int(full_tokens)
|
||||
if need_tokens <= self.available_size():
|
||||
for allocator in self._flush_targets():
|
||||
allocator.flush_for_allocation()
|
||||
if self.can_reserve(full_tokens, swa_tokens):
|
||||
return True
|
||||
return _relieve_for_alloc(self, need_tokens)
|
||||
_float_open_short_side(
|
||||
self.swa_attn_allocator,
|
||||
{
|
||||
self.full_attn_allocator: -(-full_tokens // self.page_size),
|
||||
self.swa_attn_allocator: -(-swa_tokens // self.page_size),
|
||||
self.mamba_allocator: 0,
|
||||
},
|
||||
)
|
||||
return self.can_reserve(full_tokens, swa_tokens)
|
||||
|
||||
def _compute_available_size(self) -> int:
|
||||
"""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
|
||||
either side but only ONE per batch alloc. Feasibility is monotone in N, so
|
||||
binary search; the order matches the alloc path (full takes the high band).
|
||||
"""
|
||||
def _fits_page_demand(self, full_pages: int, swa_pages: int) -> bool:
|
||||
"""Price FULL first, then SWA in one contiguous band on the float grid."""
|
||||
fa, sa = self.full_attn_allocator, self.swa_attn_allocator
|
||||
e_f = fa.entry_bytes_per_page
|
||||
# full is grow-down: its chain gap IS the high band.
|
||||
b_high = fa._current_gap_bytes()
|
||||
h_f = len(fa._free_phys_pages) if fa.lazy_compaction else 0
|
||||
h_s = sa._hole_pages()
|
||||
r_f = fa.num_pages - fa.min_page_index - fa._allocated_pages()
|
||||
r_s = sa.num_pages - sa.min_page_index - sa._allocated_pages()
|
||||
|
||||
def feasible(n: int) -> bool:
|
||||
if n > h_f + r_f or n > h_s + r_s:
|
||||
return False
|
||||
ext_f = max(0, n - h_f)
|
||||
if ext_f * e_f > b_high:
|
||||
return False
|
||||
ext_s = max(0, n - h_s)
|
||||
# On the float's page grid, never in raw bytes: a byte budget
|
||||
# credits a page `take_physical_pages` cannot yield.
|
||||
full_low_after = fa._byte_low_frontier() - ext_f * e_f
|
||||
if sa._is_frontier_transparent():
|
||||
room = sa.pages_in_band(
|
||||
low_byte=sa._chain_high_frontier_below_bytes(),
|
||||
high_byte=full_low_after,
|
||||
)
|
||||
return ext_s <= room
|
||||
p_low = sa.pages_in_band(
|
||||
if full_pages > h_f + r_f or swa_pages > h_s + r_s:
|
||||
return False
|
||||
full_bytes = max(0, full_pages - h_f) * fa.entry_bytes_per_page
|
||||
if full_bytes > fa._current_gap_bytes():
|
||||
return False
|
||||
ext_s = max(0, swa_pages - h_s)
|
||||
full_low_after = fa._byte_low_frontier() - full_bytes
|
||||
if sa._is_frontier_transparent():
|
||||
room = sa.pages_in_band(
|
||||
low_byte=sa._chain_high_frontier_below_bytes(),
|
||||
high_byte=sa._byte_low_frontier(),
|
||||
)
|
||||
p_high = sa.pages_in_band(
|
||||
low_byte=sa._byte_high_frontier(),
|
||||
high_byte=full_low_after,
|
||||
)
|
||||
return ext_s <= max(p_low, p_high)
|
||||
return ext_s <= room
|
||||
p_low = sa.pages_in_band(
|
||||
low_byte=sa._chain_high_frontier_below_bytes(),
|
||||
high_byte=sa._byte_low_frontier(),
|
||||
)
|
||||
p_high = sa.pages_in_band(
|
||||
low_byte=sa._byte_high_frontier(),
|
||||
high_byte=full_low_after,
|
||||
)
|
||||
return ext_s <= max(p_low, p_high)
|
||||
|
||||
def _compute_available_size(self) -> int:
|
||||
"""Joint TOKENS for equal FULL/SWA demand, using the same page predicate
|
||||
as tail allocation. FULL takes the high band before SWA binds its pages.
|
||||
"""
|
||||
fa, sa = self.full_attn_allocator, self.swa_attn_allocator
|
||||
h_f = len(fa._free_phys_pages) if fa.lazy_compaction else 0
|
||||
h_s = sa._hole_pages()
|
||||
r_f = fa.num_pages - fa.min_page_index - fa._allocated_pages()
|
||||
r_s = sa.num_pages - sa.min_page_index - sa._allocated_pages()
|
||||
lo_n, hi_n = 0, min(h_f + r_f, h_s + r_s)
|
||||
while lo_n < hi_n:
|
||||
mid = (lo_n + hi_n + 1) // 2
|
||||
if feasible(mid):
|
||||
if self._fits_page_demand(mid, mid):
|
||||
lo_n = mid
|
||||
else:
|
||||
hi_n = mid - 1
|
||||
|
||||
@@ -307,6 +307,30 @@ class UnifiedMambaTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||
)
|
||||
return self.full_attn_allocator.translate_kv_loc(kv_indices.to(torch.int64))
|
||||
|
||||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||||
"""Retraction backup for the FULL + mamba pair.
|
||||
|
||||
`Req.offload_kv_cache` hands over `req_to_token` rows, which hold
|
||||
VIRTUAL ids here; both unified full pools index their host copy by
|
||||
PHYSICAL ids. The mamba side is already slot-addressed and is
|
||||
translated by the pool.
|
||||
"""
|
||||
return self._kvcache.get_cpu_copy(
|
||||
self.full_attn_allocator.translate_kv_loc(indices.to(torch.int64)),
|
||||
mamba_indices=mamba_indices,
|
||||
req_pool_index=req_pool_index,
|
||||
)
|
||||
|
||||
def load_cpu_copy(
|
||||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||||
):
|
||||
return self._kvcache.load_cpu_copy(
|
||||
kv_cache_cpu,
|
||||
self.full_attn_allocator.translate_kv_loc(indices.to(torch.int64)),
|
||||
mamba_indices=mamba_indices,
|
||||
req_pool_index=req_pool_index,
|
||||
)
|
||||
|
||||
def _move_gate_targets(self):
|
||||
"""Every member a compaction gate must cover. The mamba end is gated
|
||||
even where its state is not itself transferred: the gate is about the
|
||||
|
||||
@@ -265,7 +265,11 @@ def install_move_gate(
|
||||
|
||||
|
||||
class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||
"""Allocator for one sub-pool over a `UnifiedKVPool`."""
|
||||
"""Allocator for one sub-pool over a `UnifiedKVPool`.
|
||||
|
||||
``need_sort`` applies to transfer-facing physical ids, not virtual ids.
|
||||
Physical free pages are sorted during compaction.
|
||||
"""
|
||||
|
||||
# Capacity-bearing state: any rebind bumps `_capacity_epoch`, invalidating
|
||||
# the chain's capacity memos (see `_CapacityField`).
|
||||
@@ -1207,11 +1211,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||
if not _relieve_for_alloc(self, need_tokens):
|
||||
return None
|
||||
bs = len(prefix_lens)
|
||||
if self.need_sort and extend_num_tokens // self.page_size + bs + 1 > len(
|
||||
self.free_virtual_ids
|
||||
):
|
||||
self.merge_and_sort_free()
|
||||
|
||||
# Snapshot the virtual pages the kernel will consume, to bind them
|
||||
# to physical pages afterward.
|
||||
if num_new_pages > 0:
|
||||
@@ -1274,9 +1273,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
||||
if need_tokens > self.available_size():
|
||||
if not _relieve_for_alloc(self, need_tokens):
|
||||
return None
|
||||
if self.need_sort and bs > len(self.free_virtual_ids):
|
||||
self.merge_and_sort_free()
|
||||
|
||||
# Most decode steps reuse the prefix's tail page -> num_new_pages == 0.
|
||||
if num_new_pages > 0:
|
||||
new_virtual_pages = self.free_virtual_ids[:num_new_pages].clone()
|
||||
|
||||
@@ -154,19 +154,11 @@ def resolve_decode_retraction_backup(*, tp_worker: BaseTpWorker) -> str:
|
||||
if tp_worker.is_hybrid_swa
|
||||
else None
|
||||
)
|
||||
# Host-pool retraction transfers full and sliding-window components
|
||||
# only, so a model with recurrent state stays on cpu_tensor.
|
||||
#
|
||||
# The unified pool is excluded for the same reason hierarchical cache is
|
||||
# (see `handle_unified_memory_pool`): the host-transfer path indexes the
|
||||
# device buffers with the ids it is handed, and under the unified pool
|
||||
# those are VIRTUAL. It also cannot be sized from `kv_cache.size`, which
|
||||
# is a KERNEL-FACING row count (`num_pages * 2 * layer_num * page_size`)
|
||||
# rather than a token capacity -- gpt-oss-20b reports 85M "tokens" and
|
||||
# asks for 418 GB of host memory per component.
|
||||
# Host-pool retraction does not address unified page envelopes or
|
||||
# recurrent state, so those configurations stay on cpu_tensor.
|
||||
supports_host_pool = (
|
||||
not uses_ssm_state(tp_worker.model_runner.model_config)
|
||||
and not memory.enable_unified_memory
|
||||
not memory.enable_unified_memory
|
||||
and not uses_ssm_state(tp_worker.model_runner.model_config)
|
||||
and (
|
||||
isinstance(kv_cache, MHATokenToKVPool)
|
||||
or (isinstance(kv_cache, SWAKVPool) and full_tokens_per_layer > 0)
|
||||
|
||||
@@ -868,32 +868,7 @@ class KVCacheConfigurator:
|
||||
assert not self.use_mla_backend, (
|
||||
"unified memory pool does not support MLA-SWA hybrid yet"
|
||||
)
|
||||
# Mirror the non-shared path's extra_max_context_len computation.
|
||||
extra_max_context_len = 4
|
||||
if get_spec().speculative_num_draft_tokens is not None:
|
||||
extra_max_context_len += get_spec().speculative_num_draft_tokens
|
||||
if get_disagg().disaggregation_mode == "decode":
|
||||
# A decode node hands out request rows to PREALLOCATED transfers on
|
||||
# top of its running set, so it needs the extra-slot pool (and the
|
||||
# `pre_alloc_size` the scheduler's invariant checker reads). Mirrors
|
||||
# `_build_req_to_token_pool`'s decode branch; the mamba composite
|
||||
# already takes `decode_pre_alloc_size` the same way.
|
||||
from sglang.srt.disaggregation.decode import DecodeReqToTokenPool
|
||||
|
||||
req_to_token_pool = DecodeReqToTokenPool(
|
||||
size=max_num_reqs,
|
||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||
device=self.device,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
pre_alloc_size=get_disagg().disaggregation_decode_extra_slots,
|
||||
)
|
||||
else:
|
||||
req_to_token_pool = ReqToTokenPool(
|
||||
size=max_num_reqs,
|
||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||
device=self.device,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
)
|
||||
req_to_token_pool = self._build_req_to_token_pool(max_num_reqs=max_num_reqs)
|
||||
|
||||
head_num = self.model_config.get_num_kv_heads(
|
||||
get_parallel().attn_tp_size, get_parallel().attn_dcp_size
|
||||
|
||||
@@ -624,29 +624,29 @@ class UnifiedMHATokenToKVPool(MHATokenToKVPool):
|
||||
env[tgt_pages] = env[src_pages]
|
||||
|
||||
def get_contiguous_buf_infos(self):
|
||||
"""PD-transfer registration: ONE entry, the raw buffer, addressed as
|
||||
``raw_ptr + physical_page_id * page_envelope_bytes``.
|
||||
"""Register the raw buffer as physical page envelopes for PD transfer.
|
||||
|
||||
Same whole-envelope contract as `UnifiedMLATokenToKVPool`: the transfer
|
||||
item is one page across ALL layers and both K and V, because the
|
||||
per-layer views overlap inside the envelope and index in kernel-facing
|
||||
ids. A peer must therefore build an identical spec -- enforced on the
|
||||
wire by `_validate_envelope_kv_layout`.
|
||||
Full and SWA expose the same allocation with different envelope sizes;
|
||||
the transfer backend preserves both logical entries while deduplicating
|
||||
the underlying memory registration.
|
||||
"""
|
||||
# The address formula omits the anchor; a nonzero one would mis-address.
|
||||
assert self._unified_buffer.anchor_bytes(self._sub_pool_name) == 0
|
||||
raw = self._unified_buffer._raw
|
||||
return [raw.data_ptr()], [raw.numel()], [self._page_bytes]
|
||||
|
||||
def get_cpu_copy(self, indices, mamba_indices=None):
|
||||
raise NotImplementedError(
|
||||
"CPU offloading is unsupported under the unified layout."
|
||||
)
|
||||
def _physical_to_kernel_indices(self, indices: torch.Tensor) -> torch.Tensor:
|
||||
return (indices // self.page_size) * (
|
||||
self.page_size * self.kernel_page_blocks
|
||||
) + indices % self.page_size
|
||||
|
||||
def load_cpu_copy(self, kv_cache_cpu, indices, mamba_indices=None):
|
||||
raise NotImplementedError(
|
||||
"CPU offloading is unsupported under the unified layout."
|
||||
)
|
||||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||||
"""Translate physical host-pool ids for the page-major parent path."""
|
||||
return super().get_cpu_copy(self._physical_to_kernel_indices(indices))
|
||||
|
||||
def load_cpu_copy(
|
||||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||||
):
|
||||
super().load_cpu_copy(kv_cache_cpu, self._physical_to_kernel_indices(indices))
|
||||
|
||||
def set_kv_buffer_prefix_valid(self, *args, **kwargs):
|
||||
raise NotImplementedError(
|
||||
@@ -738,6 +738,22 @@ class UnifiedMLATokenToKVPool(MLATokenToKVPool):
|
||||
raw = self._unified_buffer._raw
|
||||
return [raw.data_ptr()], [raw.numel()], [self._page_bytes]
|
||||
|
||||
def _physical_to_kernel_indices(self, indices: torch.Tensor) -> torch.Tensor:
|
||||
"""Physical TOKEN ids -> the kernel-facing ids this class's `kv_buffer`
|
||||
views are indexed by; the formula is the one in the class docstring."""
|
||||
return (indices // self.page_size) * (
|
||||
self.page_size * self.kernel_page_blocks
|
||||
) + indices % self.page_size
|
||||
|
||||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||||
"""Translate physical host-pool ids for the page-major parent path."""
|
||||
return super().get_cpu_copy(self._physical_to_kernel_indices(indices))
|
||||
|
||||
def load_cpu_copy(
|
||||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||||
):
|
||||
super().load_cpu_copy(kv_cache_cpu, self._physical_to_kernel_indices(indices))
|
||||
|
||||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||||
"""Relocate whole page envelopes.
|
||||
|
||||
@@ -1676,7 +1692,7 @@ class UnifiedSWAKVPool(SWAKVPool):
|
||||
swa_cpu = None
|
||||
if bool(valid.any().item()):
|
||||
swa_cpu = self.swa_kv_pool.get_cpu_copy(swa_phys[valid])
|
||||
return {"full": full_cpu, "swa": swa_cpu}
|
||||
return {"full": full_cpu, "swa": swa_cpu, "swa_mask": valid.cpu()}
|
||||
|
||||
def load_cpu_copy(
|
||||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||||
@@ -1689,7 +1705,14 @@ class UnifiedSWAKVPool(SWAKVPool):
|
||||
if kv_cache_cpu.get("swa") is not None:
|
||||
assert self._swa_allocator is not None
|
||||
swa_phys = self._virt_tokens_to_phys_tokens(indices, self._swa_allocator)
|
||||
self.swa_kv_pool.load_cpu_copy(kv_cache_cpu["swa"], swa_phys)
|
||||
old_swa_mask = kv_cache_cpu["swa_mask"].to(indices.device)
|
||||
assert old_swa_mask.shape == indices.shape
|
||||
row_mask = (swa_phys >= 0)[old_swa_mask].cpu()
|
||||
swa_phys = swa_phys[old_swa_mask][row_mask.to(indices.device)]
|
||||
if swa_phys.numel() == 0:
|
||||
return
|
||||
swa_cpu = self._filter_swa_cpu_copy(kv_cache_cpu["swa"], row_mask)
|
||||
self.swa_kv_pool.load_cpu_copy(swa_cpu, swa_phys)
|
||||
|
||||
|
||||
class UnifiedSWAPoolBundle(NamedTuple):
|
||||
|
||||
@@ -341,6 +341,9 @@ class ServerArgs:
|
||||
# _handle_page_major_kv_layout); the model-family gate is enforced at pool
|
||||
# construction in model_runner_kv_cache_mixin._init_pools.
|
||||
|
||||
def _unified_memory_pd_transfer_backends(self) -> set[str]:
|
||||
return {"mooncake"}
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: argparse.ArgumentParser):
|
||||
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""Allocator double for scheduler tests that admit requests.
|
||||
|
||||
`DecodePreallocQueue` asks the allocator to price a preallocation rather than
|
||||
doing the arithmetic itself, so a bare `MagicMock` returns a truthy `Mock` and
|
||||
the admission decision under test stops being made anywhere. Binding the real
|
||||
separate-buffer implementations keeps the arithmetic live while leaving the
|
||||
per-test stubs (`size_swa`, `swa_available_size`, ...) in charge of the state.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
|
||||
|
||||
# Bound on the double, so each reads the stubs the caller set on it.
|
||||
_SEPARATE_BUFFER_METHODS = {
|
||||
"prealloc_fits_assumes_reclaim": BaseTokenToKVPoolAllocator.prealloc_fits_assumes_reclaim,
|
||||
"prealloc_ceiling_fits": BaseTokenToKVPoolAllocator.prealloc_ceiling_fits,
|
||||
"prealloc_fits": BaseTokenToKVPoolAllocator.prealloc_fits,
|
||||
"reclaim_for_prealloc": SWATokenToKVPoolAllocator.reclaim_for_prealloc,
|
||||
"swa_capacity_and_available": SWATokenToKVPoolAllocator.swa_capacity_and_available,
|
||||
}
|
||||
|
||||
|
||||
def bind_separate_buffer_capacity(allocator) -> None:
|
||||
"""Make `allocator` price capacity like a pool whose sides own their own
|
||||
buffers. Call on any allocator double a `DecodePreallocQueue` will read."""
|
||||
for name, impl in _SEPARATE_BUFFER_METHODS.items():
|
||||
setattr(
|
||||
allocator,
|
||||
name,
|
||||
(lambda impl: lambda *args, **kwargs: impl(allocator, *args, **kwargs))(
|
||||
impl
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def separate_buffer_allocator_double(**attrs) -> MagicMock:
|
||||
"""A `MagicMock` allocator that prices capacity as separate buffers."""
|
||||
allocator = MagicMock(**attrs)
|
||||
bind_separate_buffer_capacity(allocator)
|
||||
return allocator
|
||||
@@ -0,0 +1,127 @@
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.unified_memory_pool import init_unified_swa_pools
|
||||
from sglang.srt.runtime_context import publish, reset_context
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
|
||||
class TestUnifiedSWATailAllocation(CustomTestCase):
|
||||
def setUp(self):
|
||||
reset_context()
|
||||
self.addCleanup(reset_context)
|
||||
publish(
|
||||
ServerArgs(model_path="dummy", enable_unified_memory=True), role="tokenizer"
|
||||
)
|
||||
|
||||
def test_extend_binds_only_new_tail_pages(self):
|
||||
"""PD tail allocation must leave new FULL-only pages unbound in SWA,
|
||||
while preserving an existing partial page and binding the trailing KV."""
|
||||
for page_size in (4, 16):
|
||||
for prefix_len, seq_len, tail_len in (
|
||||
(page_size, 5 * page_size, 0),
|
||||
(page_size, 5 * page_size, page_size),
|
||||
(page_size, 5 * page_size, 2 * page_size),
|
||||
(page_size + 2, 5 * page_size, page_size),
|
||||
(page_size + 2, 5 * page_size, 4 * page_size - 2),
|
||||
(page_size + 2, 2 * page_size - 1, page_size - 3),
|
||||
(page_size + 2, 5 * page_size - 1, page_size),
|
||||
(0, None, 1),
|
||||
):
|
||||
with self.subTest(
|
||||
page_size=page_size,
|
||||
prefix_len=prefix_len,
|
||||
seq_len=seq_len,
|
||||
tail_len=tail_len,
|
||||
):
|
||||
bundle = init_unified_swa_pools(
|
||||
device="cuda",
|
||||
kv_cache_dtype=torch.float16,
|
||||
head_num=1,
|
||||
head_dim=8,
|
||||
v_head_dim=8,
|
||||
swa_head_num=1,
|
||||
swa_head_dim=8,
|
||||
swa_v_head_dim=8,
|
||||
page_size=page_size,
|
||||
start_layer=0,
|
||||
end_layer=2,
|
||||
full_attention_layer_ids=[0],
|
||||
swa_attention_layer_ids=[1],
|
||||
total_bytes=1 << 16,
|
||||
enable_memory_saver=False,
|
||||
need_sort=False,
|
||||
)
|
||||
allocator = bundle.token_to_kv_pool_allocator
|
||||
if seq_len is None:
|
||||
seq_len = allocator.available_size() + page_size
|
||||
self.assertFalse(allocator.can_reserve(seq_len, seq_len))
|
||||
prefix_capacity = -(-prefix_len // page_size) * page_size
|
||||
prefix = allocator.alloc(prefix_capacity)[:prefix_len]
|
||||
prefix_swa = allocator.translate_swa_indices_for_transfer(
|
||||
prefix
|
||||
).clone()
|
||||
prefix_cpu = torch.tensor([prefix_len], dtype=torch.int64)
|
||||
seq_cpu = torch.tensor([seq_len], dtype=torch.int64)
|
||||
extended = allocator.alloc_extend_swa_tail(
|
||||
prefix_lens=prefix_cpu.cuda(),
|
||||
prefix_lens_cpu=prefix_cpu,
|
||||
seq_lens=seq_cpu.cuda(),
|
||||
seq_lens_cpu=seq_cpu,
|
||||
last_loc=(
|
||||
prefix[-1:]
|
||||
if prefix_len
|
||||
else torch.tensor([-1], device="cuda")
|
||||
),
|
||||
extend_num_tokens=seq_len - prefix_len,
|
||||
swa_tail_len=tail_len,
|
||||
)
|
||||
self.assertIsNotNone(extended)
|
||||
self.assertEqual(extended.numel(), seq_len - prefix_len)
|
||||
tokens = torch.cat((prefix, extended))
|
||||
full_phys = allocator.translate_kv_indices_for_transfer(tokens)
|
||||
swa_phys = allocator.translate_swa_indices_for_transfer(tokens)
|
||||
self.assertTrue(bool((full_phys > 0).all()))
|
||||
self.assertTrue(torch.equal(swa_phys[:prefix_len], prefix_swa))
|
||||
tail_start = seq_len - tail_len
|
||||
new_pages = torch.unique(tokens[prefix_capacity:] // page_size)
|
||||
tail_pages = torch.unique(tokens[tail_start:] // page_size)
|
||||
full_only_pages = new_pages[~torch.isin(new_pages, tail_pages)]
|
||||
self.assertTrue(
|
||||
bool(
|
||||
(allocator.swa_v2p_page_table[full_only_pages] == -1).all()
|
||||
)
|
||||
)
|
||||
if tail_len:
|
||||
pages = allocator.swa_v2p_page_table[
|
||||
tokens[tail_start:] // page_size
|
||||
]
|
||||
self.assertTrue(bool((pages > 0).all()))
|
||||
expected = pages * page_size + tokens[tail_start:] % page_size
|
||||
self.assertTrue(torch.equal(swa_phys[tail_start:], expected))
|
||||
self.assertEqual(
|
||||
allocator.swa_attn_allocator.allocated_count(),
|
||||
prefix_capacity
|
||||
+ torch.isin(new_pages, tail_pages).sum().item() * page_size,
|
||||
)
|
||||
if prefix_len < prefix_capacity:
|
||||
reused_tokens = min(prefix_capacity, seq_len) - prefix_len
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
extended[:reused_tokens],
|
||||
prefix[-1]
|
||||
+ torch.arange(1, reused_tokens + 1, device="cuda"),
|
||||
)
|
||||
)
|
||||
allocator.free(tokens)
|
||||
self.assertEqual(allocator.full_attn_allocator.allocated_count(), 0)
|
||||
self.assertEqual(allocator.swa_attn_allocator.allocated_count(), 0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -17,6 +17,9 @@ from sglang.srt.managers.scheduler import Scheduler
|
||||
from sglang.srt.runtime_context import get_context, publish, reset_context
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.separate_buffer_allocator_double import (
|
||||
bind_separate_buffer_capacity,
|
||||
)
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
|
||||
@@ -71,7 +74,8 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
queue.retracted_queue = reqs.copy()
|
||||
queue.num_reserved_decode_tokens = 0
|
||||
queue.req_to_token_pool = SimpleNamespace(available_size=lambda: len(reqs))
|
||||
queue.token_to_kv_pool_allocator = SimpleNamespace(page_size=page_size)
|
||||
queue.token_to_kv_pool_allocator = MagicMock(page_size=page_size)
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
queue.tree_cache = MagicMock()
|
||||
queue.scheduler = SimpleNamespace(
|
||||
sliding_window_size=2047,
|
||||
@@ -81,6 +85,9 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
queue._swa_aware_allocatable_token_budgets = MagicMock(
|
||||
return_value=(physical_available, physical_available)
|
||||
)
|
||||
queue._allocatable_token_budgets = MagicMock(
|
||||
side_effect=lambda **_: physical_available
|
||||
)
|
||||
queue._swa_tail_allocatable_token_budget = MagicMock(
|
||||
side_effect=lambda **_: physical_available
|
||||
)
|
||||
@@ -120,6 +127,10 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
queue.retracted_queue = []
|
||||
queue._resolve_pending_reqs = MagicMock()
|
||||
queue._uses_swa_tail_prealloc = MagicMock(return_value=False)
|
||||
# `_uses_swa_reservation` consults the allocator once tail prealloc is
|
||||
# off, so this abort path needs one even though it never allocates.
|
||||
queue.token_to_kv_pool_allocator = MagicMock()
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
queue._allocatable_token_budgets = MagicMock(return_value=0)
|
||||
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
|
||||
|
||||
@@ -175,6 +186,10 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
queue._resolve_pending_reqs = MagicMock()
|
||||
queue._update_handshake_waiters = MagicMock()
|
||||
queue._uses_swa_tail_prealloc = MagicMock(return_value=False)
|
||||
# `_uses_swa_reservation` consults the allocator once tail prealloc is
|
||||
# off, so this abort path needs one even though it never allocates.
|
||||
queue.token_to_kv_pool_allocator = MagicMock()
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
queue._allocatable_token_budgets = MagicMock(return_value=0)
|
||||
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
|
||||
|
||||
@@ -234,6 +249,9 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
)
|
||||
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
|
||||
queue._pre_alloc = MagicMock()
|
||||
queue.token_to_kv_pool_allocator = MagicMock()
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
queue.tree_cache = MagicMock()
|
||||
queue.req_to_token_pool = MagicMock()
|
||||
queue.req_to_token_pool.available_size.return_value = 1
|
||||
# Non-hybrid pools have no mamba allocator; MagicMock would otherwise
|
||||
|
||||
@@ -7,6 +7,7 @@ from types import SimpleNamespace
|
||||
import numpy as np
|
||||
|
||||
from sglang.srt.disaggregation.ascend.conn import AscendKVManager
|
||||
from sglang.srt.disaggregation.base.conn import StateType
|
||||
from sglang.srt.disaggregation.common.conn import CommonKVManager
|
||||
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
|
||||
from sglang.srt.disaggregation.prefill import _transfer_start_layer
|
||||
@@ -15,6 +16,8 @@ from sglang.srt.disaggregation.utils import (
|
||||
build_transfer_entry_pairs,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
|
||||
from sglang.srt.runtime_context import get_memory, publish, reset_context
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -71,6 +74,7 @@ class TestTransferStartLayer(CustomTestCase):
|
||||
|
||||
class _RecordingKVManager:
|
||||
get_mha_kv_ptrs_with_pp = CommonKVManager.get_mha_kv_ptrs_with_pp
|
||||
get_mla_kv_ptrs_with_pp = CommonKVManager.get_mla_kv_ptrs_with_pp
|
||||
|
||||
def __init__(self, *, prefill_start_layer: int, pp_size: int):
|
||||
self.is_mla_backend = False
|
||||
@@ -143,6 +147,50 @@ class TestHybridSendUsesLayerIdPairing(CustomTestCase):
|
||||
self._run_case(model_full_ids=ids, stage_full_ids=ids[:5], start_offset=0)
|
||||
|
||||
|
||||
class TestSingleRegionSWATransfer(CustomTestCase):
|
||||
def test_one_region_full_generates_transfer_block(self):
|
||||
publish(ServerArgs(model_path="dummy"), role="tokenizer")
|
||||
self.addCleanup(reset_context)
|
||||
manager = _RecordingKVManager(prefill_start_layer=0, pp_size=1)
|
||||
manager.kv_args.kv_data_ptrs = [1000]
|
||||
manager.kv_args.kv_item_lens = [64]
|
||||
manager.kv_args.kv_layer_ids = []
|
||||
manager._validate_envelope_kv_layout = (
|
||||
MooncakeKVManager._validate_envelope_kv_layout.__get__(manager)
|
||||
)
|
||||
manager._send_kvcache_generic = MooncakeKVManager._send_kvcache_generic.__get__(
|
||||
manager
|
||||
)
|
||||
with get_memory().override(enable_unified_memory=True):
|
||||
rc = MooncakeKVManager.send_kvcache(
|
||||
manager,
|
||||
mooncake_session_id="session",
|
||||
prefill_kv_indices=np.array([3, 4], dtype=np.int32),
|
||||
dst_kv_ptrs=[2000],
|
||||
dst_kv_indices=np.array([7, 8], dtype=np.int32),
|
||||
dst_kv_item_len=64,
|
||||
executor=None,
|
||||
)
|
||||
self.assertEqual(rc, 0)
|
||||
self.assertEqual(manager.blocks, [(1192, 2448, 128)])
|
||||
|
||||
def test_one_region_swa_generates_transfer_block(self):
|
||||
manager = _RecordingKVManager(prefill_start_layer=0, pp_size=1)
|
||||
rc = MooncakeKVManager._send_kvcache_generic(
|
||||
manager,
|
||||
mooncake_session_id="session",
|
||||
src_data_ptrs=[1000],
|
||||
dst_data_ptrs=[2000],
|
||||
item_lens=[64],
|
||||
prefill_data_indices=np.array([3, 4], dtype=np.int32),
|
||||
dst_data_indices=np.array([7, 8], dtype=np.int32),
|
||||
executor=None,
|
||||
state_type=StateType.SWA,
|
||||
)
|
||||
self.assertEqual(rc, 0)
|
||||
self.assertEqual(manager.blocks, [(1000 + 3 * 64, 2000 + 7 * 64, 2 * 64)])
|
||||
|
||||
|
||||
class _RecordingAscendManager:
|
||||
def __init__(self):
|
||||
self.is_hybrid_mla_backend = True
|
||||
|
||||
@@ -23,6 +23,9 @@ from sglang.srt.managers.scheduler import Scheduler # noqa: E402
|
||||
from sglang.srt.runtime_context import get_context, publish, reset_context # noqa: E402
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.separate_buffer_allocator_double import (
|
||||
bind_separate_buffer_capacity,
|
||||
)
|
||||
|
||||
register_cpu_ci(est_time=12, suite="base-a-test-cpu")
|
||||
|
||||
@@ -161,11 +164,14 @@ class TestDecodePreallocQueuePriority(unittest.TestCase):
|
||||
queue.req_to_metadata_buffer_idx_allocator.alloc.side_effect = iter(range(100))
|
||||
|
||||
queue.token_to_kv_pool_allocator = MagicMock()
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
queue.token_to_kv_pool_allocator.page_size = 1
|
||||
queue.token_to_kv_pool_allocator.available_size.return_value = 1000
|
||||
queue.token_to_kv_pool = MagicMock()
|
||||
queue.transfer_queue = SimpleNamespace(queue=[], enable_staging=False)
|
||||
queue.kv_manager = SimpleNamespace(kv_args=SimpleNamespace(state_types=[]))
|
||||
queue.kv_manager = SimpleNamespace(
|
||||
kv_args=SimpleNamespace(state_types=[]),
|
||||
)
|
||||
queue.tree_cache = MagicMock()
|
||||
|
||||
scheduler = MagicMock()
|
||||
|
||||
@@ -47,6 +47,9 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey
|
||||
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType
|
||||
from sglang.srt.utils.common import Range
|
||||
from sglang.test.separate_buffer_allocator_double import (
|
||||
bind_separate_buffer_capacity,
|
||||
)
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
|
||||
@@ -128,6 +131,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
server_args=SimpleNamespace(),
|
||||
)
|
||||
queue.token_to_kv_pool_allocator = MagicMock(page_size=64)
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
|
||||
tail_len = queue._swa_tail_len(895)
|
||||
|
||||
@@ -146,6 +150,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
queue._need_space_for_single_req = MagicMock(return_value=0)
|
||||
queue._active_req_count = MagicMock(return_value=1)
|
||||
queue.token_to_kv_pool_allocator = MagicMock()
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
queue.token_to_kv_pool_allocator.size_swa = 256
|
||||
queue.token_to_kv_pool_allocator.swa_available_size.return_value = 0
|
||||
queue.tree_cache = MagicMock()
|
||||
@@ -162,6 +167,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
def test_reclaim_swa_tail_capacity_page_rounds(self):
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue.token_to_kv_pool_allocator = MagicMock(page_size=64)
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
queue.token_to_kv_pool_allocator.swa_available_size.side_effect = [64, 192]
|
||||
queue.tree_cache = MagicMock()
|
||||
|
||||
@@ -175,6 +181,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
def test_reclaim_swa_tail_capacity_fails_before_allocation(self):
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue.token_to_kv_pool_allocator = MagicMock(page_size=64)
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
queue.token_to_kv_pool_allocator.swa_available_size.side_effect = [64, 128]
|
||||
queue.tree_cache = MagicMock()
|
||||
|
||||
@@ -471,6 +478,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
queue._update_handshake_waiters = MagicMock()
|
||||
queue._uses_swa_tail_prealloc = MagicMock(return_value=True)
|
||||
queue._swa_tail_len = MagicMock(return_value=8)
|
||||
queue._prealloc_required_tokens = MagicMock(return_value=(8, 8))
|
||||
queue._swa_aware_allocatable_token_budgets = MagicMock(return_value=(8, 8))
|
||||
queue._swa_tail_allocatable_token_budget = MagicMock(return_value=8)
|
||||
queue._match_prefix_and_lock = MagicMock(
|
||||
@@ -497,6 +505,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
queue.req_to_metadata_buffer_idx_allocator.available_size.return_value = 1
|
||||
queue.token_to_kv_pool = MagicMock()
|
||||
queue.token_to_kv_pool_allocator = MagicMock()
|
||||
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
|
||||
queue.token_to_kv_pool_allocator.page_size = 4
|
||||
|
||||
running_batch = MagicMock()
|
||||
@@ -539,7 +548,7 @@ class TestDecodeLockRefScenarios(CustomTestCase):
|
||||
skip_swa=True,
|
||||
)
|
||||
self.assertFalse(req.swa_prefix_lock_released)
|
||||
queue._swa_tail_len.assert_called_once_with(8)
|
||||
queue._swa_tail_len.assert_called_with(8)
|
||||
queue._allocatable_token_budgets.assert_called_once()
|
||||
|
||||
def test_hicache_restore_commit_hands_over_lock_with_receipt(self):
|
||||
|
||||
@@ -18,7 +18,10 @@ from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator
|
||||
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.test_utils import CustomTestCase
|
||||
from sglang.test.separate_buffer_allocator_double import (
|
||||
separate_buffer_allocator_double,
|
||||
)
|
||||
from sglang.test.test_utils import CustomTestCase, enter_override
|
||||
|
||||
register_cpu_ci(est_time=12, suite="base-a-test-cpu")
|
||||
|
||||
@@ -102,7 +105,10 @@ def _make_prealloc_queue(
|
||||
"""Build a minimal DecodePreallocQueue for _check_if_req_exceed_kv_capacity."""
|
||||
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
|
||||
queue.max_total_num_tokens = max_total_num_tokens
|
||||
queue.token_to_kv_pool_allocator = SimpleNamespace(size_swa=10**9)
|
||||
queue.num_reserved_decode_tokens = 0
|
||||
queue.token_to_kv_pool_allocator = separate_buffer_allocator_double(
|
||||
page_size=1, size_swa=10**9
|
||||
)
|
||||
# Disable the SWA-tail branch; this test only exercises the pool-length gate.
|
||||
queue._uses_swa_tail_prealloc = MagicMock(return_value=False)
|
||||
|
||||
@@ -128,6 +134,10 @@ def _make_req(rid: str, prompt_len: int):
|
||||
|
||||
|
||||
class TestCheckIfReqExceedKvCapacity(CustomTestCase):
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
enter_override(self, get_context().override_server_args())
|
||||
|
||||
def test_hisparse_admits_beyond_device_pool_up_to_host_backed_size(self):
|
||||
"""Core regression: request longer than device-only
|
||||
`max_total_num_tokens` but within HiSparse host-backed
|
||||
|
||||
@@ -40,6 +40,7 @@ from sglang.srt.mem_cache.allocator.unified_sub_pool import (
|
||||
MultiEndedAllocator,
|
||||
)
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||
from sglang.srt.mem_cache.common import kv_to_page_indices
|
||||
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_memory_pool import (
|
||||
@@ -2522,6 +2523,29 @@ class TestSWACompositeKernelIdSurface(unittest.TestCase):
|
||||
expected = v2p_swa[v // self.PS] * (self.PS * mult) + v % self.PS
|
||||
self.assertTrue(torch.equal(a.translate_loc_from_full_to_swa(v), expected))
|
||||
|
||||
def test_swa_transfer_page_is_physical_not_kernel_scaled(self):
|
||||
mult = 2 * self.SWA_L
|
||||
a = self._build()
|
||||
v = a.alloc(3 * self.PS)
|
||||
self.assertIsNotNone(v)
|
||||
|
||||
physical_pages = a.swa_attn_allocator.virtual_to_physical[
|
||||
v[:: self.PS] // self.PS
|
||||
]
|
||||
physical_tokens = a.swa_attn_allocator.translate_kv_loc(v)
|
||||
transfer_tokens = a.translate_swa_indices_for_transfer(v)
|
||||
self.assertTrue(torch.equal(transfer_tokens, physical_tokens))
|
||||
self.assertEqual(
|
||||
kv_to_page_indices(transfer_tokens, self.PS).tolist(),
|
||||
physical_pages.tolist(),
|
||||
)
|
||||
|
||||
kernel_tokens = a.translate_loc_from_full_to_swa(v)
|
||||
self.assertEqual(
|
||||
kv_to_page_indices(kernel_tokens, self.PS).tolist(),
|
||||
(physical_pages * mult).tolist(),
|
||||
)
|
||||
|
||||
def test_swa_kernel_tombstone_still_lands_on_sink(self):
|
||||
"""The scaled stride must not break the tombstone clamp: a tombstoned
|
||||
page's ids (v2p == -1 -> -stride + offset, negative for every in-page
|
||||
|
||||
@@ -1,10 +1,13 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
|
||||
import torch
|
||||
from test_unified_byte_budget_sizing import _swa_factory
|
||||
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||
|
||||
@@ -57,5 +60,38 @@ class TestSWACpuCopyFilter(unittest.TestCase):
|
||||
)
|
||||
|
||||
|
||||
class TestUnifiedSWATransfers(CustomTestCase):
|
||||
def test_cpu_copy_round_trip_with_request_index_and_swa_tombstone(self):
|
||||
bundle = _swa_factory(page_size=4)
|
||||
allocator = bundle.token_to_kv_pool_allocator
|
||||
pool = bundle.token_to_kv_pool
|
||||
indices = allocator.alloc(12)
|
||||
allocator.free_swa(indices[:4])
|
||||
raw = bundle.unified_memory_pool._raw
|
||||
raw.copy_(torch.arange(raw.numel()).remainder(251).to(torch.uint8))
|
||||
full_pages = allocator.translate_kv_indices_for_transfer(indices)[::4] // 4
|
||||
swa_pages = allocator.translate_swa_indices_for_transfer(indices[4:])[::4] // 4
|
||||
full_buffer, swa_buffer = (
|
||||
raw[: side.num_pages * side.entry_bytes_per_page].view(
|
||||
side.num_pages, side.entry_bytes_per_page
|
||||
)
|
||||
for side in (allocator.full_attn_allocator, allocator.swa_attn_allocator)
|
||||
)
|
||||
expected_full = full_buffer[full_pages].clone()
|
||||
expected_swa = swa_buffer[swa_pages].clone()
|
||||
|
||||
# The buffers are CPU tensors; no device synchronization is needed.
|
||||
with mock.patch(
|
||||
"sglang.srt.mem_cache.memory_pool.current_platform.synchronize"
|
||||
):
|
||||
saved = pool.get_cpu_copy(indices, req_pool_index=0)
|
||||
raw.zero_()
|
||||
pool.load_cpu_copy(saved, indices, req_pool_index=0)
|
||||
|
||||
self.assertTrue(torch.equal(full_buffer[full_pages], expected_full))
|
||||
self.assertTrue(torch.equal(swa_buffer[swa_pages], expected_swa))
|
||||
self.assertEqual(saved["swa_mask"].tolist(), [False] * 4 + [True] * 8)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -344,17 +344,9 @@ class TestUnifiedMHATokenToKVPool(unittest.TestCase):
|
||||
"envelope move did not relocate exactly the named pages",
|
||||
)
|
||||
|
||||
def test_transfer_entry_points_fail_loud(self):
|
||||
"""The entry points that assume per-layer buffers indexed by TOKEN id
|
||||
would silently mis-index against the row space (or hit a missing-attr
|
||||
AttributeError), so each must raise. `get_contiguous_buf_infos` is NOT
|
||||
among them: PD addresses this pool as whole page envelopes, pinned by
|
||||
`test_pd_registration_is_one_whole_envelope` below."""
|
||||
def test_prefix_valid_entry_point_fails_loud(self):
|
||||
"""Prefix-valid writes still assume token-major buffer indexing."""
|
||||
_, pool = _make_pool_and_kv(1)
|
||||
with self.assertRaises(NotImplementedError):
|
||||
pool.get_cpu_copy(torch.tensor([1]))
|
||||
with self.assertRaises(NotImplementedError):
|
||||
pool.load_cpu_copy(None, torch.tensor([1]))
|
||||
with self.assertRaises(NotImplementedError):
|
||||
pool.set_kv_buffer_prefix_valid()
|
||||
|
||||
|
||||
@@ -26,9 +26,13 @@ from sglang.test.ci.ci_register import register_cpu_ci
|
||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.allocator.unified_mamba import (
|
||||
UnifiedMambaTokenToKVPoolAllocator,
|
||||
)
|
||||
from sglang.srt.mem_cache.allocator.unified_sub_pool import MultiEndedAllocator
|
||||
from sglang.srt.mem_cache.layout.page_major import build_mla_views
|
||||
from sglang.srt.mem_cache.unified_memory_pool import (
|
||||
@@ -37,6 +41,7 @@ from sglang.srt.mem_cache.unified_memory_pool import (
|
||||
UnifiedKVPool,
|
||||
UnifiedMLATokenToKVPool,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
_DEV = "cpu"
|
||||
|
||||
@@ -219,6 +224,48 @@ class TestUnifiedMLATokenToKVPool(unittest.TestCase):
|
||||
k[7] = 2.5
|
||||
self.assertTrue(torch.all(v[7] == 2.5))
|
||||
|
||||
def test_cpu_copy_round_trips_through_physical_ids(self):
|
||||
"""REGRESSION: the host copy for decode retraction is addressed by
|
||||
PHYSICAL token ids, but this pool's `kv_buffer` views are indexed by
|
||||
kernel-facing ids. Without the rewrite the parent read a different row
|
||||
and the restore silently returned other tokens' KV."""
|
||||
for ps in (1, 4):
|
||||
with self.subTest(page_size=ps):
|
||||
pool, kv_pool = self._make(ps=ps)
|
||||
phys = torch.tensor([0, 1, ps, ps + 1], dtype=torch.int64)
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
kv_pool._physical_to_kernel_indices(phys),
|
||||
torch.tensor(
|
||||
[_kernel_id(int(t), ps, _L) for t in phys],
|
||||
dtype=torch.int64,
|
||||
),
|
||||
)
|
||||
)
|
||||
for layer in range(_L):
|
||||
kv_pool.get_key_buffer(layer)[
|
||||
kv_pool._physical_to_kernel_indices(phys)
|
||||
] = float(layer + 1)
|
||||
|
||||
with (
|
||||
get_parallel().override(dcp_enabled=False),
|
||||
mock.patch(
|
||||
"sglang.srt.mem_cache.memory_pool.current_platform.synchronize"
|
||||
),
|
||||
):
|
||||
saved = kv_pool.get_cpu_copy(phys)
|
||||
pool._raw.zero_()
|
||||
kv_pool.load_cpu_copy(saved, phys)
|
||||
|
||||
for layer in range(_L):
|
||||
restored = kv_pool.get_key_buffer(layer)[
|
||||
kv_pool._physical_to_kernel_indices(phys)
|
||||
]
|
||||
self.assertTrue(
|
||||
torch.all(restored == float(layer + 1)),
|
||||
f"layer {layer} did not round-trip at page_size {ps}",
|
||||
)
|
||||
|
||||
def test_move_kv_cache_moves_page_envelopes(self):
|
||||
"""Whole page envelopes relocate, in raw bytes and (at ps=4) as read
|
||||
back through the per-layer views at the destination kernel ids."""
|
||||
@@ -334,5 +381,58 @@ class TestTranslateKvLocForKernel(unittest.TestCase):
|
||||
self.assertTrue(torch.all(x == no_out))
|
||||
|
||||
|
||||
class _RecordingHybridPool:
|
||||
"""Stands in for `UnifiedHybridLinearKVPool`, recording the ids it is handed."""
|
||||
|
||||
def __init__(self, full_kv_pool, mamba_pool):
|
||||
self.full_kv_pool = full_kv_pool
|
||||
self.mamba_pool = mamba_pool
|
||||
self.seen = None
|
||||
|
||||
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||||
self.seen = indices.clone()
|
||||
return {"full": None}
|
||||
|
||||
def load_cpu_copy(
|
||||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||||
):
|
||||
self.seen = indices.clone()
|
||||
|
||||
|
||||
class TestMambaAllocatorCpuCopyIsPhysical(unittest.TestCase):
|
||||
"""REGRESSION: decode retraction calls the allocator's `get_cpu_copy` with
|
||||
`req_to_token` rows, which hold VIRTUAL ids. This composite inherited the
|
||||
raising base, and a plain delegate would have been just as wrong -- the
|
||||
unified pools read those ids as PHYSICAL."""
|
||||
|
||||
def _build(self, ps=1):
|
||||
pool, _, _ = _make_unified(page_size=ps)
|
||||
kvcache = _RecordingHybridPool(
|
||||
_FakeKVCache(pool.max_slots("full")),
|
||||
_FakeKVCache(pool.max_slots("mamba")),
|
||||
)
|
||||
with get_parallel().override(dcp_enabled=False, attn_dcp_size=1):
|
||||
allocator = UnifiedMambaTokenToKVPoolAllocator(
|
||||
unified_buffer=pool, kvcache=kvcache, device=_DEV, page_size=ps
|
||||
)
|
||||
return allocator, kvcache
|
||||
|
||||
def test_pool_is_handed_physical_token_ids(self):
|
||||
alloc, kvcache = self._build()
|
||||
virtual = alloc.alloc(4)
|
||||
self.assertIsNotNone(virtual)
|
||||
virtual = virtual.to(torch.int64)
|
||||
physical = alloc.full_attn_allocator.translate_kv_loc(virtual)
|
||||
# Not identity here, so a delegate that passed the virtual ids straight
|
||||
# through would read and restore other tokens' rows.
|
||||
self.assertFalse(torch.equal(physical, virtual))
|
||||
|
||||
alloc.get_cpu_copy(virtual, req_pool_index=0)
|
||||
self.assertTrue(torch.equal(kvcache.seen, physical))
|
||||
|
||||
alloc.load_cpu_copy({"full": None}, virtual, req_pool_index=0)
|
||||
self.assertTrue(torch.equal(kvcache.seen, physical))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -202,6 +202,58 @@ class TestUnifiedTriPool(unittest.TestCase):
|
||||
self.assertEqual(allocator.swa_attn_allocator.allocated_count(), 0)
|
||||
self.assertEqual(allocator.available_size(), before)
|
||||
|
||||
def test_pd_short_tail_fits_beyond_joint_capacity(self):
|
||||
for lazy in (False, True):
|
||||
for tail_len in (0, 5):
|
||||
with self.subTest(lazy=lazy, tail_len=tail_len):
|
||||
_, allocator, _, _ = self._build(page_size=4, lazy_compaction=lazy)
|
||||
full = allocator.full_attn_allocator
|
||||
length = allocator.available_size() + 4
|
||||
self.assertFalse(allocator.can_reserve(length, length))
|
||||
self.assertTrue(allocator.can_reserve(length, tail_len))
|
||||
prefix = torch.tensor([0], dtype=torch.int64)
|
||||
seq = torch.tensor([length], dtype=torch.int64)
|
||||
with patch.object(
|
||||
full,
|
||||
"alloc_extend",
|
||||
side_effect=lambda *a, **kw: full.alloc(length),
|
||||
):
|
||||
virtual = allocator.alloc_extend_swa_tail(
|
||||
prefix,
|
||||
prefix,
|
||||
seq,
|
||||
seq,
|
||||
torch.tensor([-1]),
|
||||
length,
|
||||
tail_len,
|
||||
)
|
||||
self.assertIsNotNone(virtual)
|
||||
self.assertEqual(full.allocated_count(), length)
|
||||
self.assertEqual(
|
||||
allocator.swa_attn_allocator.allocated_count(),
|
||||
-(-tail_len // 4) * 4,
|
||||
)
|
||||
self.assertEqual(allocator.verify_byte_accounting(), [])
|
||||
allocator.free(virtual)
|
||||
self.assertEqual(full.allocated_count(), 0)
|
||||
self.assertEqual(allocator.swa_attn_allocator.allocated_count(), 0)
|
||||
|
||||
def test_pd_tail_rejects_full_capacity_shortfall(self):
|
||||
_, allocator, _, _ = self._build(page_size=4)
|
||||
full = allocator.full_attn_allocator
|
||||
length = full.available_size() + 4
|
||||
prefix = torch.tensor([0], dtype=torch.int64)
|
||||
seq = torch.tensor([length], dtype=torch.int64)
|
||||
with patch.object(full, "alloc_extend") as extend:
|
||||
self.assertIsNone(
|
||||
allocator.alloc_extend_swa_tail(
|
||||
prefix, prefix, seq, seq, torch.tensor([-1]), length, 0
|
||||
)
|
||||
)
|
||||
extend.assert_not_called()
|
||||
self.assertEqual(full.allocated_count(), 0)
|
||||
self.assertEqual(allocator.swa_attn_allocator.allocated_count(), 0)
|
||||
|
||||
def test_empty_float_is_transparent_to_the_ends(self):
|
||||
_, allocator, _, _ = self._build()
|
||||
fa = allocator.full_attn_allocator
|
||||
@@ -1408,5 +1460,64 @@ class TestFloatHoleCreditIsPerSide(unittest.TestCase):
|
||||
self.assertEqual(flt._byte_accounting_violations(), [])
|
||||
|
||||
|
||||
class TestPreallocIsPricedOnTheSharedGrid(unittest.TestCase):
|
||||
"""REGRESSION: PD admission compared FULL and SWA against per-side token
|
||||
budgets, but each side's `available_size` credits the peer's drainable
|
||||
holes, so a pair that each side can host alone can be jointly infeasible.
|
||||
Such a pair was admitted and then refused inside `alloc_extend_swa_tail`."""
|
||||
|
||||
def _build(self, **kw):
|
||||
return TestUnifiedTriPool._build(self, **kw)
|
||||
|
||||
def test_a_pair_each_side_can_host_alone_is_still_refused(self):
|
||||
# page_size 1 leaves no slack between the per-side and joint views;
|
||||
# the double-count only has room to show on a paged grid.
|
||||
_, allocator, _, _ = self._build(page_size=4)
|
||||
full_demand = allocator.full_available_size()
|
||||
swa_demand = allocator.swa_available_size()
|
||||
self.assertGreater(min(full_demand, swa_demand), 0)
|
||||
# Each side alone reports room for its own half ...
|
||||
self.assertLessEqual(full_demand, allocator.full_available_size())
|
||||
self.assertLessEqual(swa_demand, allocator.swa_available_size())
|
||||
# ... yet the two draw on the same bytes, so the grid refuses the pair.
|
||||
self.assertFalse(
|
||||
allocator._fits_page_demand(
|
||||
-(-full_demand // allocator.page_size),
|
||||
-(-swa_demand // allocator.page_size),
|
||||
)
|
||||
)
|
||||
self.assertFalse(
|
||||
allocator.prealloc_fits(
|
||||
MagicMock(),
|
||||
full_demand,
|
||||
swa_demand,
|
||||
full_budget_tokens=full_demand,
|
||||
swa_budget_tokens=swa_demand,
|
||||
)
|
||||
)
|
||||
|
||||
def test_the_scheduler_budget_still_binds(self):
|
||||
_, allocator, _, _ = self._build()
|
||||
page_size = allocator.page_size
|
||||
self.assertTrue(
|
||||
allocator.prealloc_fits(
|
||||
MagicMock(),
|
||||
page_size,
|
||||
page_size,
|
||||
full_budget_tokens=page_size,
|
||||
swa_budget_tokens=page_size,
|
||||
)
|
||||
)
|
||||
self.assertFalse(
|
||||
allocator.prealloc_fits(
|
||||
MagicMock(),
|
||||
page_size,
|
||||
page_size,
|
||||
full_budget_tokens=page_size - 1,
|
||||
swa_budget_tokens=page_size,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user