From 81bfd4bfd38ad7d725a2e0c88e5ead3577f327e7 Mon Sep 17 00:00:00 2001 From: laoyao0822 Date: Sat, 6 Jun 2026 02:03:16 +0800 Subject: [PATCH] Reduce CP shared-KV batch overhead without reverting bs1 planning MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- ...cp_hicache_reactive_host_free_room_plan.md | 46 ++++ ...hared_kv_bs_gt1_parallel_workstreams_zh.md | 77 ++++++ .../attention/nsa/cp_shared_kv_runtime.py | 224 +++++++++++++----- .../srt/layers/attention/nsa/nsa_indexer.py | 12 +- .../srt/layers/attention/nsa_backend.py | 12 +- .../sglang/srt/managers/cache_controller.py | 28 +++ python/sglang/srt/mem_cache/hiradix_cache.py | 67 ++++-- .../test_cp_hicache_load_back_owner_lanes.py | 39 +++ .../mem_cache/test_cp_hicache_metadata.py | 7 +- .../mem_cache/test_cp_shared_kv_runtime.py | 124 +++++++++- 10 files changed, 544 insertions(+), 92 deletions(-) diff --git a/docs/advanced_features/nsa_prefill_cp_hicache_reactive_host_free_room_plan.md b/docs/advanced_features/nsa_prefill_cp_hicache_reactive_host_free_room_plan.md index cfaf1efba..979084640 100644 --- a/docs/advanced_features/nsa_prefill_cp_hicache_reactive_host_free_room_plan.md +++ b/docs/advanced_features/nsa_prefill_cp_hicache_reactive_host_free_room_plan.md @@ -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. diff --git a/docs/advanced_features/nsa_prefill_cp_shared_kv_bs_gt1_parallel_workstreams_zh.md b/docs/advanced_features/nsa_prefill_cp_shared_kv_bs_gt1_parallel_workstreams_zh.md index 8deec955d..6d0132f05 100644 --- a/docs/advanced_features/nsa_prefill_cp_shared_kv_bs_gt1_parallel_workstreams_zh.md +++ b/docs/advanced_features/nsa_prefill_cp_shared_kv_bs_gt1_parallel_workstreams_zh.md @@ -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。 diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py index 63c3f3916..d1c9d583b 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py @@ -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." diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 2b0eb52d6..47373ffc6 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -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]), ) diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index d55b83f99..41fcc5fbb 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -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), diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index bde98a514..55cc30565 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -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): diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index d085022f5..895b5d8e4 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -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() diff --git a/test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py b/test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py index d94628afd..565c658cb 100644 --- a/test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py +++ b/test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py @@ -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 diff --git a/test/registered/unit/mem_cache/test_cp_hicache_metadata.py b/test/registered/unit/mem_cache/test_cp_hicache_metadata.py index dbe5b2f2a..40969ed87 100644 --- a/test/registered/unit/mem_cache/test_cp_hicache_metadata.py +++ b/test/registered/unit/mem_cache/test_cp_hicache_metadata.py @@ -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) diff --git a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py index f90231e41..27fe81142 100644 --- a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py @@ -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)