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:
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user