Make CP HiCache residency owner-lane deterministic

CP shared KV cannot treat capacity as a scalar token count: cache-hit load-back and fresh extend allocation both have to preserve the logical page owner pattern or later direct writes, HiCache reload, and prefix materialization can read the wrong lane. This change moves the critical paths to owner-lane plans, makes owner-lane exhaustion recoverable during prefill scheduling, and routes shared-KV prefix prefetch through prefetch-stream-safe KV getters so HiCache layer-load waits do not attach to the forward stream.

Constraint: CP shared KV correctness depends on page owner lane preservation across allocation, backup, load, eviction, and prefix materialization.

Constraint: Avoid adding CP/global collectives for capacity agreement; derive capacity from deterministic local owner-lane state.

Rejected: Keep SGLANG_DISABLE_TAI_OWNER_SELECT fallback | legacy allocation can silently break owner-lane invariants.

Rejected: Scalar total-token eviction for CP HiCache load-back | total capacity can be sufficient while the required owner lane is exhausted.

Confidence: medium

Scope-risk: broad

Directive: Do not reintroduce silent legacy fallback in owner-lane paths; unexpected owner-lane failure must be warning-level fail-closed or recoverable capacity wait.

Tested: Remote g0034 container PYTHONPATH=python python -m pytest test/registered/unit/mem_cache/test_alloc_pages_with_owners.py test/registered/unit/mem_cache/test_cp_shared_kv_layout.py test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py -q -> 95 passed.

Tested: Local py_compile for modified runtime/cache/scheduler modules.

Not-tested: Full CUDA ETE performance trace for cache-hit overlap and MTP accept-rate impact.

Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
laoyao0822
2026-05-28 05:54:23 +08:00
parent 2c94b8de23
commit ff33446787
12 changed files with 2913 additions and 136 deletions

View File

@@ -0,0 +1,351 @@
# NSA Prefill CP / CP Shared KV Batch Size > 1 Exploration
Date: 2026-05-28
Scope:
- NSA prefill context parallelism.
- `--nsa-prefill-cp-mode in-seq-split`.
- CP shared KV.
- HiCache L2/L1 load path.
- MTP / EAGLE draft KV integration.
This note records why true `ForwardBatch(batch_size > 1)` is not currently a
safe near-term throughput lever for the CP shared-KV path, and why the safer
near-term direction is multi-slot overlap where each slot keeps
`batch_size == 1`.
## Motivation
In the cache-hit-heavy workload, observed prefix cache hit can be above 90%.
The remaining extend length is often only 2K-10K tokens. A single request then
does not always provide enough work to fill the GPU. It is natural to ask
whether prefill should group multiple short cache-hit extends into one batch.
SGLang's generic prefill scheduler and `ScheduleBatch.prepare_for_extend()` can
construct a multi-request extend batch. The blocker is the specialized NSA
prefill CP + CP shared-KV path layered underneath it.
## Terminology
### True batch
One `ForwardBatch` contains multiple sequences:
```text
ForwardBatch(batch_size=N)
```
All NSA metadata, page tables, top-k transforms, materialize remaps, direct
writes, and MTP draft paths must be batch-aware.
### Multi-slot overlap
Multiple independent request slots are scheduled together, but each slot owns
its own `ScheduleBatch`, `ModelWorkerBatch`, and `ForwardBatch`:
```text
slot0: ForwardBatch(batch_size=1)
slot1: ForwardBatch(batch_size=1)
...
```
This is the Phase9 direction. It keeps current single-request CP invariants
while allowing communication/materialize from one slot to overlap useful work
from another slot.
## Current evidence: true batch > 1 is not supported on the CP shared-KV path
### 1. In-seq CP metadata is single-request by construction
`prepare_input_dp_with_cp_dsa()` builds the zigzag split for one full-length
request.
Relevant code:
- `python/sglang/srt/layers/attention/nsa/utils.py`
- `prepare_input_dp_with_cp_dsa(...)`
- comment: `# just support batch = 1`
- `bs_per_cp_group = 1`
- computes one `split_list`, one `zigzag_index`, one pair of
`kv_len_prev/next`, and one pair of `actual_seq_q_prev/next`.
Implication:
For true batch > 1, concatenating multiple requests and passing their total
token count into this function would treat them as one long sequence. That does
not preserve per-request causal boundaries, per-request prefix/suffix splits,
or per-request owner-lane page semantics.
### 2. Page-aligned split only activates for batch size 1
The shared-KV path depends on page-aligned in-seq split so direct writes land in
owner lanes and later materialize/load paths can reproduce the same logical
page ownership.
Relevant code:
- `python/sglang/srt/layers/attention/nsa/utils.py`
- `_build_in_seq_split_for_forward_batch(...)`
- page-aligned split requires:
```python
len(forward_batch.extend_seq_lens_cpu) == 1
len(forward_batch.extend_prefix_lens_cpu) == 1
```
If those conditions fail, the code falls back to token-balanced split metadata.
Implication:
True batch > 1 loses the page-aligned split contract required by the current
CP shared-KV direct-write path.
### 3. Compute-owner allocation supports only batch size 1
CP shared KV relies on owner-lane-aware allocation: newly allocated logical
pages should come from the lane owned by the CP rank that computes the page.
Relevant code:
- `python/sglang/srt/mem_cache/allocator.py`
- `alloc_extend_compute_owner(...)`
- raises if `len(prefix_lens_cpu) != 1` or `len(seq_lens_cpu) != 1`.
- `python/sglang/srt/mem_cache/common.py`
- `alloc_paged_token_slots_extend(...)`
- only builds `page_compute_owners` when `len(prefix_lens_cpu) == 1`.
- batch > 1 records reason `multi_batch` and falls back to legacy allocation.
Implication:
True batch > 1 either fails the owner-lane allocator or falls back to legacy
allocation. That fallback weakens the invariant that direct write, HiCache
backup/load, and later CP materialize assume: logical page owner must match the
physical owner lane.
### 4. Direct-write path depends on single-request page-aligned metadata
The MLA KV and NSA index direct-write paths get this rank's local output locs
through:
- `python/sglang/srt/layers/attention/nsa/utils.py`
- `get_cp_shared_kv_local_out_cache_loc(...)`
- checks `metadata.page_aligned`.
- checks `sum(metadata.split_list) == out_cache_loc.numel()`.
- uses `cp_split_and_rebuild_1d(...)` with the single-request CP metadata.
- validates local loc ownership by `layout.owned_by_this_rank(...)`.
Callers:
- `python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py`
- `_maybe_write_cp_shared_local_mla_kv(...)`
- `python/sglang/srt/layers/attention/nsa/nsa_indexer.py`
- `_store_cp_shared_local_index_k_cache(...)`
Implication:
If true batch > 1 does not produce per-request page-aligned metadata and
owner-lane-aligned `out_cache_loc`, direct write should fall back. This removes
a major reason CP shared KV is fast.
### 5. Phase8 prefix prefetch is explicitly batch size 1
The current L1 shared-KV prefix prefetch path is conservative and only supports
one request.
Relevant code:
- `python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py`
- `CpSharedKVMlaPrefetcher.maybe_create(...)`
- rejects `forward_batch.batch_size != 1`.
- requires `len(extend_prefix_lens_cpu) == 1`.
- same file:
- `CpSharedKVIndexPrefetcher.maybe_create(...)`
- same batch-size and prefix-lens restrictions.
Implication:
If true batch > 1 is used in a cache-hit-heavy workload, the exact path we need
to hide -- prefix materialize/reduce from L1 shared KV -- is disabled. The code
falls back to synchronous materialize in the attention path.
### 6. NSA top-k in-seq CP pair path has single-batch assumptions
Relevant code:
- `python/sglang/srt/layers/attention/nsa/nsa_indexer.py`
- `_get_topk_in_seq_cp_pair(...)`
- TODO: `support mutil-batch`
- uses one `kv_len_prev/next` pair and one `actual_seq_q_prev/next` pair from
`forward_batch.nsa_cp_metadata`.
Implication:
True batch > 1 would require per-request CP pair metadata and top-k outputs
that preserve each sequence's independent query/key ranges.
### 7. MTP / EAGLE draft path depends on the same single-request CP metadata
Relevant code:
- `python/sglang/srt/models/deepseek_nextn.py`
- draft model local path uses `cp_split_and_rebuild_1d(...)`,
`cp_split_and_rebuild_data(...)`, and `cp_collect_last_token_hidden(...)`.
- `python/sglang/srt/layers/attention/nsa/utils.py`
- `_in_seq_collect_last_token(...)` only uses the exact in-seq owner/offset
calculation when `bs == 1`.
Implication:
Even if target-model CP shared KV were made batch-aware, MTP/EAGLE would still
need per-request draft hidden state slicing and per-request last-token
collection semantics before true batch > 1 can be declared safe.
### 8. HiCache load is closer to batch-capable, but it inherits owner semantics
The CP HiCache load path can merge load ops across nodes:
- `python/sglang/srt/managers/cache_controller.py`
- `load_cp(nodes_to_load, ...)`
- collects `page_owners` across nodes.
- allocates fresh device pages through `alloc_pages_with_owners(...)`.
- appends target and draft load ops.
- `start_loading()` merges load queues and performs per-layer load.
This part is not the primary batch-size blocker.
However, it assumes every backed-up node carries valid CP `page_owners` and
owned-position metadata. If the original true-batch allocation/split path does
not preserve owner-lane semantics, HiCache load can faithfully reproduce an
already-wrong owner pattern.
## Why true batch may not improve this workload immediately
For cache-hit-heavy short extends, batching increases query work, but it can
also increase or re-expose prefix materialize/reduce work:
```text
cache-hit prefix pages
-> L2 host to L1 device load
-> L1 shared-KV materialize/reduce
-> top-k / attention
```
If true batch disables Phase8 prefetch and direct write, the critical path may
become worse even though more requests are grouped.
The throughput bottleneck is likely not only "too few query tokens". It is also
"attention-visible shared-KV prefix preparation is not fully hidden".
## Recommended near-term direction
Use multi-slot overlap before true batch > 1.
Rationale:
1. Keeps current per-slot `batch_size == 1` invariants.
2. Preserves page-aligned split, owner-lane allocation, direct write, and
Phase8 prefetch eligibility.
3. Allows slot A's attention/MLP to hide slot B's CP materialize/prefetch/load
work.
4. Avoids redesigning `NSAContextParallelMetadata`, top-k, MTP hidden slicing,
and direct-write remaps in one large step.
This matches the Phase9 plan:
- `docs/advanced_features/nsa_prefill_cp_phase9_two_request_overlap_plan.md`
- distinguishes multi-slot from true `ForwardBatch(batch_size=N)`.
- explicitly keeps each initial slot at `batch_size == 1`.
- defers true batch support until batch-aware NSA CP metadata exists.
## Requirements for future true batch > 1 support
True batch support should be treated as a separate phase with explicit
fail-fast gates until all requirements below are implemented.
### P0: hard guard / loud fallback
When `uses_cp_shared_kv` and `batch_size > 1`, the system should not silently
fall back to a slow or semantically weaker path.
Acceptable behavior during development:
- fail fast for unsupported combinations; or
- emit a loud warning with exact fallback reason and expected performance loss.
### P1: batch-aware NSA CP metadata
`NSAContextParallelMetadata` needs per-request fields:
- per-request `split_list`;
- per-request `zigzag_index`;
- per-request `kv_len_prev/next`;
- per-request `actual_seq_q_prev/next`;
- per-request page-aligned segment ranges;
- a flattened representation usable by kernels without losing request
boundaries.
### P2: batch-aware page-aligned split and owner-lane allocation
For each request in the batch:
1. compute page-aligned split independently;
2. build page owners for that request's suffix;
3. concatenate page-owner plans in the same order as `out_cache_loc`;
4. allocate pages with matching owner lanes;
5. preserve per-request offsets for direct write and materialize remap.
### P3: batch-aware direct write
`get_cp_shared_kv_local_out_cache_loc(...)` cannot treat `out_cache_loc` as one
single in-seq split. It needs to produce local locs by request, then concatenate
them in local compute order while preserving owner-lane validation.
### P4: batch-aware Phase8 L1 shared-KV prefetch
The MLA and index prefetchers need:
- per-request prefix page ranges;
- per-request prefix page alignment checks;
- a batched slot-remap layout that preserves the current dense-page semantics;
- deterministic collective ordering across ranks.
### P5: batch-aware top-k CP pair path
`_get_topk_in_seq_cp_pair(...)` needs per-request CP pair metadata rather than
one global `kv_len_prev/next` pair.
### P6: batch-aware MTP / EAGLE local path
Draft model CP local input, spec hidden states, and last-token hidden
collection need per-request CP split and gather semantics.
### P7: tests
Minimum tests before enabling true batch > 1:
- two cache-hit requests with different prefix lengths;
- page-aligned and non-page-aligned suffixes;
- short radix-hit suffix with replicated compute path;
- HiCache host hit + load-back;
- draft/MTP enabled with target/draft KV mirrored;
- compare batch > 1 output against sequential batch-size-1 execution;
- verify no direct-write fallback under supported cases.
## Practical conclusion
For the current high-cache-hit workload, true `ForwardBatch(batch_size > 1)` is
not the right first optimization. The current CP shared-KV path has hard and
soft batch-size-1 assumptions in split metadata, owner-lane allocation,
direct-write, prefix prefetch, top-k, and draft MTP.
The lower-risk path is:
```text
multi-slot CP overlap first
-> each slot keeps batch_size == 1
-> preserve shared-KV correctness invariants
-> overlap CP materialize/load/prefix-prefetch with peer slot compute
-> only later add true batch-aware CP metadata
```

File diff suppressed because it is too large Load Diff

View File

@@ -45,6 +45,50 @@ def _index_prefetch_fallback_log(reason: str, message: str, *args) -> None:
)
def _prefetch_pool_get_key_buffer(
*,
token_to_kv_pool: Any,
layer_id: int,
stream: torch.cuda.Stream,
path: str,
) -> torch.Tensor:
getter = getattr(token_to_kv_pool, "get_key_buffer_for_prefetch", None)
if getter is not None:
return getter(layer_id, stream)
if getattr(token_to_kv_pool, "layer_transfer_counter", None) is not None:
logger.warning(
"[CP_SHARED_KV_FALLBACK][%s_prefetch] "
"reason=prefetch_safe_getter_unavailable layer_id=%s pool=%s "
"has_layer_transfer_counter=True",
path,
layer_id,
type(token_to_kv_pool).__name__,
)
return token_to_kv_pool.get_key_buffer(layer_id)
def _prefetch_pool_get_index_buffer(
*,
token_to_kv_pool: Any,
layer_id: int,
stream: torch.cuda.Stream,
) -> torch.Tensor:
getter = getattr(
token_to_kv_pool, "get_index_k_with_scale_buffer_for_prefetch", None
)
if getter is not None:
return getter(layer_id, stream)
if getattr(token_to_kv_pool, "layer_transfer_counter", None) is not None:
_index_prefetch_fallback_log(
"prefetch_safe_getter_unavailable",
"pool has active layer transfer counter but no prefetch-safe index "
"getter. layer_id=%s pool=%s",
layer_id,
type(token_to_kv_pool).__name__,
)
return token_to_kv_pool.get_index_k_with_scale_buffer(layer_id=layer_id)
def _is_cuda_stream_capturing() -> bool:
if not torch.cuda.is_available():
return False
@@ -235,9 +279,15 @@ class CpSharedKVMlaPrefetcher:
)
return None
prefetch_stream = stream if stream is not None else torch.cuda.Stream()
try:
first_layer_id = int(getattr(token_to_kv_pool, "start_layer", 0))
kv_cache = token_to_kv_pool.get_key_buffer(first_layer_id)
kv_cache = _prefetch_pool_get_key_buffer(
token_to_kv_pool=token_to_kv_pool,
layer_id=first_layer_id,
stream=prefetch_stream,
path="mla",
)
remap = get_or_build_shared_token_kv_slot_remap(
forward_batch,
kv_cache=kv_cache,
@@ -276,7 +326,7 @@ class CpSharedKVMlaPrefetcher:
dense_num_pages=remap.dense_num_pages,
owned_prefix_pages=owned_prefix_pages,
owned_total_pages=owned_total_pages,
stream=stream,
stream=prefetch_stream,
)
def _layer_in_pool(self, token_to_kv_pool: Any, layer_id: int) -> bool:
@@ -447,7 +497,12 @@ class CpSharedKVMlaPrefetcher:
return
try:
kv_cache = token_to_kv_pool.get_key_buffer(next_layer_id)
kv_cache = _prefetch_pool_get_key_buffer(
token_to_kv_pool=token_to_kv_pool,
layer_id=next_layer_id,
stream=self.stream,
path="mla",
)
except Exception:
logger.exception(
"Failed to get next-layer KV cache for CP shared KV MLA prefetch."
@@ -461,31 +516,61 @@ class CpSharedKVMlaPrefetcher:
return
try:
dense_kv_cache = kv_cache.new_zeros(
(self.dense_num_pages * self.page_size, *kv_cache.shape[1:])
)
self._log_next_layer(
next_layer_id,
"start_prefix_begin next_layer=%s start_slot=0 end_slot=%s "
"dense_rows=%s",
next_layer_id,
self.prefix_pages,
int(dense_kv_cache.shape[0]),
)
materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=self.slot_logical_pages,
layout=self.layout,
page_size=self.page_size,
start_slot=0,
end_slot=self.prefix_pages,
)
current_stream = torch.cuda.current_stream()
self.stream.wait_stream(current_stream)
prefix_rows = slot_range_to_token_slice(
self.page_size,
0,
self.prefix_pages,
)
with torch.cuda.stream(self.stream):
dense_kv_cache = kv_cache.new_zeros(
(self.dense_num_pages * self.page_size, *kv_cache.shape[1:])
)
self._log_next_layer(
next_layer_id,
"start_prefix_begin next_layer=%s start_slot=0 end_slot=%s "
"dense_rows=%s",
next_layer_id,
self.prefix_pages,
int(dense_kv_cache.shape[0]),
)
materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=self.slot_logical_pages,
layout=self.layout,
page_size=self.page_size,
start_slot=0,
end_slot=self.prefix_pages,
)
event = _all_reduce_materialized_buffer_async(
dense_kv_cache[prefix_rows],
cp_size=self.layout.cp_size,
stream=self.stream,
nvtx_source="mla.prefetch_prefix",
nvtx_layer_id=next_layer_id,
nvtx_cp_rank=self.layout.cp_rank,
nvtx_rows=(prefix_rows.start, prefix_rows.stop),
)
if event is None:
self.disabled = True
logger.warning(
"[CP_SHARED_KV_FALLBACK][mla_prefetch] "
"reason=async_reduce_unavailable layer_id=%s cp_rank=%s "
"cp_size=%s prefix_pages=%s total_slots=%s",
next_layer_id,
self.layout.cp_rank,
self.layout.cp_size,
self.prefix_pages,
self.total_slots,
)
self._log_next_layer(
next_layer_id,
"start_disable reason=async_reduce_unavailable next_layer=%s",
next_layer_id,
)
return
except Exception:
logger.exception("Failed to start CP shared KV MLA prefix prefetch.")
self.disabled = True
@@ -500,12 +585,14 @@ class CpSharedKVMlaPrefetcher:
layer_id=next_layer_id,
dense_kv_cache=dense_kv_cache,
prefix_rows=prefix_rows,
event=event,
)
self.handles[next_layer_id] = handle
self.pending_attention_handle = handle
self._log_next_layer(
next_layer_id,
"start next_layer=%s prefix_pages=%s prefix_rows=%s dense_rows=%s",
"start next_layer=%s prefix_pages=%s prefix_rows=%s dense_rows=%s "
"reduce_enqueued=True",
next_layer_id,
self.prefix_pages,
prefix_rows.stop - prefix_rows.start,
@@ -776,10 +863,13 @@ class CpSharedKVIndexPrefetcher:
)
return None
prefetch_stream = stream if stream is not None else torch.cuda.Stream()
try:
first_layer_id = int(getattr(token_to_kv_pool, "start_layer", 0))
page_buffer = token_to_kv_pool.get_index_k_with_scale_buffer(
layer_id=first_layer_id
page_buffer = _prefetch_pool_get_index_buffer(
token_to_kv_pool=token_to_kv_pool,
layer_id=first_layer_id,
stream=prefetch_stream,
)
remap = get_or_build_shared_paged_buffer_slot_remap(
forward_batch,
@@ -822,7 +912,7 @@ class CpSharedKVIndexPrefetcher:
dense_num_pages=remap.dense_num_pages,
owned_prefix_pages=owned_prefix_pages,
owned_total_pages=owned_total_pages,
stream=stream,
stream=prefetch_stream,
)
def _layer_in_pool(self, token_to_kv_pool: Any, layer_id: int) -> bool:
@@ -990,8 +1080,10 @@ class CpSharedKVIndexPrefetcher:
return
try:
page_buffer = token_to_kv_pool.get_index_k_with_scale_buffer(
layer_id=next_layer_id
page_buffer = _prefetch_pool_get_index_buffer(
token_to_kv_pool=token_to_kv_pool,
layer_id=next_layer_id,
stream=self.stream,
)
except Exception:
logger.exception(
@@ -1006,26 +1098,56 @@ class CpSharedKVIndexPrefetcher:
return
try:
dense_page_buffer = page_buffer.new_zeros(
(self.dense_num_pages, *page_buffer.shape[1:])
)
self._log_next_layer(
next_layer_id,
"index_start_prefix_begin next_layer=%s start_slot=0 "
"end_slot=%s dense_pages=%s",
next_layer_id,
self.prefix_pages,
int(dense_page_buffer.shape[0]),
)
materialize_local_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=self.slot_logical_pages,
layout=self.layout,
start_slot=0,
end_slot=self.prefix_pages,
)
current_stream = torch.cuda.current_stream()
self.stream.wait_stream(current_stream)
prefix_rows = slot_range_to_page_slice(0, self.prefix_pages)
with torch.cuda.stream(self.stream):
dense_page_buffer = page_buffer.new_zeros(
(self.dense_num_pages, *page_buffer.shape[1:])
)
self._log_next_layer(
next_layer_id,
"index_start_prefix_begin next_layer=%s start_slot=0 "
"end_slot=%s dense_pages=%s",
next_layer_id,
self.prefix_pages,
int(dense_page_buffer.shape[0]),
)
materialize_local_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=self.slot_logical_pages,
layout=self.layout,
start_slot=0,
end_slot=self.prefix_pages,
)
event = _all_reduce_materialized_buffer_async(
dense_page_buffer[prefix_rows],
cp_size=self.layout.cp_size,
stream=self.stream,
nvtx_source="index.prefetch_prefix",
nvtx_layer_id=next_layer_id,
nvtx_cp_rank=self.layout.cp_rank,
nvtx_rows=(prefix_rows.start, prefix_rows.stop),
)
if event is None:
self.disabled = True
_index_prefetch_fallback_log(
"async_reduce_unavailable",
"async reduce unavailable. layer_id=%s cp_rank=%s cp_size=%s "
"prefix_pages=%s total_slots=%s",
next_layer_id,
self.layout.cp_rank,
self.layout.cp_size,
self.prefix_pages,
self.total_slots,
)
self._log_next_layer(
next_layer_id,
"index_start_disable reason=async_reduce_unavailable next_layer=%s",
next_layer_id,
)
return
except Exception:
logger.exception("Failed to start CP shared KV index prefix prefetch.")
self.disabled = True
@@ -1040,12 +1162,14 @@ class CpSharedKVIndexPrefetcher:
layer_id=next_layer_id,
dense_page_buffer=dense_page_buffer,
prefix_rows=prefix_rows,
event=event,
)
self.handles[next_layer_id] = handle
self.pending_attention_handle = handle
self._log_next_layer(
next_layer_id,
"index_start next_layer=%s prefix_pages=%s dense_pages=%s",
"index_start next_layer=%s prefix_pages=%s dense_pages=%s "
"reduce_enqueued=True",
next_layer_id,
self.prefix_pages,
int(dense_page_buffer.shape[0]),

View File

@@ -61,6 +61,10 @@ class LayerLoadingEvent:
def wait(self, layer_index: int):
device_module.current_stream().wait_event(self.load_events[layer_index])
def wait_on_stream(self, layer_index: int, stream):
assert 0 <= layer_index < self._num_layers
stream.wait_event(self.load_events[layer_index])
@property
def finish_event(self):
return self.load_events[-1]
@@ -101,6 +105,12 @@ class LayerDoneCounter:
for consumer_index in self.consumer_indices:
self.events[consumer_index].wait(threshold)
def wait_until_on_stream(self, threshold: int, stream) -> None:
if not self.consumer_indices:
return
for consumer_index in self.consumer_indices:
self.events[consumer_index].wait_on_stream(threshold, stream)
def reset(self):
self.producer_index = -1
self.consumer_index = -1

View File

@@ -179,8 +179,9 @@ from sglang.srt.managers.scheduler_update_weights_mixin import (
)
from sglang.srt.managers.session_controller import SessionController
from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.mem_cache.common import KVCapacityWaitError, release_kv_cache
from sglang.srt.mem_cache.radix_cache import RadixCache
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
from sglang.srt.model_executor.forward_batch_info import ForwardMode, PPProxyTensors
@@ -2417,6 +2418,8 @@ class Scheduler(
if len(can_run_list) == 0:
return None
waiting_queue_before_prepare = list(self.waiting_queue)
chunked_req_before_prepare = self.chunked_req
can_run_set = set(can_run_list)
self.waiting_queue = [x for x in self.waiting_queue if x not in can_run_set]
if adder.preempt_list:
@@ -2456,7 +2459,20 @@ class Scheduler(
self.tree_cache.ready_to_load_host_cache()
)
new_batch.prepare_for_extend()
try:
new_batch.prepare_for_extend()
except KVCapacityWaitError as exc:
self._release_prefill_adder_locks(
can_run_list, skip_req=chunked_req_before_prepare
)
self.waiting_queue = waiting_queue_before_prepare
self.chunked_req = chunked_req_before_prepare
self.running_batch.batch_is_full = True
logger.warning(
"[CP_SHARED_KV_CAPACITY_WAIT] prefill allocation deferred: %s",
exc,
)
return None
# Record prefill stats for logging after forward
new_batch.prefill_stats = PrefillStats.from_adder(
@@ -2485,6 +2501,38 @@ class Scheduler(
return new_batch
def _release_prefill_adder_locks(
self, reqs: List[Any], *, skip_req: Optional[Any] = None
) -> None:
"""Release lock refs acquired by PrefillAdder for a failed prefill batch.
`PrefillAdder.add_one_req` persists one `inc_lock_ref` per staged
request after its short-lived scheduling lock context exits. If KV
allocation later reports a recoverable capacity wait, the batch never
runs, so those persistent refs must be dropped before the requests are
returned to the waiting queue.
"""
if getattr(self.tree_cache, "disable", False):
return
use_swa_params = self.tree_cache.supports_swa() and self.tree_cache.is_tree_cache()
for req in reqs:
if req is skip_req:
continue
last_node = getattr(req, "last_node", None)
if last_node is None:
continue
if use_swa_params:
params = DecLockRefParams(
swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)
)
self.tree_cache.dec_lock_ref(last_node, params)
if hasattr(req, "swa_uuid_for_lock"):
req.swa_uuid_for_lock = None
else:
self.tree_cache.dec_lock_ref(last_node)
def update_running_batch(self, batch: ScheduleBatch) -> Optional[ScheduleBatch]:
"""Update the current running decoding batch."""
initial_bs = batch.batch_size()

View File

@@ -78,6 +78,7 @@ class EvictParams:
num_tokens: int
swa_num_tokens: int = 0
mamba_num: int = 0
owner_lane_deficits: Optional[list[int]] = None
@dataclasses.dataclass

View File

@@ -32,6 +32,27 @@ MAMBA_STATE_PER_REQ_NO_CACHE = 1
logger = logging.getLogger(__name__)
class KVCapacityWaitError(RuntimeError):
"""Recoverable KV-capacity pressure with CP owner-lane diagnostics."""
def __init__(
self,
message: str,
*,
required_by_owner: list[int] | None = None,
available_by_owner: list[int] | None = None,
deficit_by_owner: list[int] | None = None,
):
self.required_by_owner = required_by_owner
self.available_by_owner = available_by_owner
self.deficit_by_owner = deficit_by_owner
super().__init__(
f"{message}; required_by_owner={required_by_owner} "
f"available_by_owner={available_by_owner} "
f"deficit_by_owner={deficit_by_owner}"
)
def _log_cp_shared_kv_alloc_fallback(
reason: str,
message: str,
@@ -341,7 +362,12 @@ def _evict_for_compute_owner_lanes(
before_available,
evictable_size,
)
evict_result = tree_cache.evict(EvictParams(num_tokens=evict_tokens))
evict_result = tree_cache.evict(
EvictParams(
num_tokens=evict_tokens,
owner_lane_deficits=[int(v) for v in deficits],
)
)
after_available = allocator.available_size()
evicted_tokens = getattr(evict_result, "num_tokens_evicted", 0)
logger.info(
@@ -367,7 +393,7 @@ def alloc_paged_token_slots_extend(
# Over estimate the number of tokens: assume each request needs a new page.
allocator = tree_cache.token_to_kv_pool_allocator
num_tokens = extend_num_tokens + len(seq_lens_cpu) * allocator.page_size
evict_result = evict_from_tree_cache(tree_cache, num_tokens)
evict_result = None
# logger.info(
# "[MemCache-alloc] alloc_paged_token_slots_extend: extend_num_tokens=%d batch_size=%d num_tokens=%d page_size=%d "
@@ -420,11 +446,11 @@ def alloc_paged_token_slots_extend(
"multi_batch" if len(prefix_lens_cpu) != 1 else "unknown"
)
state = None
if backup_state:
state = allocator.backup_state()
if page_compute_owners is not None:
state = None
if backup_state:
state = allocator.backup_state()
out_cache_loc = alloc_extend_compute_owner(
prefix_lens,
prefix_lens_cpu,
@@ -460,10 +486,10 @@ def alloc_paged_token_slots_extend(
required, available, deficits = compute_owner_lane_stats(
page_compute_owners
)
_log_cp_shared_kv_alloc_fallback(
"owner_lane_exhausted",
"failed to allocate pages from compute-owner lanes; "
"falling back to legacy page allocation. extend_num_tokens=%s page_size=%s "
logger.warning(
"[CP_SHARED_KV_FAIL_FAST][owner_lane_exhausted] "
"failed to allocate pages from compute-owner lanes after eviction; "
"raising recoverable capacity wait. extend_num_tokens=%s page_size=%s "
"required_by_owner=%s available_by_owner=%s deficit_by_owner=%s",
extend_num_tokens,
allocator.page_size,
@@ -471,15 +497,19 @@ def alloc_paged_token_slots_extend(
available,
deficits,
)
out_cache_loc = allocator.alloc_extend(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
raise KVCapacityWaitError(
"owner_lane_exhausted: failed to allocate pages from CP "
"compute-owner lanes after owner-lane eviction",
required_by_owner=required,
available_by_owner=available,
deficit_by_owner=deficits,
)
else:
evict_result = evict_from_tree_cache(tree_cache, num_tokens)
state = None
if backup_state:
state = allocator.backup_state()
if alloc_extend_compute_owner is not None:
_log_cp_shared_kv_alloc_fallback(
compute_owner_unavailable_reason or "compute_owner_not_available",
@@ -568,6 +598,7 @@ def alloc_for_extend(
extend_lens_device = extend_lens_cpu.to(batch.device, non_blocking=True)
# Allocate req slots
newly_allocated_reqs = [req for req in batch.reqs if req.req_pool_idx is None]
req_pool_indices = alloc_req_slots(
batch.req_to_token_pool, batch.reqs, batch.tree_cache
)
@@ -575,23 +606,29 @@ def alloc_for_extend(
req_pool_indices_device = req_pool_indices_cpu.to(batch.device, non_blocking=True)
# Allocate KV cache (throws exception on failure)
if batch.tree_cache.page_size == 1:
out_cache_loc = alloc_token_slots(batch.tree_cache, batch.extend_num_tokens)
else:
# Paged allocation - build last_loc
last_loc = [
(t[-1:] if len(t) > 0 else torch.tensor([-1], device=batch.device))
for t in prefix_tensors
]
out_cache_loc = alloc_paged_token_slots_extend(
tree_cache=batch.tree_cache,
prefix_lens=prefix_lens_device,
prefix_lens_cpu=prefix_lens_cpu,
seq_lens=batch.seq_lens,
seq_lens_cpu=batch.seq_lens_cpu,
last_loc=torch.cat(last_loc),
extend_num_tokens=batch.extend_num_tokens,
)
try:
if batch.tree_cache.page_size == 1:
out_cache_loc = alloc_token_slots(batch.tree_cache, batch.extend_num_tokens)
else:
# Paged allocation - build last_loc
last_loc = [
(t[-1:] if len(t) > 0 else torch.tensor([-1], device=batch.device))
for t in prefix_tensors
]
out_cache_loc = alloc_paged_token_slots_extend(
tree_cache=batch.tree_cache,
prefix_lens=prefix_lens_device,
prefix_lens_cpu=prefix_lens_cpu,
seq_lens=batch.seq_lens,
seq_lens_cpu=batch.seq_lens_cpu,
last_loc=torch.cat(last_loc),
extend_num_tokens=batch.extend_num_tokens,
)
except KVCapacityWaitError:
for req in newly_allocated_reqs:
if req.req_pool_idx is not None:
batch.req_to_token_pool.free(req)
raise
# Write to req_to_token_pool
write_cache_indices(

View File

@@ -33,7 +33,6 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.common import _evict_for_compute_owner_lanes
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
from sglang.srt.mem_cache.memory_pool import (
MHATokenToKVPool,
@@ -288,6 +287,22 @@ class CpHiCacheEvictionPlan:
remaining_deficit: Tuple[int, ...]
@dataclass(frozen=True)
class CpLoadBackPlan:
page_owners: List[int]
required_by_owner: List[int]
available_by_owner: List[int]
deficit_by_owner: List[int]
host_hit_len: int
@dataclass(frozen=True)
class CpLoadBackEvictionPlan:
victims: Tuple[TreeNode, ...]
planned_freed_by_owner: Tuple[int, ...]
remaining_deficit_by_owner: Tuple[int, ...]
class HiCachePendingBackupSplit(Exception):
def __init__(self, node: TreeNode):
self.node = node
@@ -869,6 +884,249 @@ class HiRadixCache(RadixCache):
remaining_deficit=tuple(deficits),
)
def _build_cp_load_back_plan(
self, nodes_to_load: List[TreeNode], *, node_id: int
) -> CpLoadBackPlan:
allocator = self.token_to_kv_pool_allocator
compute_owner_lane_stats = getattr(allocator, "compute_owner_lane_stats", None)
if compute_owner_lane_stats is None:
raise RuntimeError(
"CP HiCache load-back requires an allocator with "
f"compute_owner_lane_stats: node_id={node_id}"
)
page_owners: List[int] = []
host_hit_len = 0
for node in nodes_to_load:
metadata = getattr(node, "cp_hicache", None)
if metadata is None:
raise RuntimeError(
"CP HiCache load-back missing cp_hicache metadata: "
f"load_node_id={getattr(node, 'id', '?')} root_node_id={node_id}"
)
node_host_len = int(self._node_host_len(node))
if node_host_len % self.page_size != 0:
raise RuntimeError(
"CP HiCache load-back requires page-aligned host lengths: "
f"load_node_id={getattr(node, 'id', '?')} "
f"host_len={node_host_len} page_size={self.page_size}"
)
if int(getattr(metadata, "logical_len", node_host_len)) != node_host_len:
raise RuntimeError(
"CP HiCache load-back metadata length mismatch: "
f"load_node_id={getattr(node, 'id', '?')} "
f"host_len={node_host_len} metadata_len={metadata.logical_len}"
)
page_owners.extend(int(owner) for owner in metadata.page_owners.tolist())
host_hit_len += node_host_len
expected_pages = host_hit_len // self.page_size
if len(page_owners) != expected_pages:
raise RuntimeError(
"CP HiCache load-back page owner count mismatch: "
f"node_id={node_id} host_hit_len={host_hit_len} "
f"page_size={self.page_size} page_owners={len(page_owners)}"
)
required, available, deficits = compute_owner_lane_stats(page_owners)
return CpLoadBackPlan(
page_owners=list(page_owners),
required_by_owner=[int(v) for v in required],
available_by_owner=[int(v) for v in available],
deficit_by_owner=[int(v) for v in deficits],
host_hit_len=host_hit_len,
)
def _refresh_cp_load_back_plan(self, plan: CpLoadBackPlan) -> CpLoadBackPlan:
required, available, deficits = (
self.token_to_kv_pool_allocator.compute_owner_lane_stats(plan.page_owners)
)
return CpLoadBackPlan(
page_owners=plan.page_owners,
required_by_owner=[int(v) for v in required],
available_by_owner=[int(v) for v in available],
deficit_by_owner=[int(v) for v in deficits],
host_hit_len=plan.host_hit_len,
)
def _cp_device_leaf_is_load_back_victim(self, node: TreeNode) -> bool:
if node == getattr(self, "root_node", None):
return False
if getattr(node, "lock_ref", 0) > 0:
return False
if getattr(node, "value", None) is None:
return False
parent = getattr(node, "parent", None)
if parent is None:
return False
try:
parent_key = self.get_child_key_fn(node.key)
except Exception:
return False
if getattr(parent, "children", {}).get(parent_key) is not node:
return False
if self._is_pinned(node):
return False
return True
def _cp_load_back_node_owner_page_counts(
self, node: TreeNode, cp_size: int
) -> Tuple[int, ...]:
value = getattr(node, "value", None)
if value is None or int(value.numel()) == 0:
return tuple(0 for _ in range(cp_size))
if int(value.numel()) % self.page_size != 0:
raise RuntimeError(
"CP HiCache load-back eviction requires page-aligned device values: "
f"node_id={getattr(node, 'id', '?')} value_len={value.numel()} "
f"page_size={self.page_size}"
)
first_locs = value[:: self.page_size]
logical_pages = torch.div(first_locs, self.page_size, rounding_mode="floor")
owners = torch.remainder(logical_pages - 1, cp_size)
return tuple(int((owners == owner).sum().item()) for owner in range(cp_size))
def _plan_cp_load_back_owner_lane_evictions(
self, plan: CpLoadBackPlan
) -> CpLoadBackEvictionPlan:
deficits = [max(0, int(v)) for v in plan.deficit_by_owner]
cp_size = len(deficits)
planned_freed = [0 for _ in range(cp_size)]
victims: List[TreeNode] = []
used_node_ids = set()
while any(v > 0 for v in deficits):
best_node = None
best_counts = None
best_score = None
for node in list(getattr(self, "evictable_leaves", set())):
node_id = getattr(node, "id", None)
if node_id in used_node_ids:
continue
if not self._cp_device_leaf_is_load_back_victim(node):
continue
counts = self._cp_load_back_node_owner_page_counts(node, cp_size)
contribution = sum(
min(int(count), int(deficit))
for count, deficit in zip(counts, deficits)
)
if contribution <= 0:
continue
score = (
-int(contribution),
self.eviction_strategy.get_priority(node),
int(node_id or 0),
)
if best_score is None or score < best_score:
best_score = score
best_node = node
best_counts = counts
if best_node is None or best_counts is None:
break
victims.append(best_node)
used_node_ids.add(getattr(best_node, "id", None))
for owner, count in enumerate(best_counts):
planned_freed[owner] += int(count)
deficits[owner] = max(0, deficits[owner] - int(count))
return CpLoadBackEvictionPlan(
victims=tuple(victims),
planned_freed_by_owner=tuple(planned_freed),
remaining_deficit_by_owner=tuple(deficits),
)
def _evict_cp_load_back_owner_lanes(
self, plan: CpLoadBackPlan, *, node_id: int
) -> CpLoadBackPlan:
eviction_plan = self._plan_cp_load_back_owner_lane_evictions(plan)
if any(v > 0 for v in eviction_plan.remaining_deficit_by_owner):
logger.warning(
"[CP_HICACHE_FALLBACK][cp_load_back_owner_lane_eviction_insufficient] "
"node_id=%d required_by_owner=%s available_by_owner=%s "
"deficit_by_owner=%s planned_freed_by_owner=%s "
"remaining_deficit_by_owner=%s evictable_size=%d protected_size=%d "
"victims=%s allocator_state=%s",
node_id,
plan.required_by_owner,
plan.available_by_owner,
plan.deficit_by_owner,
eviction_plan.planned_freed_by_owner,
eviction_plan.remaining_deficit_by_owner,
int(getattr(self, "evictable_size_", 0)),
int(getattr(self, "protected_size_", 0)),
[getattr(node, "id", None) for node in eviction_plan.victims],
self.token_to_kv_pool_allocator.allocator_state_str(),
)
return self._refresh_cp_load_back_plan(plan)
if len(eviction_plan.victims) == 0:
return plan
num_evicted = 0
write_back_nodes: List[TreeNode] = []
for victim in eviction_plan.victims:
victim_id = getattr(victim, "id", None)
if not self._cp_device_leaf_is_load_back_victim(victim):
raise RuntimeError(
"CP HiCache load-back owner-lane eviction plan diverged "
f"before reservation: node_id={node_id} victim_id={victim_id}"
)
if victim.pin_expiry > 0 and time.monotonic() > victim.pin_expiry:
self._clear_pin(victim)
if not self._node_backuped(victim):
if self.cache_controller.write_policy == "write_back":
logger.warning(
"[CP_HICACHE_FALLBACK][cp_load_back_owner_lane_write_back] "
"node_id=%d victim_id=%s reason=unbacked_write_back_victim",
node_id,
victim_id,
)
written = self.write_backup(victim, write_back=True)
if written > 0:
write_back_nodes.append(victim)
continue
num_evicted += self._evict_regular(victim)
else:
num_evicted += self._evict_backuped(victim)
for child in victim.parent.children.values():
if child in write_back_nodes:
continue
if not child.evicted:
break
else:
self._update_leaf_status(victim.parent)
if write_back_nodes:
self.writing_check(write_back=True)
for victim in write_back_nodes:
if self._node_backuped(victim):
num_evicted += self._evict_backuped(victim)
refreshed = self._refresh_cp_load_back_plan(plan)
logger.info(
"[HiCache-load] owner-lane device eviction before CP load-back: "
"node_id=%d victims=%s num_evicted=%d required_by_owner=%s "
"before_available_by_owner=%s before_deficit_by_owner=%s "
"after_available_by_owner=%s after_deficit_by_owner=%s "
"planned_freed_by_owner=%s",
node_id,
[getattr(node, "id", None) for node in eviction_plan.victims],
num_evicted,
plan.required_by_owner,
plan.available_by_owner,
plan.deficit_by_owner,
refreshed.available_by_owner,
refreshed.deficit_by_owner,
eviction_plan.planned_freed_by_owner,
)
return refreshed
def _cp_capacity_debug_enabled(self) -> bool:
return bool(
getattr(self, "_cp_hicache_capacity_debug", False)
@@ -2470,14 +2728,28 @@ class HiRadixCache(RadixCache):
result = self.inc_lock_ref(ancester_node)
delta = result.delta
# load it all or not at all
host_hit_len = sum(self._node_host_len(n) for n in nodes_to_load)
# load it all or not at all. The scalar length remains only a
# coarse upper bound; final CP admission is the owner-lane vector
# check below.
try:
load_back_plan = self._build_cp_load_back_plan(
nodes_to_load, node_id=last_hit_node.id
)
except Exception:
self.dec_lock_ref(ancester_node)
raise
host_hit_len = load_back_plan.host_hit_len
logger.info(
"[HiCache-load] load_back CP: node_id=%d nodes_to_load=%d host_hit_len=%d threshold=%d",
"[HiCache-load] load_back CP: node_id=%d nodes_to_load=%d "
"host_hit_len=%d threshold=%d required_by_owner=%s "
"available_by_owner=%s deficit_by_owner=%s",
last_hit_node.id,
len(nodes_to_load),
host_hit_len,
self.load_back_threshold,
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
)
if host_hit_len < self.load_back_threshold or (
host_hit_len > mem_quota + delta if mem_quota is not None else False
@@ -2491,32 +2763,48 @@ class HiRadixCache(RadixCache):
)
return None
if any(v > 0 for v in load_back_plan.deficit_by_owner):
load_back_plan = self._evict_cp_load_back_owner_lanes(
load_back_plan, node_id=last_hit_node.id
)
if any(v > 0 for v in load_back_plan.deficit_by_owner):
self.dec_lock_ref(ancester_node)
logger.warning(
"[CP_HICACHE_FALLBACK][cp_load_back_owner_lane_capacity_failed] "
"node_id=%d host_hit_len=%d nodes_to_load=%d "
"required_by_owner=%s available_by_owner=%s "
"deficit_by_owner=%s evictable_size=%d protected_size=%d "
"ongoing_load_back_count=%d allocator_state=%s",
last_hit_node.id,
host_hit_len,
len(nodes_to_load),
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
int(getattr(self, "evictable_size_", 0)),
int(getattr(self, "protected_size_", 0)),
len(getattr(self, "ongoing_load_back", {})),
self.token_to_kv_pool_allocator.allocator_state_str(),
)
return None
device_indices = self.cache_controller.load_cp(
nodes_to_load, node_id=last_hit_node.id
)
if device_indices is None:
logger.info(
"[HiCache-load] load_back CP retry with lane-aware eviction: "
"node_id=%d tokens_needed=%d",
failed_plan = self._refresh_cp_load_back_plan(load_back_plan)
logger.warning(
"[CP_HICACHE_FALLBACK][cp_load_back_preflight_mismatch] "
"node_id=%d host_hit_len=%d required_by_owner=%s "
"available_by_owner=%s deficit_by_owner=%s "
"allocator_state=%s",
last_hit_node.id,
host_hit_len,
)
# Lane-aware eviction: alloc_pages_with_owners failed because
# at least one owner lane is short of free pages. Targeted
# eviction frees pages from THOSE specific lanes, leaving
# other lanes untouched. Mirrors cold prefill's retry path
# in common.py:alloc_paged_token_slots_extend.
_retry_page_owners: List[int] = []
for _node in nodes_to_load:
# page_owners is CPU int8; tolist() returns list[int].
_retry_page_owners.extend(_node.cp_hicache.page_owners.tolist())
_evict_for_compute_owner_lanes(
tree_cache=self,
allocator=self.token_to_kv_pool_allocator,
page_compute_owners=_retry_page_owners,
)
device_indices = self.cache_controller.load_cp(
nodes_to_load, node_id=last_hit_node.id
failed_plan.required_by_owner,
failed_plan.available_by_owner,
failed_plan.deficit_by_owner,
self.token_to_kv_pool_allocator.allocator_state_str(),
)
self.dec_lock_ref(ancester_node)
if device_indices is None:

View File

@@ -746,6 +746,32 @@ class KVCache(abc.ABC):
def register_layer_transfer_counter(self, layer_transfer_counter: LayerDoneCounter):
self.layer_transfer_counter = layer_transfer_counter
def wait_layer_transfer_on_stream(self, layer_id: int, stream) -> None:
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until_on_stream(
layer_id - self.start_layer, stream
)
def get_key_buffer_for_prefetch(self, layer_id: int, stream) -> torch.Tensor:
self.wait_layer_transfer_on_stream(layer_id, stream)
raw_getter = getattr(self, "_get_key_buffer", None)
if raw_getter is None:
raise NotImplementedError(
f"{type(self).__name__} does not expose a non-blocking key-buffer getter"
)
return raw_getter(layer_id)
def get_index_k_with_scale_buffer_for_prefetch(
self, layer_id: int, stream
) -> torch.Tensor:
self.wait_layer_transfer_on_stream(layer_id, stream)
raw_getter = getattr(self, "_get_index_k_with_scale_buffer", None)
if raw_getter is None:
raise NotImplementedError(
f"{type(self).__name__} does not expose a non-blocking index-buffer getter"
)
return raw_getter(layer_id)
def register_layer_backup_notifier(self, notifier):
self.layer_backup_notifiers.append(notifier)
@@ -1379,6 +1405,10 @@ class HybridLinearKVPool(KVCache):
layer_id = self._transfer_full_attention_id(layer_id)
return self.full_kv_pool.get_key_buffer(layer_id)
def get_key_buffer_for_prefetch(self, layer_id: int, stream):
layer_id = self._transfer_full_attention_id(layer_id)
return self.full_kv_pool.get_key_buffer_for_prefetch(layer_id, stream)
def get_value_buffer(self, layer_id: int):
layer_id = self._transfer_full_attention_id(layer_id)
return self.full_kv_pool.get_value_buffer(layer_id)
@@ -1552,25 +1582,29 @@ class MLATokenToKVPool(KVCache):
return _copy_buf_infos(self._contiguous_buf_infos)
def get_key_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
def _get_key_buffer(self, layer_id: int):
if self.store_dtype != self.dtype:
return self.kv_buffer[layer_id - self.start_layer].view(self.dtype)
return self.kv_buffer[layer_id - self.start_layer]
def get_value_buffer(self, layer_id: int):
def get_key_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self._get_key_buffer(layer_id)
def _get_value_buffer(self, layer_id: int):
if self.store_dtype != self.dtype:
return self.kv_buffer[layer_id - self.start_layer][
..., : self.kv_lora_rank
].view(self.dtype)
return self.kv_buffer[layer_id - self.start_layer][..., : self.kv_lora_rank]
def get_value_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self._get_value_buffer(layer_id)
def get_kv_buffer(self, layer_id: int):
return self.get_key_buffer(layer_id), self.get_value_buffer(layer_id)
@@ -1724,10 +1758,7 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
del self.kv_scale_buffer
self._contiguous_buf_infos = None
def get_key_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
def _get_key_buffer(self, layer_id: int):
if self.store_dtype != self.dtype:
cache_k_nope_fp4 = self.kv_buffer[layer_id - self.start_layer].view(
torch.uint8
@@ -1743,6 +1774,11 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
return self.kv_buffer[layer_id - self.start_layer]
def get_key_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self._get_key_buffer(layer_id)
def set_kv_buffer(
self,
layer: RadixAttention,
@@ -1893,10 +1929,13 @@ class NSATokenToKVPool(MLATokenToKVPool):
]
self._finalize_allocation_log(size)
def _get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
return self.index_k_with_scale_buffer[layer_id - self.start_layer]
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self.index_k_with_scale_buffer[layer_id - self.start_layer]
return self._get_index_k_with_scale_buffer(layer_id)
def get_index_k_continuous(
self,

View File

@@ -0,0 +1,318 @@
import functools
import json
import sys
import types
import unittest
import torch
if "pybase64" not in sys.modules:
pybase64_stub = types.ModuleType("pybase64")
pybase64_stub.b64encode = lambda *args, **kwargs: b""
pybase64_stub.b64decode = lambda *args, **kwargs: b""
sys.modules["pybase64"] = pybase64_stub
if "orjson" not in sys.modules:
orjson_stub = types.ModuleType("orjson")
orjson_stub.loads = lambda data, *args, **kwargs: json.loads(
data.decode() if isinstance(data, (bytes, bytearray)) else data
)
orjson_stub.dumps = lambda obj, *args, **kwargs: json.dumps(obj).encode()
sys.modules["orjson"] = orjson_stub
sgl_kernel_stub = sys.modules.setdefault("sgl_kernel", types.ModuleType("sgl_kernel"))
sgl_kernel_stub.__file__ = getattr(sgl_kernel_stub, "__file__", "sgl_kernel_stub.py")
sgl_kernel_stub.__path__ = getattr(sgl_kernel_stub, "__path__", [])
if not hasattr(sgl_kernel_stub, "__getattr__"):
def _sgl_kernel_getattr(name):
if name.startswith("__"):
raise AttributeError(name)
fn = lambda *args, **kwargs: None
setattr(sgl_kernel_stub, name, fn)
return fn
sgl_kernel_stub.__getattr__ = _sgl_kernel_getattr
for _name in (
"sgl_per_token_group_quant_8bit",
"sgl_per_token_group_quant_fp8",
"sgl_per_token_quant_fp8",
"fp8_blockwise_scaled_mm",
"fp8_scaled_mm",
"silu_and_mul",
):
if not hasattr(sgl_kernel_stub, _name):
setattr(sgl_kernel_stub, _name, lambda *args, **kwargs: None)
quantization_stub = sys.modules.setdefault(
"sgl_kernel.quantization", types.ModuleType("sgl_kernel.quantization")
)
for _name in (
"ggml_dequantize",
"ggml_moe_a8",
"ggml_moe_a8_vec",
"ggml_moe_get_block_size",
"ggml_mul_mat_a8",
"ggml_mul_mat_vec_a8",
):
if not hasattr(quantization_stub, _name):
setattr(quantization_stub, _name, lambda *args, **kwargs: None)
kvcacheio_stub = sys.modules.setdefault(
"sgl_kernel.kvcacheio", types.ModuleType("sgl_kernel.kvcacheio")
)
for _name in (
"transfer_kv_all_layer",
"transfer_kv_all_layer_direct_lf_pf",
"transfer_kv_all_layer_lf_pf",
"transfer_kv_all_layer_lf_ph",
"transfer_kv_all_layer_mla",
"transfer_kv_all_layer_mla_lf_pf",
"transfer_kv_direct",
"transfer_kv_per_layer",
"transfer_kv_per_layer_direct_pf_lf",
"transfer_kv_per_layer_mla",
"transfer_kv_per_layer_mla_pf_lf",
"transfer_kv_per_layer_pf_lf",
"transfer_kv_per_layer_ph_lf",
):
if not hasattr(kvcacheio_stub, _name):
setattr(kvcacheio_stub, _name, lambda *args, **kwargs: None)
_sgl_kernel_lib = torch.library.Library("sgl_kernel", "FRAGMENT")
for _schema in (
"sgl_per_token_group_quant_8bit(Tensor input, Tensor(a!) output_q, Tensor(b!) output_s, int group_size, float eps, float fp8_min, float fp8_max, bool scale_ue8m0) -> ()",
"sgl_per_token_group_quant_fp8(Tensor input, Tensor(a!) output_q, Tensor(b!) output_s, int group_size, float eps, float fp8_min, float fp8_max, bool scale_ue8m0) -> ()",
"sgl_per_token_quant_fp8(Tensor input, Tensor(a!) output_q, Tensor(b!) output_s) -> ()",
"fp8_scaled_mm(Tensor mat_a, Tensor mat_b, Tensor scales_a, Tensor scales_b, ScalarType out_dtype, Tensor? bias=None) -> Tensor",
"fp8_blockwise_scaled_mm(Tensor mat_a, Tensor mat_b, Tensor scales_a, Tensor scales_b, ScalarType out_dtype) -> Tensor",
):
try:
_sgl_kernel_lib.define(_schema)
except RuntimeError as exc:
if "already" not in str(exc).lower() and "duplicate" not in str(exc).lower():
raise
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.hiradix_cache import (
CpHiCacheNodeMetadata,
HiRadixCache,
)
from sglang.srt.mem_cache.radix_cache import RadixKey, TreeNode, get_child_key
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
class _FakeLayout:
cp_size = 4
cp_rank = 0
class _FakeEvictionStrategy:
def get_priority(self, node):
return getattr(node, "priority", 0)
class _FakeController:
write_policy = "write_through"
has_draft_hicache = False
cp_shared_kv_layout = _FakeLayout()
def __init__(self, allocator):
self.allocator = allocator
self.load_calls = 0
self.evicted_device_indices = []
self.ack_load_queue = []
self.force_load_none = False
def load_cp(self, nodes_to_load, node_id=-1):
page_owners = []
for node in nodes_to_load:
page_owners.extend(node.cp_hicache.page_owners.tolist())
_, _, deficits = self.allocator.compute_owner_lane_stats(page_owners)
if any(deficits):
raise AssertionError(
f"load_cp called before owner-lane deficits were evicted: {deficits}"
)
self.load_calls += 1
if self.force_load_none:
return None
return self.allocator.alloc_pages_with_owners(page_owners)
def evict_device(self, indices):
self.evicted_device_indices.append(indices.clone())
self.allocator.free(indices)
return int(indices.numel())
def _make_allocator(page_size=4, cp_size=4):
allocator = CPSharedPagedTokenToKVPoolAllocator(
logical_size=page_size * 16,
physical_size=page_size * 4,
page_size=page_size,
dtype=torch.bfloat16,
device="cpu",
kvcache=None,
need_sort=False,
cp_size=cp_size,
cp_rank=0,
)
allocator.release_pages = torch.empty((0,), dtype=torch.int64)
return allocator
def _metadata(page_owners, page_size=4):
logical_len = len(page_owners) * page_size
owned_positions = []
for page_idx, owner in enumerate(page_owners):
if owner == 0:
owned_positions.extend(
range(page_idx * page_size, (page_idx + 1) * page_size)
)
return CpHiCacheNodeMetadata(
logical_len=logical_len,
owned_positions=torch.tensor(owned_positions, dtype=torch.int64),
host_indices=torch.arange(len(owned_positions), dtype=torch.int64),
page_owners=torch.tensor(page_owners, dtype=torch.int8),
page_size=page_size,
)
def _make_node(node_id, key_start, page_owners, *, value=None, priority=0):
page_size = 4
node = TreeNode(id=node_id, priority=priority)
node.key = RadixKey(token_ids=list(range(key_start, key_start + len(page_owners) * page_size)))
node.value = value
node.host_len = len(page_owners) * page_size
node.cp_hicache = _metadata(page_owners, page_size=page_size)
return node
def _attach_child(cache, parent, child):
child.parent = parent
parent.children[cache.get_child_key_fn(child.key)] = child
def _make_cache(allocator):
cache = HiRadixCache.__new__(HiRadixCache)
cache.disable = False
cache._uses_cp_hicache = True
cache.token_to_kv_pool_allocator = allocator
cache.page_size = allocator.page_size
cache.cache_controller = _FakeController(allocator)
cache.root_node = TreeNode(id=0, priority=-999)
cache.root_node.key = RadixKey(token_ids=[])
cache.root_node.value = []
cache.root_node.lock_ref = 1
cache.evictable_leaves = set()
cache.evictable_host_leaves = set()
cache.evictable_size_ = 0
cache.protected_size_ = 0
cache.ongoing_load_back = {}
cache.ongoing_write_through = {}
cache.pending_host_backups = {}
cache.load_back_threshold = 0
cache.metrics_collector = None
cache.eviction_strategy = _FakeEvictionStrategy()
cache.get_child_key_fn = functools.partial(get_child_key, page_size=allocator.page_size)
cache.enable_kv_cache_events = False
return cache
class TestCpHiCacheLoadBackOwnerLanes(CustomTestCase):
def test_load_back_plan_reports_owner_lane_vectors(self):
allocator = _make_allocator()
allocator.free_pages = torch.tensor([1, 2], dtype=torch.int64)
allocator.release_pages = torch.tensor([4], dtype=torch.int64)
cache = _make_cache(allocator)
node = _make_node(10, 100, [0, 1, 1, 3])
plan = cache._build_cp_load_back_plan([node], node_id=node.id)
self.assertEqual(plan.page_owners, [0, 1, 1, 3])
self.assertEqual(plan.required_by_owner, [1, 2, 0, 1])
self.assertEqual(plan.available_by_owner, [1, 1, 0, 1])
self.assertEqual(plan.deficit_by_owner, [0, 1, 0, 0])
self.assertEqual(plan.host_hit_len, 16)
def test_load_back_plan_fails_closed_without_cp_metadata(self):
allocator = _make_allocator()
cache = _make_cache(allocator)
node = TreeNode(id=11)
node.host_len = 4
node.cp_hicache = None
with self.assertRaisesRegex(RuntimeError, "missing cp_hicache metadata"):
cache._build_cp_load_back_plan([node], node_id=node.id)
def test_load_back_evicts_owner_lane_deficit_before_allocating(self):
allocator = _make_allocator()
# Only owner lane 1 is initially available. The load-back target needs
# lanes [0, 1], so calling load_cp before targeted eviction is a bug.
allocator.free_pages = torch.tensor([2], dtype=torch.int64)
cache = _make_cache(allocator)
victim = _make_node(
20,
200,
[0],
value=torch.tensor([4, 5, 6, 7], dtype=torch.int64),
priority=0,
)
_attach_child(cache, cache.root_node, victim)
cache.evictable_leaves.add(victim)
cache.evictable_size_ = len(victim.key)
target = _make_node(21, 300, [0, 1], value=None, priority=10)
_attach_child(cache, cache.root_node, target)
loaded = cache.load_back(target, mem_quota=100)
self.assertIsNotNone(loaded)
self.assertEqual(cache.cache_controller.load_calls, 1)
self.assertEqual(len(cache.cache_controller.evicted_device_indices), 1)
self.assertEqual(cache.cache_controller.evicted_device_indices[0].tolist(), [4, 5, 6, 7])
self.assertEqual(target.value.tolist(), loaded.tolist())
self.assertIn(target.id, cache.ongoing_load_back)
def test_load_back_failure_leaves_node_unassigned_and_unlocked(self):
allocator = _make_allocator()
allocator.free_pages = torch.tensor([1, 2], dtype=torch.int64)
cache = _make_cache(allocator)
cache.cache_controller.force_load_none = True
target = _make_node(30, 400, [0, 1], value=None)
_attach_child(cache, cache.root_node, target)
loaded = cache.load_back(target, mem_quota=100)
self.assertIsNone(loaded)
self.assertIsNone(target.value)
self.assertEqual(target.lock_ref, 0)
self.assertNotIn(target.id, cache.ongoing_load_back)
def test_load_back_lock_is_released_only_by_loading_ack(self):
class ReadyEvent:
def query(self):
return True
allocator = _make_allocator()
allocator.free_pages = torch.tensor([1, 2], dtype=torch.int64)
cache = _make_cache(allocator)
target = _make_node(31, 500, [0, 1], value=None)
_attach_child(cache, cache.root_node, target)
loaded = cache.load_back(target, mem_quota=100)
self.assertIsNotNone(loaded)
self.assertEqual(target.lock_ref, 1)
self.assertIn(target.id, cache.ongoing_load_back)
cache.cache_controller.ack_load_queue.append((None, ReadyEvent(), [target.id]))
cache.loading_check()
self.assertEqual(target.lock_ref, 0)
self.assertNotIn(target.id, cache.ongoing_load_back)
self.assertEqual(cache.cache_controller.ack_load_queue, [])
if __name__ == "__main__":
unittest.main()

View File

@@ -4,6 +4,18 @@ from unittest.mock import patch
import numpy as np
import torch
_sgl_kernel_lib = torch.library.Library("sgl_kernel", "FRAGMENT")
try:
_sgl_kernel_lib.define(
"moe_fused_gate(Tensor input_tensor, Tensor? bias, int num_expert_group, "
"int topk_group, int topk, int num_fused_shared_experts, "
"float routed_scaling_factor, bool apply_routed_scaling_factor_on_output) "
"-> (Tensor, Tensor)"
)
except RuntimeError as exc:
if "already" not in str(exc).lower() and "duplicate" not in str(exc).lower():
raise
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -435,6 +447,75 @@ class TestCPSharedPagedAllocator(CustomTestCase):
self.assertEqual(allocator.calls, 1)
self.assertEqual(out.numel(), page_size * 8)
def test_compute_owner_alloc_skips_aggregate_evict_before_owner_attempt(self):
from types import SimpleNamespace
from sglang.srt.mem_cache import common
page_size = 64
class FakeAllocator:
def __init__(self):
self.page_size = page_size
self.cp_size = 4
self.calls = 0
def available_size(self):
# Aggregate capacity looks insufficient, but the owner-aware
# allocator can satisfy the request from the right lanes.
# Aggregate eviction before this attempt would evict unrelated
# lanes and destroy cache locality.
return 0
def alloc_extend_compute_owner(
self,
_prefix_lens,
_prefix_lens_cpu,
_seq_lens,
_seq_lens_cpu,
_last_loc,
extend_num_tokens,
_page_compute_owners,
):
self.calls += 1
return torch.arange(extend_num_tokens, dtype=torch.int64)
def alloc_extend(self, *_args, **_kwargs):
raise AssertionError("legacy allocation should not be used")
class FakeTreeCache:
def __init__(self, allocator):
self.token_to_kv_pool_allocator = allocator
def is_chunk_cache(self):
return False
def evict(self, _params):
raise AssertionError(
"aggregate eviction should not run before owner-aware allocation"
)
allocator = FakeAllocator()
server_args = SimpleNamespace(
enable_nsa_prefill_cp_shared_kv=True,
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
with patch.object(common, "get_global_server_args", return_value=server_args):
out = common.alloc_paged_token_slots_extend(
tree_cache=FakeTreeCache(allocator),
prefix_lens=torch.tensor([0], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
seq_lens=torch.tensor([page_size * 8], dtype=torch.int64),
seq_lens_cpu=torch.tensor([page_size * 8], dtype=torch.int64),
last_loc=torch.tensor([-1], dtype=torch.int64),
extend_num_tokens=page_size * 8,
)
self.assertEqual(allocator.calls, 1)
self.assertEqual(out.numel(), page_size * 8)
def test_compute_owner_lane_eviction_recovers_exhausted_owner_lane(self):
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
@@ -508,6 +589,253 @@ class TestCPSharedPagedAllocator(CustomTestCase):
)
self.assertIsNotNone(locs)
def test_compute_owner_lane_eviction_passes_deficits_to_tree_cache(self):
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
from sglang.srt.mem_cache.common import _evict_for_compute_owner_lanes
page_size = 64
cp_size = 4
allocator = CPSharedPagedTokenToKVPoolAllocator(
logical_size=page_size * 16,
physical_size=page_size * 4,
page_size=page_size,
dtype=torch.bfloat16,
device="cpu",
kvcache=None,
need_sort=False,
cp_size=cp_size,
cp_rank=0,
)
allocator.free_pages = torch.empty((0,), dtype=torch.int64)
allocator.release_pages = torch.empty((0,), dtype=torch.int64)
class FakeTreeCache:
def __init__(self):
self.owner_deficits = []
def is_chunk_cache(self):
return False
def evictable_size(self):
return page_size * cp_size
def evict(self, params):
deficits = getattr(params, "owner_lane_deficits", None)
self.owner_deficits.append(deficits)
if deficits == [1, 0, 0, 0]:
allocator.free(torch.tensor([page_size], dtype=torch.int64))
else:
allocator.free(torch.tensor([page_size * 2], dtype=torch.int64))
return EvictResult(num_tokens_evicted=page_size)
tree_cache = FakeTreeCache()
_evict_for_compute_owner_lanes(
tree_cache=tree_cache,
allocator=allocator,
page_compute_owners=[0],
)
self.assertEqual(tree_cache.owner_deficits[0], [1, 0, 0, 0])
locs = allocator.alloc_extend_compute_owner(
prefix_lens=torch.tensor([0], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
seq_lens=torch.tensor([page_size], dtype=torch.int64),
seq_lens_cpu=torch.tensor([page_size], dtype=torch.int64),
last_loc=torch.tensor([-1], dtype=torch.int64),
extend_num_tokens=page_size,
page_compute_owners=[0],
)
self.assertIsNotNone(locs)
def test_compute_owner_capacity_wait_reports_owner_lane_deficits(self):
from types import SimpleNamespace
from sglang.srt.mem_cache import common
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
page_size = 64
class FakeAllocator:
def __init__(self):
self.page_size = page_size
self.cp_size = 4
def available_size(self):
return self.page_size * 4
def allocator_state_str(self):
return "allocator_state_for_test"
def alloc_extend_compute_owner(self, *_args, **_kwargs):
return None
def compute_owner_lane_stats(self, _page_compute_owners):
return [2, 2, 2, 2], [2, 2, 2, 0], [0, 0, 0, 2]
def alloc_extend(self, *_args, **_kwargs):
raise AssertionError("legacy allocation should not be used")
class FakeTreeCache:
def __init__(self):
self.token_to_kv_pool_allocator = FakeAllocator()
self.evict_params = []
def is_chunk_cache(self):
return False
def evictable_size(self):
return page_size * 8
def evict(self, params):
self.evict_params.append(params)
return EvictResult(num_tokens_evicted=0)
def pretty_print(self):
raise AssertionError("recoverable capacity wait should not dump tree")
server_args = SimpleNamespace(
enable_nsa_prefill_cp_shared_kv=True,
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
tree_cache = FakeTreeCache()
with patch.object(common, "get_global_server_args", return_value=server_args):
with self.assertRaises(common.KVCapacityWaitError) as cm:
common.alloc_paged_token_slots_extend(
tree_cache=tree_cache,
prefix_lens=torch.tensor([0], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
seq_lens=torch.tensor([page_size * 8], dtype=torch.int64),
seq_lens_cpu=torch.tensor([page_size * 8], dtype=torch.int64),
last_loc=torch.tensor([-1], dtype=torch.int64),
extend_num_tokens=page_size * 8,
)
err = cm.exception
self.assertEqual(err.required_by_owner, [2, 2, 2, 2])
self.assertEqual(err.available_by_owner, [2, 2, 2, 0])
self.assertEqual(err.deficit_by_owner, [0, 0, 0, 2])
self.assertIn("owner_lane_exhausted", str(err))
self.assertIn("deficit_by_owner=[0, 0, 0, 2]", str(err))
self.assertIn(
[0, 0, 0, 2],
[params.owner_lane_deficits for params in tree_cache.evict_params],
)
def test_alloc_for_extend_releases_req_slots_on_recoverable_capacity_wait(self):
from types import SimpleNamespace
from sglang.srt.mem_cache import common
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
req_to_token_pool = ReqToTokenPool(
size=1,
max_context_len=128,
device="cpu",
enable_memory_saver=False,
)
req = SimpleNamespace(
req_pool_idx=None,
is_chunked=0,
kv_committed_len=0,
prefix_indices=torch.empty((0,), dtype=torch.int64),
)
batch = SimpleNamespace(
maybe_evict_swa=lambda: None,
reqs=[req],
prefix_lens=[0],
extend_lens=[64],
device="cpu",
req_to_token_pool=req_to_token_pool,
tree_cache=SimpleNamespace(page_size=64),
seq_lens=torch.tensor([64], dtype=torch.int64),
seq_lens_cpu=torch.tensor([64], dtype=torch.int64),
extend_num_tokens=64,
)
wait_error = common.KVCapacityWaitError(
"owner_lane_exhausted for test",
required_by_owner=[1, 0],
available_by_owner=[0, 0],
deficit_by_owner=[1, 0],
)
with patch.object(
common,
"alloc_paged_token_slots_extend",
side_effect=wait_error,
):
with self.assertRaises(common.KVCapacityWaitError):
common.alloc_for_extend(batch)
self.assertIsNone(req.req_pool_idx)
self.assertEqual(req_to_token_pool.available_size(), 1)
def test_scheduler_capacity_wait_rollback_releases_adder_locks(self):
from types import SimpleNamespace
from sglang.srt.managers.scheduler import Scheduler
class FakeTreeCache:
def __init__(self):
self.dec_calls = []
def supports_swa(self):
return False
def is_tree_cache(self):
return True
def dec_lock_ref(self, node, params=None):
self.dec_calls.append((node, params))
scheduler = object.__new__(Scheduler)
scheduler.tree_cache = FakeTreeCache()
req = SimpleNamespace(last_node="scheduled-node", swa_uuid_for_lock=None)
chunked_req = SimpleNamespace(last_node="chunked-node", swa_uuid_for_lock=None)
scheduler._release_prefill_adder_locks(
[req, chunked_req],
skip_req=chunked_req,
)
self.assertEqual(
[(node, params) for node, params in scheduler.tree_cache.dec_calls],
[("scheduled-node", None)],
)
def test_scheduler_capacity_wait_rollback_releases_swa_lock_uuid(self):
from types import SimpleNamespace
from sglang.srt.managers.scheduler import Scheduler
class FakeTreeCache:
def __init__(self):
self.dec_calls = []
def supports_swa(self):
return True
def is_tree_cache(self):
return True
def dec_lock_ref(self, node, params=None):
self.dec_calls.append((node, params))
scheduler = object.__new__(Scheduler)
scheduler.tree_cache = FakeTreeCache()
req = SimpleNamespace(last_node="scheduled-node", swa_uuid_for_lock=17)
scheduler._release_prefill_adder_locks([req])
self.assertEqual(scheduler.tree_cache.dec_calls[0][0], "scheduled-node")
self.assertEqual(
scheduler.tree_cache.dec_calls[0][1].swa_uuid_for_lock,
17,
)
self.assertIsNone(req.swa_uuid_for_lock)
if __name__ == "__main__":
unittest.main()

View File

@@ -1,9 +1,30 @@
import unittest
import sys
import types
from types import SimpleNamespace
from unittest.mock import patch
import torch
_sgl_kernel_lib = torch.library.Library("sgl_kernel", "FRAGMENT")
try:
_sgl_kernel_lib.define(
"moe_fused_gate(Tensor input_tensor, Tensor? bias, int num_expert_group, "
"int topk_group, int topk, int num_fused_shared_experts, "
"float routed_scaling_factor, bool apply_routed_scaling_factor_on_output) "
"-> (Tensor, Tensor)"
)
except RuntimeError as exc:
if "already" not in str(exc).lower() and "duplicate" not in str(exc).lower():
raise
flash_attn_stub = sys.modules.setdefault(
"sgl_kernel.flash_attn", types.ModuleType("sgl_kernel.flash_attn")
)
for _name in ("flash_attn_varlen_func", "flash_attn_with_kvcache"):
if not hasattr(flash_attn_stub, _name):
setattr(flash_attn_stub, _name, lambda *args, **kwargs: None)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
@@ -14,7 +35,7 @@ def _identity_all_reduce(buffer, *args, **kwargs):
class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
def test_mla_prefetch_materializes_on_current_stream_and_defers_reduce_to_launch(
def test_mla_prefetch_materializes_and_reduces_on_prefetch_stream(
self,
):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
@@ -48,8 +69,13 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
def __init__(self):
self.kv_cache = torch.zeros((32, 1, 2), dtype=torch.float32)
self.prefetch_getter_streams = []
def get_key_buffer(self, layer_id):
raise AssertionError("prefetch must not call blocking get_key_buffer")
def get_key_buffer_for_prefetch(self, layer_id, stream):
self.prefetch_getter_streams.append((layer_id, stream.name))
return self.kv_cache
active_stream = ["current"]
@@ -89,23 +115,28 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
"_all_reduce_materialized_buffer_async",
side_effect=record_reduce,
):
pool = FakePool()
prefetcher.start_next_layer_prefix(
next_layer_id=1,
token_to_kv_pool=FakePool(),
token_to_kv_pool=pool,
)
self.assertEqual(calls, [("materialize", "current")])
self.assertEqual(prefetch_stream.waited, [])
self.assertEqual(pool.prefetch_getter_streams, [(1, "prefetch")])
self.assertEqual(
calls,
[("materialize", "prefetch"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
prefetcher.launch_pending_reduce()
self.assertEqual(
calls,
[("materialize", "current"), ("reduce", "prefetch", "prefetch")],
[("materialize", "prefetch"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
def test_index_prefetch_materializes_on_current_stream_and_defers_reduce_to_launch(
def test_index_prefetch_materializes_and_reduces_on_prefetch_stream(
self,
):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
@@ -139,8 +170,15 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
def __init__(self):
self.page_buffer = torch.zeros((16, 3), dtype=torch.uint8)
self.prefetch_getter_streams = []
def get_index_k_with_scale_buffer(self, layer_id):
raise AssertionError(
"prefetch must not call blocking get_index_k_with_scale_buffer"
)
def get_index_k_with_scale_buffer_for_prefetch(self, layer_id, stream):
self.prefetch_getter_streams.append((layer_id, stream.name))
return self.page_buffer
active_stream = ["current"]
@@ -179,22 +217,55 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
"_all_reduce_materialized_buffer_async",
side_effect=record_reduce,
):
pool = FakePool()
prefetcher.start_next_layer_prefix(
next_layer_id=1,
token_to_kv_pool=FakePool(),
token_to_kv_pool=pool,
)
self.assertEqual(calls, [("materialize", "current")])
self.assertEqual(prefetch_stream.waited, [])
self.assertEqual(pool.prefetch_getter_streams, [(1, "prefetch")])
self.assertEqual(
calls,
[("materialize", "prefetch"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
prefetcher.launch_pending_reduce()
self.assertEqual(
calls,
[("materialize", "current"), ("reduce", "prefetch", "prefetch")],
[("materialize", "prefetch"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
def test_mla_pool_prefetch_getter_orders_layer_transfer_on_prefetch_stream(self):
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
class FakeCounter:
def __init__(self):
self.calls = []
def wait_until(self, threshold):
raise AssertionError("prefetch getter must not wait on current stream")
def wait_until_on_stream(self, threshold, stream):
self.calls.append((threshold, stream.name))
class FakeStream:
name = "prefetch"
pool = MLATokenToKVPool.__new__(MLATokenToKVPool)
pool.start_layer = 2
pool.layer_transfer_counter = FakeCounter()
pool.store_dtype = torch.float32
pool.dtype = torch.float32
pool.kv_buffer = [torch.ones((4, 1), dtype=torch.float32) for _ in range(3)]
out = pool.get_key_buffer_for_prefetch(3, FakeStream())
self.assertIs(out, pool.kv_buffer[1])
self.assertEqual(pool.layer_transfer_counter.calls, [(1, "prefetch")])
def test_all_reduce_uses_group_fast_path_for_float_buffers(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
@@ -1782,12 +1853,12 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
self.assertIs(dense_buffer, fallback_buffer)
self.assertIs(dense_pages, fallback_pages)
logger.info.assert_called_once()
logger.warning.assert_called_once()
self.assertIn(
"CP shared KV index prefetch fallback",
logger.info.call_args.args[0],
"[CP_SHARED_KV_FALLBACK][index_prefetch]",
logger.warning.call_args.args[0],
)
self.assertIn("consume_miss", logger.info.call_args.args[1])
self.assertIn("consume_miss", logger.warning.call_args.args[1])
def test_index_prefetch_first_layer_miss_does_not_log_fallback(self):
from sglang.srt.layers.attention.nsa import nsa_indexer
@@ -1835,7 +1906,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
logical_page_table=logical_pages,
)
logger.info.assert_not_called()
logger.warning.assert_not_called()
def test_index_prefetch_create_skip_logs_fallback_when_enabled(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
@@ -1847,7 +1918,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
), patch.object(
prefetch.torch.cuda, "is_available", return_value=False
), patch.object(
prefetch.logger, "info"
prefetch.logger, "warning"
) as logger:
result = prefetch.CpSharedKVIndexPrefetcher.maybe_create(
forward_batch=SimpleNamespace(),
@@ -1858,7 +1929,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
self.assertIsNone(result)
logger.assert_called_once()
self.assertIn(
"CP shared KV index prefetch fallback",
"[CP_SHARED_KV_FALLBACK][index_prefetch]",
logger.call_args.args[0],
)
self.assertIn("cuda_unavailable_or_stream_capturing", logger.call_args.args[1])