From 142e7a5a645a43053b1f3f7289d2d387ef34d0b4 Mon Sep 17 00:00:00 2001 From: laoyao0822 Date: Fri, 12 Jun 2026 03:07:47 +0800 Subject: [PATCH] Keep CP shared-KV fast paths off dense current collectives CP shared-KV cache-hit batches should compose long prefix pages and short current pages through page-slot IPC instead of falling back to dense all_reduce. Wire the runtime and prefetch consume paths to the TAI current-staging helpers, fail fast when the configured CUDA fast path cannot run, and document the bs>1 cache-hit benchmark evidence. Constraint: bs>1 prefill must preserve the page-slot contract across fp8/bf16 and zero-lane current tails. Rejected: Silent all_reduce fallback | hides correctness and performance regressions in production. Confidence: medium Scope-risk: moderate Directive: Any future fallback in CP shared-KV CUDA fast paths must be explicit warning/fail-fast and covered by runtime tests. Tested: Local py_compile cp_shared_kv_runtime.py and cp_shared_kv_prefetch.py; remote PYTHONPATH=python pytest -q test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py (144 passed, 21 warnings, 2 subtests passed); remote TAI IPC benchmark fp8 bs>1 cache-hit matrix recorded in docs. Not-tested: Full ETE mixed replay after replacing all current collectives with IPC. (cherry picked from commit 8aa3b4ce59e0ebef5da5b0d07499a5f1d9785997) --- ...d_kv_bs_gt1_l1_prefetch_zero_sm_plan_zh.md | 170 +- ...d_kv_ipc_collective_replacement_plan_zh.md | 271 +++ .../attention/nsa/cp_shared_kv_prefetch.py | 1164 ++++++------ .../attention/nsa/cp_shared_kv_runtime.py | 1589 ++++++----------- .../mem_cache/test_cp_shared_kv_runtime.py | 1132 +++++++++--- 5 files changed, 2471 insertions(+), 1855 deletions(-) create mode 100644 docs/advanced_features/nsa_prefill_cp_shared_kv_ipc_collective_replacement_plan_zh.md diff --git a/docs/advanced_features/nsa_prefill_cp_shared_kv_bs_gt1_l1_prefetch_zero_sm_plan_zh.md b/docs/advanced_features/nsa_prefill_cp_shared_kv_bs_gt1_l1_prefetch_zero_sm_plan_zh.md index c093b398e..c2037cfc8 100644 --- a/docs/advanced_features/nsa_prefill_cp_shared_kv_bs_gt1_l1_prefetch_zero_sm_plan_zh.md +++ b/docs/advanced_features/nsa_prefill_cp_shared_kv_bs_gt1_l1_prefetch_zero_sm_plan_zh.md @@ -442,6 +442,23 @@ L2->L1 load finished on owner rank ### P1. 补 bs>1 prefetch plan 单测 +**状态(2026-06-12):已完成 MLA baseline。** + +已新增: + +- `test_mla_prefetch_create_batch_uses_exact_prefix_and_current_spans` +- `test_mla_prefetch_batch_consume_reduces_exact_current_spans` + +RED 证据:旧代码在 `batch_size=2` 时 `maybe_create()` 直接返回 `None`,且 +`CpSharedKVMlaPrefetcher.__init__()` 不接受 `prefix_slot_spans/current_slot_spans`。 + +GREEN 证据:远端 `cjy-glm5-new` 容器内 +`test_cp_shared_kv_runtime.py` 全文件通过: + +```text +132 passed, 21 warnings, 2 subtests passed +``` + **文件:** - 修改:`test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py` @@ -479,6 +496,31 @@ PYTHONPATH=python python -m pytest -q \ ### P2. 改造 MLA prefetcher 支持 bs>1 spans +**状态(2026-06-12):已完成第一版 spans baseline。** + +当前实现: + +1. `CpSharedKVMlaPrefetcher.maybe_create()` 不再以 `batch_size != 1` 为 + skip 条件。 +2. create 阶段基于 `metadata.real_page_table`、`extend_prefix_lens_cpu`、 + `extend_seq_lens_cpu` 构造: + - `prefix_slot_spans` + - `current_slot_spans` + - `prefix_page_count` + - `current_page_count` +3. `start_next_layer_prefix()` 只 materialize/reduce `prefix_slot_spans`, + 不再把 batch flattened page table 当成 `[0:prefix_pages)`。 +4. `consume_prefix_with_current()` 只 reduce `current_slot_spans`,避免把 + row gap / 其他 request prefix 一起 reduce。 +5. `consume()` 的 legacy full-materialize suffix 路径也改为使用 + `current_slot_spans`,避免 bs>1 bounding suffix。 + +当前限制: + +- MLA prefix spans 仍走现有 materialize + async all-reduce baseline;还没有 + 接入 spans-list TAI IPC 或 0SM CE。 +- index prefetcher 仍未改造,继续由 P3 处理。 + **文件:** - 修改:`python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py` @@ -490,7 +532,8 @@ PYTHONPATH=python python -m pytest -q \ 2. create 阶段构造 `prefix_slot_spans/current_slot_spans`。 3. `start_next_layer_prefix()` 只 materialize prefix spans。 4. `consume_prefix_with_current()` 只 fill/reduce current spans。 -5. `consume()` 如果仍存在 legacy suffix path,bs>1 下 fail-fast,避免错误 bounding suffix。 +5. `consume()` 如果仍存在 legacy suffix path,必须使用 `current_slot_spans` + 或 fail-fast,不能回到错误 bounding suffix。 **第一版允许:** @@ -503,6 +546,46 @@ PYTHONPATH=python python -m pytest -q \ ### P3. 改造 index prefetcher 支持 bs>1 spans +**状态(2026-06-12):已完成第一版 spans baseline。** + +已新增: + +- `test_index_prefetch_create_batch_uses_exact_prefix_and_current_spans` +- `test_index_prefetch_batch_consume_reduces_exact_current_spans` + +RED 证据:旧代码在 `batch_size=2` 时以 +`[CP_SHARED_KV_FALLBACK][index_prefetch] reason=batch_size` 返回 `None`, +且 `CpSharedKVIndexPrefetcher.__init__()` 不接受 +`prefix_slot_spans/current_slot_spans`。 + +当前实现: + +1. `CpSharedKVIndexPrefetcher.maybe_create()` 不再以 `batch_size != 1` + 为 skip 条件。 +2. create 阶段复用 MLA 同一套 deterministic spans: + - `prefix_slot_spans` + - `current_slot_spans` + - `prefix_page_count` + - `current_page_count` +3. `start_next_layer_prefix()` 只 materialize/reduce index prefix spans。 +4. `consume_prefix_with_current()` 只 fill/reduce index current spans。 +5. `consume()` 的 legacy suffix 路径也改为 `current_slot_spans`, + 不再使用 batch bounding suffix。 + +远端验证: + +```text +test_cp_shared_kv_runtime.py +134 passed, 21 warnings, 2 subtests passed +``` + +当前限制: + +- index prefetch 仍使用现有 materialize + async all-reduce baseline; + spans-list TAI IPC / 0SM CE 留给后续 P5/P6。 +- active index layer / index skip 的 runtime hook 当前沿用已有 + `nsa_backend.py` 调用点;本阶段没有修改 skip 参数语义。 + **文件:** - 修改:`python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py` @@ -547,6 +630,19 @@ PYTHONPATH=python python -m pytest -q \ - forward stream 上同步等待 L2->L1。 - 用新的 all-reduce 确认所有 rank load 完成。 +**2026-06-12 P4 实施记录:** + +- 当前 HiCache load 已通过 `LayerDoneCounter` 暴露 per-layer readiness: + - `LayerDoneCounter.wait_until_on_stream(layer_id - start_layer, stream)` + - `MLATokenToKVPool.get_key_buffer_for_prefetch(layer_id, stream)` + - `MLATokenToKVPool.get_index_k_with_scale_buffer_for_prefetch(layer_id, stream)` +- 发现并修正一个 ordering 疏漏:P1-P3 后 prefix materialize 仍在 current stream 上执行,但 `start_next_layer_prefix()` 把 L2->L1 readiness wait 绑定到了 prefetch stream。这样只能保护后续 reduce,不能保护实际读取 L1 raw pages 的 materialize。 +- 修正策略:MLA/index 的 `start_next_layer_prefix()` 先取得 `current_stream`,把它传给 prefetch-safe getter,使 L2->L1 ready event 挂到实际读取 KV/index page 的 stream;随后仍保持 `prefetch_stream.wait_stream(current_stream)`,reduce 在 prefetch stream 上异步提交。 +- 没有新增 collective;没有把 forward stream 改成 CPU 同步等待。 +- 单测锁住: + - `test_mla_prefetch_waits_l2_l1_on_materialize_stream_and_reduces_on_prefetch_stream` + - `test_index_prefetch_waits_l2_l1_on_materialize_stream_and_reduces_on_prefetch_stream` + ### P5. TAI SM IPC spans baseline **文件:** @@ -735,52 +831,48 @@ P7 GSM8K/replay/Nsight 验证 不要先做 P6 再修 P2/P3。0SM kernel 只解决 transport 代价,不解决 bs>1 prefix/current slot 语义;如果语义仍是 scalar prefix,kernel 越快只会越快地产生错误。 -## 2026-06-12 追加:symm compose 与旧 IPC kernel 的对比口径 +--- -`symm-syh` 分支的 compose/symm kernel 可能比旧 CUDA IPC current-staging kernel 单次更快,但不能只比较 kernel elapsed time。需要把同步点作为一等指标,否则可能出现“单 kernel 快,但每层/每 buffer 多一次同步,ETE 更慢”的误判。 +## 9. P5 当前实现状态:bs>1 prefix/suffix IPC spans -对比 benchmark / Nsight trace 必须至少拆分以下事件: +已补齐一个 TAI SM IPC baseline kernel,用于先替代 bs>1 prefix/suffix 上的 +“local materialize + all_reduce”: -1. **IPC capability agreement** - - `_agreed_tai_ipc_peer_ptrs()` 里的 group agreement / all-reduce。 - - 应确认是每个 pool tensor 一次,还是每个 forward/layer 重复触发。 +```text +materialize_cuda_ipc_peer_pages_slot_indices( + peer_ptrs, + dst, + owner_ranks, + src_page_indices, + dst_page_indices, + page_nbytes, +) +``` -2. **symm barrier** - - `cp_symm_barrier()` 调用次数、耗时、等待方差。 - - 需要按 token KV / index buffer 分开统计。 - - 当前 tai-kernel 实现是 1 个 block / 1 个 warp 的 CUDA spin barrier, - 不是 copy-engine/0SM 路径;它占用很少 SM,但会在当前 stream 上形成 - 明确同步点。判断 symm 是否优于旧 IPC 时,必须把这个 barrier 的次数和 - rank 间等待方差算进去。 +语义: -3. **symm mega gather** - - `materialize_cuda_ipc_peer_pages_slot_dense()` 在 symm combined ptr table 上的耗时。 - - 统计 prefix pages、current pages、request 数、dense pages。 +```text +dst[dst_page_indices[i]] = peer_ptrs[owner_ranks[i]][src_page_indices[i]] +``` -4. **compact current reduce** - - symm 未开启或不可用时的 `_reduce_current_pages_compact()`。 - - 这是 compact current collective,不是 dense full fallback;但仍会同步/占用通信资源,需要单独计数。 +与旧 `slot_dense` kernel 的区别: -5. **dense full fallback reduce** - - `_all_reduce_materialized_buffer(... v2_full ...)`。 - - 生产 CUDA + TAI materialize 开启时不允许静默发生;应 fail-fast 暴露 `CP_SHARED_KV_FAIL_FAST][compose_v2]`。 +1. 旧 kernel 只能写 `slot i -> dense page i+1`,等价于只支持从 slot 0 开始的一段连续 prefix。 +2. 新 kernel 显式传入 `dst_page_indices`,支持 bs>1 的多个 prefix spans 和 suffix spans。 +3. invalid owner/src 会 zero-fill 对应 dst page;未被 descriptor 指向的 dst page 保持原值。 +4. 该 kernel 仍是 SM copy baseline,不是最终 0SM copy-engine queue。 -6. **CPU descriptor / plan 成本** - - `get_or_build_compose_plan()` cache hit/miss。 - - per-forward 是否复用 descriptor;不要把一次性 build 成本误算到每层 steady state。 +SGLang 接入位置: -建议 benchmark 矩阵: +1. `cp_shared_kv_runtime.py` + - 新增 token/index span IPC helper。 + - partial-current prefix 支持多 span IPC,prefix 不再因为 bs>1 退回 all_reduce。 +2. `cp_shared_kv_prefetch.py` + - MLA/index prefix prefetch 先尝试 span IPC;成功时在 current stream record event,不再 enqueue prefix all_reduce。 + - full-cache-hit suffix consume 先尝试 span IPC;成功时跳过 suffix all_reduce。 -- dtype:bf16 / fp8_e4m3; -- batch size:1, 2, 5, 10; -- extend:1k, 2k, 10k, 40k, 65k; -- cached/prefix:100k, 200k, 300k; -- case:cache-hit partial-current、current-only、multi-request shared prefix; -- mode:legacy dense fallback(只作为基线,不允许生产静默)、old CUDA IPC current staging、symm staging、symm prefetch。 +尚未覆盖: -结论标准:保留默认路径必须同时满足: - -- 无 dense full fallback; -- 同步点数量不高于旧 IPC 路径,或同步耗时被更少 kernel/更高带宽抵消; -- ETE replay 在 cache-hit-heavy 短 extend 场景提升,而不是只在 micro benchmark 提升; -- GSM8K cache-hit 二轮精度不掉点。 +1. partial-current 的 current rows 仍是本 rank 当前 forward 产生的临时 buffer,source layout 不是长期 L1 page buffer;不能直接复用 peer page IPC。 +2. 该 current rows all_reduce 需要单独设计 owner-aware current-source IPC/fused compose kernel。 +3. 0SM CE path 仍未实现,本阶段只是先消除 bs>1 prefix/suffix 上不必要的 collective。 diff --git a/docs/advanced_features/nsa_prefill_cp_shared_kv_ipc_collective_replacement_plan_zh.md b/docs/advanced_features/nsa_prefill_cp_shared_kv_ipc_collective_replacement_plan_zh.md new file mode 100644 index 000000000..288ed36ea --- /dev/null +++ b/docs/advanced_features/nsa_prefill_cp_shared_kv_ipc_collective_replacement_plan_zh.md @@ -0,0 +1,271 @@ +# NSA Prefill CP shared-KV:用自研 IPC collective 替换 materialize all_reduce + +## 目标 + +彻底移除 CP shared-KV materialize 热路径上的 NCCL/Gloo `all_reduce`: + +1. prefix / suffix / full-cache-hit:继续使用 L1 page buffer 上的 IPC page gather,失败必须显式 warning / fail-fast,不能静默回退。 +2. partial-current / current reuse:新增 current staging + ready flag IPC collective,替换当前 dense page fill 后的 slot-range `all_reduce`。 +3. bs>1:descriptor 必须一次覆盖 batch 内多个 request 的 slot spans,不允许 per-request 循环发射 kernel。 +4. fp8 / bf16:token KV 与 index page buffer 两条路径都要支持。 +5. benchmark:必须覆盖 all_reduce baseline、现有 IPC prefix/suffix、current staging IPC,并输出 CPU submit、GPU elapsed、有效带宽、kernel launch 数。 + +## 当前 all_reduce 分类 + +### 已可用 IPC 替换的路径 + +- `materialize_shared_token_kv_buffer` / `materialize_shared_paged_buffer`:full materialize fallback all_reduce。 +- `materialize_prefix_and_reuse_current_*` 的 prefix spans:已有 `_try_tai_ipc_materialize_*_page_slot_spans_into`。 +- `CpSharedKV*Mla/IndexPrefetcher.consume()` 的 suffix spans:已有 IPC span gather。 +- `start_next_layer_prefix()` 的 prefix prefetch:已有 IPC span gather。 + +这些路径的源数据是长期存在的 L1 `kv_cache` / `page_buffer`,IPC handle 可以按 storage 缓存,只需在分配/扩容后重新 open。 + +### 仍依赖 all_reduce 的路径 + +- `materialize_prefix_and_reuse_current_kv_page_slots()`:`fill_current_kv_page_slots_and_remap_locs()` 后,对 current slot spans 做 `_all_reduce_materialized_buffer_range()`。 +- `materialize_prefix_and_reuse_current_index_page_slots()`:`fill_current_index_page_slots()` 后,对 current page spans 做 `_all_reduce_materialized_buffer_range()`。 +- `CpSharedKVMlaPrefetcher.consume_prefix_with_current()` / `CpSharedKVIndexPrefetcher.consume_prefix_with_current()`:prefetched prefix + current fill 后仍 reduce current slot spans。 + +这些路径的源是每层 forward 产生的临时 `current_kv_cache` / `current_index_k` / `current_index_scale`。不能直接对临时 tensor 做 per-layer IPC handle all_gather,否则只是把 all_reduce 换成另一个高频 collective。 + +## 设计选择 + +### 方案 A:临时 tensor IPC handle all_gather(拒绝) + +每层对 `current_*` tensor open IPC handle,然后 peer-read current rows。 + +拒绝原因: +- data_ptr/shape 每层/每 batch 变化,handle cache 命中率低。 +- 仍需要高频 `all_gather` 交换 handle/offset。 +- CPU submit 和同步开销不可控,违背“彻底干掉 collective”的目标。 + +### 方案 B:persistent current staging + ready flag(采用) + +每个 CP rank 维护长期 CUDA staging buffer 和 ready counter buffer: + +1. current fill kernel 同时把本 rank owner-lane current pages 写入本 rank staging buffer,布局与 dense slot page 对齐。 +2. publish 完成后在同 stream 写 ready seq:`__threadfence_system()` 后 store seq。 +3. 所有 rank 用 IPC peer ptrs 读取 owner rank 的 staging pages,gather 到本地 dense buffer;gather kernel 在读取每个 owner 前等待 `peer_ready[owner] >= seq`。 +4. descriptor 以 dense slot page 为单位,跨 batch request 合并为一个 owner/page/slot list,一次 kernel launch 完成多个 request。 + +采用原因: +- IPC handle 只在 staging buffer 分配/扩容时交换,热路径无 NCCL/Gloo collective。 +- current 数据仍按 page 最小单位发布,符合当前 page-aligned cache 合同。 +- bs>1 可以复用 slot span descriptor,一次 launch 覆盖多个 request。 +- 可与现有 TAI current fill kernel 融合,避免重复 remap/row-mask 逻辑。 + +## staging buffer 合同 + +### token KV staging + +- 形状语义:flat bytes,容量至少覆盖 `dense_num_pages * page_size * kv_row_bytes`。 +- 写入地址:`dense_slot_page * page_nbytes + row_offset * row_nbytes`。 +- 每次 publish 只保证 valid current rows 正确;为了避免 stale tail,publish kernel 需要对 touched current slot pages 的 tail slack 清零,或 SGLang 必须保证返回 locs 不引用 tail slack。第一版建议在 publish kernel 内按 touched page 清零,优先正确性。 + +### index staging + +- 形状语义:flat bytes,容量至少覆盖 `dense_num_pages * index_page_bytes`。 +- 写入地址:`dense_slot_page * page_bytes`,内部包含 K rows + scale rows。 +- current K/scale 的 valid rows 写入对应 row offset;tail slack 清零。 + +### ready flag + +- 每 rank 一个 `uint64/int64` counter buffer,通过 IPC peer ptrs 打开。 +- 每次 current publish 使用递增 seq。 +- gather kernel 对需要读取的 owner 执行 device-side wait,直到 `peer_ready[owner] >= seq`。 +- wait kernel 需要 watchdog/iteration bound,debug build 可 fail-fast;生产第一版可以保留有限 spin + error flag,避免死锁静默挂住。 + +## descriptor 合同 + +输入:`current_slot_spans` / `slot_logical_pages` / layout / physical capacity。 + +输出: +- `owner_ranks[num_pages]` +- `src_page_indices[num_pages]`:对 staging 来说等于 dense slot page id;对 persistent L1 prefix/suffix 来说是 physical page index。 +- `dst_slot_indices[num_pages]`:dense buffer 1-based slot index,保持现有 kernel 合同。 + +bs>1 要求: +- span list 可以覆盖多个 request。 +- descriptor 构造只按 merged slot spans 生成一次,不允许 request loop + 多次 kernel。 +- 如果某些 request 没 current page,descriptor 为空时直接成功。 + +## TAI kernel 阶段 + +### P1:RED tests / benchmark skeleton + +- SGLang unit:partial-current compose 在 cp_size>1 时应调用 IPC current gather helper,不应调用 `_all_reduce_materialized_buffer_range`。 +- tai-kernel CUDA test:声明期望 API `publish_current_*_to_staging` 和 `materialize_cuda_ipc_peer_pages_slot_indices_wait_ready`,先验证缺失失败。 +- benchmark skeleton:同一输入比较 `local fill + all_reduce` 与 `publish + IPC gather`。 + +### P2:current publish kernel + +- token:扩展/新增 TAI op,复用 `cp_fill_current_kv_page_slots` remap 逻辑,同时写 staging 和 ready seq。 +- index:扩展/新增 TAI op,复用 `cp_fill_current_index_page_slots`,同时写 staging 和 ready seq。 + +### P3:ready-wait IPC gather kernel + +- 基于现有 `materialize_cuda_ipc_peer_pages_slot_indices` 增加 ready peer ptrs + seq 参数。 +- 支持 page_nbytes 变长,owner/dst descriptor 长度变长。 +- 支持 self-rank peer ptr 快路径,便于单机 CUDA unit test。 + +### P4:SGLang runtime 接入 + +- 新增 runtime helper:`_try_tai_ipc_materialize_current_token_kv_page_slot_spans_into`。 +- 新增 runtime helper:`_try_tai_ipc_materialize_current_paged_buffer_page_slot_spans_into`。 +- current compose:先 fill+publish,再 IPC gather current spans;失败时 warning/fail-fast,不再静默 all_reduce。 +- prefetch consume_prefix_with_current:同样走 current staging IPC。 + +### P5:严格化 fallback + +- prefix/suffix/full materialize:IPC 失败在生产 CP shared-KV fast path 下 fail-fast;仅在显式 debug env 下允许 fallback,用 warning 标记。 +- all_reduce helper 保留给非 CP shared-KV 或测试 reference,不在 fast path 默认触发。 + +### P6:验证 + +- CPU unit:descriptor / fallback contract / no-allreduce path。 +- CUDA unit:self-rank staging roundtrip、multi-process cp=2/8 roundtrip、fp8/bf16、bs=1/5/10。 +- benchmark:4k/16k/40k/80k/160k prefix+current,bs=1/5/10,报告 CPU/GPU 时间和有效带宽。 +- ETE:GSM8K 两轮 cache-hit 精度不掉点;mixed replay 不出现 detokenizer hang / all_reduce collective mismatch。 + +## 风险与约束 + +- device-side ready wait 如果某 rank 没有 publish 会死等;必须确保所有 rank 都按同一 seq 进入 publish/gather,即使本 rank current rows 为空也要 publish ready。 +- staging buffer 扩容会触发一次 IPC handle exchange;必须高水位缓存,不能每层分配。 +- current staging tail slack 不能污染 attention/index;第一版应清零 touched current pages,后续再优化成 valid-locs 完全约束。 +- 一次 kernel 同时 publish+peer-gather 在跨进程场景没有全局同步,容易死锁;第一版采用 publish kernel + gather kernel 两阶段。 + +## 当前结论 + +先实现 staging+ready 的 SM IPC collective,彻底移除 current compose all_reduce。0SM/copy-engine 版本后续单独做;当前优先解决 correctness、CPU collective overhead、bs>1 一次 launch。 + +## 2026-06-11 实现与 benchmark 更新 + +### 已完成 + +- TAI 已新增 `publish_cuda_ipc_slot_pages_and_mark_ready`:把本 rank 已经填好的 dense slot pages 复制到 persistent staging,并写 ready seq。 +- TAI 已新增 `materialize_cuda_ipc_peer_pages_slot_indices_wait_ready`:按 owner/src/dst descriptor 等待 peer ready 后从 IPC peer staging 复制到本地 dense buffer。 +- SGLang runtime 的以下 current compose 已接入 current-staging IPC helper: + - token KV:`materialize_prefix_and_reuse_current_kv_page_slots()`。 + - index page buffer:`materialize_prefix_and_reuse_current_index_page_slots()`。 + - MLA/index prefetch `consume_prefix_with_current()`。 +- bs>1 descriptor 合同已覆盖:一次 launch 可以用 flattened slot-page descriptor 覆盖多个 request spans;benchmark 用 `--current-batch-requests 5` 验证该形态。 + +### 与原计划的差异 + +P2 暂时没有把 publish 融入 `cp_fill_current_*` row-fill kernel,而是采用: + +1. current fill 先写 dense buffer; +2. `publish_cuda_ipc_slot_pages_and_mark_ready` 再按 page 复制本 rank owner pages 到 staging; +3. wait-ready IPC gather 从 peer staging 拉取所有 current pages。 + +这样 correctness 风险低,接入面小,但 current path 多一次 page copy + 一个额外 kernel。benchmark 也证明小 current span 下该版本不一定优于 NCCL/Gloo all_reduce;后续要进一步追性能,需要做“current fill 同时写 staging + mark ready”的融合版本。 + +### 远端 benchmark 证据 + +环境:`g0034` / `cjy-glm5-new` / `torchrun --nproc_per_node=8` / `benchmark/nsa_prefill/benchmark_cp_shared_kv_ipc_gather.py`。 + +命令示例: + +```bash +PYTHONPATH=python torchrun --standalone --nproc_per_node=8 \ + benchmark/nsa_prefill/benchmark_cp_shared_kv_ipc_gather.py \ + --tokens 16384 32768 65536 98304 122880 \ + --dtype uint8 --kv-dim 656 \ + --warmup 3 --repeat 8 \ + --include-current-staging --current-batch-requests 5 --no-check +``` + +FP8/uint8 MLA page(`kv_dim=656`)current-staging IPC vs current all_reduce: + +| tokens | pages | all_reduce p50 | current IPC p50 | 结论 | +| --- | ---: | ---: | ---: | --- | +| 16k | 256 | 0.184 ms | 0.199 ms | 当前两阶段 IPC 小幅变慢 | +| 32k | 512 | 0.227 ms | 0.217 ms | 基本持平/略快 | +| 65k | 1024 | 0.337 ms | 0.253 ms | IPC 快约 25% | +| 98k | 1536 | 0.450 ms | 0.307 ms | IPC 快约 32% | +| 122k | 1920 | 0.527 ms | 0.362 ms | IPC 快约 31% | + +BF16 MLA page(`kv_dim=576`)current-staging IPC vs current all_reduce: + +| tokens | pages | all_reduce p50 | current IPC p50 | 结论 | +| --- | ---: | ---: | ---: | --- | +| 16k | 256 | 0.175 ms | 0.175 ms | 持平 | +| 65k | 1024 | 0.397 ms | 0.414 ms | 两阶段 IPC 小幅变慢 | +| 122k | 1920 | 0.671 ms | 0.594 ms | IPC 快约 11% | + +同一 benchmark 中 prefix/L1 persistent IPC 的 `cuda_ipc_peer_pages_materialize_slot_dense` 仍明显优于 dense all_reduce,例如 FP8:65k tokens 从 0.334 ms 降到 0.164 ms,122k tokens 从 0.517 ms 降到 0.274 ms。这说明 persistent L1 prefix/suffix 场景适合直接 IPC;current 临时数据场景的瓶颈在 publish 阶段,下一步应优先融合 current fill + staging publish。 + +### 下一步必须处理的点 + +1. fast path fallback 已收窄:当 tensor 在 CUDA 上且 `SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE=1` 时,prefix/current IPC 失败会 fail-fast,不再静默走 all_reduce;CPU reference/TAI 显式关闭路径仍保留测试 fallback。 +2. wait-ready kernel 现在有 `max_spins`,但没有 device error flag;超过 spin 后可能复制 stale data。生产化前应增加 error flag 或明确 fail-fast 检测,避免 silent corruption。 +3. current publish 未融合,短 current span 可能回退性能。应新增 TAI fused current fill+publish op:fill dense 的同时写 staging,随后只用一个 tiny mark-ready kernel,再 wait-gather。 + +## 2026-06-11 bs>1 cache-hit compose benchmark 补充 + +### 为什么补这个 benchmark + +之前的 IPC benchmark 只按“总 pages”测 current staging 或 prefix gather,不能覆盖真实线上 cache-hit 形态: + +- 每个 request 有很长 cached prefix(100k-300k tokens)。 +- 每个 request 的 extend/current 较短(10k-65k tokens)。 +- prefill bs>1 时一次 batch 内有 2-10 个 request。 +- 部分短 current tail 只落在少数 CP owner lanes,上游 zero-lane rank 仍必须 publish ready,不能退出 collective 合同。 + +因此新增 `--include-cache-hit-compose` / `--cache-hit-only`,显式构造 request-major dense slot layout: + +```text +page0(dummy) | +req0 cached pages | req0 current pages | +req1 cached pages | req1 current pages | ... +``` + +对比两条路径: + +1. `cache_hit_dense_all_reduce_full`:每 rank 构造本 owner pages 的 dense buffer,然后对整个 dense slot buffer 做 all_reduce。 +2. `cache_hit_ipc_prefix_current_compose`:cached prefix 从 persistent compact owner staging 走 IPC slot-index materialize;current/extend 从 dense current staging publish + ready-wait IPC gather。 + +注意:current owner 分布按完整 request page positions 计算,所以短 extend 可能只触达少数 owner lanes。这是预期现象,不再要求每个 rank 都拥有 current page;zero-lane rank 仍会 publish ready,避免 peer wait 死锁。 + +### 远端命令 + +```bash +PYTHONPATH=python torchrun --standalone --nproc_per_node=8 \ + benchmark/nsa_prefill/benchmark_cp_shared_kv_ipc_gather.py \ + --cache-hit-only \ + --cache-hit-cached-tokens 102400 204800 307200 \ + --cache-hit-extend-tokens 10240 32768 65536 \ + --cache-hit-batch-requests 2 5 10 \ + --dtype uint8 --kv-dim 656 \ + --warmup 2 --repeat 5 --no-check +``` + +### FP8/uint8 结果摘要 + +环境:`g0034` / `cjy-glm5-new` / 8 ranks / `kv_dim=656` / page size 64。 + +| bs | cached/req | extend/req | dense all_reduce p50 | IPC compose p50 | 收益 | +| ---: | ---: | ---: | ---: | ---: | ---: | +| 2 | 100k | 10k | 0.858 ms | 0.686 ms | 1.25x | +| 2 | 200k | 65k | 1.960 ms | 1.503 ms | 1.30x | +| 2 | 300k | 65k | 2.669 ms | 1.999 ms | 1.34x | +| 5 | 100k | 10k | 2.043 ms | 1.525 ms | 1.34x | +| 5 | 200k | 65k | 4.792 ms | 3.517 ms | 1.36x | +| 5 | 300k | 65k | 6.605 ms | 4.798 ms | 1.38x | +| 10 | 100k | 10k | 4.036 ms | 2.951 ms | 1.37x | +| 10 | 200k | 65k | 9.441 ms | 6.999 ms | 1.35x | +| 10 | 300k | 65k | 13.023 ms | 9.552 ms | 1.36x | + +完整矩阵结论:在 fp8 cache-hit bs>1 形态下,IPC compose 对 dense all_reduce 稳定约 1.25x-1.6x;cached 越长、bs 越大收益越稳定。这个 benchmark 比单独 current staging 更贴近实际,因为真实 cache-hit 主要成本来自长 cached prefix 的 materialize,而这部分 persistent IPC 收益明显。 + +### BF16 spot-check + +命令只测代表性 case:`bs=5 cached=200k extend=10k/65k dtype=bf16 kv_dim=656`。 + +| bs | cached/req | extend/req | dense all_reduce p50 | IPC compose p50 | 结论 | +| ---: | ---: | ---: | ---: | ---: | --- | +| 5 | 200k | 10k | 5.929 ms | 5.595 ms | 略快 | +| 5 | 200k | 65k | 7.387 ms | 7.366 ms | 基本持平 | + +BF16 下收益不明显,原因是两阶段 current publish 的额外 copy 被放大;当前线上 GLM5 使用 fp8 KV cache,因此优先级仍然是把 fp8 bs>1 fast path 接稳。后续若要兼顾 BF16,需要 fused current fill+publish,减少 current staging 的额外 HBM copy。 diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py index a3f5a3587..668b50ce9 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py @@ -7,21 +7,11 @@ from typing import Any, Optional import torch -from sglang.srt.layers.attention.nsa.cp_shared_kv_compose import ( - cp_shared_kv_compose_symm_enabled, - get_compose_staging, - get_or_build_compose_plan, -) from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( - _agreed_tai_ipc_peer_ptrs, _all_reduce_materialized_buffer_async, _all_reduce_materialized_buffer_range, - _page_nbytes_from_page_tensor, - _symm_begin_current_staging, - _symm_staging_ready_or_register, - _token_kv_page_nbytes, - build_current_loc_remap, - build_current_page_mask, + build_batch_current_slot_spans, + build_batch_prefix_slot_spans, cp_shared_kv_debug_enabled, cp_shared_kv_mla_prefetch_enabled, cp_shared_kv_mla_prefetch_log, @@ -38,11 +28,16 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( get_or_build_shared_token_kv_slot_remap, materialize_local_paged_buffer_page_slots_into, materialize_local_token_kv_page_slots_into, - maybe_build_current_page_writer_ranks, remap_logical_pages_to_slot_dense_pages, remap_logical_locs_to_slot_dense_locs_optimized, slot_range_to_page_slice, slot_range_to_token_slice, + _raise_tai_ipc_materialize_required, + _should_fail_fast_tai_ipc_materialize, + _try_tai_ipc_materialize_current_paged_buffer_page_slot_spans_into, + _try_tai_ipc_materialize_current_token_kv_page_slot_spans_into, + _try_tai_ipc_materialize_paged_buffer_page_slot_spans_into, + _try_tai_ipc_materialize_token_kv_page_slot_spans_into, ) from sglang.srt.layers.attention.nsa.utils import is_nsa_prefill_cp_in_seq_split from sglang.srt.layers.dp_attention import get_attention_cp_group @@ -51,72 +46,6 @@ from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout logger = logging.getLogger(__name__) -class _PrefetchSymmState: - """Layer-invariant symm-exchange descriptors for one prefetch batch. - - ``num_current_pages`` / ``staging_page_inverse`` make this object a valid - ``plan`` argument for ``_symm_begin_current_staging``. The gather covers - ALL current pages (this rank's own come from its own staging — the fill - wrote them there, not into the dense buffer), so unlike the sync mega - gather no prefix descriptors are involved: the prefix was already - materialized by the prefetch stream. - """ - - __slots__ = ( - "num_current_pages", - "staging_page_inverse", - "writer_ranks", - "staging_slots", - "current_dense_pages", - "page_nbytes", - "staging_current_rows", - "mixed_locs", - "dense_pages", - ) - - def __init__(self, **kwargs) -> None: - for name in self.__slots__: - setattr(self, name, kwargs.get(name)) - - -def _seed_symm_ipc_agreement(prefetcher: Any, pool_buffer: torch.Tensor) -> None: - """Seed the group-agreed IPC capability for this layer's pool tensor - BEFORE any per-rank hit/miss divergence. - - The sync-compose fallback's first use of a pool tensor issues collectives - (capability MIN all-reduce + IPC handle all-gather); if only a lone miss - rank ran them the CP group would deadlock. Seeding here is uniform — the - consume call itself is batch-logical — and idempotent (cached per pool - tensor).""" - - if ( - prefetcher.symm_writers is not None - and prefetcher.layout.cp_size > 1 - and cp_shared_kv_compose_symm_enabled() - ): - _agreed_tai_ipc_peer_ptrs(pool_buffer, prefetcher.layout) - - -def _prefetch_symm_active(prefetcher: Any, device: torch.device) -> bool: - """Rank-uniform gate for the prefetch-path symm current exchange. - - Every condition is batch-logical or group-agreed: ``symm_writers`` comes - from ``maybe_build_current_page_writer_ranks`` (env + batch metadata), - and ``registered`` only flips inside a collective registration. A - prefetch hit/miss divergence across ranks stays barrier-safe because the - sync-compose fallback also runs exactly one begin_round + barrier per - (layer, kind). - """ - - return ( - prefetcher.symm_writers is not None - and prefetcher.layout.cp_size > 1 - and prefetcher.prefix_pages < prefetcher.total_slots - and cp_shared_kv_compose_symm_enabled() - and get_compose_staging(device).registered - ) - - def _prefetch_log(message: str, *args) -> None: cp_shared_kv_mla_prefetch_log(message, *args) @@ -209,6 +138,115 @@ def _debug_owned_pages_count( return -1 +def _slot_spans_page_count(spans: list[tuple[int, int]]) -> int: + return sum(max(0, int(end) - int(start)) for start, end in spans) + + +def _slot_spans_to_token_slices( + page_size: int, + spans: list[tuple[int, int]], +) -> list[slice]: + return [ + slot_range_to_token_slice(page_size, int(start), int(end)) + for start, end in spans + if int(end) > int(start) + ] + + +def _slot_spans_to_page_slices(spans: list[tuple[int, int]]) -> list[slice]: + return [ + slot_range_to_page_slice(int(start), int(end)) + for start, end in spans + if int(end) > int(start) + ] + + +def _slot_span_logical_pages( + slot_logical_pages: torch.Tensor, + spans: list[tuple[int, int]], +) -> torch.Tensor: + pieces = [ + slot_logical_pages[int(start) : int(end)] + for start, end in spans + if int(end) > int(start) + ] + if not pieces: + return slot_logical_pages.new_empty((0,)) + return torch.cat(pieces, dim=0) + + +def _materialize_local_token_kv_page_slot_spans_into( + *, + kv_cache: torch.Tensor, + dense_kv_cache: torch.Tensor, + slot_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + page_size: int, + spans: list[tuple[int, int]], +) -> None: + for start_slot, end_slot in spans: + materialize_local_token_kv_page_slots_into( + kv_cache=kv_cache, + dense_kv_cache=dense_kv_cache, + slot_logical_pages=slot_logical_pages, + layout=layout, + page_size=page_size, + start_slot=int(start_slot), + end_slot=int(end_slot), + ) + + +def _materialize_local_paged_buffer_page_slot_spans_into( + *, + page_buffer: torch.Tensor, + dense_page_buffer: torch.Tensor, + slot_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + spans: list[tuple[int, int]], +) -> None: + for start_slot, end_slot in spans: + materialize_local_paged_buffer_page_slots_into( + page_buffer=page_buffer, + dense_page_buffer=dense_page_buffer, + slot_logical_pages=slot_logical_pages, + layout=layout, + start_slot=int(start_slot), + end_slot=int(end_slot), + ) + + +def _all_reduce_materialized_buffer_ranges_async( + *, + dense_kv_cache: torch.Tensor, + row_slices: list[slice], + cp_size: int, + stream: torch.cuda.Stream, + nvtx_source: str, + nvtx_layer_id: int, + nvtx_cp_rank: int, +) -> torch.cuda.Event | None: + event: torch.cuda.Event | None = None + for rows in row_slices: + event = _all_reduce_materialized_buffer_async( + dense_kv_cache[rows], + cp_size=cp_size, + stream=stream, + nvtx_source=nvtx_source, + nvtx_layer_id=nvtx_layer_id, + nvtx_cp_rank=nvtx_cp_rank, + nvtx_rows=(rows.start, rows.stop), + ) + if event is None: + return None + return event + + +def _record_event_on_stream(stream: torch.cuda.Stream) -> torch.cuda.Event: + event = torch.cuda.Event() + event.record(stream) + return event + + def _debug_handle_keys( layer_id: int, handles: dict[int, Any], @@ -365,6 +403,7 @@ class CpSharedKVMlaPrefetchHandle: dense_kv_cache: torch.Tensor prefix_rows: slice event: Optional[torch.cuda.Event] = None + prefix_row_spans: tuple[slice, ...] | None = None @dataclass @@ -373,6 +412,7 @@ class CpSharedKVIndexPrefetchHandle: dense_page_buffer: torch.Tensor prefix_rows: slice event: Optional[torch.cuda.Event] = None + prefix_row_spans: tuple[slice, ...] | None = None @dataclass(frozen=True) @@ -406,6 +446,8 @@ class CpSharedKVMlaPrefetcher: layout: CpSharedKVLayout, page_size: int, prefix_pages: int, + prefix_slot_spans: list[tuple[int, int]] | None = None, + current_slot_spans: list[tuple[int, int]] | None = None, slot_logical_pages: torch.Tensor, page_inverse: torch.Tensor, slot_sorted_logical_pages_by_row: torch.Tensor | None = None, @@ -414,8 +456,6 @@ class CpSharedKVMlaPrefetcher: owned_prefix_pages: int = -1, owned_total_pages: int = -1, stream: Optional[torch.cuda.Stream] = None, - slot_remap: Any = None, - symm_writers: Optional[list] = None, ) -> None: self.layout = layout self.page_size = page_size @@ -428,14 +468,22 @@ class CpSharedKVMlaPrefetcher: self.owned_prefix_pages = owned_prefix_pages self.owned_total_pages = owned_total_pages self.total_slots = int(slot_logical_pages.numel()) + self.prefix_slot_spans = ( + [(0, int(prefix_pages))] if prefix_slot_spans is None and prefix_pages > 0 + else list(prefix_slot_spans or []) + ) + self.current_slot_spans = ( + [(int(prefix_pages), self.total_slots)] + if current_slot_spans is None and prefix_pages < self.total_slots + else list(current_slot_spans or []) + ) + self.prefix_page_count = _slot_spans_page_count(self.prefix_slot_spans) + self.current_page_count = _slot_spans_page_count(self.current_slot_spans) self.stream = stream if stream is not None else torch.cuda.Stream() self.handles: dict[int, CpSharedKVMlaPrefetchHandle] = {} self.pending_attention_handle: Optional[CpSharedKVMlaPrefetchHandle] = None self.disabled = False self._cpu_timing = _PrefetchCpuTiming() - self.slot_remap = slot_remap - self.symm_writers = symm_writers - self._symm_state: Optional[_PrefetchSymmState] = None @classmethod def maybe_create( @@ -473,7 +521,8 @@ class CpSharedKVMlaPrefetcher: if not topk_transform_is_paged: _prefetch_log("create_skip reason=not_paged_topk") return None - if int(getattr(forward_batch, "batch_size", 0)) != 1: + batch_size = int(getattr(forward_batch, "batch_size", 0)) + if batch_size <= 0: _prefetch_log( "create_skip reason=batch_size batch_size=%s", getattr(forward_batch, "batch_size", None), @@ -491,70 +540,116 @@ class CpSharedKVMlaPrefetcher: return None extend_prefix_lens_cpu = getattr(forward_batch, "extend_prefix_lens_cpu", None) - if extend_prefix_lens_cpu is None or len(extend_prefix_lens_cpu) != 1: + if extend_prefix_lens_cpu is None or len(extend_prefix_lens_cpu) != batch_size: _prefetch_log("create_skip reason=bad_prefix_lens_metadata") return None page_size = int(getattr(token_to_kv_pool, "page_size", 1)) if page_size <= 1: _prefetch_log("create_skip reason=bad_page_size page_size=%s", page_size) return None - extend_prefix_len = int(extend_prefix_lens_cpu[0]) - if extend_prefix_len <= 0 or extend_prefix_len % page_size != 0: + prefix_lens = [int(prefix_len) for prefix_len in extend_prefix_lens_cpu] + bad_prefix_lens = [ + prefix_len + for prefix_len in prefix_lens + if prefix_len < 0 or prefix_len % page_size != 0 + ] + if bad_prefix_lens or not any(prefix_len > 0 for prefix_len in prefix_lens): _mla_prefetch_fallback_log( "prefix_not_page_aligned", "prefix length is zero or not page-aligned. " - "prefix_len=%s page_size=%s", - extend_prefix_len, + "prefix_lens=%s page_size=%s", + prefix_lens, page_size, ) return None - prefix_pages = extend_prefix_len // page_size real_page_table = getattr(metadata, "real_page_table", None) page_table_1 = getattr(metadata, "page_table_1", None) if real_page_table is None or page_table_1 is None: _prefetch_log("create_skip reason=missing_page_tables") return None - if prefix_pages <= 0 or prefix_pages > int(real_page_table.numel()): + try: + prefix_slot_spans = build_batch_prefix_slot_spans( + logical_pages=real_page_table, + prefix_lens_cpu=prefix_lens, + page_size=page_size, + ) + except ValueError as exc: + _mla_prefetch_fallback_log( + "bad_prefix_slot_spans", + "failed to build prefix slot spans. error=%s prefix_lens=%s " + "page_table_shape=%s page_size=%s", + exc, + prefix_lens, + tuple(real_page_table.shape), + page_size, + ) + return None + prefix_page_count = _slot_spans_page_count(prefix_slot_spans) + if prefix_page_count <= 0 or prefix_page_count > int(real_page_table.numel()): _prefetch_log( "create_skip reason=prefix_pages_out_of_range prefix_pages=%s real_pages=%s", - prefix_pages, + prefix_page_count, int(real_page_table.numel()), ) return None min_prefix_pages = cp_shared_kv_mla_prefetch_min_prefix_pages( layout.cp_size, page_size=page_size ) - if prefix_pages < min_prefix_pages: + if prefix_page_count < min_prefix_pages: _prefetch_log( "create_skip reason=prefix_below_min cp_rank=%s cp_size=%s " "prefix_pages=%s min_prefix_pages=%s prefix_len=%s page_size=%s", layout.cp_rank, layout.cp_size, - prefix_pages, + prefix_page_count, min_prefix_pages, - extend_prefix_len, + sum(prefix_lens), page_size, ) return None extend_seq_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None) - if extend_seq_lens_cpu is None or len(extend_seq_lens_cpu) != 1: + if extend_seq_lens_cpu is None or len(extend_seq_lens_cpu) != batch_size: _prefetch_log("create_skip reason=bad_extend_lens_metadata") return None - extend_len = int(extend_seq_lens_cpu[0]) + extend_lens = [int(extend_len) for extend_len in extend_seq_lens_cpu] + if any(extend_len < 0 for extend_len in extend_lens): + _prefetch_log("create_skip reason=bad_extend_lens_metadata") + return None + try: + current_slot_spans = build_batch_current_slot_spans( + logical_pages=real_page_table, + prefix_lens_cpu=prefix_lens, + extend_lens_cpu=extend_lens, + page_size=page_size, + ) + except ValueError as exc: + _mla_prefetch_fallback_log( + "bad_current_slot_spans", + "failed to build current slot spans. error=%s prefix_lens=%s " + "extend_lens=%s page_table_shape=%s page_size=%s", + exc, + prefix_lens, + extend_lens, + tuple(real_page_table.shape), + page_size, + ) + return None + current_page_count = _slot_spans_page_count(current_slot_spans) + total_extend_len = sum(extend_lens) min_async_extend_tokens = cp_shared_kv_mla_prefetch_min_async_extend_tokens( cp_size=layout.cp_size, page_size=page_size ) - if extend_len < min_async_extend_tokens: + if total_extend_len < min_async_extend_tokens: _prefetch_log( "create_skip reason=extend_below_min cp_rank=%s cp_size=%s " "extend_len=%s min_async_extend_tokens=%s prefix_pages=%s page_size=%s", layout.cp_rank, layout.cp_size, - extend_len, + total_extend_len, min_async_extend_tokens, - prefix_pages, + prefix_page_count, page_size, ) return None @@ -593,44 +688,29 @@ class CpSharedKVMlaPrefetcher: logger.exception("Failed to initialize CP shared KV MLA prefetcher.") return None - symm_writers = None - if cp_shared_kv_compose_symm_enabled() and layout.cp_size > 1: - symm_writers = maybe_build_current_page_writer_ranks( - forward_batch=forward_batch, - prefix_lens_cpu=extend_prefix_lens_cpu, - extend_lens_cpu=extend_seq_lens_cpu, - page_size=page_size, - layout=layout, - ) - if symm_writers is not None: - # Collective registration at a batch-uniform point: with a - # prefetcher active the sync compose only runs on per-rank - # misses, so its lazy first-compose registration would - # diverge. Must NOT be swallowed — a half-registered group - # is a hang, not a fallback. - _symm_staging_ready_or_register( - layout=layout, kv_cache=kv_cache, page_size=page_size - ) - + prefix_pages = prefix_lens[0] // page_size if batch_size == 1 else 0 owned_prefix_pages = _debug_owned_pages_count( - layout, remap.slot_logical_pages[:prefix_pages] + layout, _slot_span_logical_pages(remap.slot_logical_pages, prefix_slot_spans) ) owned_total_pages = _debug_owned_pages_count(layout, remap.slot_logical_pages) _prefetch_log( "create cp_rank=%s cp_size=%s prefix_pages=%s total_slots=%s " - "owned_prefix_pages=%s owned_total_pages=%s dense_pages=%s page_size=%s " + "prefix_slot_spans=%s current_slot_spans=%s owned_prefix_pages=%s " + "owned_total_pages=%s dense_pages=%s page_size=%s " "min_async_extend_tokens=%s extend_len=%s", layout.cp_rank, layout.cp_size, - prefix_pages, + prefix_page_count, int(remap.slot_logical_pages.numel()), + prefix_slot_spans, + current_slot_spans, owned_prefix_pages, owned_total_pages, remap.dense_num_pages, page_size, min_async_extend_tokens, - extend_len, + total_extend_len, ) create_total_ms = _cpu_timing_ms(create_cpu) _prefetch_log( @@ -642,7 +722,7 @@ class CpSharedKVMlaPrefetcher: create_total_ms, get_ms, remap_ms, - prefix_pages, + prefix_page_count, int(remap.slot_logical_pages.numel()), remap.dense_num_pages, ) @@ -651,6 +731,8 @@ class CpSharedKVMlaPrefetcher: layout=layout, page_size=page_size, prefix_pages=prefix_pages, + prefix_slot_spans=prefix_slot_spans, + current_slot_spans=current_slot_spans, slot_logical_pages=remap.slot_logical_pages, page_inverse=remap.page_inverse, slot_sorted_logical_pages_by_row=remap.slot_sorted_logical_pages_by_row, @@ -659,90 +741,8 @@ class CpSharedKVMlaPrefetcher: owned_prefix_pages=owned_prefix_pages, owned_total_pages=owned_total_pages, stream=prefetch_stream, - slot_remap=remap, - symm_writers=symm_writers, ) - def _get_or_build_symm_state( - self, - *, - kv_cache: torch.Tensor, - logical_locs: torch.Tensor, - current_locs: torch.Tensor, - loc_req_id: torch.Tensor, - current_req_id: torch.Tensor, - ) -> _PrefetchSymmState: - state = self._symm_state - if state is not None: - return state - plan = get_or_build_compose_plan( - slot_remap=self.slot_remap, - layout=self.layout, - physical_page_capacity=kv_cache.shape[0] // self.page_size, - prefix_spans=[(0, self.prefix_pages)], - current_spans=[(self.prefix_pages, self.total_slots)], - kind="token_kv", - current_page_writer_ranks=self.symm_writers, - ) - num_current = int(plan.num_current_pages) - writer_ranks = ( - plan.symm_all_owner_ranks.index_select(0, plan.current_dense_pages - 1) - - int(self.layout.cp_size) - ).contiguous() - # Staging slot of current page i is i + 1 (row 0 = dummy page). - staging_slots = torch.arange( - 1, num_current + 1, dtype=torch.long, device=writer_ranks.device - ) - staging_current_rows = remap_logical_locs_to_slot_dense_locs_optimized( - current_locs.reshape(-1), - page_inverse=plan.staging_page_inverse, - page_size=self.page_size, - loc_req_id=current_req_id.reshape(-1), - ).to(torch.long) - if cp_shared_kv_debug_enabled() and staging_current_rows.numel() > 0: - # index_copy_ has no skip semantics for stray -1s (debug-only: - # the min() syncs). - if int(staging_current_rows.min().item()) < 0: - raise RuntimeError( - "[CP_SHARED_KV_FAIL_FAST][prefetch_symm] current locs " - "map outside the staging page inverse" - ) - # mixed_locs is layer-invariant; the fused fill that used to produce - # it builds masks sized by its target buffer, which is now the - # staging — so compute it once here in logical space instead. - logical_locs = filter_locs_mappable_to_physical_pool( - logical_locs=logical_locs, - layout=self.layout, - physical_token_capacity=kv_cache.shape[0], - ) - dense_locs = remap_logical_locs_to_slot_dense_locs_optimized( - logical_locs, - page_inverse=self.page_inverse, - page_size=self.page_size, - loc_req_id=loc_req_id, - ) - current_mask, _ = build_current_loc_remap(logical_locs, current_locs) - current_page_mask = build_current_page_mask( - logical_locs, current_locs, page_size=self.page_size - ) - mixed_locs = torch.where( - current_page_mask & (~current_mask), - torch.full_like(dense_locs, -1), - dense_locs, - ) - state = _PrefetchSymmState( - num_current_pages=num_current, - staging_page_inverse=plan.staging_page_inverse, - writer_ranks=writer_ranks, - staging_slots=staging_slots, - current_dense_pages=plan.current_dense_pages, - page_nbytes=_token_kv_page_nbytes(kv_cache, self.page_size), - staging_current_rows=staging_current_rows, - mixed_locs=mixed_locs, - ) - self._symm_state = state - return state - def _layer_in_pool(self, token_to_kv_pool: Any, layer_id: int) -> bool: start_layer = int(getattr(token_to_kv_pool, "start_layer", 0)) kv_buffer = getattr(token_to_kv_pool, "kv_buffer", None) @@ -806,49 +806,65 @@ class CpSharedKVMlaPrefetcher: torch.cuda.current_stream().wait_event(handle.event) wait_ms = _cpu_timing_ms(wait_cpu) dense_kv_cache = handle.dense_kv_cache - suffix_slots = self.total_slots - self.prefix_pages + suffix_spans = self.current_slot_spans + suffix_slots = _slot_spans_page_count(suffix_spans) suffix_ms = 0.0 - if self.prefix_pages < self.total_slots: + if suffix_spans: self._log_layer( layer_id, - "consume_suffix_begin layer=%s start_slot=%s end_slot=%s " + "consume_suffix_begin layer=%s suffix_spans=%s " "suffix_slots=%s", layer_id, - self.prefix_pages, - self.total_slots, + suffix_spans, suffix_slots, ) suffix_cpu = _cpu_timing_start() - materialize_local_token_kv_page_slots_into( - kv_cache=kv_cache, - dense_kv_cache=dense_kv_cache, - slot_logical_pages=self.slot_logical_pages, - layout=self.layout, - page_size=self.page_size, - start_slot=self.prefix_pages, - end_slot=self.total_slots, - ) - suffix_rows = slot_range_to_token_slice( - self.page_size, - self.prefix_pages, - self.total_slots, - ) - _all_reduce_materialized_buffer_range( - dense_kv_cache, - self.layout.cp_size, - suffix_rows.start, - suffix_rows.stop, - nvtx_source="mla.consume_suffix", - nvtx_layer_id=layer_id, - nvtx_cp_rank=self.layout.cp_rank, + materialized_suffix_by_ipc = ( + _try_tai_ipc_materialize_token_kv_page_slot_spans_into( + kv_cache=kv_cache, + dense_kv_cache=dense_kv_cache, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + page_size=self.page_size, + spans=suffix_spans, + ) ) + if not materialized_suffix_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_kv_cache): + _raise_tai_ipc_materialize_required( + "mla_prefetch_suffix_ipc_unavailable", + cp_rank=self.layout.cp_rank, + cp_size=self.layout.cp_size, + spans=suffix_spans, + dense_shape=tuple(dense_kv_cache.shape), + ) + if not materialized_suffix_by_ipc: + _materialize_local_token_kv_page_slot_spans_into( + kv_cache=kv_cache, + dense_kv_cache=dense_kv_cache, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + page_size=self.page_size, + spans=suffix_spans, + ) + if not materialized_suffix_by_ipc: + for suffix_rows in _slot_spans_to_token_slices( + self.page_size, suffix_spans + ): + _all_reduce_materialized_buffer_range( + dense_kv_cache, + self.layout.cp_size, + suffix_rows.start, + suffix_rows.stop, + nvtx_source="mla.consume_suffix", + nvtx_layer_id=layer_id, + nvtx_cp_rank=self.layout.cp_rank, + ) self._log_layer( layer_id, - "consume_suffix_done layer=%s rows=%s:%s", + "consume_suffix_done layer=%s suffix_spans=%s", layer_id, - suffix_rows.start, - suffix_rows.stop, + suffix_spans, ) suffix_ms = _cpu_timing_ms(suffix_cpu) @@ -918,7 +934,6 @@ class CpSharedKVMlaPrefetcher: in ``logical_locs`` are remapped to the appended current KV rows. """ - _seed_symm_ipc_agreement(self, kv_cache) if self.disabled: self._log_layer( layer_id, @@ -968,85 +983,15 @@ class CpSharedKVMlaPrefetcher: dense_kv_cache = handle.dense_kv_cache remap_cpu = _cpu_timing_start() - if loc_req_id is None: - loc_req_id = torch.zeros_like(logical_locs, dtype=torch.long) - if current_req_id is None: - current_req_id = torch.zeros_like(current_locs, dtype=torch.long) - - if _prefetch_symm_active(self, dense_kv_cache.device): - # Symm current exchange: fill current rows straight into this - # round's staging span, barrier, gather ALL current pages from - # the stagings (this rank's own included) into the prefetched - # dense buffer. Zero NCCL; descriptors and mixed_locs are - # layer-invariant and cached on the prefetcher. - from tai_kernel.nsa_prefill.ipc import ( - cp_symm_barrier, - gather_cuda_ipc_peer_pages, - ) - - state = self._get_or_build_symm_state( - kv_cache=kv_cache, - logical_locs=logical_locs, - current_locs=current_locs, - loc_req_id=loc_req_id, - current_req_id=current_req_id, - ) - staging = get_compose_staging(dense_kv_cache.device) - parity, staging_span = _symm_begin_current_staging( - staging=staging, - plan=state, - kind="token_kv", - layer_id=layer_id, - page_nbytes=state.page_nbytes, - ) - num_rows = int(state.staging_current_rows.numel()) - if num_rows > 0: - # Byte view on both sides: index_copy_ has no fp8 CUDA - # kernel, and the copy is dtype-agnostic (whole token rows). - staging_rows = staging_span.view( - (state.num_current_pages + 1) * self.page_size, - state.page_nbytes // self.page_size, - ) - staging_rows.index_copy_( - 0, - state.staging_current_rows, - current_kv_cache[:num_rows] - .reshape(num_rows, -1) - .view(torch.uint8), - ) - cp_symm_barrier( - staging.flag_ptrs, self_rank=int(self.layout.cp_rank) - ) - gather_cuda_ipc_peer_pages( - staging.peer_region_ptrs("token_kv", parity), - dense_kv_cache, - state.writer_ranks, - state.staging_slots, - state.current_dense_pages, - page_nbytes=state.page_nbytes, - ) - remap_ms = _cpu_timing_ms(remap_cpu) - total_ms = _cpu_timing_ms(consume_cpu) - self._log_layer( - layer_id, - "consume_prefix_current_hit layer=%s prefix_pages=%s " - "dense_rows=%s current_rows=%s symm=1 total_ms=%.3f " - "wait_ms=%.3f remap_ms=%.3f", - layer_id, - self.prefix_pages, - int(dense_kv_cache.shape[0]), - int(current_kv_cache.shape[0]), - total_ms, - wait_ms, - remap_ms, - ) - return dense_kv_cache, state.mixed_locs - logical_locs = filter_locs_mappable_to_physical_pool( logical_locs=logical_locs, layout=self.layout, physical_token_capacity=kv_cache.shape[0], ) + if loc_req_id is None: + loc_req_id = torch.zeros_like(logical_locs, dtype=torch.long) + if current_req_id is None: + current_req_id = torch.zeros_like(current_locs, dtype=torch.long) dense_locs = remap_logical_locs_to_slot_dense_locs_optimized( logical_locs, page_inverse=self.page_inverse, @@ -1064,21 +1009,41 @@ class CpSharedKVMlaPrefetcher: current_req_id=current_req_id, mask_non_current_in_current_pages=True, ) - if self.layout.cp_size > 1 and self.prefix_pages < self.total_slots: - current_rows = slot_range_to_token_slice( - self.page_size, - self.prefix_pages, - self.total_slots, - ) - _all_reduce_materialized_buffer_range( - mixed_kv_cache, - self.layout.cp_size, - current_rows.start, - current_rows.stop, - nvtx_source="mla.prefetch_current", - nvtx_layer_id=layer_id, - nvtx_cp_rank=self.layout.cp_rank, + if self.layout.cp_size > 1: + current_materialized_by_ipc = ( + _try_tai_ipc_materialize_current_token_kv_page_slot_spans_into( + dense_kv_cache=mixed_kv_cache, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + page_size=self.page_size, + spans=self.current_slot_spans, + ) ) + if ( + not current_materialized_by_ipc + and self.current_slot_spans + and _should_fail_fast_tai_ipc_materialize(mixed_kv_cache) + ): + _raise_tai_ipc_materialize_required( + "mla_prefetch_current_ipc_unavailable", + cp_rank=self.layout.cp_rank, + cp_size=self.layout.cp_size, + spans=self.current_slot_spans, + dense_shape=tuple(mixed_kv_cache.shape), + ) + if not current_materialized_by_ipc: + for current_rows in _slot_spans_to_token_slices( + self.page_size, self.current_slot_spans + ): + _all_reduce_materialized_buffer_range( + mixed_kv_cache, + self.layout.cp_size, + current_rows.start, + current_rows.stop, + nvtx_source="mla.prefetch_current", + nvtx_layer_id=layer_id, + nvtx_cp_rank=self.layout.cp_rank, + ) remap_ms = _cpu_timing_ms(remap_cpu) total_ms = _cpu_timing_ms(consume_cpu) self._log_layer( @@ -1139,12 +1104,13 @@ class CpSharedKVMlaPrefetcher: return start_cpu = _cpu_timing_start() + current_stream = torch.cuda.current_stream() get_cpu = _cpu_timing_start() try: kv_cache = _prefetch_pool_get_key_buffer( token_to_kv_pool=token_to_kv_pool, layer_id=next_layer_id, - stream=self.stream, + stream=current_stream, path="mla", ) get_ms = _cpu_timing_ms(get_cpu) @@ -1161,46 +1127,63 @@ class CpSharedKVMlaPrefetcher: return try: - current_stream = torch.cuda.current_stream() - prefix_rows = slot_range_to_token_slice( - self.page_size, - 0, - self.prefix_pages, + prefix_row_spans = _slot_spans_to_token_slices( + self.page_size, self.prefix_slot_spans ) + prefix_rows = prefix_row_spans[0] if prefix_row_spans else slice(0, 0) materialize_cpu = _cpu_timing_start() dense_kv_cache = kv_cache.new_zeros( (self.dense_num_pages * self.page_size, *kv_cache.shape[1:]) ) self._log_next_layer( next_layer_id, - "start_prefix_begin next_layer=%s start_slot=0 end_slot=%s " + "start_prefix_begin next_layer=%s prefix_slot_spans=%s " "dense_rows=%s", next_layer_id, - self.prefix_pages, + self.prefix_slot_spans, int(dense_kv_cache.shape[0]), ) - materialize_local_token_kv_page_slots_into( + materialized_by_ipc = _try_tai_ipc_materialize_token_kv_page_slot_spans_into( kv_cache=kv_cache, dense_kv_cache=dense_kv_cache, slot_logical_pages=self.slot_logical_pages, layout=self.layout, page_size=self.page_size, - start_slot=0, - end_slot=self.prefix_pages, + spans=self.prefix_slot_spans, ) + if not materialized_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_kv_cache): + _raise_tai_ipc_materialize_required( + "mla_prefetch_prefix_ipc_unavailable", + cp_rank=self.layout.cp_rank, + cp_size=self.layout.cp_size, + spans=self.prefix_slot_spans, + dense_shape=tuple(dense_kv_cache.shape), + ) + if not materialized_by_ipc: + _materialize_local_token_kv_page_slot_spans_into( + kv_cache=kv_cache, + dense_kv_cache=dense_kv_cache, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + page_size=self.page_size, + spans=self.prefix_slot_spans, + ) materialize_ms = _cpu_timing_ms(materialize_cpu) reduce_cpu = _cpu_timing_start() - self.stream.wait_stream(current_stream) - with torch.cuda.stream(self.stream): - event = _all_reduce_materialized_buffer_async( - dense_kv_cache[prefix_rows], - cp_size=self.layout.cp_size, - stream=self.stream, - nvtx_source="mla.prefetch_prefix", - nvtx_layer_id=next_layer_id, - nvtx_cp_rank=self.layout.cp_rank, - nvtx_rows=(prefix_rows.start, prefix_rows.stop), - ) + if materialized_by_ipc: + event = _record_event_on_stream(current_stream) + else: + self.stream.wait_stream(current_stream) + with torch.cuda.stream(self.stream): + event = _all_reduce_materialized_buffer_ranges_async( + dense_kv_cache=dense_kv_cache, + row_slices=prefix_row_spans, + cp_size=self.layout.cp_size, + stream=self.stream, + nvtx_source="mla.prefetch_prefix", + nvtx_layer_id=next_layer_id, + nvtx_cp_rank=self.layout.cp_rank, + ) reduce_enqueue_ms = _cpu_timing_ms(reduce_cpu) if event is None: self.disabled = True @@ -1256,6 +1239,7 @@ class CpSharedKVMlaPrefetcher: dense_kv_cache=dense_kv_cache, prefix_rows=prefix_rows, event=event, + prefix_row_spans=tuple(prefix_row_spans), ) self.handles[next_layer_id] = handle self.pending_attention_handle = handle @@ -1281,19 +1265,27 @@ class CpSharedKVMlaPrefetcher: ) return - prefix_rows = handle.prefix_rows + prefix_row_spans = list(handle.prefix_row_spans or (handle.prefix_rows,)) try: + if _should_fail_fast_tai_ipc_materialize(handle.dense_kv_cache): + _raise_tai_ipc_materialize_required( + "mla_prefetch_deferred_prefix_ipc_unavailable", + cp_rank=self.layout.cp_rank, + cp_size=self.layout.cp_size, + dense_shape=tuple(handle.dense_kv_cache.shape), + layer_id=handle.layer_id, + ) current_stream = torch.cuda.current_stream() self.stream.wait_stream(current_stream) with torch.cuda.stream(self.stream): - event = _all_reduce_materialized_buffer_async( - handle.dense_kv_cache[prefix_rows], + event = _all_reduce_materialized_buffer_ranges_async( + dense_kv_cache=handle.dense_kv_cache, + row_slices=prefix_row_spans, cp_size=self.layout.cp_size, stream=self.stream, nvtx_source="mla.prefetch_prefix", nvtx_layer_id=handle.layer_id, nvtx_cp_rank=self.layout.cp_rank, - nvtx_rows=(prefix_rows.start, prefix_rows.stop), ) if event is None: self.disabled = True @@ -1309,10 +1301,9 @@ class CpSharedKVMlaPrefetcher: handle.event = event self._log_next_layer( handle.layer_id, - "start_prefix_reduce_enqueued next_layer=%s rows=%s:%s", + "start_prefix_reduce_enqueued next_layer=%s row_spans=%s", handle.layer_id, - prefix_rows.start, - prefix_rows.stop, + [(rows.start, rows.stop) for rows in prefix_row_spans], ) except Exception: logger.exception("Failed to launch CP shared KV MLA prefix prefetch reduce.") @@ -1368,6 +1359,8 @@ class CpSharedKVIndexPrefetcher: *, layout: CpSharedKVLayout, prefix_pages: int, + prefix_slot_spans: list[tuple[int, int]] | None = None, + current_slot_spans: list[tuple[int, int]] | None = None, slot_logical_pages: torch.Tensor, page_inverse: torch.Tensor, slot_sorted_logical_pages_by_row: torch.Tensor | None = None, @@ -1376,8 +1369,6 @@ class CpSharedKVIndexPrefetcher: owned_prefix_pages: int = -1, owned_total_pages: int = -1, stream: Optional[torch.cuda.Stream] = None, - slot_remap: Any = None, - symm_writers: Optional[list] = None, ) -> None: self.layout = layout self.prefix_pages = prefix_pages @@ -1389,14 +1380,22 @@ class CpSharedKVIndexPrefetcher: self.owned_prefix_pages = owned_prefix_pages self.owned_total_pages = owned_total_pages self.total_slots = int(slot_logical_pages.numel()) + self.prefix_slot_spans = ( + [(0, int(prefix_pages))] if prefix_slot_spans is None and prefix_pages > 0 + else list(prefix_slot_spans or []) + ) + self.current_slot_spans = ( + [(int(prefix_pages), self.total_slots)] + if current_slot_spans is None and prefix_pages < self.total_slots + else list(current_slot_spans or []) + ) + self.prefix_page_count = _slot_spans_page_count(self.prefix_slot_spans) + self.current_page_count = _slot_spans_page_count(self.current_slot_spans) self.stream = stream if stream is not None else torch.cuda.Stream() self.handles: dict[int, CpSharedKVIndexPrefetchHandle] = {} self.pending_attention_handle: Optional[CpSharedKVIndexPrefetchHandle] = None self.disabled = False self._cpu_timing = _PrefetchCpuTiming() - self.slot_remap = slot_remap - self.symm_writers = symm_writers - self._symm_state: Optional[_PrefetchSymmState] = None @classmethod def maybe_create( @@ -1454,7 +1453,8 @@ class CpSharedKVIndexPrefetcher: "topk transform is not PAGED.", ) return None - if int(getattr(forward_batch, "batch_size", 0)) != 1: + batch_size = int(getattr(forward_batch, "batch_size", 0)) + if batch_size <= 0: _index_prefetch_fallback_log( "batch_size", "batch size is not supported. batch_size=%s", @@ -1479,11 +1479,12 @@ class CpSharedKVIndexPrefetcher: return None extend_prefix_lens_cpu = getattr(forward_batch, "extend_prefix_lens_cpu", None) - if extend_prefix_lens_cpu is None or len(extend_prefix_lens_cpu) != 1: + if extend_prefix_lens_cpu is None or len(extend_prefix_lens_cpu) != batch_size: _index_prefetch_fallback_log( "bad_prefix_lens_metadata", - "extend_prefix_lens_cpu is missing or not single-batch. value=%s", + "extend_prefix_lens_cpu is missing or does not match batch. value=%s batch_size=%s", extend_prefix_lens_cpu, + batch_size, ) return None page_size = int(getattr(token_to_kv_pool, "page_size", 1)) @@ -1494,16 +1495,20 @@ class CpSharedKVIndexPrefetcher: page_size, ) return None - extend_prefix_len = int(extend_prefix_lens_cpu[0]) - if extend_prefix_len <= 0 or extend_prefix_len % page_size != 0: + prefix_lens = [int(prefix_len) for prefix_len in extend_prefix_lens_cpu] + bad_prefix_lens = [ + prefix_len + for prefix_len in prefix_lens + if prefix_len < 0 or prefix_len % page_size != 0 + ] + if bad_prefix_lens or not any(prefix_len > 0 for prefix_len in prefix_lens): _index_prefetch_fallback_log( "prefix_not_page_aligned", - "prefix length is zero or not page-aligned. prefix_len=%s page_size=%s", - extend_prefix_len, + "prefix length is zero or not page-aligned. prefix_lens=%s page_size=%s", + prefix_lens, page_size, ) return None - prefix_pages = extend_prefix_len // page_size real_page_table = getattr(metadata, "real_page_table", None) page_table_1 = getattr(metadata, "page_table_1", None) @@ -1513,47 +1518,89 @@ class CpSharedKVIndexPrefetcher: "metadata is missing real_page_table or page_table_1.", ) return None - if prefix_pages <= 0 or prefix_pages > int(real_page_table.numel()): + try: + prefix_slot_spans = build_batch_prefix_slot_spans( + logical_pages=real_page_table, + prefix_lens_cpu=prefix_lens, + page_size=page_size, + ) + except ValueError as exc: + _index_prefetch_fallback_log( + "bad_prefix_slot_spans", + "failed to build prefix slot spans. error=%s prefix_lens=%s " + "page_table_shape=%s page_size=%s", + exc, + prefix_lens, + tuple(real_page_table.shape), + page_size, + ) + return None + prefix_page_count = _slot_spans_page_count(prefix_slot_spans) + if prefix_page_count <= 0 or prefix_page_count > int(real_page_table.numel()): _index_prefetch_fallback_log( "prefix_pages_out_of_range", "prefix pages are outside real page table. prefix_pages=%s real_pages=%s", - prefix_pages, + prefix_page_count, int(real_page_table.numel()), ) return None min_prefix_pages = cp_shared_kv_mla_prefetch_min_prefix_pages( layout.cp_size, page_size=page_size ) - if prefix_pages < min_prefix_pages: + if prefix_page_count < min_prefix_pages: _prefetch_log( "index_create_skip reason=prefix_below_min cp_rank=%s cp_size=%s " "prefix_pages=%s min_prefix_pages=%s prefix_len=%s page_size=%s", layout.cp_rank, layout.cp_size, - prefix_pages, + prefix_page_count, min_prefix_pages, - extend_prefix_len, + sum(prefix_lens), page_size, ) return None extend_seq_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None) - if extend_seq_lens_cpu is None or len(extend_seq_lens_cpu) != 1: + if extend_seq_lens_cpu is None or len(extend_seq_lens_cpu) != batch_size: _prefetch_log("index_create_skip reason=bad_extend_lens_metadata") return None - extend_len = int(extend_seq_lens_cpu[0]) + extend_lens = [int(extend_len) for extend_len in extend_seq_lens_cpu] + if any(extend_len < 0 for extend_len in extend_lens): + _prefetch_log("index_create_skip reason=bad_extend_lens_metadata") + return None + try: + current_slot_spans = build_batch_current_slot_spans( + logical_pages=real_page_table, + prefix_lens_cpu=prefix_lens, + extend_lens_cpu=extend_lens, + page_size=page_size, + ) + except ValueError as exc: + _index_prefetch_fallback_log( + "bad_current_slot_spans", + "failed to build current slot spans. error=%s prefix_lens=%s " + "extend_lens=%s page_table_shape=%s page_size=%s", + exc, + prefix_lens, + extend_lens, + tuple(real_page_table.shape), + page_size, + ) + return None + current_page_count = _slot_spans_page_count(current_slot_spans) + total_extend_len = sum(extend_lens) min_extend_tokens = cp_shared_kv_mla_prefetch_min_async_extend_tokens( cp_size=layout.cp_size, page_size=page_size ) - if extend_len < min_extend_tokens: + if total_extend_len < min_extend_tokens: _prefetch_log( "index_create_skip reason=extend_below_min cp_rank=%s cp_size=%s " "extend_len=%s min_extend_tokens=%s prefix_pages=%s page_size=%s", layout.cp_rank, layout.cp_size, - extend_len, + total_extend_len, min_extend_tokens, - prefix_pages, + prefix_page_count, page_size, ) return None @@ -1596,47 +1643,27 @@ class CpSharedKVIndexPrefetcher: logger.exception("Failed to initialize CP shared KV index prefetcher.") return None - symm_writers = None - if cp_shared_kv_compose_symm_enabled() and layout.cp_size > 1: - symm_writers = maybe_build_current_page_writer_ranks( - forward_batch=forward_batch, - prefix_lens_cpu=extend_prefix_lens_cpu, - extend_lens_cpu=extend_seq_lens_cpu, - page_size=page_size, - layout=layout, - ) - if symm_writers is not None: - # Registration sizing needs the token-KV page bytes, so fetch - # the key buffer; a no-op when the MLA prefetcher (created - # first) already registered. Collective — must not be - # swallowed (see the MLA twin). - _symm_staging_ready_or_register( - layout=layout, - kv_cache=_prefetch_pool_get_key_buffer( - token_to_kv_pool=token_to_kv_pool, - layer_id=first_layer_id, - stream=prefetch_stream, - path="index_symm_register", - ), - page_size=page_size, - ) - + prefix_pages = prefix_lens[0] // page_size if batch_size == 1 else 0 owned_prefix_pages = _debug_owned_pages_count( - layout, remap.slot_logical_pages[:prefix_pages] + layout, _slot_span_logical_pages(remap.slot_logical_pages, prefix_slot_spans) ) owned_total_pages = _debug_owned_pages_count(layout, remap.slot_logical_pages) _prefetch_log( "index_create cp_rank=%s cp_size=%s prefix_pages=%s total_slots=%s " - "owned_prefix_pages=%s owned_total_pages=%s dense_pages=%s page_size=%s", + "prefix_slot_spans=%s current_slot_spans=%s owned_prefix_pages=%s " + "owned_total_pages=%s dense_pages=%s page_size=%s current_pages=%s", layout.cp_rank, layout.cp_size, - prefix_pages, + prefix_page_count, int(remap.slot_logical_pages.numel()), + prefix_slot_spans, + current_slot_spans, owned_prefix_pages, owned_total_pages, remap.dense_num_pages, page_size, + current_page_count, ) create_total_ms = _cpu_timing_ms(create_cpu) _prefetch_log( @@ -1648,7 +1675,7 @@ class CpSharedKVIndexPrefetcher: create_total_ms, get_ms, remap_ms, - prefix_pages, + prefix_page_count, int(remap.slot_logical_pages.numel()), remap.dense_num_pages, ) @@ -1656,6 +1683,8 @@ class CpSharedKVIndexPrefetcher: return cls( layout=layout, prefix_pages=prefix_pages, + prefix_slot_spans=prefix_slot_spans, + current_slot_spans=current_slot_spans, slot_logical_pages=remap.slot_logical_pages, page_inverse=remap.page_inverse, slot_sorted_logical_pages_by_row=remap.slot_sorted_logical_pages_by_row, @@ -1664,55 +1693,8 @@ class CpSharedKVIndexPrefetcher: owned_prefix_pages=owned_prefix_pages, owned_total_pages=owned_total_pages, stream=prefetch_stream, - slot_remap=remap, - symm_writers=symm_writers, ) - def _get_or_build_symm_state( - self, - *, - dense_page_buffer: torch.Tensor, - logical_pages: torch.Tensor, - ) -> _PrefetchSymmState: - state = self._symm_state - if state is not None: - return state - plan = get_or_build_compose_plan( - slot_remap=self.slot_remap, - layout=self.layout, - physical_page_capacity=None, - prefix_spans=[(0, self.prefix_pages)], - current_spans=[(self.prefix_pages, self.total_slots)], - kind="index", - current_page_writer_ranks=self.symm_writers, - ) - num_current = int(plan.num_current_pages) - writer_ranks = ( - plan.symm_all_owner_ranks.index_select(0, plan.current_dense_pages - 1) - - int(self.layout.cp_size) - ).contiguous() - # Staging slot of current page i is i + 1 (row 0 = dummy page). - staging_slots = torch.arange( - 1, num_current + 1, dtype=torch.long, device=writer_ranks.device - ) - # The returned dense-pages remap is layer-invariant too. - dense_pages = remap_logical_pages_to_slot_dense_pages( - logical_pages, - page_inverse=self.page_inverse, - page_req_id=build_page_table_row_req_id(logical_pages), - ) - state = _PrefetchSymmState( - num_current_pages=num_current, - staging_page_inverse=plan.staging_page_inverse, - writer_ranks=writer_ranks, - staging_slots=staging_slots, - current_dense_pages=plan.current_dense_pages, - page_nbytes=_page_nbytes_from_page_tensor(dense_page_buffer), - dense_pages=dense_pages, - ) - self._symm_state = state - return state - def _layer_in_pool(self, token_to_kv_pool: Any, layer_id: int) -> bool: start_layer = int(getattr(token_to_kv_pool, "start_layer", 0)) kv_buffer = getattr(token_to_kv_pool, "kv_buffer", None) @@ -1775,47 +1757,61 @@ class CpSharedKVIndexPrefetcher: torch.cuda.current_stream().wait_event(handle.event) wait_ms = _cpu_timing_ms(wait_cpu) dense_page_buffer = handle.dense_page_buffer - suffix_slots = self.total_slots - self.prefix_pages + suffix_spans = self.current_slot_spans + suffix_slots = _slot_spans_page_count(suffix_spans) suffix_ms = 0.0 - if self.prefix_pages < self.total_slots: + if suffix_spans: self._log_layer( layer_id, - "index_consume_suffix_begin layer=%s start_slot=%s end_slot=%s " + "index_consume_suffix_begin layer=%s suffix_spans=%s " "suffix_slots=%s", layer_id, - self.prefix_pages, - self.total_slots, + suffix_spans, suffix_slots, ) suffix_cpu = _cpu_timing_start() - materialize_local_paged_buffer_page_slots_into( - page_buffer=page_buffer, - dense_page_buffer=dense_page_buffer, - slot_logical_pages=self.slot_logical_pages, - layout=self.layout, - start_slot=self.prefix_pages, - end_slot=self.total_slots, - ) - suffix_rows = slot_range_to_page_slice( - self.prefix_pages, - self.total_slots, - ) - _all_reduce_materialized_buffer_range( - dense_page_buffer, - self.layout.cp_size, - suffix_rows.start, - suffix_rows.stop, - nvtx_source="index.consume_suffix", - nvtx_layer_id=layer_id, - nvtx_cp_rank=self.layout.cp_rank, + materialized_suffix_by_ipc = ( + _try_tai_ipc_materialize_paged_buffer_page_slot_spans_into( + page_buffer=page_buffer, + dense_page_buffer=dense_page_buffer, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + spans=suffix_spans, + ) ) + if not materialized_suffix_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_page_buffer): + _raise_tai_ipc_materialize_required( + "index_prefetch_suffix_ipc_unavailable", + cp_rank=self.layout.cp_rank, + cp_size=self.layout.cp_size, + spans=suffix_spans, + dense_shape=tuple(dense_page_buffer.shape), + ) + if not materialized_suffix_by_ipc: + _materialize_local_paged_buffer_page_slot_spans_into( + page_buffer=page_buffer, + dense_page_buffer=dense_page_buffer, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + spans=suffix_spans, + ) + if not materialized_suffix_by_ipc: + for suffix_rows in _slot_spans_to_page_slices(suffix_spans): + _all_reduce_materialized_buffer_range( + dense_page_buffer, + self.layout.cp_size, + suffix_rows.start, + suffix_rows.stop, + nvtx_source="index.consume_suffix", + nvtx_layer_id=layer_id, + nvtx_cp_rank=self.layout.cp_rank, + ) self._log_layer( layer_id, - "index_consume_suffix_done layer=%s rows=%s:%s", + "index_consume_suffix_done layer=%s suffix_spans=%s", layer_id, - suffix_rows.start, - suffix_rows.stop, + suffix_spans, ) suffix_ms = _cpu_timing_ms(suffix_cpu) @@ -1872,10 +1868,7 @@ class CpSharedKVIndexPrefetcher: page_size: int, index_head_dim: int, current_req_id: torch.Tensor | None = None, - pool_page_buffer: torch.Tensor | None = None, ) -> Optional[tuple[torch.Tensor, torch.Tensor]]: - if pool_page_buffer is not None: - _seed_symm_ipc_agreement(self, pool_page_buffer) if self.disabled: self._log_layer( layer_id, @@ -1935,72 +1928,6 @@ class CpSharedKVIndexPrefetcher: remap_cpu = _cpu_timing_start() if current_req_id is None: current_req_id = torch.zeros_like(current_locs, dtype=torch.long) - - if _prefetch_symm_active(self, dense_page_buffer.device): - # Symm current exchange (see the MLA twin): fill into the - # staging, barrier, gather all current pages into the prefetched - # dense page buffer. Zero NCCL. - from tai_kernel.nsa_prefill.ipc import ( - cp_symm_barrier, - gather_cuda_ipc_peer_pages, - ) - - state = self._get_or_build_symm_state( - dense_page_buffer=dense_page_buffer, - logical_pages=logical_pages, - ) - staging = get_compose_staging(dense_page_buffer.device) - parity, staging_span = _symm_begin_current_staging( - staging=staging, - plan=state, - kind="index", - layer_id=layer_id, - page_nbytes=state.page_nbytes, - ) - if state.num_current_pages > 0: - staging_pages = staging_span.view( - dense_page_buffer.dtype - ).reshape( - state.num_current_pages + 1, *dense_page_buffer.shape[1:] - ) - fill_current_index_page_slots( - dense_page_buffer=staging_pages, - current_index_k=current_index_k, - current_index_scale=current_index_scale, - current_locs=current_locs, - page_inverse=state.staging_page_inverse, - page_size=page_size, - index_head_dim=index_head_dim, - current_req_id=current_req_id, - ) - cp_symm_barrier( - staging.flag_ptrs, self_rank=int(self.layout.cp_rank) - ) - gather_cuda_ipc_peer_pages( - staging.peer_region_ptrs("index", parity), - dense_page_buffer, - state.writer_ranks, - state.staging_slots, - state.current_dense_pages, - page_nbytes=state.page_nbytes, - ) - remap_ms = _cpu_timing_ms(remap_cpu) - total_ms = _cpu_timing_ms(consume_cpu) - self._log_layer( - layer_id, - "index_consume_prefix_current_hit layer=%s prefix_pages=%s " - "dense_pages=%s current_rows=%s symm=1 total_ms=%.3f " - "wait_ms=%.3f remap_ms=%.3f", - layer_id, - self.prefix_pages, - int(dense_page_buffer.shape[0]), - int(current_index_k.shape[0]), - total_ms, - wait_ms, - remap_ms, - ) - return dense_page_buffer, state.dense_pages - dense_pages = remap_logical_pages_to_slot_dense_pages( logical_pages, page_inverse=self.page_inverse, @@ -2016,20 +1943,38 @@ class CpSharedKVIndexPrefetcher: index_head_dim=index_head_dim, current_req_id=current_req_id, ) - if self.layout.cp_size > 1 and self.prefix_pages < self.total_slots: - current_pages = slot_range_to_page_slice( - self.prefix_pages, - self.total_slots, - ) - _all_reduce_materialized_buffer_range( - dense_page_buffer, - self.layout.cp_size, - current_pages.start, - current_pages.stop, - nvtx_source="index.prefetch_current", - nvtx_layer_id=layer_id, - nvtx_cp_rank=self.layout.cp_rank, + if self.layout.cp_size > 1: + current_materialized_by_ipc = ( + _try_tai_ipc_materialize_current_paged_buffer_page_slot_spans_into( + dense_page_buffer=dense_page_buffer, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + spans=self.current_slot_spans, + ) ) + if ( + not current_materialized_by_ipc + and self.current_slot_spans + and _should_fail_fast_tai_ipc_materialize(dense_page_buffer) + ): + _raise_tai_ipc_materialize_required( + "index_prefetch_current_ipc_unavailable", + cp_rank=self.layout.cp_rank, + cp_size=self.layout.cp_size, + spans=self.current_slot_spans, + dense_shape=tuple(dense_page_buffer.shape), + ) + if not current_materialized_by_ipc: + for current_pages in _slot_spans_to_page_slices(self.current_slot_spans): + _all_reduce_materialized_buffer_range( + dense_page_buffer, + self.layout.cp_size, + current_pages.start, + current_pages.stop, + nvtx_source="index.prefetch_current", + nvtx_layer_id=layer_id, + nvtx_cp_rank=self.layout.cp_rank, + ) remap_ms = _cpu_timing_ms(remap_cpu) total_ms = _cpu_timing_ms(consume_cpu) self._log_layer( @@ -2092,12 +2037,13 @@ class CpSharedKVIndexPrefetcher: return start_cpu = _cpu_timing_start() + current_stream = torch.cuda.current_stream() get_cpu = _cpu_timing_start() try: page_buffer = _prefetch_pool_get_index_buffer( token_to_kv_pool=token_to_kv_pool, layer_id=next_layer_id, - stream=self.stream, + stream=current_stream, ) get_ms = _cpu_timing_ms(get_cpu) except Exception: @@ -2113,41 +2059,61 @@ class CpSharedKVIndexPrefetcher: return try: - current_stream = torch.cuda.current_stream() - prefix_rows = slot_range_to_page_slice(0, self.prefix_pages) + prefix_row_spans = _slot_spans_to_page_slices(self.prefix_slot_spans) + prefix_rows = prefix_row_spans[0] if prefix_row_spans else slice(0, 0) materialize_cpu = _cpu_timing_start() dense_page_buffer = page_buffer.new_zeros( (self.dense_num_pages, *page_buffer.shape[1:]) ) self._log_next_layer( next_layer_id, - "index_start_prefix_begin next_layer=%s start_slot=0 " - "end_slot=%s dense_pages=%s", + "index_start_prefix_begin next_layer=%s prefix_slot_spans=%s " + "dense_pages=%s", next_layer_id, - self.prefix_pages, + self.prefix_slot_spans, int(dense_page_buffer.shape[0]), ) - materialize_local_paged_buffer_page_slots_into( - page_buffer=page_buffer, - dense_page_buffer=dense_page_buffer, - slot_logical_pages=self.slot_logical_pages, - layout=self.layout, - start_slot=0, - end_slot=self.prefix_pages, + materialized_by_ipc = ( + _try_tai_ipc_materialize_paged_buffer_page_slot_spans_into( + page_buffer=page_buffer, + dense_page_buffer=dense_page_buffer, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + spans=self.prefix_slot_spans, + ) ) + if not materialized_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_page_buffer): + _raise_tai_ipc_materialize_required( + "index_prefetch_prefix_ipc_unavailable", + cp_rank=self.layout.cp_rank, + cp_size=self.layout.cp_size, + spans=self.prefix_slot_spans, + dense_shape=tuple(dense_page_buffer.shape), + ) + if not materialized_by_ipc: + _materialize_local_paged_buffer_page_slot_spans_into( + page_buffer=page_buffer, + dense_page_buffer=dense_page_buffer, + slot_logical_pages=self.slot_logical_pages, + layout=self.layout, + spans=self.prefix_slot_spans, + ) materialize_ms = _cpu_timing_ms(materialize_cpu) reduce_cpu = _cpu_timing_start() - self.stream.wait_stream(current_stream) - with torch.cuda.stream(self.stream): - event = _all_reduce_materialized_buffer_async( - dense_page_buffer[prefix_rows], - cp_size=self.layout.cp_size, - stream=self.stream, - nvtx_source="index.prefetch_prefix", - nvtx_layer_id=next_layer_id, - nvtx_cp_rank=self.layout.cp_rank, - nvtx_rows=(prefix_rows.start, prefix_rows.stop), - ) + if materialized_by_ipc: + event = _record_event_on_stream(current_stream) + else: + self.stream.wait_stream(current_stream) + with torch.cuda.stream(self.stream): + event = _all_reduce_materialized_buffer_ranges_async( + dense_kv_cache=dense_page_buffer, + row_slices=prefix_row_spans, + cp_size=self.layout.cp_size, + stream=self.stream, + nvtx_source="index.prefetch_prefix", + nvtx_layer_id=next_layer_id, + nvtx_cp_rank=self.layout.cp_rank, + ) reduce_enqueue_ms = _cpu_timing_ms(reduce_cpu) if event is None: self.disabled = True @@ -2203,6 +2169,7 @@ class CpSharedKVIndexPrefetcher: dense_page_buffer=dense_page_buffer, prefix_rows=prefix_rows, event=event, + prefix_row_spans=tuple(prefix_row_spans), ) self.handles[next_layer_id] = handle self.pending_attention_handle = handle @@ -2227,19 +2194,27 @@ class CpSharedKVIndexPrefetcher: ) return - prefix_rows = handle.prefix_rows + prefix_row_spans = list(handle.prefix_row_spans or (handle.prefix_rows,)) try: + if _should_fail_fast_tai_ipc_materialize(handle.dense_page_buffer): + _raise_tai_ipc_materialize_required( + "index_prefetch_deferred_prefix_ipc_unavailable", + cp_rank=self.layout.cp_rank, + cp_size=self.layout.cp_size, + dense_shape=tuple(handle.dense_page_buffer.shape), + layer_id=handle.layer_id, + ) current_stream = torch.cuda.current_stream() self.stream.wait_stream(current_stream) with torch.cuda.stream(self.stream): - event = _all_reduce_materialized_buffer_async( - handle.dense_page_buffer[prefix_rows], + event = _all_reduce_materialized_buffer_ranges_async( + dense_kv_cache=handle.dense_page_buffer, + row_slices=prefix_row_spans, cp_size=self.layout.cp_size, stream=self.stream, nvtx_source="index.prefetch_prefix", nvtx_layer_id=handle.layer_id, nvtx_cp_rank=self.layout.cp_rank, - nvtx_rows=(prefix_rows.start, prefix_rows.stop), ) if event is None: self.disabled = True @@ -2255,10 +2230,9 @@ class CpSharedKVIndexPrefetcher: handle.event = event self._log_next_layer( handle.layer_id, - "index_start_prefix_reduce_enqueued next_layer=%s rows=%s:%s", + "index_start_prefix_reduce_enqueued next_layer=%s row_spans=%s", handle.layer_id, - prefix_rows.start, - prefix_rows.stop, + [(rows.start, rows.stop) for rows in prefix_row_spans], ) except Exception: logger.exception( 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 8db9457ad..1eeb7797b 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 @@ -9,14 +9,6 @@ from typing import Any import torch from sglang.srt.environ import envs -from sglang.srt.layers.attention.nsa.cp_shared_kv_compose import ( - acquire_dense_buffer, - compute_staging_capacity_pages, - cp_shared_kv_compose_symm_enabled, - cp_shared_kv_compose_v2_enabled, - get_compose_staging, - get_or_build_compose_plan, -) from sglang.srt.layers.attention.nsa.utils import ( cp_shared_kv_bs_gt1_timing_start, get_cp_shared_kv_local_out_cache_loc, @@ -36,6 +28,8 @@ _TAI_INDEX_MQA_PREPARE_FALLBACK_LOG_COUNTS: dict[str, int] = {} _CURRENT_REUSE_FALLBACK_LOG_COUNTS: dict[str, int] = {} _SLOT_REMAP_CACHE_LOG_COUNTS: dict[str, int] = {} _TAI_IPC_PEER_PTR_CACHE: dict[tuple[object, ...], torch.Tensor] = {} +_TAI_IPC_CURRENT_STAGING_CACHE: dict[tuple[object, ...], "_TaiIpcCurrentStagingState"] = {} +_TAI_IPC_RETIRED_CURRENT_STAGING: list[torch.Tensor] = [] _MLA_PREFETCH_LOG_PROBE_LAYER = 2 _MLA_PREFETCH_DEFAULT_MIN_PREFIX_TOKENS = max( int(envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_TOKENS.get()), @@ -45,6 +39,16 @@ _SORT_NVTX_ENABLED = envs.SGLANG_DEBUG_SORT_NVTX.get() _MATERIALIZE_NVTX_ENABLED = envs.SGLANG_CP_SHARED_KV_MATERIALIZE_NVTX.get() +@dataclass +class _TaiIpcCurrentStagingState: + staging: torch.Tensor + ready: torch.Tensor + peer_ptrs: torch.Tensor + ready_peer_ptrs: torch.Tensor + capacity_nbytes: int + ready_seq: int = 0 + + def cp_shared_kv_debug_enabled() -> bool: return envs.SGLANG_DEBUG_CP_SHARED_KV.get() @@ -449,9 +453,13 @@ def _load_tai_ipc_kernels(): from tai_kernel.nsa_prefill import ipc as tai_ipc required = ( + "allocate_cuda_ipc_buffer", "get_cuda_ipc_mem_handle_with_offset", "open_cuda_ipc_mem_handles_with_offsets", "materialize_cuda_ipc_peer_pages_slot_dense", + "materialize_cuda_ipc_peer_pages_slot_indices", + "materialize_cuda_ipc_peer_pages_slot_indices_wait_ready", + "publish_cuda_ipc_slot_pages_and_mark_ready", ) missing = [name for name in required if not hasattr(tai_ipc, name)] if missing: @@ -832,24 +840,19 @@ def _log_tai_index_mqa_prepare_fallback( ) -def _should_fail_fast_compose_v2_dense_fallback(dense_tensor: torch.Tensor) -> bool: - """Whether compose_v2 may fall back to dense full-buffer collectives. - CPU/unit-test paths can still use the simple fallback. In production CUDA - runs with TAI materialize enabled, a prefix-IPC miss would silently turn a - bs>1 cache-hit compose into a dense all_reduce over the whole buffer. That - hides both correctness-contract drift and severe performance regressions, - so fail fast instead. - """ +def _should_fail_fast_tai_ipc_materialize(dense_tensor: torch.Tensor) -> bool: return bool(dense_tensor.is_cuda and cp_shared_kv_tai_materialize_enabled()) -def _raise_compose_v2_dense_fallback_required(reason: str, **details: Any) -> None: - detail_str = " ".join(f"{key}={value}" for key, value in details.items()) - message = f"[CP_SHARED_KV_FAIL_FAST][compose_v2] reason={reason}" - if detail_str: - message = f"{message} {detail_str}" +def _raise_tai_ipc_materialize_required(reason: str, **details: Any) -> None: + detail_text = " ".join(f"{key}={value}" for key, value in details.items()) + message = ( + f"[CP_SHARED_KV_FAIL_FAST][tai_ipc_materialize] reason={reason}" + + (f" {detail_text}" if detail_text else "") + ) + logger.error(message) raise RuntimeError(message) @@ -2298,53 +2301,118 @@ def _get_or_open_tai_ipc_peer_ptrs( return None -_TAI_IPC_CAPABILITY_AGREEMENT: dict[tuple[object, ...], bool] = {} +def _next_power_of_two_nbytes(nbytes: int) -> int: + nbytes = max(int(nbytes), 1) + return 1 << (nbytes - 1).bit_length() -def _agreed_tai_ipc_peer_ptrs( - tensor: torch.Tensor, +def _tai_ipc_current_staging_cache_key( + *, + kind: str, + device: torch.device, layout: CpSharedKVLayout, -) -> tuple[Any, torch.Tensor] | None: - """Probe peer-IPC capability and AGREE on it across the CP group. + cp_group: Any, +) -> tuple[object, ...]: + return ( + kind, + str(device), + int(layout.cp_size), + int(layout.cp_rank), + getattr(cp_group, "unique_name", None), + ) - The probe alone is per-rank; if it diverged (one rank's open failing), - ranks would issue different collectives downstream and deadlock with a - shape mismatch. A one-time MIN-agreement per pool tensor pins every - rank to the same path. Called at uniform points only (the compose path - runs identically on all ranks). - """ - state = _get_or_open_tai_ipc_peer_ptrs(tensor, layout) +def _get_or_create_tai_ipc_current_staging( + *, + kind: str, + dense_tensor: torch.Tensor, + layout: CpSharedKVLayout, + required_nbytes: int, +) -> tuple[Any, _TaiIpcCurrentStagingState] | None: if layout.cp_size <= 1: - return state - if not torch.distributed.is_initialized(): - # Single-process context (unit tests / tools): no group to agree - # with. Production CP always has torch.distributed initialized. - return state + return None + if not dense_tensor.is_cuda or not dense_tensor.is_contiguous(): + _log_tai_ipc_materialize_fallback( + "current_dense_tensor_unsupported", + "CP shared KV current IPC staging requires a contiguous CUDA dense " + "tensor; falling back to current-slot collective. kind=%s " + "device=%s contiguous=%s shape=%s", + kind, + dense_tensor.device, + dense_tensor.is_contiguous(), + tuple(dense_tensor.shape), + limit=4, + ) + return None + + kernels = _load_tai_ipc_kernels() + if kernels is None: + return None + cp_group = get_attention_cp_group() - key = _tai_ipc_peer_ptr_cache_key(tensor, layout, cp_group) - agreed = _TAI_IPC_CAPABILITY_AGREEMENT.get(key) - if agreed is None: - flag = torch.tensor( - [1 if state is not None else 0], - dtype=torch.int32, - device=tensor.device, + if int(getattr(cp_group, "world_size", layout.cp_size)) != int(layout.cp_size): + _log_tai_ipc_materialize_fallback( + "current_group_size_mismatch", + "CP shared KV current IPC staging cp group size mismatch; " + "falling back to current-slot collective. layout_cp_size=%s " + "group_world_size=%s", + layout.cp_size, + getattr(cp_group, "world_size", None), + limit=1, ) - torch.distributed.all_reduce( - flag, - op=torch.distributed.ReduceOp.MIN, - group=cp_group.device_group, + return None + + cache_key = _tai_ipc_current_staging_cache_key( + kind=kind, + device=dense_tensor.device, + layout=layout, + cp_group=cp_group, + ) + required_nbytes = max(int(required_nbytes), 1) + state = _TAI_IPC_CURRENT_STAGING_CACHE.get(cache_key) + if state is not None and state.capacity_nbytes >= required_nbytes: + return kernels, state + + if state is not None: + # Keep retired buffers alive because CUDA IPC peer pointer caches name + # their allocations. Growth should be rare due power-of-two capacity. + _TAI_IPC_RETIRED_CURRENT_STAGING.extend([state.staging, state.ready]) + + capacity_nbytes = _next_power_of_two_nbytes(required_nbytes) + try: + staging = kernels.allocate_cuda_ipc_buffer( + capacity_nbytes, device=dense_tensor.device ) - agreed = bool(int(flag.item()) == 1) - _TAI_IPC_CAPABILITY_AGREEMENT[key] = agreed - if state is not None and not agreed: - logger.warning( - "[CP_SHARED_KV_FALLBACK][compose_v2] peer IPC works on this " - "rank but not on every CP rank; the whole group uses the " - "collective fallback. cp_rank=%s", - layout.cp_rank, - ) - return state if agreed else None + ready = torch.zeros((1,), dtype=torch.int64, device=dense_tensor.device) + staging_state = _get_or_open_tai_ipc_peer_ptrs(staging, layout) + ready_state = _get_or_open_tai_ipc_peer_ptrs(ready, layout) + if staging_state is None or ready_state is None: + return None + _, peer_ptrs = staging_state + _, ready_peer_ptrs = ready_state + state = _TaiIpcCurrentStagingState( + staging=staging, + ready=ready, + peer_ptrs=peer_ptrs, + ready_peer_ptrs=ready_peer_ptrs, + capacity_nbytes=capacity_nbytes, + ) + _TAI_IPC_CURRENT_STAGING_CACHE[cache_key] = state + return kernels, state + except Exception as exc: + _log_tai_ipc_materialize_fallback( + "current_staging_setup_failed", + "CP shared KV current IPC staging setup failed; falling back to " + "current-slot collective. kind=%s cp_rank=%s cp_size=%s " + "required_nbytes=%s error=%s", + kind, + layout.cp_rank, + layout.cp_size, + required_nbytes, + exc, + limit=4, + ) + return None def _page_nbytes_from_page_tensor(tensor: torch.Tensor) -> int: @@ -2370,19 +2438,14 @@ def _try_tai_ipc_materialize_token_kv_page_slots_into( if start_slot == end_slot: return True if start_slot != 0: - _log_tai_ipc_materialize_fallback( - "start_slot_nonzero", - "CP shared KV tai IPC token materialize only supports slot ranges " - "starting at zero; falling back to local materialize plus " - "collective. cp_rank=%s cp_size=%s start_slot=%s end_slot=%s " - "page_size=%s", - layout.cp_rank, - layout.cp_size, - start_slot, - end_slot, - page_size, + return _try_tai_ipc_materialize_token_kv_page_slot_spans_into( + kv_cache=kv_cache, + dense_kv_cache=dense_kv_cache, + slot_logical_pages=slot_logical_pages, + layout=layout, + page_size=page_size, + spans=[(start_slot, end_slot)], ) - return False if not dense_kv_cache.is_cuda or not dense_kv_cache.is_contiguous(): _log_tai_ipc_materialize_fallback( "dense_token_tensor_unsupported", @@ -2442,6 +2505,207 @@ def _try_tai_ipc_materialize_token_kv_page_slots_into( return False +def _slot_spans_to_cuda_slot_indices( + spans: list[tuple[int, int]], + *, + total_slots: int, + device: torch.device, +) -> torch.Tensor: + merged = _merge_slot_spans(spans) + for start_slot, end_slot in merged: + if start_slot < 0 or end_slot < start_slot or end_slot > total_slots: + raise ValueError( + "Invalid CP shared KV IPC slot span: " + f"start_slot={start_slot} end_slot={end_slot} total_slots={total_slots}" + ) + if not merged: + return torch.empty((0,), dtype=torch.long, device=device) + return torch.cat( + [ + torch.arange(start_slot, end_slot, dtype=torch.long, device=device) + for start_slot, end_slot in merged + ] + ).contiguous() + + +def _build_current_staging_ipc_descriptors( + *, + slot_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + spans: list[tuple[int, int]], + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + flat_slot_logical_pages = _contiguous_for_tai( + slot_logical_pages.reshape(-1).to(device=device) + ) + slot_indices = _slot_spans_to_cuda_slot_indices( + spans, + total_slots=int(flat_slot_logical_pages.numel()), + device=device, + ) + if slot_indices.numel() == 0: + empty = torch.empty((0,), dtype=torch.long, device=device) + return empty, empty, empty + logical_pages = _contiguous_for_tai(flat_slot_logical_pages.index_select(0, slot_indices)) + valid = logical_pages > 0 + owner_ranks = torch.remainder(logical_pages - 1, int(layout.cp_size)).to(torch.long) + owner_ranks = torch.where(valid, owner_ranks, torch.full_like(owner_ranks, -1)) + dense_page_indices = (slot_indices + 1).to(torch.long).contiguous() + src_page_indices = torch.where( + valid, + dense_page_indices, + torch.full_like(dense_page_indices, -1), + ).contiguous() + return owner_ranks.contiguous(), src_page_indices, dense_page_indices + + +def _try_tai_ipc_materialize_token_kv_page_slot_spans_into( + *, + kv_cache: torch.Tensor, + dense_kv_cache: torch.Tensor, + slot_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + page_size: int, + spans: list[tuple[int, int]], +) -> bool: + if not spans: + return True + if not dense_kv_cache.is_cuda or not dense_kv_cache.is_contiguous(): + _log_tai_ipc_materialize_fallback( + "dense_token_tensor_unsupported", + "CP shared KV tai IPC token span materialize requires a contiguous " + "CUDA dense tensor; falling back to local materialize plus collective. " + "device=%s contiguous=%s dense_shape=%s spans=%s", + dense_kv_cache.device, + dense_kv_cache.is_contiguous(), + tuple(dense_kv_cache.shape), + spans, + limit=4, + ) + return False + + ipc_state = _get_or_open_tai_ipc_peer_ptrs(kv_cache, layout) + if ipc_state is None: + return False + kernels, peer_ptrs = ipc_state + + flat_slot_logical_pages = _contiguous_for_tai( + slot_logical_pages.reshape(-1).to(device=dense_kv_cache.device) + ) + try: + slot_indices = _slot_spans_to_cuda_slot_indices( + spans, + total_slots=int(flat_slot_logical_pages.numel()), + device=dense_kv_cache.device, + ) + if slot_indices.numel() == 0: + return True + slot_logical_pages_range = _contiguous_for_tai( + flat_slot_logical_pages.index_select(0, slot_indices) + ) + owner_ranks, src_page_indices = build_cp_shared_kv_ipc_page_descriptors( + slot_logical_pages_range, + layout, + physical_page_capacity=kv_cache.shape[0] // page_size, + ) + kernels.materialize_cuda_ipc_peer_pages_slot_indices( + peer_ptrs, + dense_kv_cache, + owner_ranks, + src_page_indices, + (slot_indices + 1).contiguous(), + page_nbytes=_token_kv_page_nbytes(kv_cache, page_size), + ) + return True + except Exception as exc: + _log_tai_ipc_materialize_fallback( + "token_span_kernel_failed", + "CP shared KV tai IPC token span materialize failed; falling back " + "to collective materialize. cp_rank=%s cp_size=%s spans=%s " + "page_size=%s kv_shape=%s dense_shape=%s error=%s", + layout.cp_rank, + layout.cp_size, + spans, + page_size, + tuple(kv_cache.shape), + tuple(dense_kv_cache.shape), + exc, + ) + return False + + +def _try_tai_ipc_materialize_current_token_kv_page_slot_spans_into( + *, + dense_kv_cache: torch.Tensor, + slot_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + page_size: int, + spans: list[tuple[int, int]], +) -> bool: + """Materialize peer current KV slots through persistent IPC staging.""" + if not spans: + return True + page_nbytes = _token_kv_page_nbytes(dense_kv_cache, page_size) + required_nbytes = int(dense_kv_cache.shape[0]) * _page_nbytes_from_page_tensor( + dense_kv_cache + ) + staging_state = _get_or_create_tai_ipc_current_staging( + kind="token", + dense_tensor=dense_kv_cache, + layout=layout, + required_nbytes=required_nbytes, + ) + if staging_state is None: + return False + kernels, state = staging_state + try: + owner_ranks, src_page_indices, dense_page_indices = ( + _build_current_staging_ipc_descriptors( + slot_logical_pages=slot_logical_pages, + layout=layout, + spans=spans, + device=dense_kv_cache.device, + ) + ) + if dense_page_indices.numel() == 0: + return True + state.ready_seq += 1 + ready_seq = int(state.ready_seq) + kernels.publish_cuda_ipc_slot_pages_and_mark_ready( + dense_kv_cache, + state.staging, + dense_page_indices, + state.ready, + ready_seq=ready_seq, + page_nbytes=page_nbytes, + ) + kernels.materialize_cuda_ipc_peer_pages_slot_indices_wait_ready( + state.peer_ptrs, + state.ready_peer_ptrs, + dense_kv_cache, + owner_ranks, + src_page_indices, + dense_page_indices, + ready_seq=ready_seq, + page_nbytes=page_nbytes, + ) + return True + except Exception as exc: + _log_tai_ipc_materialize_fallback( + "current_token_kernel_failed", + "CP shared KV current token IPC materialize failed; falling back " + "to current-slot collective. cp_rank=%s cp_size=%s spans=%s " + "page_size=%s dense_shape=%s error=%s", + layout.cp_rank, + layout.cp_size, + spans, + page_size, + tuple(dense_kv_cache.shape), + exc, + ) + return False + + def _try_tai_ipc_materialize_paged_buffer_page_slots_into( *, page_buffer: torch.Tensor, @@ -2454,17 +2718,13 @@ def _try_tai_ipc_materialize_paged_buffer_page_slots_into( if start_slot == end_slot: return True if start_slot != 0: - _log_tai_ipc_materialize_fallback( - "paged_start_slot_nonzero", - "CP shared KV tai IPC paged materialize only supports slot ranges " - "starting at zero; falling back to local materialize plus " - "collective. cp_rank=%s cp_size=%s start_slot=%s end_slot=%s", - layout.cp_rank, - layout.cp_size, - start_slot, - end_slot, + return _try_tai_ipc_materialize_paged_buffer_page_slot_spans_into( + page_buffer=page_buffer, + dense_page_buffer=dense_page_buffer, + slot_logical_pages=slot_logical_pages, + layout=layout, + spans=[(start_slot, end_slot)], ) - return False if not dense_page_buffer.is_cuda or not dense_page_buffer.is_contiguous(): _log_tai_ipc_materialize_fallback( "dense_paged_tensor_unsupported", @@ -2522,6 +2782,147 @@ def _try_tai_ipc_materialize_paged_buffer_page_slots_into( return False +def _try_tai_ipc_materialize_paged_buffer_page_slot_spans_into( + *, + page_buffer: torch.Tensor, + dense_page_buffer: torch.Tensor, + slot_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + spans: list[tuple[int, int]], +) -> bool: + if not spans: + return True + if not dense_page_buffer.is_cuda or not dense_page_buffer.is_contiguous(): + _log_tai_ipc_materialize_fallback( + "dense_paged_tensor_unsupported", + "CP shared KV tai IPC paged span materialize requires a contiguous " + "CUDA dense tensor; falling back to local materialize plus collective. " + "device=%s contiguous=%s dense_shape=%s spans=%s", + dense_page_buffer.device, + dense_page_buffer.is_contiguous(), + tuple(dense_page_buffer.shape), + spans, + limit=4, + ) + return False + + ipc_state = _get_or_open_tai_ipc_peer_ptrs(page_buffer, layout) + if ipc_state is None: + return False + kernels, peer_ptrs = ipc_state + + flat_slot_logical_pages = _contiguous_for_tai( + slot_logical_pages.reshape(-1).to(device=dense_page_buffer.device) + ) + try: + slot_indices = _slot_spans_to_cuda_slot_indices( + spans, + total_slots=int(flat_slot_logical_pages.numel()), + device=dense_page_buffer.device, + ) + if slot_indices.numel() == 0: + return True + slot_logical_pages_range = _contiguous_for_tai( + flat_slot_logical_pages.index_select(0, slot_indices) + ) + owner_ranks, src_page_indices = build_cp_shared_kv_ipc_page_descriptors( + slot_logical_pages_range, + layout, + physical_page_capacity=page_buffer.shape[0], + ) + kernels.materialize_cuda_ipc_peer_pages_slot_indices( + peer_ptrs, + dense_page_buffer, + owner_ranks, + src_page_indices, + (slot_indices + 1).contiguous(), + page_nbytes=_page_nbytes_from_page_tensor(page_buffer), + ) + return True + except Exception as exc: + _log_tai_ipc_materialize_fallback( + "paged_span_kernel_failed", + "CP shared KV tai IPC paged span materialize failed; falling back " + "to collective materialize. cp_rank=%s cp_size=%s spans=%s " + "page_shape=%s dense_shape=%s error=%s", + layout.cp_rank, + layout.cp_size, + spans, + tuple(page_buffer.shape), + tuple(dense_page_buffer.shape), + exc, + ) + return False + + +def _try_tai_ipc_materialize_current_paged_buffer_page_slot_spans_into( + *, + dense_page_buffer: torch.Tensor, + slot_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + spans: list[tuple[int, int]], +) -> bool: + """Materialize peer current index/page slots through persistent IPC staging.""" + if not spans: + return True + page_nbytes = _page_nbytes_from_page_tensor(dense_page_buffer) + required_nbytes = int(dense_page_buffer.shape[0]) * page_nbytes + staging_state = _get_or_create_tai_ipc_current_staging( + kind="paged", + dense_tensor=dense_page_buffer, + layout=layout, + required_nbytes=required_nbytes, + ) + if staging_state is None: + return False + kernels, state = staging_state + try: + owner_ranks, src_page_indices, dense_page_indices = ( + _build_current_staging_ipc_descriptors( + slot_logical_pages=slot_logical_pages, + layout=layout, + spans=spans, + device=dense_page_buffer.device, + ) + ) + if dense_page_indices.numel() == 0: + return True + state.ready_seq += 1 + ready_seq = int(state.ready_seq) + kernels.publish_cuda_ipc_slot_pages_and_mark_ready( + dense_page_buffer, + state.staging, + dense_page_indices, + state.ready, + ready_seq=ready_seq, + page_nbytes=page_nbytes, + ) + kernels.materialize_cuda_ipc_peer_pages_slot_indices_wait_ready( + state.peer_ptrs, + state.ready_peer_ptrs, + dense_page_buffer, + owner_ranks, + src_page_indices, + dense_page_indices, + ready_seq=ready_seq, + page_nbytes=page_nbytes, + ) + return True + except Exception as exc: + _log_tai_ipc_materialize_fallback( + "current_paged_kernel_failed", + "CP shared KV current paged IPC materialize failed; falling back " + "to current-slot collective. cp_rank=%s cp_size=%s spans=%s " + "dense_shape=%s error=%s", + layout.cp_rank, + layout.cp_size, + spans, + tuple(dense_page_buffer.shape), + exc, + ) + return False + + def is_current_only_extend_batch(forward_batch) -> bool: """Return whether an extend batch has no cached/history tokens. @@ -2807,51 +3208,6 @@ def build_batch_prefix_slot_span( return (start_slot, end_slot) -def get_or_build_batch_slot_spans( - forward_batch, - *, - logical_pages: torch.Tensor, - prefix_lens_cpu, - extend_lens_cpu, - page_size: int, - want_prefix: bool, -) -> tuple[list[tuple[int, int]] | None, list[tuple[int, int]]]: - """Per-batch cache for the layer-invariant slot-span builders. - - The builders read ``logical_pages`` only for its SHAPE; together with the - batch-scoped ``prefix/extend`` lens that makes the spans identical for - every layer of a forward — rebuilding the per-request Python loops per - layer was part of the measured pre-attention CPU gap. - """ - - key = (tuple(logical_pages.shape), int(page_size), bool(want_prefix)) - cache = getattr(forward_batch, "_cp_batch_slot_spans_cache", None) - if cache is None: - cache = {} - forward_batch._cp_batch_slot_spans_cache = cache - hit = cache.get(key) - if hit is not None: - return hit - prefix_spans = ( - build_batch_prefix_slot_spans( - logical_pages=logical_pages, - prefix_lens_cpu=prefix_lens_cpu, - page_size=page_size, - ) - if want_prefix - else None - ) - current_spans = build_batch_current_slot_spans( - logical_pages=logical_pages, - prefix_lens_cpu=prefix_lens_cpu, - extend_lens_cpu=extend_lens_cpu, - page_size=page_size, - ) - result = (prefix_spans, current_spans) - cache[key] = result - return result - - def build_batch_prefix_slot_spans( *, logical_pages: torch.Tensor, @@ -4450,578 +4806,6 @@ def materialize_local_token_kv_page_slots_into( dense_range.copy_(torch.where(owned_view, gathered, zero)) -def build_batch_current_page_writer_ranks( - *, - prefix_lens_cpu, - extend_lens_cpu, - page_size: int, - cp_size: int, -) -> list[int]: - """Per-current-page compute owner (writer) in batch current-slot order. - - Valid ONLY for the page-aligned in-seq split, where each current page is - written by exactly one rank. Order matches - ``build_batch_current_slot_spans`` (request-major ascending slots). - """ - - from sglang.srt.mem_cache.cp_shared_kv_compute_owner import ( - build_in_seq_page_compute_owners, - ) - - writers: list[int] = [] - for prefix_len, extend_len in zip(prefix_lens_cpu, extend_lens_cpu): - writers.extend( - int(owner) - for owner in build_in_seq_page_compute_owners( - extend_len=int(extend_len), - extend_prefix_len=int(prefix_len), - page_size=page_size, - cp_size=cp_size, - ) - ) - return writers - - -def maybe_build_current_page_writer_ranks( - *, - forward_batch, - prefix_lens_cpu, - extend_lens_cpu, - page_size: int, - layout: CpSharedKVLayout, -) -> list[int] | None: - """Caller-side gate for the symm current-page exchange. - - Returns writer ranks only when the symm path is enabled AND the batch - uses the page-aligned in-seq split (single writer per current page); - otherwise None, which keeps compose on the collective current exchange. - - The list is batch-invariant (a pure function of the batch lens), so it is - cached on the forward batch — rebuilding ~bs x current-pages ints per - layer per call site would sit on the launch-critical path for nothing. - """ - - if not cp_shared_kv_compose_symm_enabled(): - return None - metadata = getattr(forward_batch, "nsa_cp_metadata", None) - if metadata is None or not getattr(metadata, "page_aligned", False): - return None - if prefix_lens_cpu is None or extend_lens_cpu is None: - return None - cache_key = (int(page_size), int(layout.cp_size)) - cached_key = getattr(forward_batch, "cp_shared_kv_current_writer_key", None) - if cached_key == cache_key: - return forward_batch.cp_shared_kv_current_writer_ranks - writers = build_batch_current_page_writer_ranks( - prefix_lens_cpu=prefix_lens_cpu, - extend_lens_cpu=extend_lens_cpu, - page_size=page_size, - cp_size=int(layout.cp_size), - ) - forward_batch.cp_shared_kv_current_writer_key = cache_key - forward_batch.cp_shared_kv_current_writer_ranks = writers - return writers - - -def _symm_staging_ready_or_register( - *, - layout: CpSharedKVLayout, - kv_cache: torch.Tensor, - page_size: int, -) -> bool: - """Register the compact symm staging on first use (collective; uniform - call point). - - Only the token-KV compose registers (it knows the exact KV page bytes); - the index compose uses symm once registration has happened. - """ - - staging = get_compose_staging(kv_cache.device) - if staging.registered: - return True - cp_group = get_attention_cp_group() - # NSA index page bytes are model constants (index_head_dim 128, - # quant_block 128 -> head + 4 scale bytes per token). - index_page_nbytes = page_size * (128 + 4) - kv_page_nbytes = _token_kv_page_nbytes(kv_cache, page_size) - staging.register( - cp_group=cp_group, - cp_rank=int(layout.cp_rank), - cp_size=int(layout.cp_size), - capacity_pages=compute_staging_capacity_pages( - kv_pool_tokens=int(kv_cache.shape[0]), - page_size=page_size, - cp_size=int(layout.cp_size), - kv_page_nbytes=kv_page_nbytes, - index_page_nbytes=index_page_nbytes, - ), - kv_page_nbytes=kv_page_nbytes, - index_page_nbytes=index_page_nbytes, - ) - return staging.registered - - -class _SymmTokenFillState: - """Layer-invariant pieces of the symm token-KV compose, built once per - batch and anchored on the (batch-lifetime) compose plan. - - ``staging_current_rows`` maps each current KV row to its row in the - compact staging (via the staging page inverse), so the per-layer fill is - a single ``index_copy_``. ``mixed_locs`` is the page-slack-masked loc - remap the fused fill kernel used to recompute every layer — it depends - only on the batch's locs, never on buffer contents. - """ - - __slots__ = ("staging_current_rows", "mixed_locs") - - def __init__(self, staging_current_rows, mixed_locs): - self.staging_current_rows = staging_current_rows - self.mixed_locs = mixed_locs - - -def _get_or_build_symm_token_fill_state( - *, - plan, - slot_remap, - layout: CpSharedKVLayout, - kv_cache: torch.Tensor, - logical_locs: torch.Tensor, - current_locs: torch.Tensor, - loc_req_id: torch.Tensor, - current_req_id: torch.Tensor, - page_size: int, -) -> "_SymmTokenFillState": - state = getattr(plan, "_symm_token_fill_state", None) - if state is not None: - return state - logical_locs = filter_locs_mappable_to_physical_pool( - logical_locs=logical_locs, - layout=layout, - physical_token_capacity=kv_cache.shape[0], - ) - dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs( - logical_locs, - slot_remap=slot_remap, - page_size=page_size, - loc_req_id=loc_req_id, - ) - # Same masking as the fused fill kernel / torch reference: page-slack - # rows inside current pages become -1 (invisible to attention). - current_mask, _ = build_current_loc_remap(logical_locs, current_locs) - current_page_mask = build_current_page_mask( - logical_locs, current_locs, page_size=page_size - ) - mixed_locs = torch.where( - current_page_mask & (~current_mask), - torch.full_like(dense_locs, -1), - dense_locs, - ) - staging_current_rows = remap_logical_locs_to_slot_dense_locs_optimized( - current_locs.reshape(-1), - page_inverse=plan.staging_page_inverse, - page_size=page_size, - loc_req_id=current_req_id.reshape(-1), - ).to(torch.long) - state = _SymmTokenFillState(staging_current_rows, mixed_locs) - object.__setattr__(plan, "_symm_token_fill_state", state) - return state - - -def _symm_begin_current_staging( - *, - staging, - plan, - kind: str, - layer_id: int, - page_nbytes: int, -) -> tuple[int, torch.Tensor]: - """Open this round's staging span for the current-row fill. - - Validates the batch against the registered staging layout (both checks - are batch-logical, hence rank-uniform), advances the round parity, and - returns the zeroed span — zeroed so that tail-slack rows inside current - pages stay zero in every peer's gathered copy, exactly like the sentinel - zero-fill did on the dense buffer. - """ - - if int(plan.num_current_pages) > staging.capacity_pages: - raise RuntimeError( - "[CP_SHARED_KV_FAIL_FAST][compose_symm] batch current pages " - f"exceed the symm staging capacity: pages={plan.num_current_pages} " - f"capacity={staging.capacity_pages}. Raise " - "SGLANG_CP_SHARED_KV_SYMM_HEAP_MB or lower " - "--cp-shared-kv-prefill-max-total-extend-tokens." - ) - if page_nbytes != staging.page_nbytes(kind): - raise RuntimeError( - "[CP_SHARED_KV_FAIL_FAST][compose_symm] page bytes diverge from " - f"the registered staging layout: kind={kind} page_nbytes=" - f"{page_nbytes} registered={staging.page_nbytes(kind)}" - ) - if plan.staging_page_inverse is None: - raise RuntimeError( - "[CP_SHARED_KV_FAIL_FAST][compose_symm] compose plan has no " - "staging page inverse (slot_remap.page_inverse missing?)" - ) - parity = staging.begin_round(int(layer_id), kind) - # +1: staging row 0 is the (unused) dummy page; zeroing it keeps any - # accidental slot-0 read deterministic. - span = staging.buffer(kind, parity)[ - : (int(plan.num_current_pages) + 1) * page_nbytes - ] - span.zero_() - return parity, span - - -def _symm_barrier_and_gather_all( - *, - kernels, - pool_peer_ptrs: torch.Tensor, - dense_buffer: torch.Tensor, - plan, - layout: CpSharedKVLayout, - staging, - kind: str, - parity: int, - page_nbytes: int, -) -> None: - """Step B compose: barrier, then ONE slot-dense gather for everything. - - The current rows were already written into this rank's staging span (the - fill kernels were pointed there via ``plan.staging_page_inverse``), so - after the barrier a single gather against the concatenated - ``[pool peer ptrs | staging peer ptrs]`` table materializes prefix pages - (owner < cp_size, from the KV pools), current pages (owner = cp_size + - writer, from the stagings — including this rank's own), and zero-fills - the remaining sentinel slots. No publish copy, no second gather. - - The barrier runs even with zero current pages — barrier COUNTS must - match across ranks. - - INVARIANT: every gate on the path to this call (compose env flags, - page_aligned metadata, staging.registered, the agreed IPC capability, - prefetch absence) MUST be rank-uniform. A per-rank divergence here - desyncs the barrier counting and hangs the CP group — never add a - per-rank condition without routing it through a group agreement first - (see _agreed_tai_ipc_peer_ptrs).""" - - from tai_kernel.nsa_prefill.ipc import cp_symm_barrier - - cp_symm_barrier(staging.flag_ptrs, self_rank=int(layout.cp_rank)) - kernels.materialize_cuda_ipc_peer_pages_slot_dense( - staging.combined_ptr_table(pool_peer_ptrs, kind, parity), - dense_buffer, - plan.symm_all_owner_ranks, - plan.symm_all_src_pages, - page_nbytes=page_nbytes, - ) - - -def _resolve_partial_current_spans( - *, - current_slot_spans: list[tuple[int, int]] | None, - prefix_slot_span: tuple[int, int] | None, - prefix_slot_spans: list[tuple[int, int]] | None, - prefix_pages: int, - total_slots: int, -) -> list[tuple[int, int]]: - """Replicate the legacy current-span defaulting (incl. its error).""" - - if current_slot_spans is 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." - ) - current_slot_spans = ( - [(int(prefix_pages), total_slots)] - if int(prefix_pages) < total_slots - else [] - ) - return _merge_slot_spans(current_slot_spans) - - -def _reduce_current_pages_compact( - dense_buffer: torch.Tensor, - dense_num_pages: int, - current_dense_pages: torch.Tensor, - layout: CpSharedKVLayout, - *, - nvtx_source: str, - layer_id: int | None, -) -> None: - """One collective over the compact current pages, scattered back in place. - - Used when prefix rows must NOT be reduced (after an IPC prefix gather they - hold real values on every rank). Operates on a uint8 byte view: every - byte of a current page is writer-exclusive (current rows are written by - exactly one rank, non-current rows are zero on all ranks), so the byte sum - has no carries and is exact for any element dtype. ``view`` (not - ``reshape``) keeps failure loud if the buffer is ever non-viewable — - a silent copy here would drop the scatter-back. - """ - - paged = dense_buffer.view(torch.uint8).view(int(dense_num_pages), -1) - compact = paged.index_select(0, current_dense_pages) - _all_reduce_materialized_buffer( - compact, - layout.cp_size, - nvtx_source=nvtx_source, - nvtx_layer_id=layer_id, - nvtx_cp_rank=layout.cp_rank, - ) - paged.index_copy_(0, current_dense_pages, compact) - - -def _compose_token_kv_partial_current_v2( - *, - kv_cache: torch.Tensor, - logical_locs: torch.Tensor, - current_kv_cache: torch.Tensor, - current_locs: torch.Tensor, - slot_remap: SharedTokenKVSlotRemap, - layout: CpSharedKVLayout, - page_size: int, - prefix_spans: list[tuple[int, int]], - current_spans: list[tuple[int, int]], - loc_req_id: torch.Tensor, - current_req_id: torch.Tensor, - layer_id: int | None, - nvtx_source: str, - timing_start: object, - current_page_writer_ranks: list[int] | None = None, -) -> tuple[torch.Tensor, torch.Tensor]: - """Step A/B compose for the MLA token-KV dense buffer. - - Fast path: one IPC slot-dense gather covers ALL prefix spans (slots - outside them are -1 sentinels, zero-filled by the kernel); current pages - are then shared either via barrier + symm-heap IPC gather (Step B, zero - NCCL) or via one collective over the compact current pages. Fallback - (no peer IPC): local materialize of all prefix spans (no per-span - collectives) + ONE sum-all-reduce over the whole dense buffer, exact - because every row is writer-exclusive at reduce time. - """ - - use_symm = ( - cp_shared_kv_compose_symm_enabled() - and current_page_writer_ranks is not None - and layer_id is not None - ) - plan = get_or_build_compose_plan( - slot_remap=slot_remap, - layout=layout, - physical_page_capacity=kv_cache.shape[0] // page_size, - prefix_spans=prefix_spans, - current_spans=current_spans, - kind="token_kv", - current_page_writer_ranks=( - current_page_writer_ranks if use_symm else None - ), - ) - dense_rows = int(slot_remap.dense_num_pages) * page_size - - # Path selection is a capability decision, not error handling: the - # group-agreed probe (cached per pool tensor) tells us whether peer-IPC - # is available everywhere. After a successful probe, a failing gather - # is a bug — let it raise. - ipc_state = ( - _agreed_tai_ipc_peer_ptrs(kv_cache, layout) - if plan.num_prefix_slots > 0 - else None - ) - if use_symm: - # Symm requires the prefix IPC capability (same transport) and the - # registered staging; registration is collective and happens here, - # at the uniform first-use point. - use_symm = ipc_state is not None and _symm_staging_ready_or_register( - layout=layout, - kv_cache=kv_cache, - page_size=page_size, - ) - gathered = ipc_state is not None - kv_page_nbytes = _token_kv_page_nbytes(kv_cache, page_size) - if gathered: - kernels, peer_ptrs = ipc_state - dense_kv_cache = acquire_dense_buffer( - device=kv_cache.device, - shape=(dense_rows, *kv_cache.shape[1:]), - dtype=kv_cache.dtype, - layer_id=layer_id, - kind="token_kv", - ) - if not use_symm: - kernels.materialize_cuda_ipc_peer_pages_slot_dense( - peer_ptrs, - dense_kv_cache, - plan.prefix_owner_ranks, - plan.prefix_src_pages, - page_nbytes=kv_page_nbytes, - ) - else: - dense_kv_cache = kv_cache.new_zeros((dense_rows, *kv_cache.shape[1:])) - if prefix_spans and _should_fail_fast_compose_v2_dense_fallback( - dense_kv_cache - ): - _raise_compose_v2_dense_fallback_required( - "token_kv_dense_fallback", - cp_rank=layout.cp_rank, - cp_size=layout.cp_size, - layer_id=layer_id, - prefix_spans=prefix_spans, - current_spans=current_spans, - dense_shape=tuple(dense_kv_cache.shape), - kv_dtype=kv_cache.dtype, - ) - 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, - ) - - if use_symm: - # Step B: write current rows straight into this round's staging span - # (one index_copy; the loc remap and slack masking are layer-invariant - # and cached on the plan), then barrier + ONE slot-dense gather for - # prefix AND current pages. The prefix-only gather above was skipped - # — the mega gather covers it. - staging = get_compose_staging(kv_cache.device) - parity, staging_span = _symm_begin_current_staging( - staging=staging, - plan=plan, - kind="token_kv", - layer_id=layer_id, - page_nbytes=kv_page_nbytes, - ) - fill_state = _get_or_build_symm_token_fill_state( - plan=plan, - slot_remap=slot_remap, - layout=layout, - kv_cache=kv_cache, - logical_locs=logical_locs, - current_locs=current_locs, - loc_req_id=loc_req_id, - current_req_id=current_req_id, - page_size=page_size, - ) - num_rows = int(fill_state.staging_current_rows.numel()) - if num_rows > 0: - # Byte view on both sides: index_copy_ has no fp8 CUDA kernel, - # and the copy is dtype-agnostic anyway (whole token rows). - staging_rows = staging_span.view( - (int(plan.num_current_pages) + 1) * page_size, - kv_page_nbytes // page_size, - ) - staging_rows.index_copy_( - 0, - fill_state.staging_current_rows, - current_kv_cache[:num_rows] - .reshape(num_rows, -1) - .view(torch.uint8), - ) - _symm_barrier_and_gather_all( - kernels=kernels, - pool_peer_ptrs=peer_ptrs, - dense_buffer=dense_kv_cache, - plan=plan, - layout=layout, - staging=staging, - kind="token_kv", - parity=parity, - page_nbytes=kv_page_nbytes, - ) - mixed_kv_cache = dense_kv_cache - mixed_locs = fill_state.mixed_locs - log_cp_shared_kv_bs_gt1_timing( - "mla_partial_current_compose_v2", - timing_start, - "cp_rank=%s layer=%s cp_size=%s total_slots=%s dense_rows=%s " - "prefix_span_pages=%s current_pages=%s symm=1 kv_dtype=%s", - layout.cp_rank, - layer_id, - layout.cp_size, - plan.total_slots, - int(mixed_kv_cache.shape[0]), - plan.num_prefix_slots, - plan.num_current_pages, - kv_cache.dtype, - ) - return mixed_kv_cache, mixed_locs - - logical_locs = filter_locs_mappable_to_physical_pool( - logical_locs=logical_locs, - layout=layout, - physical_token_capacity=kv_cache.shape[0], - ) - dense_locs = remap_logical_locs_to_shared_token_slot_dense_locs( - logical_locs, - slot_remap=slot_remap, - page_size=page_size, - loc_req_id=loc_req_id, - ) - mixed_kv_cache, mixed_locs, _ = fill_current_kv_page_slots_and_remap_locs( - dense_kv_cache=dense_kv_cache, - materialized_dense_locs=dense_locs, - current_kv_cache=current_kv_cache, - logical_locs=logical_locs, - current_locs=current_locs, - page_inverse=slot_remap.page_inverse, - page_size=page_size, - current_req_id=current_req_id, - mask_non_current_in_current_pages=True, - ) - - if gathered: - if plan.num_current_pages > 0: - _reduce_current_pages_compact( - mixed_kv_cache, - int(slot_remap.dense_num_pages), - plan.current_dense_pages, - layout, - nvtx_source=f"{nvtx_source}.v2_current_compact", - layer_id=layer_id, - ) - elif prefix_spans: - # Unreduced prefix rows present: one reduce must cover them too. - _all_reduce_materialized_buffer( - mixed_kv_cache, - layout.cp_size, - nvtx_source=f"{nvtx_source}.v2_full", - nvtx_layer_id=layer_id, - nvtx_cp_rank=layout.cp_rank, - ) - elif plan.num_current_pages > 0: - _reduce_current_pages_compact( - mixed_kv_cache, - int(slot_remap.dense_num_pages), - plan.current_dense_pages, - layout, - nvtx_source=f"{nvtx_source}.v2_current_compact", - layer_id=layer_id, - ) - - log_cp_shared_kv_bs_gt1_timing( - "mla_partial_current_compose_v2", - timing_start, - "cp_rank=%s layer=%s cp_size=%s total_slots=%s dense_rows=%s " - "prefix_span_pages=%s current_pages=%s gathered_by_ipc=%s kv_dtype=%s", - layout.cp_rank, - layer_id, - layout.cp_size, - plan.total_slots, - int(mixed_kv_cache.shape[0]), - plan.num_prefix_slots, - plan.num_current_pages, - gathered, - kv_cache.dtype, - ) - return mixed_kv_cache, mixed_locs - - def materialize_prefix_and_reuse_current_kv_page_slots( *, kv_cache: torch.Tensor, @@ -5040,7 +4824,6 @@ def materialize_prefix_and_reuse_current_kv_page_slots( logical_locs_row_ids: torch.Tensor | None = None, layer_id: int | None = None, nvtx_source: str = "mla.partial_current_sync", - current_page_writer_ranks: list[int] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """Synchronously compose prefix materialization with current KV rows. @@ -5111,45 +4894,26 @@ def materialize_prefix_and_reuse_current_kv_page_slots( else [] ) - if cp_shared_kv_compose_v2_enabled() and layout.cp_size > 1: - return _compose_token_kv_partial_current_v2( - kv_cache=kv_cache, - logical_locs=logical_locs, - current_kv_cache=current_kv_cache, - current_locs=current_locs, - slot_remap=slot_remap, - layout=layout, - page_size=page_size, - prefix_spans=prefix_spans, - current_spans=_resolve_partial_current_spans( - current_slot_spans=current_slot_spans, - prefix_slot_span=prefix_slot_span, - prefix_slot_spans=prefix_slot_spans, - prefix_pages=prefix_pages, - total_slots=total_slots, - ), - loc_req_id=loc_req_id, - current_req_id=current_req_id, - layer_id=layer_id, - nvtx_source=nvtx_source, - timing_start=timing_start, - current_page_writer_ranks=current_page_writer_ranks, - ) - dense_kv_cache = kv_cache.new_zeros( (slot_remap.dense_num_pages * page_size, *kv_cache.shape[1:]) ) 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( + if prefix_spans: + materialized_by_ipc = _try_tai_ipc_materialize_token_kv_page_slot_spans_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, + spans=prefix_spans, + ) + if not materialized_by_ipc and prefix_spans and _should_fail_fast_tai_ipc_materialize(dense_kv_cache): + _raise_tai_ipc_materialize_required( + "token_prefix_ipc_unavailable", + cp_rank=layout.cp_rank, + cp_size=layout.cp_size, + spans=prefix_spans, + dense_shape=tuple(dense_kv_cache.shape), ) if not materialized_by_ipc: for prefix_start_slot, prefix_end_slot in prefix_spans: @@ -5201,6 +4965,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots( mask_non_current_in_current_pages=True, ) if layout.cp_size > 1: + current_materialized_by_ipc = False if current_slot_spans is None: if prefix_slot_span is not None or prefix_slot_spans is not None: raise ValueError( @@ -5212,23 +4977,45 @@ def materialize_prefix_and_reuse_current_kv_page_slots( if int(prefix_pages) < total_slots else [] ) - for current_start_slot, current_end_slot in _merge_slot_spans( - current_slot_spans + merged_current_spans_for_reduce = _merge_slot_spans(current_slot_spans) + if merged_current_spans_for_reduce: + current_materialized_by_ipc = ( + _try_tai_ipc_materialize_current_token_kv_page_slot_spans_into( + dense_kv_cache=mixed_kv_cache, + slot_logical_pages=slot_remap.slot_logical_pages, + layout=layout, + page_size=page_size, + spans=merged_current_spans_for_reduce, + ) + ) + if ( + not current_materialized_by_ipc + and merged_current_spans_for_reduce + and _should_fail_fast_tai_ipc_materialize(mixed_kv_cache) ): - current_rows = slot_range_to_token_slice( - page_size, - current_start_slot, - current_end_slot, - ) - _all_reduce_materialized_buffer_range( - mixed_kv_cache, - layout.cp_size, - current_rows.start, - current_rows.stop, - nvtx_source=f"{nvtx_source}.current", - nvtx_layer_id=layer_id, - nvtx_cp_rank=layout.cp_rank, + _raise_tai_ipc_materialize_required( + "token_current_ipc_unavailable", + cp_rank=layout.cp_rank, + cp_size=layout.cp_size, + spans=merged_current_spans_for_reduce, + dense_shape=tuple(mixed_kv_cache.shape), ) + if not current_materialized_by_ipc: + for current_start_slot, current_end_slot in merged_current_spans_for_reduce: + current_rows = slot_range_to_token_slice( + page_size, + current_start_slot, + current_end_slot, + ) + _all_reduce_materialized_buffer_range( + mixed_kv_cache, + layout.cp_size, + current_rows.start, + current_rows.stop, + nvtx_source=f"{nvtx_source}.current", + nvtx_layer_id=layer_id, + nvtx_cp_rank=layout.cp_rank, + ) merged_current_spans = ( _merge_slot_spans(current_slot_spans) if current_slot_spans is not None @@ -5255,212 +5042,6 @@ def materialize_prefix_and_reuse_current_kv_page_slots( return mixed_kv_cache, mixed_locs -def _compose_index_partial_current_v2( - *, - page_buffer: torch.Tensor, - current_index_k: torch.Tensor, - current_index_scale: torch.Tensor, - current_locs: torch.Tensor, - slot_remap: SharedPagedBufferSlotRemap, - layout: CpSharedKVLayout, - page_size: int, - index_head_dim: int, - prefix_spans: list[tuple[int, int]], - current_spans: list[tuple[int, int]], - current_req_id: torch.Tensor, - layer_id: int | None, - nvtx_source: str, - timing_start: object, - current_page_writer_ranks: list[int] | None = None, -) -> tuple[torch.Tensor, torch.Tensor]: - """Step A/B compose for the indexer page buffer (see token-KV variant). - - The index compose never registers the symm staging itself (the token-KV - compose does, knowing the exact KV page bytes); it uses symm only once - registration already happened — deterministic across ranks because the - layer order is identical everywhere. - """ - - use_symm = ( - cp_shared_kv_compose_symm_enabled() - and current_page_writer_ranks is not None - and layer_id is not None - and get_compose_staging(page_buffer.device).registered - ) - plan = get_or_build_compose_plan( - slot_remap=slot_remap, - layout=layout, - physical_page_capacity=page_buffer.shape[0], - prefix_spans=prefix_spans, - current_spans=current_spans, - kind="index", - current_page_writer_ranks=( - current_page_writer_ranks if use_symm else None - ), - ) - dense_num_pages = int(slot_remap.dense_num_pages) - - ipc_state = ( - _agreed_tai_ipc_peer_ptrs(page_buffer, layout) - if plan.num_prefix_slots > 0 - else None - ) - gathered = ipc_state is not None - use_symm = use_symm and gathered - if gathered: - kernels, peer_ptrs = ipc_state - dense_page_buffer = acquire_dense_buffer( - device=page_buffer.device, - shape=(dense_num_pages, *page_buffer.shape[1:]), - dtype=page_buffer.dtype, - layer_id=layer_id, - kind="index", - ) - if not use_symm: - kernels.materialize_cuda_ipc_peer_pages_slot_dense( - peer_ptrs, - dense_page_buffer, - plan.prefix_owner_ranks, - plan.prefix_src_pages, - page_nbytes=_page_nbytes_from_page_tensor(page_buffer), - ) - else: - dense_page_buffer = page_buffer.new_zeros( - (dense_num_pages, *page_buffer.shape[1:]) - ) - if prefix_spans and _should_fail_fast_compose_v2_dense_fallback( - dense_page_buffer - ): - _raise_compose_v2_dense_fallback_required( - "index_dense_fallback", - cp_rank=layout.cp_rank, - cp_size=layout.cp_size, - layer_id=layer_id, - prefix_spans=prefix_spans, - current_spans=current_spans, - dense_shape=tuple(dense_page_buffer.shape), - index_dtype=page_buffer.dtype, - ) - 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, - ) - - if use_symm: - # Step B: fill current index rows straight into this round's staging - # span (the unchanged fill kernel — its write destinations all come - # from the page inverse, which the plan remapped to staging slots), - # then barrier + ONE slot-dense gather for prefix AND current pages. - index_page_nbytes = _page_nbytes_from_page_tensor(page_buffer) - staging = get_compose_staging(page_buffer.device) - parity, staging_span = _symm_begin_current_staging( - staging=staging, - plan=plan, - kind="index", - layer_id=layer_id, - page_nbytes=index_page_nbytes, - ) - if plan.num_current_pages > 0: - staging_pages = staging_span.view(page_buffer.dtype).reshape( - int(plan.num_current_pages) + 1, *page_buffer.shape[1:] - ) - fill_current_index_page_slots( - dense_page_buffer=staging_pages, - current_index_k=current_index_k, - current_index_scale=current_index_scale, - current_locs=current_locs, - page_inverse=plan.staging_page_inverse, - page_size=page_size, - index_head_dim=index_head_dim, - current_req_id=current_req_id, - ) - _symm_barrier_and_gather_all( - kernels=kernels, - pool_peer_ptrs=peer_ptrs, - dense_buffer=dense_page_buffer, - plan=plan, - layout=layout, - staging=staging, - kind="index", - parity=parity, - page_nbytes=index_page_nbytes, - ) - log_cp_shared_kv_bs_gt1_timing( - "index_partial_current_compose_v2", - timing_start, - "cp_rank=%s layer=%s cp_size=%s total_slots=%s dense_pages=%s " - "prefix_span_pages=%s current_pages=%s symm=1", - layout.cp_rank, - layer_id, - layout.cp_size, - plan.total_slots, - dense_num_pages, - plan.num_prefix_slots, - plan.num_current_pages, - ) - return dense_page_buffer, slot_remap.dense_pages - - dense_page_buffer = fill_current_index_page_slots( - dense_page_buffer=dense_page_buffer, - current_index_k=current_index_k, - current_index_scale=current_index_scale, - current_locs=current_locs, - page_inverse=slot_remap.page_inverse, - page_size=page_size, - index_head_dim=index_head_dim, - current_req_id=current_req_id, - ) - - if gathered: - if plan.num_current_pages > 0: - _reduce_current_pages_compact( - dense_page_buffer, - dense_num_pages, - plan.current_dense_pages, - layout, - nvtx_source=f"{nvtx_source}.v2_current_compact", - layer_id=layer_id, - ) - elif prefix_spans: - _all_reduce_materialized_buffer( - dense_page_buffer, - layout.cp_size, - nvtx_source=f"{nvtx_source}.v2_full", - nvtx_layer_id=layer_id, - nvtx_cp_rank=layout.cp_rank, - ) - elif plan.num_current_pages > 0: - _reduce_current_pages_compact( - dense_page_buffer, - dense_num_pages, - plan.current_dense_pages, - layout, - nvtx_source=f"{nvtx_source}.v2_current_compact", - layer_id=layer_id, - ) - - log_cp_shared_kv_bs_gt1_timing( - "index_partial_current_compose_v2", - timing_start, - "cp_rank=%s layer=%s cp_size=%s total_slots=%s dense_pages=%s " - "prefix_span_pages=%s current_pages=%s gathered_by_ipc=%s", - layout.cp_rank, - layer_id, - layout.cp_size, - plan.total_slots, - dense_num_pages, - plan.num_prefix_slots, - plan.num_current_pages, - gathered, - ) - return dense_page_buffer, slot_remap.dense_pages - - def materialize_prefix_and_reuse_current_index_page_slots( *, page_buffer: torch.Tensor, @@ -5478,7 +5059,6 @@ def materialize_prefix_and_reuse_current_index_page_slots( current_slot_spans: list[tuple[int, int]] | None = None, layer_id: int | None = None, nvtx_source: str = "index.partial_current_sync", - current_page_writer_ranks: list[int] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """Synchronously compose prefix index materialization with current index rows.""" timing_start = cp_shared_kv_bs_gt1_timing_start() @@ -5536,44 +5116,25 @@ def materialize_prefix_and_reuse_current_index_page_slots( else [] ) - if cp_shared_kv_compose_v2_enabled() and layout.cp_size > 1: - return _compose_index_partial_current_v2( - page_buffer=page_buffer, - current_index_k=current_index_k, - current_index_scale=current_index_scale, - current_locs=current_locs, - slot_remap=slot_remap, - layout=layout, - page_size=page_size, - index_head_dim=index_head_dim, - prefix_spans=prefix_spans, - current_spans=_resolve_partial_current_spans( - current_slot_spans=current_slot_spans, - prefix_slot_span=prefix_slot_span, - prefix_slot_spans=prefix_slot_spans, - prefix_pages=prefix_pages, - total_slots=total_slots, - ), - current_req_id=current_req_id, - layer_id=layer_id, - nvtx_source=nvtx_source, - timing_start=timing_start, - current_page_writer_ranks=current_page_writer_ranks, - ) - dense_page_buffer = page_buffer.new_zeros( (slot_remap.dense_num_pages, *page_buffer.shape[1:]) ) 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( + if prefix_spans: + materialized_by_ipc = _try_tai_ipc_materialize_paged_buffer_page_slot_spans_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, + spans=prefix_spans, + ) + if not materialized_by_ipc and prefix_spans and _should_fail_fast_tai_ipc_materialize(dense_page_buffer): + _raise_tai_ipc_materialize_required( + "index_prefix_ipc_unavailable", + cp_rank=layout.cp_rank, + cp_size=layout.cp_size, + spans=prefix_spans, + dense_shape=tuple(dense_page_buffer.shape), ) if not materialized_by_ipc: for prefix_start_slot, prefix_end_slot in prefix_spans: @@ -5606,6 +5167,7 @@ def materialize_prefix_and_reuse_current_index_page_slots( current_req_id=current_req_id, ) if layout.cp_size > 1: + current_materialized_by_ipc = False if current_slot_spans is None: if prefix_slot_span is not None or prefix_slot_spans is not None: raise ValueError( @@ -5617,22 +5179,43 @@ def materialize_prefix_and_reuse_current_index_page_slots( if int(prefix_pages) < total_slots else [] ) - for current_start_slot, current_end_slot in _merge_slot_spans( - current_slot_spans + merged_current_spans_for_reduce = _merge_slot_spans(current_slot_spans) + if merged_current_spans_for_reduce: + current_materialized_by_ipc = ( + _try_tai_ipc_materialize_current_paged_buffer_page_slot_spans_into( + dense_page_buffer=dense_page_buffer, + slot_logical_pages=slot_remap.slot_logical_pages, + layout=layout, + spans=merged_current_spans_for_reduce, + ) + ) + if ( + not current_materialized_by_ipc + and merged_current_spans_for_reduce + and _should_fail_fast_tai_ipc_materialize(dense_page_buffer) ): - current_pages = slot_range_to_page_slice( - current_start_slot, - current_end_slot, - ) - _all_reduce_materialized_buffer_range( - dense_page_buffer, - layout.cp_size, - current_pages.start, - current_pages.stop, - nvtx_source=f"{nvtx_source}.current", - nvtx_layer_id=layer_id, - nvtx_cp_rank=layout.cp_rank, + _raise_tai_ipc_materialize_required( + "index_current_ipc_unavailable", + cp_rank=layout.cp_rank, + cp_size=layout.cp_size, + spans=merged_current_spans_for_reduce, + dense_shape=tuple(dense_page_buffer.shape), ) + if not current_materialized_by_ipc: + for current_start_slot, current_end_slot in merged_current_spans_for_reduce: + current_pages = slot_range_to_page_slice( + current_start_slot, + current_end_slot, + ) + _all_reduce_materialized_buffer_range( + dense_page_buffer, + layout.cp_size, + current_pages.start, + current_pages.stop, + nvtx_source=f"{nvtx_source}.current", + nvtx_layer_id=layer_id, + nvtx_cp_rank=layout.cp_rank, + ) merged_current_spans = ( _merge_slot_spans(current_slot_spans) if current_slot_spans is not None @@ -6179,6 +5762,14 @@ def materialize_shared_token_kv_buffer( ), ) + if not materialized_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_kv_cache): + _raise_tai_ipc_materialize_required( + "token_full_ipc_unavailable", + cp_rank=layout.cp_rank, + cp_size=layout.cp_size, + dense_shape=tuple(dense_kv_cache.shape), + logical_pages=int(materialized_logical_pages.numel()), + ) if not materialized_by_ipc: dense_kv_cache = _all_reduce_materialized_buffer( dense_kv_cache, @@ -6289,6 +5880,14 @@ def materialize_shared_paged_buffer( ), ) + if not materialized_by_ipc and _should_fail_fast_tai_ipc_materialize(dense_page_buffer): + _raise_tai_ipc_materialize_required( + "index_full_ipc_unavailable", + cp_rank=layout.cp_rank, + cp_size=layout.cp_size, + dense_shape=tuple(dense_page_buffer.shape), + logical_pages=int(materialized_logical_pages.numel()), + ) if not materialized_by_ipc: dense_page_buffer = _all_reduce_materialized_buffer( dense_page_buffer, 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 bf23ad631..f09110f0b 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 @@ -155,7 +155,7 @@ class _FakeExtendForwardMode: class TestCpSharedKVRuntimeHelpers(unittest.TestCase): - def test_mla_prefetch_materializes_on_current_stream_and_reduces_on_prefetch_stream( + def test_mla_prefetch_waits_l2_l1_on_materialize_stream_and_reduces_on_prefetch_stream( self, ): from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch @@ -241,7 +241,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): token_to_kv_pool=pool, ) - self.assertEqual(pool.prefetch_getter_streams, [(1, "prefetch")]) + self.assertEqual(pool.prefetch_getter_streams, [(1, "current")]) self.assertEqual( calls, [("materialize", "current"), ("reduce", "prefetch", "prefetch")], @@ -256,7 +256,146 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): ) self.assertEqual(prefetch_stream.waited, ["current"]) - def test_index_prefetch_materializes_on_current_stream_and_reduces_on_prefetch_stream( + def test_mla_prefetch_ipc_spans_skip_prefix_reduce(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + class FakeStream: + def __init__(self, name): + self.name = name + self.waited = [] + + def wait_stream(self, stream): + self.waited.append(stream.name) + + class FakeEvent: + def __init__(self): + self.recorded = [] + + def record(self, stream): + self.recorded.append(stream.name) + + class FakePool: + start_layer = 0 + page_size = 4 + kv_buffer = [object(), object(), object()] + + def __init__(self): + self.kv_cache = torch.zeros((64, 1, 2), dtype=torch.float32) + + def get_key_buffer_for_prefetch(self, layer_id, stream): + return self.kv_cache + + current_stream = FakeStream("current") + prefetch_stream = FakeStream("prefetch") + ipc_calls = [] + + def record_ipc(**kwargs): + ipc_calls.append(kwargs["spans"]) + return True + + prefetcher = prefetch.CpSharedKVMlaPrefetcher( + layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0), + page_size=4, + prefix_pages=3, + prefix_slot_spans=[(0, 2), (4, 5)], + current_slot_spans=[(2, 4), (5, 6)], + slot_logical_pages=torch.tensor([0, 1, 2, 3, 8, 9], dtype=torch.int64), + page_inverse=torch.tensor([0, 1, 2, 3, -1, -1, -1, -1, 5, 6]), + dense_num_pages=7, + stream=prefetch_stream, + ) + + with patch.object( + prefetch.torch.cuda, "current_stream", return_value=current_stream + ), patch.object( + prefetch.torch.cuda, "Event", side_effect=FakeEvent + ), patch.object( + prefetch, + "_try_tai_ipc_materialize_token_kv_page_slot_spans_into", + side_effect=record_ipc, + ), patch.object( + prefetch, + "_all_reduce_materialized_buffer_ranges_async", + side_effect=AssertionError("IPC prefix path must not all-reduce"), + ): + prefetcher.start_next_layer_prefix( + next_layer_id=1, + token_to_kv_pool=FakePool(), + ) + + self.assertEqual(ipc_calls, [[(0, 2), (4, 5)]]) + handle = prefetcher.handles[1] + self.assertEqual(handle.event.recorded, ["current"]) + self.assertEqual(prefetch_stream.waited, []) + + def test_mla_prefetch_consume_ipc_suffix_skips_suffix_reduce(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + class FakeCurrentStream: + def __init__(self): + self.events = [] + + def wait_event(self, event): + self.events.append(event) + + fake_event = object() + current_stream = FakeCurrentStream() + dense_kv = torch.zeros((28, 1, 2), dtype=torch.float32) + prefetcher = prefetch.CpSharedKVMlaPrefetcher( + layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0), + page_size=4, + prefix_pages=3, + prefix_slot_spans=[(0, 2), (4, 5)], + current_slot_spans=[(2, 4), (5, 6)], + slot_logical_pages=torch.tensor([0, 1, 2, 3, 8, 9], dtype=torch.int64), + page_inverse=torch.tensor([0, 1, 2, 3, -1, -1, -1, -1, 5, 6]), + dense_num_pages=7, + stream=object(), + ) + prefetcher.handles[0] = prefetch.CpSharedKVMlaPrefetchHandle( + layer_id=0, + dense_kv_cache=dense_kv, + prefix_rows=slice(4, 12), + event=fake_event, + ) + ipc_calls = [] + + def record_ipc(**kwargs): + ipc_calls.append(kwargs["spans"]) + return True + + with patch.object( + prefetch.torch.cuda, "current_stream", return_value=current_stream + ), patch.object( + prefetch, + "_try_tai_ipc_materialize_token_kv_page_slot_spans_into", + side_effect=record_ipc, + ), patch.object( + prefetch, + "_materialize_local_token_kv_page_slot_spans_into", + side_effect=AssertionError("IPC suffix path must not local materialize"), + ), patch.object( + prefetch, + "_all_reduce_materialized_buffer_range", + side_effect=AssertionError("IPC suffix path must not all-reduce"), + ), patch.object( + prefetch, + "remap_logical_locs_to_slot_dense_locs_optimized", + return_value=torch.tensor([12, 20], dtype=torch.int64), + ): + result = prefetcher.consume( + layer_id=0, + kv_cache=torch.zeros((64, 1, 2), dtype=torch.float32), + logical_locs=torch.tensor([8, 20], dtype=torch.int64), + ) + + self.assertIsNotNone(result) + self.assertEqual(ipc_calls, [[(2, 4), (5, 6)]]) + self.assertEqual(current_stream.events, [fake_event]) + + def test_index_prefetch_waits_l2_l1_on_materialize_stream_and_reduces_on_prefetch_stream( self, ): from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch @@ -343,7 +482,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): token_to_kv_pool=pool, ) - self.assertEqual(pool.prefetch_getter_streams, [(1, "prefetch")]) + self.assertEqual(pool.prefetch_getter_streams, [(1, "current")]) self.assertEqual( calls, [("materialize", "current"), ("reduce", "prefetch", "prefetch")], @@ -358,6 +497,143 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): ) self.assertEqual(prefetch_stream.waited, ["current"]) + def test_index_prefetch_ipc_spans_skip_prefix_reduce(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + class FakeStream: + def __init__(self, name): + self.name = name + self.waited = [] + + def wait_stream(self, stream): + self.waited.append(stream.name) + + class FakeEvent: + def __init__(self): + self.recorded = [] + + def record(self, stream): + self.recorded.append(stream.name) + + class FakePool: + start_layer = 0 + page_size = 4 + kv_buffer = [object(), object(), object()] + + def __init__(self): + self.page_buffer = torch.zeros((64, 3), dtype=torch.uint8) + + def get_index_k_with_scale_buffer_for_prefetch(self, layer_id, stream): + return self.page_buffer + + current_stream = FakeStream("current") + prefetch_stream = FakeStream("prefetch") + ipc_calls = [] + + def record_ipc(**kwargs): + ipc_calls.append(kwargs["spans"]) + return True + + prefetcher = prefetch.CpSharedKVIndexPrefetcher( + layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0), + prefix_pages=3, + prefix_slot_spans=[(0, 2), (4, 5)], + current_slot_spans=[(2, 4), (5, 6)], + slot_logical_pages=torch.tensor([0, 1, 2, 3, 8, 9], dtype=torch.int64), + page_inverse=torch.tensor([0, 1, 2, 3, -1, -1, -1, -1, 5, 6]), + dense_num_pages=7, + stream=prefetch_stream, + ) + + with patch.object( + prefetch.torch.cuda, "current_stream", return_value=current_stream + ), patch.object( + prefetch.torch.cuda, "Event", side_effect=FakeEvent + ), patch.object( + prefetch, + "_try_tai_ipc_materialize_paged_buffer_page_slot_spans_into", + side_effect=record_ipc, + ), patch.object( + prefetch, + "_all_reduce_materialized_buffer_ranges_async", + side_effect=AssertionError("IPC prefix path must not all-reduce"), + ): + prefetcher.start_next_layer_prefix( + next_layer_id=1, + token_to_kv_pool=FakePool(), + ) + + self.assertEqual(ipc_calls, [[(0, 2), (4, 5)]]) + handle = prefetcher.handles[1] + self.assertEqual(handle.event.recorded, ["current"]) + self.assertEqual(prefetch_stream.waited, []) + + def test_index_prefetch_consume_ipc_suffix_skips_suffix_reduce(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + class FakeCurrentStream: + def __init__(self): + self.events = [] + + def wait_event(self, event): + self.events.append(event) + + fake_event = object() + current_stream = FakeCurrentStream() + dense_pages = torch.zeros((7, 3), dtype=torch.uint8) + prefetcher = prefetch.CpSharedKVIndexPrefetcher( + layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0), + prefix_pages=3, + prefix_slot_spans=[(0, 2), (4, 5)], + current_slot_spans=[(2, 4), (5, 6)], + slot_logical_pages=torch.tensor([0, 1, 2, 3, 8, 9], dtype=torch.int64), + page_inverse=torch.tensor([0, 1, 2, 3, -1, -1, -1, -1, 5, 6]), + dense_num_pages=7, + stream=object(), + ) + prefetcher.handles[0] = prefetch.CpSharedKVIndexPrefetchHandle( + layer_id=0, + dense_page_buffer=dense_pages, + prefix_rows=slice(1, 3), + event=fake_event, + ) + ipc_calls = [] + + def record_ipc(**kwargs): + ipc_calls.append(kwargs["spans"]) + return True + + with patch.object( + prefetch.torch.cuda, "current_stream", return_value=current_stream + ), patch.object( + prefetch, + "_try_tai_ipc_materialize_paged_buffer_page_slot_spans_into", + side_effect=record_ipc, + ), patch.object( + prefetch, + "_materialize_local_paged_buffer_page_slot_spans_into", + side_effect=AssertionError("IPC suffix path must not local materialize"), + ), patch.object( + prefetch, + "_all_reduce_materialized_buffer_range", + side_effect=AssertionError("IPC suffix path must not all-reduce"), + ), patch.object( + prefetch, + "remap_logical_pages_to_slot_dense_pages", + return_value=torch.tensor([[3, 6]], dtype=torch.int64), + ): + result = prefetcher.consume( + layer_id=0, + page_buffer=torch.zeros((64, 3), dtype=torch.uint8), + logical_pages=torch.tensor([[2, 9]], dtype=torch.int64), + ) + + self.assertIsNotNone(result) + self.assertEqual(ipc_calls, [[(2, 4), (5, 6)]]) + self.assertEqual(current_stream.events, [fake_event]) + def test_mla_pool_prefetch_getter_orders_layer_transfer_on_prefetch_stream(self): index_accessor_stub = types.ModuleType( "sglang.srt.layers.attention.nsa.index_buf_accessor" @@ -1130,7 +1406,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): joined = "\n".join(logs.output) self.assertIn("[CP_SHARED_KV_FALLBACK][tai_ipc_materialize]", joined) - self.assertIn("start_slot_nonzero", joined) + self.assertIn("dense_token_tensor_unsupported", joined) def test_mla_prefetch_sync_compose_paths_log_warning_in_source(self): from pathlib import Path @@ -1759,6 +2035,393 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): [(0, 2), (4, 5)], ) + def test_mla_prefetch_create_batch_uses_exact_prefix_and_current_spans(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + + class Mode: + def is_context_parallel_extend(self): + return True + + page_size = 4 + logical_pages = torch.tensor( + [ + [1, 2, 5, 0], + [9, 11, 12, 13], + ], + dtype=torch.int64, + ) + stream = object() + kv_cache = torch.zeros((128, 1, 1), dtype=torch.float32) + remap = SimpleNamespace( + slot_logical_pages=logical_pages.reshape(-1), + page_inverse=torch.zeros((2, 32), dtype=torch.int64), + slot_sorted_logical_pages_by_row=None, + slot_sorted_dense_pages_by_row=None, + dense_num_pages=9, + ) + forward_batch = SimpleNamespace( + uses_cp_shared_kv=True, + hisparse_coordinator=None, + forward_mode=Mode(), + batch_size=2, + token_to_kv_pool=SimpleNamespace(page_size=page_size, start_layer=0), + cp_shared_kv_layout=SimpleNamespace(cp_size=2, cp_rank=0), + extend_prefix_lens_cpu=[8, 4], + extend_seq_lens_cpu=[2, 7], + ) + metadata = SimpleNamespace( + real_page_table=logical_pages, + page_table_1=torch.zeros((2, 4), dtype=torch.int32), + ) + + with patch.object( + prefetch, "cp_shared_kv_mla_prefetch_enabled", return_value=True + ), patch.object( + prefetch, "cp_shared_kv_debug_enabled", return_value=False + ), patch.object( + prefetch.torch.cuda, "is_available", return_value=True + ), patch.object( + prefetch, "_is_cuda_stream_capturing", return_value=False + ), patch.object( + prefetch, "is_nsa_prefill_cp_in_seq_split", return_value=True + ), patch.object( + prefetch, + "cp_shared_kv_mla_prefetch_min_prefix_pages", + return_value=0, + ), patch.object( + prefetch, + "cp_shared_kv_mla_prefetch_min_async_extend_tokens", + return_value=0, + ), patch.object( + prefetch, + "get_attention_cp_group", + return_value=SimpleNamespace(pynccl_comm=object()), + ), patch.object( + prefetch.torch.cuda, "Stream", return_value=stream + ), patch.object( + prefetch, "_prefetch_pool_get_key_buffer", return_value=kv_cache + ), patch.object( + prefetch, + "get_or_build_shared_token_kv_slot_remap", + return_value=remap, + ): + prefetcher = prefetch.CpSharedKVMlaPrefetcher.maybe_create( + forward_batch=forward_batch, + metadata=metadata, + topk_transform_is_paged=True, + ) + + self.assertIsNotNone(prefetcher) + self.assertEqual(prefetcher.prefix_slot_spans, [(0, 2), (4, 5)]) + self.assertEqual(prefetcher.current_slot_spans, [(2, 3), (5, 7)]) + self.assertEqual(prefetcher.prefix_page_count, 3) + self.assertEqual(prefetcher.current_page_count, 3) + + def test_mla_prefetch_batch_consume_reduces_exact_current_spans(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + 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=2, cp_rank=0) + logical_pages = torch.tensor( + [ + [1, 2, 5, 0], + [9, 11, 12, 13], + ], + dtype=torch.int64, + ) + slot_remap = runtime.build_shared_token_kv_slot_remap( + kv_cache=torch.zeros((128, 1, 1), dtype=torch.float32), + logical_locs=None, + remap_logical_pages=logical_pages, + layout=layout, + page_size=page_size, + ) + prefetcher = prefetch.CpSharedKVMlaPrefetcher( + layout=layout, + page_size=page_size, + prefix_pages=0, + prefix_slot_spans=[(0, 2), (4, 5)], + current_slot_spans=[(2, 3), (5, 7)], + slot_logical_pages=slot_remap.slot_logical_pages, + page_inverse=slot_remap.page_inverse, + slot_sorted_logical_pages_by_row=slot_remap.slot_sorted_logical_pages_by_row, + slot_sorted_dense_pages_by_row=slot_remap.slot_sorted_dense_pages_by_row, + dense_num_pages=slot_remap.dense_num_pages, + stream=object(), + ) + fake_event = object() + dense_kv = torch.zeros( + (slot_remap.dense_num_pages * page_size, 1, 1), dtype=torch.float32 + ) + prefetcher.handles[1] = prefetch.CpSharedKVMlaPrefetchHandle( + layer_id=1, + dense_kv_cache=dense_kv, + prefix_rows=slice(0, 0), + event=fake_event, + ) + prefetcher.pending_attention_handle = prefetcher.handles[1] + current_kv = torch.arange(100, 104, dtype=torch.float32).view(4, 1, 1) + current_locs = torch.tensor([20, 21, 44, 45], dtype=torch.int64) + current_req_id = torch.tensor([0, 0, 1, 1], dtype=torch.int64) + logical_locs = torch.tensor( + [ + [4, 8, 20, 21, -1, -1], + [36, 44, 45, -1, -1, -1], + ], + dtype=torch.int64, + ) + loc_req_id = torch.tensor( + [ + [0, 0, 0, 0, 0, 0], + [1, 1, 1, 1, 1, 1], + ], + dtype=torch.int64, + ) + reduce_ranges = [] + + class FakeCurrentStream: + def __init__(self): + self.events = [] + + def wait_event(self, event): + self.events.append(event) + + current_stream = FakeCurrentStream() + + def record_range_reduce(buffer, cp_size, start_row, end_row, **kwargs): + reduce_ranges.append((start_row, end_row, kwargs.get("nvtx_source"))) + return buffer + + with patch.object( + prefetch.torch.cuda, "current_stream", return_value=current_stream + ), patch.object( + prefetch, + "_all_reduce_materialized_buffer_range", + side_effect=record_range_reduce, + ): + mixed_kv, mixed_locs = prefetcher.consume_prefix_with_current( + layer_id=1, + kv_cache=torch.zeros((128, 1, 1), dtype=torch.float32), + logical_locs=logical_locs, + current_kv_cache=current_kv, + current_locs=current_locs, + loc_req_id=loc_req_id, + current_req_id=current_req_id, + ) + + self.assertEqual(current_stream.events, [fake_event]) + self.assertTrue(torch.equal(mixed_kv[12:14], current_kv[:2])) + self.assertTrue(torch.equal(mixed_kv[24:26], current_kv[2:])) + self.assertEqual( + mixed_locs.tolist(), + [[4, 8, 12, 13, -1, -1], [20, 24, 25, -1, -1, -1]], + ) + self.assertEqual( + reduce_ranges, + [ + (12, 16, "mla.prefetch_current"), + (24, 32, "mla.prefetch_current"), + ], + ) + + def test_index_prefetch_create_batch_uses_exact_prefix_and_current_spans(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + class Mode: + def is_context_parallel_extend(self): + return True + + page_size = 4 + logical_pages = torch.tensor( + [ + [1, 2, 5, 0], + [9, 11, 12, 13], + ], + dtype=torch.int64, + ) + stream = object() + page_buffer = torch.zeros((128, 32), dtype=torch.uint8) + remap = SimpleNamespace( + slot_logical_pages=logical_pages.reshape(-1), + page_inverse=torch.zeros((2, 32), dtype=torch.int64), + slot_sorted_logical_pages_by_row=None, + slot_sorted_dense_pages_by_row=None, + dense_pages=torch.tensor([[1, 2, 3, 0], [5, 6, 7, 8]], dtype=torch.int64), + dense_num_pages=9, + ) + forward_batch = SimpleNamespace( + uses_cp_shared_kv=True, + hisparse_coordinator=None, + forward_mode=Mode(), + batch_size=2, + token_to_kv_pool=SimpleNamespace(page_size=page_size, start_layer=0), + cp_shared_kv_layout=CpSharedKVLayout( + page_size=page_size, cp_size=2, cp_rank=0 + ), + extend_prefix_lens_cpu=[8, 4], + extend_seq_lens_cpu=[2, 7], + ) + metadata = SimpleNamespace( + real_page_table=logical_pages, + page_table_1=torch.zeros((2, 4), dtype=torch.int32), + ) + + with patch.object( + prefetch, "cp_shared_kv_mla_prefetch_enabled", return_value=True + ), patch.object( + prefetch, "cp_shared_kv_debug_enabled", return_value=False + ), patch.object( + prefetch.torch.cuda, "is_available", return_value=True + ), patch.object( + prefetch, "_is_cuda_stream_capturing", return_value=False + ), patch.object( + prefetch, "is_nsa_prefill_cp_in_seq_split", return_value=True + ), patch.object( + prefetch, + "cp_shared_kv_mla_prefetch_min_prefix_pages", + return_value=0, + ), patch.object( + prefetch, + "cp_shared_kv_mla_prefetch_min_async_extend_tokens", + return_value=0, + ), patch.object( + prefetch, + "get_attention_cp_group", + return_value=SimpleNamespace(pynccl_comm=object()), + ), patch.object( + prefetch.torch.cuda, "Stream", return_value=stream + ), patch.object( + prefetch, "_prefetch_pool_get_index_buffer", return_value=page_buffer + ), patch.object( + prefetch, + "get_or_build_shared_paged_buffer_slot_remap", + return_value=remap, + ): + prefetcher = prefetch.CpSharedKVIndexPrefetcher.maybe_create( + forward_batch=forward_batch, + metadata=metadata, + topk_transform_is_paged=True, + ) + + self.assertIsNotNone(prefetcher) + self.assertEqual(prefetcher.prefix_slot_spans, [(0, 2), (4, 5)]) + self.assertEqual(prefetcher.current_slot_spans, [(2, 3), (5, 7)]) + self.assertEqual(prefetcher.prefix_page_count, 3) + self.assertEqual(prefetcher.current_page_count, 3) + + def test_index_prefetch_batch_consume_reduces_exact_current_spans(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + 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 + index_head_dim = 4 + scale_bytes = 4 + page_bytes = page_size * index_head_dim + page_size * scale_bytes + layout = CpSharedKVLayout(page_size=page_size, cp_size=2, cp_rank=0) + logical_pages = torch.tensor( + [ + [1, 2, 5, 0], + [9, 11, 12, 13], + ], + dtype=torch.int64, + ) + slot_remap = runtime.build_shared_paged_buffer_slot_remap( + torch.zeros((128, page_bytes), dtype=torch.uint8), + logical_pages, + layout, + ) + prefetcher = prefetch.CpSharedKVIndexPrefetcher( + layout=layout, + prefix_pages=0, + prefix_slot_spans=[(0, 2), (4, 5)], + current_slot_spans=[(2, 3), (5, 7)], + slot_logical_pages=slot_remap.slot_logical_pages, + page_inverse=slot_remap.page_inverse, + slot_sorted_logical_pages_by_row=slot_remap.slot_sorted_logical_pages_by_row, + slot_sorted_dense_pages_by_row=slot_remap.slot_sorted_dense_pages_by_row, + dense_num_pages=slot_remap.dense_num_pages, + stream=object(), + ) + fake_event = object() + dense_page_buffer = torch.zeros( + (slot_remap.dense_num_pages, page_bytes), dtype=torch.uint8 + ) + prefetcher.handles[1] = prefetch.CpSharedKVIndexPrefetchHandle( + layer_id=1, + dense_page_buffer=dense_page_buffer, + prefix_rows=slice(0, 0), + event=fake_event, + ) + prefetcher.pending_attention_handle = prefetcher.handles[1] + current_k = torch.tensor( + [[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12], [13, 14, 15, 16]], + dtype=torch.uint8, + ) + current_scale = torch.tensor([[1.25], [2.5], [3.5], [4.5]], dtype=torch.float32) + current_locs = torch.tensor([20, 21, 44, 45], dtype=torch.int64) + current_req_id = torch.tensor([0, 0, 1, 1], dtype=torch.int64) + reduce_ranges = [] + + class FakeCurrentStream: + def __init__(self): + self.events = [] + + def wait_event(self, event): + self.events.append(event) + + current_stream = FakeCurrentStream() + + def record_range_reduce(buffer, cp_size, start_row, end_row, **kwargs): + reduce_ranges.append((start_row, end_row, kwargs.get("nvtx_source"))) + return buffer + + with patch.object( + prefetch.torch.cuda, "current_stream", return_value=current_stream + ), patch.object( + prefetch, + "_all_reduce_materialized_buffer_range", + side_effect=record_range_reduce, + ): + mixed_buffer, dense_pages = prefetcher.consume_prefix_with_current( + layer_id=1, + logical_pages=logical_pages, + current_index_k=current_k, + current_index_scale=current_scale, + current_locs=current_locs, + page_size=page_size, + index_head_dim=index_head_dim, + current_req_id=current_req_id, + ) + + scale_offset = page_size * index_head_dim + self.assertEqual(current_stream.events, [fake_event]) + self.assertEqual(dense_pages.tolist(), [[1, 2, 3, 0], [5, 6, 7, 8]]) + self.assertTrue(torch.equal(mixed_buffer[3, 0:8], current_k[:2].reshape(-1))) + self.assertTrue(torch.equal(mixed_buffer[6, 0:8], current_k[2:].reshape(-1))) + self.assertTrue( + torch.equal( + mixed_buffer[3, scale_offset : scale_offset + 8], + current_scale[:2].contiguous().view(torch.uint8).reshape(-1), + ) + ) + self.assertTrue( + torch.equal( + mixed_buffer[6, scale_offset : scale_offset + 8], + current_scale[2:].contiguous().view(torch.uint8).reshape(-1), + ) + ) + self.assertEqual( + reduce_ranges, + [ + (3, 4, "index.prefetch_current"), + (6, 8, "index.prefetch_current"), + ], + ) + def test_valid_page_mask_prevents_stale_rectangular_tail_remap(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 @@ -2117,8 +2780,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): ) def fake_ipc_into(**kwargs): - self.assertEqual(kwargs["start_slot"], 0) - self.assertEqual(kwargs["end_slot"], 2) + self.assertEqual(kwargs["spans"], [(0, 2)]) kwargs["dense_kv_cache"][4:12] = torch.arange( 10, 18, dtype=torch.float32 ).view(8, 1, 1) @@ -2130,11 +2792,9 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): range_calls.append((start_row, end_row, kwargs.get("nvtx_source"))) return buffer - from sglang.srt.environ import envs - - with envs.SGLANG_CP_SHARED_KV_COMPOSE_V2.override(False), patch.object( + with patch.object( runtime, - "_try_tai_ipc_materialize_token_kv_page_slots_into", + "_try_tai_ipc_materialize_token_kv_page_slot_spans_into", side_effect=fake_ipc_into, ), patch.object( runtime, @@ -2160,70 +2820,40 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): self.assertTrue(torch.equal(mixed_kv[12:14], current_kv)) self.assertEqual(range_calls, [(12, 16, "mla.partial_current_sync.current")]) - def test_materialize_prefix_current_token_kv_compose_v2_ipc_and_compact_reduce(self): - """compose_v2 contract: ONE full-range sentinel IPC gather for the - prefix + ONE collective over the compact current pages (never a - per-span range reduce, never a whole-buffer reduce).""" - from sglang.srt.environ import envs + def test_materialize_current_token_kv_uses_ipc_without_all_reduce(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=2, cp_rank=0) kv_cache = torch.zeros((32, 1, 1), dtype=torch.float32) - logical_locs = torch.tensor([[4, 8, 20, 21, 22, 23]], dtype=torch.int64) + logical_locs = torch.tensor([[20, 21]], dtype=torch.int64) current_locs = torch.tensor([20, 21], dtype=torch.int64) current_kv = torch.arange(100, 102, dtype=torch.float32).view(2, 1, 1) - remap_logical_pages = torch.tensor([[1, 2, 5]], 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, + remap_logical_pages=torch.tensor([[5]], dtype=torch.int64), layout=layout, page_size=page_size, ) + helper_calls = [] - captured = {} - test_case = self - - class FakeKernels: - @staticmethod - def materialize_cuda_ipc_peer_pages_slot_dense( - peer_ptrs, dense, owner_ranks, src_pages, *, page_nbytes - ): - captured["owner_ranks"] = owner_ranks.clone() - captured["src_pages"] = src_pages.clone() - captured["page_nbytes"] = page_nbytes - # Kernel contract: every dst page is either gathered or - # zero-filled (dummy page 0 + sentinel slots). - dense.zero_() - dense[4:12] = torch.arange(10, 18, dtype=torch.float32).view( - 8, 1, 1 - ) - - reduce_calls = [] - - def record_reduce(buffer, cp_size, **kwargs): - reduce_calls.append( - (tuple(buffer.shape), kwargs.get("nvtx_source")) - ) - return buffer - - def fail_range_reduce(*args, **kwargs): - test_case.fail("compose_v2 must not issue per-span range reduces") + def fake_current_ipc(**kwargs): + helper_calls.append(kwargs["spans"]) + self.assertEqual(kwargs["page_size"], page_size) + self.assertIs(kwargs["layout"], layout) + return True with patch.object( runtime, - "_get_or_open_tai_ipc_peer_ptrs", - return_value=(FakeKernels, object()), - ), patch.object( - runtime, - "_all_reduce_materialized_buffer", - side_effect=record_reduce, + "_try_tai_ipc_materialize_current_token_kv_page_slot_spans_into", + side_effect=fake_current_ipc, + create=True, ), patch.object( runtime, "_all_reduce_materialized_buffer_range", - side_effect=fail_range_reduce, + side_effect=AssertionError("current compose must use IPC, not all_reduce"), ): mixed_kv, mixed_locs = ( runtime.materialize_prefix_and_reuse_current_kv_page_slots( @@ -2234,66 +2864,68 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): slot_remap=slot_remap, layout=layout, page_size=page_size, - prefix_pages=2, + prefix_pages=0, + current_slot_spans=[(0, 1)], + layer_id=0, ) ) - # Descriptors cover the full slot range with -1 sentinels on the - # current slot (slot 2, logical page 5). - self.assertEqual(captured["owner_ranks"].numel(), 3) - self.assertEqual(int(captured["owner_ranks"][2].item()), -1) - self.assertEqual(int(captured["src_pages"][2].item()), -1) - self.assertGreaterEqual(int(captured["owner_ranks"][0].item()), 0) - self.assertGreaterEqual(int(captured["owner_ranks"][1].item()), 0) - # Exactly one collective: the compact current pages (1 page, uint8 - # byte view of page_size * 1 * fp32 = 16 bytes). - self.assertEqual(len(reduce_calls), 1) - self.assertEqual(reduce_calls[0][0], (1, 16)) - self.assertIn("v2_current_compact", reduce_calls[0][1]) - # Composed result matches the legacy contract. - self.assertEqual(mixed_locs.tolist(), [[4, 8, 12, 13, -1, -1]]) - self.assertEqual(mixed_kv[4].item(), 10) - self.assertEqual(mixed_kv[8].item(), 14) - self.assertTrue(torch.equal(mixed_kv[12:14], current_kv)) + self.assertEqual(helper_calls, [[(0, 1)]]) + self.assertEqual(mixed_locs.tolist(), [[4, 5]]) + self.assertTrue(torch.equal(mixed_kv[4:6], current_kv)) - def test_materialize_prefix_current_token_kv_compose_v2_fails_on_dense_fallback(self): + def test_current_token_ipc_helper_uses_dense_slot_pages_for_staging(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=2, cp_rank=0) - kv_cache = torch.zeros((32, 1, 1), dtype=torch.float32) - logical_locs = torch.tensor([[4, 8, 20, 21]], dtype=torch.int64) - current_locs = torch.tensor([20, 21], dtype=torch.int64) - current_kv = torch.arange(100, 102, dtype=torch.float32).view(2, 1, 1) - slot_remap = runtime.build_shared_token_kv_slot_remap( - kv_cache=kv_cache, - logical_locs=logical_locs, - remap_logical_pages=torch.tensor([[1, 2, 5]], dtype=torch.int64), - layout=layout, - page_size=page_size, + layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0) + dense = torch.zeros((16, 1), dtype=torch.float32) + slot_logical_pages = torch.tensor([5, 6, 0], dtype=torch.int64) + state = SimpleNamespace( + staging=torch.empty((64,), dtype=torch.uint8), + ready=torch.zeros((1,), dtype=torch.int64), + peer_ptrs=torch.tensor([111, 222], dtype=torch.int64), + ready_peer_ptrs=torch.tensor([333, 444], dtype=torch.int64), + ready_seq=0, ) + calls = [] + + class FakeKernels: + def publish_cuda_ipc_slot_pages_and_mark_ready(self, *args, **kwargs): + calls.append(("publish", args, kwargs)) + + def materialize_cuda_ipc_peer_pages_slot_indices_wait_ready( + self, *args, **kwargs + ): + calls.append(("gather", args, kwargs)) with patch.object( - runtime, "_get_or_open_tai_ipc_peer_ptrs", return_value=None - ), patch.object( - runtime, "_should_fail_fast_compose_v2_dense_fallback", return_value=True - ), self.assertRaisesRegex( - RuntimeError, - r"\[CP_SHARED_KV_FAIL_FAST\]\[compose_v2\].*token_kv_dense_fallback", + runtime, + "_get_or_create_tai_ipc_current_staging", + return_value=(FakeKernels(), state), ): - 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, + ok = runtime._try_tai_ipc_materialize_current_token_kv_page_slot_spans_into( + dense_kv_cache=dense, + slot_logical_pages=slot_logical_pages, layout=layout, - page_size=page_size, - prefix_pages=2, - layer_id=3, + page_size=4, + spans=[(0, 2)], ) + self.assertTrue(ok) + self.assertEqual(state.ready_seq, 1) + publish_name, publish_args, publish_kwargs = calls[0] + self.assertEqual(publish_name, "publish") + self.assertIs(publish_args[0], dense) + self.assertTrue(torch.equal(publish_args[2], torch.tensor([1, 2]))) + self.assertEqual(publish_kwargs["ready_seq"], 1) + gather_name, gather_args, gather_kwargs = calls[1] + self.assertEqual(gather_name, "gather") + self.assertTrue(torch.equal(gather_args[3], torch.tensor([0, 1]))) + self.assertTrue(torch.equal(gather_args[4], torch.tensor([1, 2]))) + self.assertTrue(torch.equal(gather_args[5], torch.tensor([1, 2]))) + self.assertEqual(gather_kwargs["ready_seq"], 1) + def test_mla_prefetch_consume_prefix_with_current_skips_suffix_materialize(self): from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout @@ -2364,6 +2996,57 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): self.assertEqual(mixed_locs.tolist(), [[4, 12], [13, 7], [-1, -1]]) self.assertEqual(range_calls, [(12, 16, "mla.prefetch_current")]) + def test_mla_prefetch_consume_prefix_with_current_uses_ipc_without_all_reduce(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + class FakeCurrentStream: + def __init__(self): + self.events = [] + + def wait_event(self, event): + self.events.append(event) + + dense_kv = torch.arange(0, 16, dtype=torch.float32).view(16, 1, 1) + current_kv = torch.arange(100, 102, dtype=torch.float32).view(2, 1, 1) + prefetcher = prefetch.CpSharedKVMlaPrefetcher( + layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0), + page_size=4, + prefix_pages=2, + slot_logical_pages=torch.tensor([1, 2, 5], dtype=torch.int64), + page_inverse=torch.tensor([0, 1, 2, -1, -1, 3], dtype=torch.int64), + dense_num_pages=4, + stream=object(), + ) + prefetcher.handles[1] = prefetch.CpSharedKVMlaPrefetchHandle( + layer_id=1, + dense_kv_cache=dense_kv, + prefix_rows=slice(4, 12), + event=object(), + ) + + with patch.object( + prefetch.torch.cuda, "current_stream", return_value=FakeCurrentStream() + ), patch.object( + prefetch, + "_try_tai_ipc_materialize_current_token_kv_page_slot_spans_into", + return_value=True, + ), patch.object( + prefetch, + "_all_reduce_materialized_buffer_range", + side_effect=AssertionError("prefetch current must use IPC"), + ): + mixed_kv, mixed_locs = prefetcher.consume_prefix_with_current( + layer_id=1, + kv_cache=torch.zeros((64, 1, 1), dtype=torch.float32), + logical_locs=torch.tensor([[20], [21]], dtype=torch.int32), + current_kv_cache=current_kv, + current_locs=torch.tensor([20, 21], dtype=torch.int64), + ) + + self.assertTrue(torch.equal(mixed_kv[12:14], current_kv)) + self.assertEqual(mixed_locs.tolist(), [[12], [13]]) + def test_mla_prefetch_attention_window_defers_pending_event_wait(self): from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout @@ -2680,8 +3363,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): current_scale = torch.tensor([[1.25], [2.5]], dtype=torch.float32) def fake_ipc_into(**kwargs): - self.assertEqual(kwargs["start_slot"], 0) - self.assertEqual(kwargs["end_slot"], 1) + self.assertEqual(kwargs["spans"], [(0, 1)]) kwargs["dense_page_buffer"][1] = torch.arange( page_bytes, dtype=torch.uint8 ) @@ -2693,11 +3375,9 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): range_calls.append((start_row, end_row, kwargs.get("nvtx_source"))) return buffer - from sglang.srt.environ import envs - - with envs.SGLANG_CP_SHARED_KV_COMPOSE_V2.override(False), patch.object( + with patch.object( runtime, - "_try_tai_ipc_materialize_paged_buffer_page_slots_into", + "_try_tai_ipc_materialize_paged_buffer_page_slot_spans_into", side_effect=fake_ipc_into, ), patch.object( runtime, @@ -2725,9 +3405,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): self.assertTrue(torch.equal(dense_page_buffer[2, 4:8], current_k[1])) self.assertEqual(range_calls, [(2, 3, "index.partial_current_sync.current")]) - def test_materialize_prefix_current_index_compose_v2_ipc_and_compact_reduce(self): - """compose_v2 contract for the index buffer (see token-KV twin).""" - from sglang.srt.environ import envs + def test_materialize_current_index_uses_ipc_without_all_reduce(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 @@ -2737,10 +3415,9 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): page_bytes = page_size * index_head_dim + page_size * scale_bytes layout = CpSharedKVLayout(page_size=page_size, cp_size=2, cp_rank=0) page_buffer = torch.zeros((8, page_bytes), dtype=torch.uint8) - logical_pages = torch.tensor([[1, 2]], dtype=torch.int64) slot_remap = runtime.build_shared_paged_buffer_slot_remap( page_buffer, - logical_pages, + torch.tensor([[5]], dtype=torch.int64), layout, ) current_k = torch.tensor( @@ -2748,113 +3425,90 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): dtype=torch.uint8, ) current_scale = torch.tensor([[1.25], [2.5]], dtype=torch.float32) + helper_calls = [] - captured = {} - test_case = self - - class FakeKernels: - @staticmethod - def materialize_cuda_ipc_peer_pages_slot_dense( - peer_ptrs, dense, owner_ranks, src_pages, *, page_nbytes - ): - captured["owner_ranks"] = owner_ranks.clone() - captured["src_pages"] = src_pages.clone() - dense.zero_() - dense[1] = torch.arange(dense.shape[1], dtype=torch.uint8) - - reduce_calls = [] - - def record_reduce(buffer, cp_size, **kwargs): - reduce_calls.append( - (tuple(buffer.shape), kwargs.get("nvtx_source")) - ) - return buffer - - def fail_range_reduce(*args, **kwargs): - test_case.fail("compose_v2 must not issue per-span range reduces") + def fake_current_ipc(**kwargs): + helper_calls.append(kwargs["spans"]) + self.assertIs(kwargs["layout"], layout) + return True with patch.object( runtime, - "_get_or_open_tai_ipc_peer_ptrs", - return_value=(FakeKernels, object()), - ), patch.object( - runtime, - "_all_reduce_materialized_buffer", - side_effect=record_reduce, + "_try_tai_ipc_materialize_current_paged_buffer_page_slot_spans_into", + side_effect=fake_current_ipc, + create=True, ), patch.object( runtime, "_all_reduce_materialized_buffer_range", - side_effect=fail_range_reduce, + side_effect=AssertionError("current index compose must use IPC, not all_reduce"), ): dense_page_buffer, dense_pages = ( runtime.materialize_prefix_and_reuse_current_index_page_slots( page_buffer=page_buffer, current_index_k=current_k, current_index_scale=current_scale, - current_locs=torch.tensor([8, 9], dtype=torch.int64), + current_locs=torch.tensor([20, 21], dtype=torch.int64), slot_remap=slot_remap, layout=layout, page_size=page_size, index_head_dim=index_head_dim, - prefix_pages=1, - layer_id=2, + prefix_pages=0, + current_slot_spans=[(0, 1)], + layer_id=0, ) ) - # Full-range descriptors: slot 0 = prefix (real owner), slot 1 = - # current (-1 sentinel). - self.assertEqual(captured["owner_ranks"].numel(), 2) - self.assertGreaterEqual(int(captured["owner_ranks"][0].item()), 0) - self.assertEqual(int(captured["owner_ranks"][1].item()), -1) - self.assertEqual(int(captured["src_pages"][1].item()), -1) - # One compact collective over the single current page. - self.assertEqual(len(reduce_calls), 1) - self.assertEqual(reduce_calls[0][0], (1, page_bytes)) - self.assertIn("v2_current_compact", reduce_calls[0][1]) - self.assertEqual(dense_pages.tolist(), [[1, 2]]) - self.assertEqual(dense_page_buffer[1].tolist(), list(range(page_bytes))) - self.assertTrue(torch.equal(dense_page_buffer[2, 0:4], current_k[0])) - self.assertTrue(torch.equal(dense_page_buffer[2, 4:8], current_k[1])) + self.assertEqual(helper_calls, [[(0, 1)]]) + self.assertEqual(dense_pages.tolist(), [[1]]) + self.assertTrue(torch.equal(dense_page_buffer[1, 0:4], current_k[0])) + self.assertTrue(torch.equal(dense_page_buffer[1, 4:8], current_k[1])) - def test_materialize_prefix_current_index_compose_v2_fails_on_dense_fallback(self): + def test_current_index_ipc_helper_uses_dense_slot_pages_for_staging(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 - index_head_dim = 4 - scale_bytes = 4 - page_bytes = page_size * index_head_dim + page_size * scale_bytes - layout = CpSharedKVLayout(page_size=page_size, cp_size=2, cp_rank=0) - page_buffer = torch.zeros((8, page_bytes), dtype=torch.uint8) - slot_remap = runtime.build_shared_paged_buffer_slot_remap( - page_buffer, - torch.tensor([[1, 2]], dtype=torch.int64), - layout, + layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1) + dense_pages = torch.zeros((5, 32), dtype=torch.uint8) + slot_logical_pages = torch.tensor([5, 6, 0], dtype=torch.int64) + state = SimpleNamespace( + staging=torch.empty((160,), dtype=torch.uint8), + ready=torch.zeros((1,), dtype=torch.int64), + peer_ptrs=torch.tensor([111, 222], dtype=torch.int64), + ready_peer_ptrs=torch.tensor([333, 444], dtype=torch.int64), + ready_seq=6, ) - current_k = torch.tensor([[1, 2, 3, 4], [5, 6, 7, 8]], dtype=torch.uint8) - current_scale = torch.tensor([[1.25], [2.5]], dtype=torch.float32) + calls = [] + + class FakeKernels: + def publish_cuda_ipc_slot_pages_and_mark_ready(self, *args, **kwargs): + calls.append(("publish", args, kwargs)) + + def materialize_cuda_ipc_peer_pages_slot_indices_wait_ready( + self, *args, **kwargs + ): + calls.append(("gather", args, kwargs)) with patch.object( - runtime, "_get_or_open_tai_ipc_peer_ptrs", return_value=None - ), patch.object( - runtime, "_should_fail_fast_compose_v2_dense_fallback", return_value=True - ), self.assertRaisesRegex( - RuntimeError, - r"\[CP_SHARED_KV_FAIL_FAST\]\[compose_v2\].*index_dense_fallback", + runtime, + "_get_or_create_tai_ipc_current_staging", + return_value=(FakeKernels(), state), ): - runtime.materialize_prefix_and_reuse_current_index_page_slots( - page_buffer=page_buffer, - current_index_k=current_k, - current_index_scale=current_scale, - current_locs=torch.tensor([8, 9], dtype=torch.int64), - slot_remap=slot_remap, + ok = runtime._try_tai_ipc_materialize_current_paged_buffer_page_slot_spans_into( + dense_page_buffer=dense_pages, + slot_logical_pages=slot_logical_pages, layout=layout, - page_size=page_size, - index_head_dim=index_head_dim, - prefix_pages=1, - layer_id=4, + spans=[(1, 3)], ) + self.assertTrue(ok) + self.assertEqual(state.ready_seq, 7) + self.assertTrue(torch.equal(calls[0][1][2], torch.tensor([2, 3]))) + self.assertEqual(calls[0][2]["ready_seq"], 7) + self.assertTrue(torch.equal(calls[1][1][3], torch.tensor([1, -1]))) + self.assertTrue(torch.equal(calls[1][1][4], torch.tensor([2, -1]))) + self.assertTrue(torch.equal(calls[1][1][5], torch.tensor([2, 3]))) + self.assertEqual(calls[1][2]["ready_seq"], 7) + def test_index_prefetch_partial_current_compose_fills_current_page_slots(self): from sglang.srt.environ import envs from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch @@ -2927,6 +3581,69 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): self.assertEqual(prefetcher.handles, {}) self.assertIsNone(prefetcher.pending_attention_handle) + def test_index_prefetch_consume_prefix_with_current_uses_ipc_without_all_reduce(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + class FakeCurrentStream: + def __init__(self): + self.events = [] + + def wait_event(self, event): + self.events.append(event) + + page_size = 4 + index_head_dim = 4 + page_bytes = page_size * index_head_dim + page_size * 4 + dense_page_buffer = torch.zeros((3, page_bytes), dtype=torch.uint8) + prefetcher = prefetch.CpSharedKVIndexPrefetcher( + layout=CpSharedKVLayout(page_size=page_size, cp_size=2, cp_rank=0), + prefix_pages=1, + slot_logical_pages=torch.tensor([1, 20], dtype=torch.int64), + page_inverse=torch.tensor( + [-1, 1] + [-1] * 18 + [2], + dtype=torch.int64, + ), + dense_num_pages=3, + stream=object(), + ) + prefetcher.handles[1] = prefetch.CpSharedKVIndexPrefetchHandle( + layer_id=1, + dense_page_buffer=dense_page_buffer, + prefix_rows=slice(1, 2), + event=object(), + ) + current_k = torch.tensor( + [[11, 12, 13, 14], [15, 16, 17, 18]], + dtype=torch.uint8, + ) + current_scale = torch.tensor([[3.25], [4.5]], dtype=torch.float32) + + with patch.object( + prefetch.torch.cuda, "current_stream", return_value=FakeCurrentStream() + ), patch.object( + prefetch, + "_try_tai_ipc_materialize_current_paged_buffer_page_slot_spans_into", + return_value=True, + ), patch.object( + prefetch, + "_all_reduce_materialized_buffer_range", + side_effect=AssertionError("index prefetch current must use IPC"), + ): + mixed_buffer, dense_pages = prefetcher.consume_prefix_with_current( + layer_id=1, + logical_pages=torch.tensor([[1, 20]], dtype=torch.int64), + current_index_k=current_k, + current_index_scale=current_scale, + current_locs=torch.tensor([80, 81], dtype=torch.int64), + page_size=page_size, + index_head_dim=index_head_dim, + ) + + self.assertEqual(dense_pages.tolist(), [[1, 2]]) + self.assertTrue(torch.equal(mixed_buffer[2, 0:4], current_k[0])) + self.assertTrue(torch.equal(mixed_buffer[2, 4:8], current_k[1])) + def test_index_current_reuse_gate_uses_partial_current_contract(self): from pathlib import Path @@ -4194,8 +4911,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): ) with patch.object( runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce - ), patch.object( - runtime, "_all_reduce_materialized_buffer", _identity_all_reduce ), patch.object( runtime, "_try_tai_ipc_materialize_token_kv_page_slots_into", return_value=False, @@ -4228,13 +4943,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): actual_row = merged_dense[dense_page * page_size + token_offset] torch.testing.assert_close(actual_row, expected_row, atol=0, rtol=0) - @unittest.skipIf(not torch.cuda.is_available(), "CUDA required") - def test_fp8_mla_persistent_pages_bs5_cache_hit_materialize_legacy(self): - from sglang.srt.environ import envs - - with envs.SGLANG_CP_SHARED_KV_COMPOSE_V2.override(False): - self.test_fp8_mla_persistent_pages_survive_bs5_cache_hit_materialize() - @unittest.skipIf(not torch.cuda.is_available(), "CUDA required") def test_fp8_index_fused_store_persistent_pages_survive_bs5_materialize(self): from sglang.jit_kernel.fused_store_index_cache import ( @@ -4583,11 +5291,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): else: current_k = torch.empty((0, index_head_dim), dtype=torch.uint8) current_scale = torch.empty((0, 1), dtype=torch.float32) - with patch.object( - runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce - ), patch.object( - runtime, "_all_reduce_materialized_buffer", _identity_all_reduce - ): + with patch.object(runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce): dense_page_buffer, dense_pages = runtime.materialize_prefix_and_reuse_current_index_page_slots( page_buffer=page_buffer, current_index_k=current_k, @@ -4615,12 +5319,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): merged = torch.stack(rank_outputs, dim=0).sum(dim=0).to(torch.uint8) torch.testing.assert_close(merged, expected, atol=0, rtol=0) - def test_cp8_index_partial_current_rank_merged_reference_bs5_legacy(self): - from sglang.srt.environ import envs - - with envs.SGLANG_CP_SHARED_KV_COMPOSE_V2.override(False): - self.test_cp8_index_partial_current_compose_matches_rank_merged_reference_for_bs5() - def test_cp8_kv_partial_current_keeps_remote_current_valid_locs_after_reduce(self): from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime @@ -4653,11 +5351,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): ) current_kv_cache = torch.tensor([[10.0, 11.0], [12.0, 13.0]]) - with patch.object( - runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce - ), patch.object( - runtime, "_all_reduce_materialized_buffer", _identity_all_reduce - ): + with patch.object(runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce): _mixed_cache, mixed_locs = runtime.materialize_prefix_and_reuse_current_kv_page_slots( kv_cache=kv_cache, logical_locs=logical_locs, @@ -4677,12 +5371,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): "remote-rank valid current loc must remain visible after current slot all-reduce", ) - def test_cp8_kv_partial_current_remote_current_valid_locs_legacy(self): - from sglang.srt.environ import envs - - with envs.SGLANG_CP_SHARED_KV_COMPOSE_V2.override(False): - self.test_cp8_kv_partial_current_keeps_remote_current_valid_locs_after_reduce() - def test_cp8_batch_kv_partial_current_keeps_request_packed_layout(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 @@ -4787,8 +5475,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): ) with patch.object( runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce - ), patch.object( - runtime, "_all_reduce_materialized_buffer", _identity_all_reduce ): mixed_kv, mixed_locs = runtime.materialize_prefix_and_reuse_current_kv_page_slots( kv_cache=kv_cache, @@ -4816,12 +5502,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): self.assertTrue(torch.equal(actual, expected)) - def test_cp8_batch_kv_partial_current_request_packed_layout_legacy(self): - from sglang.srt.environ import envs - - with envs.SGLANG_CP_SHARED_KV_COMPOSE_V2.override(False): - self.test_cp8_batch_kv_partial_current_keeps_request_packed_layout() - @unittest.skipIf(not torch.cuda.is_available(), "CUDA required") def test_tai_batched_index_mqa_prepare_matches_getk_gets_reference_gsm8k_bs5(self): from types import SimpleNamespace