Reduce CP shared-KV batch overhead without reverting bs1 planning

CP shared-KV bs>1 exposed two separate overhead sources: HiCache load-back could thrash near capacity, and partial-current sync compose could all-reduce row-major page-table gaps between request prefixes. Keep the intended batch-plan path for bs=1, but make host/L1 free-room handling less reactive and teach the MLA/index sync compose path to use exact per-request prefix slot spans instead of one bounding span.\n\nThe exact-span path preserves the single-span IPC fast path for bs=1/single-span cases, while avoiding over-communication for heterogeneous cache-hit batches. The HiCache metadata tests cover host/L1 free-room propagation and load-back batching behavior.\n\nConstraint: bs=1 using CPSharedKVBatchPlan is the expected steady-state path and must not be treated as a regression.\nConstraint: Remote production-like validation runs inside g0034 container /sgl-workspace/sglang-tai.\nRejected: Disable batch-plan for bs=1 | user confirmed this is intended behavior and it would hide the actual bs>1 overhead.\nRejected: Keep one bounding prefix span for bs>1 | row-major page tables can include current/gap slots and inflate per-layer all-reduce work.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not replace exact prefix spans with a single row-major bounding span unless ETE data proves collective launch count dominates gap over-communication.\nTested: g0034 docker py_compile for changed Python/test files.\nTested: g0034 docker PYTHONPATH=python python -m pytest -q test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py test/registered/unit/mem_cache/test_cp_hicache_metadata.py — 241 passed, 5 warnings, 2 subtests passed.\nNot-tested: Full ETE bs>1 throughput after the exact-prefix-span change.\nNot-tested: CUDA kernel benchmark and live traffic replay.
This commit is contained in:
laoyao0822
2026-06-06 02:03:16 +08:00
parent 1b99de7459
commit 81bfd4bfd3
10 changed files with 544 additions and 92 deletions
@@ -230,3 +230,49 @@ Implications:
separate ratios because L1 hit value and host capacity pressure are different.
- Free-room alone is not sufficient for contiguous allocation. `HostKVCache` and
CP owner-lane device allocation also need contiguous-preferred selection.
## 2026-06-05 correction: load-back must not synchronously evict for L1 free room
Remote evidence from `/mnt/beegfs/cjy/log/sglang_cp_hicache_20260605_152940.log`
showed prefill slowdown and death caused by a CP0 HiCache load-back planning
straggler:
```text
node_id=677 load_back CP starts on all CP ranks at 15:48:26
CP1-7 report SUCCESS by 15:48:27
CP0 reports SUCCESS at 15:48:49
health check fires at 15:48:47 due no detokenizer heartbeat for 20s
```
Important detail: `HiCache-load load_back CP SUCCESS` is emitted before
`cache_controller.start_loading()`, so this 23s gap is not the actual L2->L1 H2D
copy. It is in the scheduler-side admission/planning path:
```text
HiRadixCache.load_back
-> _build_cp_load_back_plan
-> optional _evict_cp_load_back_owner_lanes
-> HiCacheController.load_cp
-> alloc_pages_with_owners
-> build host/device index descriptors
```
The failing node had exact capacity available but non-zero L1 free-room deficit:
```text
required_by_owner=[41,45,45,45,46,46,46,46]
available_by_owner=[221,194,197,198,237,285,290,308]
free-room deficit on lanes 1-3
```
Therefore the reactive free-room rule must be different for synchronous
load-back admission:
- exact owner-lane deficit is blocking and may trigger synchronous eviction;
- free-room deficit is advisory/observability for load-back and must not force
synchronous eviction when exact capacity already fits.
This preserves correctness while avoiding heavy eviction planning on the cache-hit
scheduler hot path. Free-room maintenance for L1 should happen on real extend
allocation pressure or a later non-blocking/proactive path, not inside a
cache-hit load-back that already has enough exact capacity.
@@ -857,6 +857,58 @@ runtime / kernel 都消费 CPSharedKVBatchPlan descriptors
### W6A 完成状态
---
## 25. 2026-06-06 bs>1 打开后性能回退排查记录
用户确认:`batch_size == 1` 也生成并消费 `CPSharedKVBatchPlan` 是预期收敛方向,
不是本轮性能问题本身。后续不要再把“bs=1 走 batch-plan”当成需要回退的 bug。
最新远端日志(`/mnt/beegfs/cjy/log/sglang_cp_hicache_20260605_164245.log`)
显示:
- 没有 `[HiCache-load] slow load_back planning` 或容量等待类日志;
- fallback 主要是 `[CP_SHARED_KV_FALLBACK][tai_ipc_materialize] reason=paged_start_slot_nonzero`,
每 rank 限频后共 64 条,说明 index/page-buffer materialize 的非 0 起始 slot 会退到
local materialize + collective;
- CP0 batch 分布仍以 bs=1 为主,少量 bs=2/3,因此不能简单用“真实 batch 数变大”
解释全部慢点;
- debug 日志本身会污染性能测试,但不是一个足够可靠的根因解释。
当前高优先级性能嫌疑:
1. **bs>1 partial-current sync compose 的 row-major prefix slot span 可能过度覆盖。**
`build_batch_prefix_slot_span()` 为了用一个 contiguous span 覆盖所有 request prefix,
在 row-major page table 下会包含 request 之间的 current/空洞 slot。cache-hit 高、
prefix 分布不均时,prefix materialize/all-reduce 的页数可能远大于真实 prefix 页数。
这会落在每层 critical path 上,尤其在 bs>1 async prefetch 还没恢复时更明显。
2. **index/top-k batch path 仍有 per-layer Python/Torch descriptor overhead。**
`_get_topk_in_seq_cp_pair_batch()` 每层构造 Python list、`torch.tensor(..., device)`、
`torch.cat`、`torch.full`,再 scatter compact topk。TAI batch kernel 已减少部分
K/S copy,但 descriptor 构造还不是 per-forward 预计算。
3. **tiny extend compute padding 会把 `< cp_size` pages 的 request 补到 `cp_size` pages。**
例如 page=64、cp=8 时,200 token 会以 512 token compute 形态参与 CP split。
这对 200-512 token 的线上短 extend 可能抵消 batch 的 compute 填充收益。
4. **bs>1 L1 prefix prefetch 仍未实现。**
`CpSharedKVMlaPrefetcher` / `CpSharedKVIndexPrefetcher` 仍要求
`forward_batch.batch_size == 1`。如果开启真实 bs>1 后 cache-hit 高,prefix 准备
会更多落到同步 compose 路径,而不是目标的 async L2->L1/L1 materialize 流水线。
下一步验证顺序:
1. 加或复用低频 perf counter,记录每个 forward 的 prefix true pages、prefix span pages、
current span pages、batch size、extend/prefix lens,不在每层刷屏;
2. microbench `_get_topk_in_seq_cp_pair_batch()` descriptor/cat/scatter 与
`split_tensor_by_cp_batch_plan()` 在 200/512/1k/2k extend、bs=1/2/4/5 下的 CPU submit
和 GPU elapsed;
3. 若证实 prefix span 过度覆盖,优先改为 compact prefix descriptor 或支持 nonzero-slot
TAI IPC materialize,避免 row-major gap 被 all-reduce;
4. 若证实 descriptor overhead 主导,先把 index batch descriptor 缓存在 `ForwardBatch`
上,避免 78 层重复构造。
已补 characterization tests,锁住现有 L2->L1 load-back batch 行为:
- `test_cp_start_loading_batches_multiple_load_cp_requests_with_draft`
@@ -1616,3 +1668,28 @@ PYTHONPATH=python python -m pytest -q \
`batch_gt1_index_q_length_mismatch`;
- 日志中的 `[CP_SHARED_KV_FALLBACK][tai_index_mqa_prepare] current_index_k must be uint8`
是另一个性能 fast-path dtype 问题,本次未修;它当前是 warning fallback,不是本次进程退出原因。
## 26. 2026-06-06 bs>1 partial-current prefix span 修正
用户确认:`bs=1` 走 batch-plan 路径是预期行为,不应作为回退或禁用目标。
本轮性能排查确认一个明确问题:bs>1 的 partial-current 同步 compose 之前使用
`build_batch_prefix_slot_span()` 把 row-major page table 中多个 request 的 prefix 区间压成一个
bounding span。这个 span 会覆盖 request 行之间的 current/gap slot,导致 index/MLA prefix
materialize 后对并非 prefix 的页一起做 all-reduce。在线上 cache-hit-heavy、prefix 长度不一致的
batch 中,这会把本应只覆盖 prefix pages 的通信量放大。
修正:
- 保留旧的 `build_batch_prefix_slot_span()` 作为单 bounding span helper。
- 新增 `build_batch_prefix_slot_spans()`,返回每个 request prefix 的精确 slot span,并只 merge 相邻区间。
- `materialize_prefix_and_reuse_current_kv_page_slots()` /
`materialize_prefix_and_reuse_current_index_page_slots()` 支持 `prefix_slot_spans`。
- MLA 与 index 的 bs>1 partial-current sync compose 使用精确 prefix spans;bs=1 或单 span 仍保留 IPC fast path。
权衡:
- 可能把一次 prefix all-reduce 变成至多 batch_size 次 range all-reduce;但这不超过逐 request 执行的 collective 次数,且避免 row-major gap 过通信。
- 这只修正同步 compose 的过通信;bs>1 仍有 per-layer Python descriptor 开销、bs>1 L1 async prefetch 未支持等潜在性能项,后续需要用 ETE/benchmark 继续确认。
验证:
- 远端 `g0034` container `/sgl-workspace/sglang-tai` py_compile 通过。
- 远端 `test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py`:113 passed, 5 warnings, 2 subtests passed。
@@ -2310,6 +2310,72 @@ def build_batch_prefix_slot_span(
return (start_slot, end_slot)
def build_batch_prefix_slot_spans(
*,
logical_pages: torch.Tensor,
prefix_lens_cpu,
page_size: int,
) -> list[tuple[int, int]]:
"""Return exact flattened page-table slot spans for batched prefix pages.
Unlike :func:`build_batch_prefix_slot_span`, this does not collapse multiple
request rows into one bounding span. Row-major page tables can have large
current/empty gaps between request prefixes; reducing those gaps is a
performance bug for cache-hit-heavy bs>1 prefill.
"""
if page_size <= 0:
raise ValueError(f"page_size must be positive, got {page_size}")
if prefix_lens_cpu is None:
raise ValueError("prefix_lens_cpu is required for batch prefix slot spans")
prefix_lens = [int(x) for x in prefix_lens_cpu]
if not prefix_lens:
return []
if logical_pages.dim() == 1:
if len(prefix_lens) != 1:
raise ValueError(
"1D logical_pages can only describe one request for prefix slot spans: "
f"batch_size={len(prefix_lens)} logical_pages_shape={tuple(logical_pages.shape)}"
)
pages_per_request = int(logical_pages.numel())
else:
if int(logical_pages.shape[0]) < len(prefix_lens):
raise ValueError(
"logical_pages has fewer rows than prefix_lens_cpu: "
f"rows={int(logical_pages.shape[0])} batch_size={len(prefix_lens)}"
)
pages_per_request = int(
logical_pages.reshape(logical_pages.shape[0], -1).shape[1]
)
spans: list[tuple[int, int]] = []
for req_id, prefix_len in enumerate(prefix_lens):
if prefix_len < 0:
raise ValueError(
f"prefix_lens_cpu contains negative prefix length: req_id={req_id} "
f"prefix_len={prefix_len}"
)
if prefix_len % page_size != 0:
raise ValueError(
"CP shared KV batch prefix slot spans require page-aligned prefixes: "
f"req_id={req_id} prefix_len={prefix_len} page_size={page_size}"
)
prefix_pages = prefix_len // page_size
if prefix_pages == 0:
continue
if prefix_pages > pages_per_request:
raise ValueError(
"prefix pages exceed per-request logical page-table width: "
f"req_id={req_id} prefix_pages={prefix_pages} "
f"pages_per_request={pages_per_request}"
)
req_start = req_id * pages_per_request
spans.append((req_start, req_start + prefix_pages))
return _merge_slot_spans(spans)
def _merge_slot_spans(spans: list[tuple[int, int]]) -> list[tuple[int, int]]:
normalized = sorted(
(int(start), int(end)) for start, end in spans if int(end) > int(start)
@@ -3402,6 +3468,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
page_size: int,
prefix_pages: int,
prefix_slot_span: tuple[int, int] | None = None,
prefix_slot_spans: list[tuple[int, int]] | None = None,
current_slot_spans: list[tuple[int, int]] | None = None,
layer_id: int | None = None,
nvtx_source: str = "mla.partial_current_sync",
@@ -3416,14 +3483,30 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
"""
total_slots = int(slot_remap.slot_logical_pages.numel())
if prefix_slot_span is None:
if prefix_slot_spans is not None and prefix_slot_span is not None:
raise ValueError(
"Specify either prefix_slot_span or prefix_slot_spans, not both."
)
if prefix_slot_spans is not None:
prefix_spans = _merge_slot_spans(prefix_slot_spans)
for prefix_start_slot, prefix_end_slot in prefix_spans:
if (
prefix_start_slot < 0
or prefix_end_slot < prefix_start_slot
or prefix_end_slot > total_slots
):
raise ValueError(
"Invalid CP shared KV partial-current prefix slot span: "
f"prefix_slot_span={(prefix_start_slot, prefix_end_slot)} "
f"total_slots={total_slots}"
)
elif prefix_slot_span is None:
if prefix_pages < 0 or prefix_pages > total_slots:
raise ValueError(
"Invalid CP shared KV partial-current prefix range: "
f"prefix_pages={prefix_pages} total_slots={total_slots}"
)
prefix_start_slot = 0
prefix_end_slot = int(prefix_pages)
prefix_spans = [(0, int(prefix_pages))] if int(prefix_pages) > 0 else []
else:
prefix_start_slot, prefix_end_slot = (
int(prefix_slot_span[0]),
@@ -3438,21 +3521,19 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
"Invalid CP shared KV partial-current prefix slot span: "
f"prefix_slot_span={prefix_slot_span} total_slots={total_slots}"
)
prefix_spans = (
[(prefix_start_slot, prefix_end_slot)]
if prefix_end_slot > prefix_start_slot
else []
)
dense_kv_cache = kv_cache.new_zeros(
(slot_remap.dense_num_pages * page_size, *kv_cache.shape[1:])
)
materialized_by_ipc = _try_tai_ipc_materialize_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=slot_remap.slot_logical_pages,
layout=layout,
page_size=page_size,
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
if not materialized_by_ipc:
materialize_local_token_kv_page_slots_into(
materialized_by_ipc = False
if len(prefix_spans) == 1:
prefix_start_slot, prefix_end_slot = prefix_spans[0]
materialized_by_ipc = _try_tai_ipc_materialize_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=slot_remap.slot_logical_pages,
@@ -3461,21 +3542,32 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
if not materialized_by_ipc:
for prefix_start_slot, prefix_end_slot in prefix_spans:
materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=slot_remap.slot_logical_pages,
layout=layout,
page_size=page_size,
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
prefix_rows = slot_range_to_token_slice(
page_size,
prefix_start_slot,
prefix_end_slot,
)
_all_reduce_materialized_buffer_range(
dense_kv_cache,
layout.cp_size,
prefix_rows.start,
prefix_rows.stop,
nvtx_source=nvtx_source,
nvtx_layer_id=layer_id,
nvtx_cp_rank=layout.cp_rank,
)
prefix_rows = slot_range_to_token_slice(
page_size,
prefix_start_slot,
prefix_end_slot,
)
_all_reduce_materialized_buffer_range(
dense_kv_cache,
layout.cp_size,
prefix_rows.start,
prefix_rows.stop,
nvtx_source=nvtx_source,
nvtx_layer_id=layer_id,
nvtx_cp_rank=layout.cp_rank,
)
logical_locs = filter_locs_mappable_to_physical_pool(
logical_locs=logical_locs,
@@ -3499,7 +3591,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
)
if layout.cp_size > 1:
if current_slot_spans is None:
if prefix_slot_span is not None:
if prefix_slot_span is not None or prefix_slot_spans is not None:
raise ValueError(
"CP shared KV batched current compose requires explicit "
"current_slot_spans to avoid reducing prefix slots twice."
@@ -3541,6 +3633,7 @@ def materialize_prefix_and_reuse_current_index_page_slots(
index_head_dim: int,
prefix_pages: int,
prefix_slot_span: tuple[int, int] | None = None,
prefix_slot_spans: list[tuple[int, int]] | None = None,
current_slot_spans: list[tuple[int, int]] | None = None,
layer_id: int | None = None,
nvtx_source: str = "index.partial_current_sync",
@@ -3548,14 +3641,30 @@ def materialize_prefix_and_reuse_current_index_page_slots(
"""Synchronously compose prefix index materialization with current index rows."""
total_slots = int(slot_remap.slot_logical_pages.numel())
if prefix_slot_span is None:
if prefix_slot_spans is not None and prefix_slot_span is not None:
raise ValueError(
"Specify either prefix_slot_span or prefix_slot_spans, not both."
)
if prefix_slot_spans is not None:
prefix_spans = _merge_slot_spans(prefix_slot_spans)
for prefix_start_slot, prefix_end_slot in prefix_spans:
if (
prefix_start_slot < 0
or prefix_end_slot < prefix_start_slot
or prefix_end_slot > total_slots
):
raise ValueError(
"Invalid CP shared KV index partial-current prefix slot span: "
f"prefix_slot_span={(prefix_start_slot, prefix_end_slot)} "
f"total_slots={total_slots}"
)
elif prefix_slot_span is None:
if prefix_pages < 0 or prefix_pages > total_slots:
raise ValueError(
"Invalid CP shared KV index partial-current prefix range: "
f"prefix_pages={prefix_pages} total_slots={total_slots}"
)
prefix_start_slot = 0
prefix_end_slot = int(prefix_pages)
prefix_spans = [(0, int(prefix_pages))] if int(prefix_pages) > 0 else []
else:
prefix_start_slot, prefix_end_slot = (
int(prefix_slot_span[0]),
@@ -3570,20 +3679,19 @@ def materialize_prefix_and_reuse_current_index_page_slots(
"Invalid CP shared KV index partial-current prefix slot span: "
f"prefix_slot_span={prefix_slot_span} total_slots={total_slots}"
)
prefix_spans = (
[(prefix_start_slot, prefix_end_slot)]
if prefix_end_slot > prefix_start_slot
else []
)
dense_page_buffer = page_buffer.new_zeros(
(slot_remap.dense_num_pages, *page_buffer.shape[1:])
)
materialized_by_ipc = _try_tai_ipc_materialize_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=slot_remap.slot_logical_pages,
layout=layout,
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
if not materialized_by_ipc:
materialize_local_paged_buffer_page_slots_into(
materialized_by_ipc = False
if len(prefix_spans) == 1:
prefix_start_slot, prefix_end_slot = prefix_spans[0]
materialized_by_ipc = _try_tai_ipc_materialize_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=slot_remap.slot_logical_pages,
@@ -3591,16 +3699,26 @@ def materialize_prefix_and_reuse_current_index_page_slots(
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
prefix_rows = slot_range_to_page_slice(prefix_start_slot, prefix_end_slot)
_all_reduce_materialized_buffer_range(
dense_page_buffer,
layout.cp_size,
prefix_rows.start,
prefix_rows.stop,
nvtx_source=nvtx_source,
nvtx_layer_id=layer_id,
nvtx_cp_rank=layout.cp_rank,
)
if not materialized_by_ipc:
for prefix_start_slot, prefix_end_slot in prefix_spans:
materialize_local_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=slot_remap.slot_logical_pages,
layout=layout,
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
prefix_rows = slot_range_to_page_slice(prefix_start_slot, prefix_end_slot)
_all_reduce_materialized_buffer_range(
dense_page_buffer,
layout.cp_size,
prefix_rows.start,
prefix_rows.stop,
nvtx_source=nvtx_source,
nvtx_layer_id=layer_id,
nvtx_cp_rank=layout.cp_rank,
)
dense_page_buffer = fill_current_index_page_slots(
dense_page_buffer=dense_page_buffer,
current_index_k=current_index_k,
@@ -3612,7 +3730,7 @@ def materialize_prefix_and_reuse_current_index_page_slots(
)
if layout.cp_size > 1:
if current_slot_spans is None:
if prefix_slot_span is not None:
if prefix_slot_span is not None or prefix_slot_spans is not None:
raise ValueError(
"CP shared KV batched index current compose requires explicit "
"current_slot_spans to avoid reducing prefix slots twice."
@@ -17,7 +17,7 @@ from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa import index_buf_accessor
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
build_batch_current_slot_spans,
build_batch_prefix_slot_span,
build_batch_prefix_slot_spans,
cp_shared_kv_debug_enabled,
cp_shared_kv_debug_log,
cp_shared_kv_mla_prefetch_enabled,
@@ -593,7 +593,7 @@ class Indexer(MultiPlatformOp):
f"current_scale_shape={tuple(current_index_kv[1].shape)} "
f"out_cache_loc_shape={tuple(current_locs.shape)}"
)
prefix_slot_span = None
prefix_slot_spans = None
current_slot_spans = build_batch_current_slot_spans(
logical_pages=logical_page_table,
prefix_lens_cpu=prefix_lens_cpu,
@@ -604,7 +604,7 @@ class Indexer(MultiPlatformOp):
prefix_pages = int(prefix_lens_cpu[0]) // page_size
else:
prefix_pages = 0
prefix_slot_span = build_batch_prefix_slot_span(
prefix_slot_spans = build_batch_prefix_slot_spans(
logical_pages=logical_page_table,
prefix_lens_cpu=prefix_lens_cpu,
page_size=page_size,
@@ -663,7 +663,7 @@ class Indexer(MultiPlatformOp):
page_size=page_size,
index_head_dim=forward_batch.token_to_kv_pool.index_head_dim,
prefix_pages=prefix_pages,
prefix_slot_span=prefix_slot_span,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
layer_id=layer_id,
)
@@ -675,13 +675,13 @@ class Indexer(MultiPlatformOp):
cp_shared_kv_mla_prefetch_log(
"index_partial_current_sync_compose cp_rank=%s layer=%s "
"prefix_lens=%s extend_lens=%s prefix_pages=%s "
"prefix_slot_span=%s current_rows=%s dense_pages=%s",
"prefix_slot_spans=%s current_rows=%s dense_pages=%s",
layout.cp_rank,
layer_id,
prefix_lens,
extend_lens,
prefix_pages,
prefix_slot_span,
prefix_slot_spans,
int(current_index_kv[0].shape[0]),
int(materialized.shape[0]),
)
@@ -16,7 +16,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_prefetch import (
)
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
build_batch_current_slot_spans,
build_batch_prefix_slot_span,
build_batch_prefix_slot_spans,
build_current_loc_remap,
cp_shared_kv_debug_enabled,
cp_shared_kv_debug_log,
@@ -2137,7 +2137,7 @@ class NativeSparseAttnBackend(
f"current_locs_shape={tuple(current_locs_for_reuse.shape)} "
f"page_size={page_size}"
)
prefix_slot_span = None
prefix_slot_spans = None
current_slot_spans = build_batch_current_slot_spans(
logical_pages=metadata.real_page_table,
prefix_lens_cpu=prefix_lens_cpu,
@@ -2148,7 +2148,7 @@ class NativeSparseAttnBackend(
prefix_pages = int(prefix_lens_cpu[0]) // page_size
else:
prefix_pages = 0
prefix_slot_span = build_batch_prefix_slot_span(
prefix_slot_spans = build_batch_prefix_slot_spans(
logical_pages=metadata.real_page_table,
prefix_lens_cpu=prefix_lens_cpu,
page_size=page_size,
@@ -2170,7 +2170,7 @@ class NativeSparseAttnBackend(
layout=forward_batch.cp_shared_kv_layout,
page_size=page_size,
prefix_pages=prefix_pages,
prefix_slot_span=prefix_slot_span,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
layer_id=layer.layer_id,
)
@@ -2184,7 +2184,7 @@ class NativeSparseAttnBackend(
cp_shared_kv_mla_prefetch_log(
"forward_partial_current_sync_compose cp_rank=%s "
"layer=%s reason=%s prefix_lens=%s extend_lens=%s "
"prefix_pages=%s prefix_slot_span=%s "
"prefix_pages=%s prefix_slot_spans=%s "
"current_rows=%s kv_rows=%s page_table_shape=%s",
forward_batch.cp_shared_kv_layout.cp_rank,
layer.layer_id,
@@ -2192,7 +2192,7 @@ class NativeSparseAttnBackend(
prefix_lens,
extend_lens,
prefix_pages,
prefix_slot_span,
prefix_slot_spans,
int(current_kv_cache.shape[0]),
int(kv_cache.shape[0]),
tuple(page_table_1.shape),
@@ -1436,6 +1436,16 @@ class HiCacheController:
return device_indices
def load_cp(self, nodes_to_load, node_id: int = -1) -> Optional[torch.Tensor]:
start_time = time.perf_counter()
stage_start_time = start_time
stage_durations_ms: List[tuple[str, float]] = []
def record_stage(stage: str) -> None:
nonlocal stage_start_time
now = time.perf_counter()
stage_durations_ms.append((stage, (now - stage_start_time) * 1000.0))
stage_start_time = now
# Reproduce the original (write-time) CP owner pattern. Each node
# carries `page_owners` (one int8 per logical page, identical on all
# CP ranks) so we can ask the allocator for a fresh device range
@@ -1455,10 +1465,12 @@ class HiCacheController:
)
# page_owners is CPU int8; .tolist() returns list[int] directly.
page_owners.extend(meta.page_owners.tolist())
record_stage("collect_page_owners")
device_indices = self.mem_pool_device_allocator.alloc_pages_with_owners(
page_owners
)
record_stage("alloc_pages_with_owners")
# Fail closed: returning None lets the caller drop to cache miss
# (cold prefill). Never proceed with a non-matching owner pattern.
if device_indices is None:
@@ -1537,6 +1549,7 @@ class HiCacheController:
host_chunks.append(meta.host_indices)
if draft_host_indices is not None:
draft_host_chunks.append(draft_host_indices)
record_stage("build_chunks")
visible_device_indices = (
torch.cat(visible_chunks)
@@ -1566,6 +1579,7 @@ class HiCacheController:
draft_host_indices = None
if self.has_draft_hicache:
draft_host_indices = torch.cat(draft_host_chunks)
record_stage("concat_indices")
try:
self._validate_cp_hicache_page_indices(
@@ -1575,6 +1589,7 @@ class HiCacheController:
self._validate_cp_hicache_page_indices(
draft_host_indices, physical_device_indices
)
record_stage("validate_indices")
except Exception:
self.mem_pool_device_allocator.free(device_indices)
raise
@@ -1595,6 +1610,19 @@ class HiCacheController:
)
)
total_duration = time.perf_counter() - start_time
if total_duration >= 1.0:
logger.warning(
"[HiCache-load] slow load_cp planning: node_id=%d duration_ms=%.3f "
"pages=%d host_indices=%d physical_indices=%d draft=%s stages_ms=%s",
node_id,
total_duration * 1000.0,
len(page_owners),
int(host_indices.numel()),
int(physical_device_indices.numel()),
draft_host_indices is not None,
[(stage, round(ms, 3)) for stage, ms in stage_durations_ms],
)
return visible_device_indices
def move_indices(self, op: CacheOperation, mem_pool_host=None):
+53 -14
View File
@@ -356,7 +356,12 @@ class CpLoadBackPlan:
page_owners: List[int]
required_by_owner: List[int]
available_by_owner: List[int]
# Exact capacity deficit. This is the only deficit allowed to block the
# synchronous load-back admission path.
deficit_by_owner: List[int]
# Advisory free-room deficit. This is logged for observability but must not
# force synchronous load-back eviction when exact capacity already fits.
free_room_deficit_by_owner: List[int]
host_hit_len: int
@@ -1017,39 +1022,36 @@ class HiRadixCache(RadixCache):
f"page_size={self.page_size} page_owners={len(page_owners)}"
)
required, available, deficits = self._cp_load_back_owner_lane_stats(page_owners)
required, available, deficits, free_room_deficits = (
self._cp_load_back_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],
free_room_deficit_by_owner=[int(v) for v in free_room_deficits],
host_hit_len=host_hit_len,
)
def _cp_load_back_owner_lane_stats(
self, page_owners: List[int]
) -> Tuple[List[int], List[int], List[int]]:
) -> Tuple[List[int], List[int], List[int], List[int]]:
allocator = self.token_to_kv_pool_allocator
target_ratio = float(getattr(self, "hicache_l1_free_room_ratio", 0.0) or 0.0)
trigger_ratio = float(
getattr(self, "hicache_l1_free_room_trigger_ratio", 0.0) or 0.0
)
free_room_stats = getattr(allocator, "compute_owner_lane_free_room_stats", None)
if free_room_stats is not None:
return free_room_stats(
page_owners,
target_ratio=target_ratio,
trigger_ratio=trigger_ratio,
)
required, available, exact_deficits = allocator.compute_owner_lane_stats(
page_owners
)
lane_capacity_pages = getattr(allocator, "compute_owner_lane_capacity_pages", None)
if lane_capacity_pages is None:
return required, available, exact_deficits
return required, available, exact_deficits, exact_deficits
return (
required,
available,
exact_deficits,
compute_owner_lane_free_room_deficits(
required=required,
available=available,
@@ -1060,14 +1062,15 @@ class HiRadixCache(RadixCache):
)
def _refresh_cp_load_back_plan(self, plan: CpLoadBackPlan) -> CpLoadBackPlan:
required, available, deficits = self._cp_load_back_owner_lane_stats(
plan.page_owners
required, available, deficits, free_room_deficits = (
self._cp_load_back_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],
free_room_deficit_by_owner=[int(v) for v in free_room_deficits],
host_hit_len=plan.host_hit_len,
)
@@ -1327,6 +1330,7 @@ class HiRadixCache(RadixCache):
required_by_owner=[],
available_by_owner=[],
deficit_by_owner=deficits,
free_room_deficit_by_owner=deficits,
host_hit_len=0,
)
eviction_plan = self._plan_cp_load_back_owner_lane_evictions(plan)
@@ -3362,6 +3366,15 @@ class HiRadixCache(RadixCache):
if self._uses_cp_hicache:
start_time = time.perf_counter()
stage_start_time = start_time
stage_durations_ms: List[Tuple[str, float]] = []
def record_stage(stage: str) -> None:
nonlocal stage_start_time
now = time.perf_counter()
stage_durations_ms.append((stage, (now - stage_start_time) * 1000.0))
stage_start_time = now
last_hit_node = node
nodes_to_load = []
while node.evicted:
@@ -3388,6 +3401,7 @@ class HiRadixCache(RadixCache):
load_back_plan = self._build_cp_load_back_plan(
nodes_to_load, node_id=last_hit_node.id
)
record_stage("build_plan")
except Exception:
self.dec_lock_ref(ancester_node)
raise
@@ -3395,7 +3409,8 @@ class HiRadixCache(RadixCache):
logger.info(
"[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",
"available_by_owner=%s deficit_by_owner=%s "
"free_room_deficit_by_owner=%s",
last_hit_node.id,
len(nodes_to_load),
host_hit_len,
@@ -3403,6 +3418,7 @@ class HiRadixCache(RadixCache):
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
load_back_plan.free_room_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
@@ -3420,6 +3436,7 @@ class HiRadixCache(RadixCache):
load_back_plan = self._evict_cp_load_back_owner_lanes(
load_back_plan, node_id=last_hit_node.id
)
record_stage("owner_lane_evict")
if any(v > 0 for v in load_back_plan.deficit_by_owner):
self.dec_lock_ref(ancester_node)
@@ -3427,7 +3444,8 @@ class HiRadixCache(RadixCache):
"[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 "
"deficit_by_owner=%s free_room_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,
@@ -3435,6 +3453,7 @@ class HiRadixCache(RadixCache):
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
load_back_plan.free_room_deficit_by_owner,
int(getattr(self, "evictable_size_", 0)),
int(getattr(self, "protected_size_", 0)),
len(getattr(self, "ongoing_load_back", {})),
@@ -3445,18 +3464,21 @@ class HiRadixCache(RadixCache):
device_indices = self.cache_controller.load_cp(
nodes_to_load, node_id=last_hit_node.id
)
record_stage("load_cp_plan")
if device_indices is None:
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 "
"free_room_deficit_by_owner=%s "
"allocator_state=%s",
last_hit_node.id,
host_hit_len,
failed_plan.required_by_owner,
failed_plan.available_by_owner,
failed_plan.deficit_by_owner,
failed_plan.free_room_deficit_by_owner,
self.token_to_kv_pool_allocator.allocator_state_str(),
)
self.dec_lock_ref(ancester_node)
@@ -3482,6 +3504,7 @@ class HiRadixCache(RadixCache):
physical_loaded_len += int(
getattr(metadata, "padded_len", host_len)
)
record_stage("assign_tree")
self.evictable_size_ += physical_loaded_len
self.inc_lock_ref(last_hit_node)
@@ -3499,6 +3522,22 @@ class HiRadixCache(RadixCache):
len(device_indices),
physical_loaded_len,
)
total_duration = time.perf_counter() - start_time
if total_duration >= 1.0:
logger.warning(
"[HiCache-load] slow CP load_back planning: node_id=%d "
"duration_ms=%.3f host_hit_len=%d required_by_owner=%s "
"available_by_owner=%s exact_deficit_by_owner=%s "
"free_room_deficit_by_owner=%s stages_ms=%s",
last_hit_node.id,
total_duration * 1000.0,
host_hit_len,
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
load_back_plan.free_room_deficit_by_owner,
[(stage, round(ms, 3)) for stage, ms in stage_durations_ms],
)
return device_indices
start_time = time.perf_counter()
@@ -359,6 +359,45 @@ class TestCpHiCacheLoadBackOwnerLanes(CustomTestCase):
self.assertEqual(target.value.tolist(), loaded.tolist())
self.assertIn(target.id, cache.ongoing_load_back)
def test_load_back_does_not_synchronously_evict_for_l1_free_room_when_exact_capacity_fits(self):
allocator = _make_allocator()
# Owner lane 0 has exactly enough capacity for the load-back target, but
# not enough to satisfy the configured free-room target. Load-back must
# not synchronously evict just to refill free room; doing so puts heavy
# eviction planning on the scheduler hot path.
allocator.free_pages = torch.tensor([1], dtype=torch.int64)
cache = _make_cache(allocator)
cache.hicache_l1_free_room_ratio = 0.5
cache.hicache_l1_free_room_trigger_ratio = 0.25
victim = _make_node(
22,
220,
[0],
value=torch.tensor([8, 9, 10, 11], 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(23, 320, [0], value=None, priority=10)
_attach_child(cache, cache.root_node, target)
plan = cache._build_cp_load_back_plan([target], node_id=target.id)
self.assertEqual(plan.deficit_by_owner, [0, 0, 0, 0])
# Existing free-room accounting reports all lanes below the watermark.
# Load-back keeps that signal advisory and does not synchronously act on it.
self.assertEqual(plan.free_room_deficit_by_owner, [2, 2, 2, 2])
loaded = cache.load_back(target, mem_quota=100)
self.assertIsNotNone(loaded)
self.assertEqual(cache.cache_controller.load_calls, 1)
self.assertEqual(cache.cache_controller.evicted_device_indices, [])
self.assertEqual(target.value.tolist(), loaded.tolist())
self.assertIn(target.id, cache.ongoing_load_back)
def test_owner_lane_evict_params_choose_deficit_contributing_victim(self):
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
@@ -2839,7 +2839,7 @@ class TestHiRadixCacheCPSplitEvict(CustomTestCase):
class TestHiRadixCacheCPLoadBack(CustomTestCase):
def test_cp_load_back_plan_uses_l1_free_room_target(self):
def test_cp_load_back_plan_reports_l1_free_room_without_blocking_exact_fit(self):
class FreeRoomAllocator:
def compute_owner_lane_stats(self, _page_owners):
return [1, 0], [0, 8], [1, 0]
@@ -2869,8 +2869,11 @@ class TestHiRadixCacheCPLoadBack(CustomTestCase):
self.assertEqual(plan.required_by_owner, [1, 0])
self.assertEqual(plan.available_by_owner, [0, 8])
# Synchronous load-back admission must only block on exact capacity.
self.assertEqual(plan.deficit_by_owner, [1, 0])
# The free-room target is still reported for observability/proactive policy.
# required=1 page, available=0, target_room=ceil(8*0.5)=4 pages.
self.assertEqual(plan.deficit_by_owner, [5, 0])
self.assertEqual(plan.free_room_deficit_by_owner, [5, 0])
def test_cp_load_back_uses_host_len_not_host_value(self):
cache = HiRadixCache.__new__(HiRadixCache)
@@ -749,7 +749,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertIn("prefix_pages=0", branch_source)
self.assertNotIn("kv_cache = current_kv_cache", branch_source)
def test_mla_partial_current_sync_uses_batch_prefix_slot_span(self):
def test_mla_partial_current_sync_uses_batch_prefix_slot_spans(self):
from pathlib import Path
source = (
@@ -762,8 +762,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
branch_source = source[branch_start:branch_end]
local_locs_source = source[method_start:branch_end]
self.assertIn("build_batch_prefix_slot_span", source)
self.assertIn("prefix_slot_span=", branch_source)
self.assertIn("build_batch_prefix_slot_spans", source)
self.assertIn("prefix_slot_spans=", branch_source)
self.assertIn("get_cp_shared_kv_local_out_cache_loc", local_locs_source)
self.assertNotIn("current_locs = forward_batch.out_cache_loc", branch_source)
self.assertNotIn("or len(prefix_lens_cpu) != 1", branch_source)
@@ -1372,7 +1372,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertTrue(torch.equal(mixed_kv[12:14], current_kv))
self.assertEqual(mixed_locs.tolist(), [[4, 8, 12, 13, -1, -1]])
def test_batch_prefix_slot_span_covers_request_prefix_pages_only(self):
def test_batch_prefix_slot_span_covers_bounding_prefix_range(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
logical_pages = torch.tensor(
@@ -1408,6 +1408,42 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
(0, 0),
)
def test_batch_prefix_slot_spans_keep_request_prefix_ranges_exact(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
logical_pages = torch.tensor(
[
[1, 2, 5],
[9, 11, 0],
],
dtype=torch.int64,
)
self.assertEqual(
runtime.build_batch_prefix_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[8, 4],
page_size=4,
),
[(0, 2), (3, 4)],
)
self.assertEqual(
runtime.build_batch_prefix_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[0, 4],
page_size=4,
),
[(3, 4)],
)
self.assertEqual(
runtime.build_batch_prefix_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[0, 0],
page_size=4,
),
[],
)
def test_batch_current_slot_spans_follow_prefix_and_extend_pages(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
@@ -1468,7 +1504,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=layout,
page_size=page_size,
)
prefix_slot_span = runtime.build_batch_prefix_slot_span(
prefix_slot_spans = runtime.build_batch_prefix_slot_spans(
logical_pages=remap_logical_pages,
prefix_lens_cpu=[8, 4],
page_size=page_size,
@@ -1493,7 +1529,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=layout,
page_size=page_size,
prefix_pages=0,
prefix_slot_span=prefix_slot_span,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
)
)
@@ -1505,6 +1541,72 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertTrue(torch.equal(mixed_kv[12:14], current_kv[:2]))
self.assertTrue(torch.equal(mixed_kv[20:22], current_kv[2:]))
def test_materialize_batch_prefix_spans_do_not_reduce_row_gaps(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
page_size = 4
layout = CpSharedKVLayout(page_size=page_size, cp_size=1, cp_rank=0)
kv_cache = torch.arange(0, 64, dtype=torch.float32).view(64, 1, 1)
logical_locs = torch.tensor(
[
[4, 8, 20, 21, -1, -1],
[36, 44, 45, -1, -1, -1],
],
dtype=torch.int64,
)
current_locs = torch.tensor([20, 21, 44, 45], dtype=torch.int64)
current_kv = torch.arange(100, 104, dtype=torch.float32).view(4, 1, 1)
remap_logical_pages = torch.tensor(
[
[1, 2, 5],
[9, 11, 0],
],
dtype=torch.int64,
)
slot_remap = runtime.build_shared_token_kv_slot_remap(
kv_cache=kv_cache,
logical_locs=logical_locs,
remap_logical_pages=remap_logical_pages,
layout=layout,
page_size=page_size,
)
prefix_slot_spans = runtime.build_batch_prefix_slot_spans(
logical_pages=remap_logical_pages,
prefix_lens_cpu=[8, 4],
page_size=page_size,
)
current_slot_spans = runtime.build_batch_current_slot_spans(
logical_pages=remap_logical_pages,
prefix_lens_cpu=[8, 4],
extend_lens_cpu=[2, 2],
page_size=page_size,
)
reduced_ranges = []
def record_all_reduce(buffer, cp_size, start, end, **kwargs):
reduced_ranges.append((start, end))
return buffer
with patch.object(
runtime, "_all_reduce_materialized_buffer_range", record_all_reduce
):
runtime.materialize_prefix_and_reuse_current_kv_page_slots(
kv_cache=kv_cache,
logical_locs=logical_locs,
current_kv_cache=current_kv,
current_locs=current_locs,
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=0,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
)
self.assertEqual(reduced_ranges[:2], [(4, 12), (16, 20)])
self.assertNotIn((4, 20), reduced_ranges)
def test_materialize_batch_prefix_span_and_reuse_current_index_page_slots(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
@@ -1524,7 +1626,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
logical_pages,
layout,
)
prefix_slot_span = runtime.build_batch_prefix_slot_span(
prefix_slot_spans = runtime.build_batch_prefix_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[8, 4],
page_size=page_size,
@@ -1555,7 +1657,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
page_size=page_size,
index_head_dim=index_head_dim,
prefix_pages=0,
prefix_slot_span=prefix_slot_span,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
layer_id=2,
)
@@ -1583,7 +1685,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
)
)
def test_index_partial_current_sync_uses_batch_prefix_slot_span(self):
def test_index_partial_current_sync_uses_batch_prefix_slot_spans(self):
from pathlib import Path
source = (
@@ -1598,8 +1700,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
branch_source = source[branch_start:branch_end]
branch_compact = "".join(branch_source.split())
self.assertIn("build_batch_prefix_slot_span", source)
self.assertIn("prefix_slot_span=", branch_source)
self.assertIn("build_batch_prefix_slot_spans", source)
self.assertIn("prefix_slot_spans=", branch_source)
self.assertIn("get_cp_shared_kv_local_out_cache_loc", branch_source)
self.assertNotIn("current_locs=forward_batch.out_cache_loc", branch_compact)
self.assertNotIn("or len(prefix_lens_cpu) != 1", branch_source)