From f8fca72635512095a9cf65eb8af1550767e107d2 Mon Sep 17 00:00:00 2001 From: laoyao0822 Date: Sun, 26 Apr 2026 04:11:17 +0800 Subject: [PATCH] Expand prefill CP KV capacity by sharding persistent NSA KV Prefill CP previously replicated NSA/MLA persistent KV on every CP rank, so CP8 consumed eight copies of KV memory while exposing only one rank of logical cache capacity. This change splits logical KV locs from per-rank physical storage, shards MLA latent KV and NSA index K/scale by deterministic page ownership, and keeps existing NSA attention kernels working through a full-view runtime materialization layer. Mooncake PD transfer now sends each prefill CP rank's owned physical pages with explicit logical page positions so non-CP decode can reconstruct full-layout KV. The implementation is guarded by an explicit server flag and startup checks, and the design documentation records the implemented scope, debug environment, and Phase 3 boundary. Constraint: Phase 2 must preserve existing NSA attention/index kernels via runtime full-view materialization Constraint: Decode side remains non-CP and receives full KV through Mooncake Rejected: Shard-aware NSA attention in this change | belongs to Phase 3 because it requires distributed topk/softmax/output contracts Rejected: Request-contiguous CP ownership | unstable under chunked prefill and tied to attention split mode Confidence: medium Scope-risk: broad Directive: Do not enable round-robin CP shared KV without wiring runtime materialization/PD transfer contracts for that split mode Directive: Keep SGLANG_DEBUG_CP_SHARED_KV disabled for perf measurements; it intentionally enables CUDA-syncing diagnostics Tested: Remote py_compile for shared-KV touched Python files in g0034 container Tested: Remote pytest selected cp_shared/shared_kv/nsa suite: 37 passed, 34 deselected Not-tested: Full GLM5 multi-node throughput/regression run after final doc update Not-tested: Phase 3 shard-aware runtime, round-robin CP mode, and non-Mooncake PD backends --- .../nsa_prefill_cp_phase2_shared_kv_design.md | 74 ++- .../srt/disaggregation/mooncake/conn.py | 149 ++++- python/sglang/srt/disaggregation/utils.py | 38 ++ python/sglang/srt/environ.py | 1 + .../attention/nsa/cp_shared_kv_runtime.py | 603 ++++++++++++++++++ .../srt/layers/attention/nsa/nsa_indexer.py | 233 +++++-- .../layers/attention/nsa/transform_index.py | 9 +- .../sglang/srt/layers/attention/nsa/utils.py | 69 +- .../srt/layers/attention/nsa_backend.py | 296 ++++++++- python/sglang/srt/mem_cache/allocator.py | 38 ++ .../srt/mem_cache/cp_shared_kv_layout.py | 85 +++ .../srt/model_executor/forward_batch_info.py | 5 + .../sglang/srt/model_executor/model_runner.py | 14 + .../model_runner_kv_cache_mixin.py | 46 +- python/sglang/srt/server_args.py | 36 +- .../test_cp_shared_kv_transfer_mapping.py | 61 ++ .../unit/layers/test_nsa_cp_utils.py | 41 ++ .../unit/layers/test_nsa_topk_transform.py | 210 ++++++ .../unit/layers/test_nsa_transform_index.py | 39 ++ .../mem_cache/test_cp_shared_kv_layout.py | 78 +++ .../mem_cache/test_cp_shared_kv_runtime.py | 435 +++++++++++++ .../unit/server_args/test_server_args.py | 19 + 22 files changed, 2489 insertions(+), 90 deletions(-) create mode 100644 python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py create mode 100644 python/sglang/srt/mem_cache/cp_shared_kv_layout.py create mode 100644 test/registered/unit/disaggregation/test_cp_shared_kv_transfer_mapping.py create mode 100644 test/registered/unit/layers/test_nsa_cp_utils.py create mode 100644 test/registered/unit/layers/test_nsa_topk_transform.py create mode 100644 test/registered/unit/layers/test_nsa_transform_index.py create mode 100644 test/registered/unit/mem_cache/test_cp_shared_kv_layout.py create mode 100644 test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py diff --git a/docs/advanced_features/nsa_prefill_cp_phase2_shared_kv_design.md b/docs/advanced_features/nsa_prefill_cp_phase2_shared_kv_design.md index 8f42641a2..5e5250cd5 100644 --- a/docs/advanced_features/nsa_prefill_cp_phase2_shared_kv_design.md +++ b/docs/advanced_features/nsa_prefill_cp_phase2_shared_kv_design.md @@ -6,6 +6,45 @@ Phase 2 的目标是先扩大 **persistent KV cache pool 的逻辑容量**。它 --- +## 0. 当前实现状态 + +当前代码已经落地 Phase 2 的首版兼容实现,启用开关为: + +```bash +--enable-nsa-prefill-cp-shared-kv +``` + +首版实现的关键语义: + +1. **persistent KV at rest 是 sharded/shared 的**: + - `req_to_token`、radix cache、scheduler admission 使用 CP group 统一的 **logical loc/page**。 + - 每个 prefill CP rank 的 `NSATokenToKVPool` 只按本 rank **physical capacity** 分配。 + - 每个 rank 只写入自己 owner 的 logical page,对应本地 physical page。 +2. **MLA latent KV 和 NSA index K/scale 同时 shard**: + - latent KV 写入前会把 `out_cache_loc` 从 logical loc 过滤为本 rank owner loc,并转换为 physical loc。 + - NSA indexer K/scale 写入使用同一 owner mapping,避免 index cache 与 latent KV 不一致。 +3. **attention runtime 仍走 full-view compatibility**: + - 现有 NSA attention/indexer kernel 仍假设可读 full KV/index view。 + - shared KV 下运行时会 materialize dense full-view buffer:每 rank 拷贝本地 owner pages,其余填 0,然后 CP group all-reduce 得到完整 runtime view。 + - 这是 Phase 2 的主要性能成本;Phase 3 才会改成 shard-aware attention/topk。 +4. **Mooncake PD transfer 已支持 prefill CP shards -> decode full KV**: + - prefill CP rank 发送本 rank owner physical pages。 + - transfer chunk 携带 `logical_page_positions`,decode 侧用这些位置选择 full-layout dst pages。 + - 因 decode 不开启 CP,prefill PD 必须设置 `SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER=1`。 +5. **debug/assert 默认不影响性能**: + - `SGLANG_DEBUG_CP_SHARED_KV=0` 时不执行 expensive checksum、CUDA tensor predicate、topk validation。 + - `SGLANG_DEBUG_CP_SHARED_KV=1` 仅用于定位 correctness 问题,不应参与性能压测。 + +当前仍未完成、属于 Phase 3 或后续工作: + +- runtime 不再 materialize full/maxlen KV。 +- NSA index/topk/attention 的 shard-aware 计算与 global topk merge。 +- `nsa_prefill_cp_mode=round-robin` 的 runtime wiring。 +- decode 侧 CP/shared KV。 +- Mooncake 之外的 PD transfer backend 完整支持。 + +--- + ## 1. 适用范围 Phase 2 首版只覆盖以下组合: @@ -608,7 +647,7 @@ PD disaggregation 下,如果 prompt KV 长度超过 decode side physical KV po ## 8. 配置与保护 -建议新增显式开关: +新增显式开关: ```text --enable-nsa-prefill-cp-shared-kv @@ -629,7 +668,7 @@ PD disaggregation 下,如果 prompt KV 长度超过 decode side physical KV po 如果任一条件不满足,启动时报明确错误。 -建议日志: +建议/当前日志: ```text CP shared KV enabled: physical_tokens_per_rank=237312, logical_tokens=1898496, cp_size=8, shard_policy=page_interleaved @@ -637,6 +676,37 @@ CP shared KV PD transfer: all CP ranks transfer enabled, backend=mooncake CP shared KV attention compatibility: full-view materialization enabled; Phase3 required for shard-aware runtime ``` +### 8.1 调试环境变量 + +```text +SGLANG_DEBUG_CP_SHARED_KV=1 +``` + +用途: + +1. 打印 CP shared KV 相关 debug 日志,包括 persistent KV shard write、NSA index materialize、Mooncake sender filtering、runtime full-view materialize 前后 checksum。 +2. 开启 shared KV 源头定位 assert,用于区分合法 `-1` sentinel 和真正异常的负 logical page/loc: + - attention token materialize path 允许 `logical_locs == -1`,这是 NSA topk/page table 的 invalid sentinel。 + - attention token materialize path 拒绝 `logical_locs < -1`。 + - token materialize 的 `remap_logical_locs` 拒绝任何 `< 0`,因为它应来自 `metadata.page_table_1` / `req_to_token`。 + - paged buffer materialize 拒绝 `logical_pages < 0`,因为 NSA index page table / real page table 不应携带 topk sentinel。 + - persistent MLA KV / NSA index write filtering 拒绝 `out_cache_loc < 0`。 + +性能影响: + +- 默认值为 `0`,正常线上路径只多一个 Python debug 分支,不执行 GPU tensor predicate,不引入 CUDA 同步。 +- 设置为 `1` 后,部分 assert 会对 CUDA tensor 执行 `torch.any()` / `min()` / `max()`,Python 控制流会触发同步;只能用于问题定位,不应参与性能压测。 +- Mooncake PD transfer 负页检查是 always-on CPU numpy 检查,用于防止异常 page list 被发送;成本相对 RDMA 传输可以忽略。 + +建议使用方式: + +```bash +SGLANG_DEBUG_CP_SHARED_KV=1 \ +python -m sglang.launch_server ... +``` + +当定位到负值源头后,应关闭该变量重新做吞吐和延迟验证。 + --- ## 9. 验证计划 diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 15e815e69..f3fe7f538 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -38,6 +38,30 @@ from sglang.srt.server_args import ServerArgs from sglang.srt.utils.network import NetworkAddress logger = logging.getLogger(__name__) +_CP_SHARED_DEBUG_COUNTS: dict[str, int] = {} + + +def _cp_shared_debug_log(key: str, message: str, *args, limit: int = 64) -> None: + if not envs.SGLANG_DEBUG_CP_SHARED_KV.get(): + return + count = _CP_SHARED_DEBUG_COUNTS.get(key, 0) + if count >= limit: + return + _CP_SHARED_DEBUG_COUNTS[key] = count + 1 + logger.info("[CP_SHARED_KV_DEBUG] " + message, *args) + + +def _np_summary(arr) -> str: + if arr is None: + return "None" + arr = np.asarray(arr) + if arr.size == 0: + return f"shape={arr.shape} dtype={arr.dtype} size=0" + head = arr.reshape(-1)[: min(8, arr.size)].tolist() + return ( + f"shape={arr.shape} dtype={arr.dtype} size={arr.size} " + f"min={int(arr.min())} max={int(arr.max())} head={head}" + ) class KVTransferError(Exception): @@ -55,7 +79,9 @@ class KVTransferError(Exception): class TransferKVChunk: room: int prefill_kv_indices: npt.NDArray[np.int32] - index_slice: slice + index_slice: Optional[slice] + logical_page_positions: Optional[npt.NDArray[np.int32]] + state_logical_page_positions: Optional[npt.NDArray[np.int32]] is_last_chunk: bool prefill_aux_index: Optional[int] state_indices: Optional[List[int]] @@ -593,6 +619,7 @@ class MooncakeKVManager(CommonKVManager): dst_state_data_ptrs: list[int], executor: concurrent.futures.ThreadPoolExecutor, target_rank_registration_info: Optional[KVArgsRegisterInfo] = None, + dst_state_indices: Optional[npt.NDArray[np.int32]] = None, ): """Send state or extra pool data with type-specific handling.""" state_type = getattr(self.kv_args, "state_type", "none") @@ -628,23 +655,27 @@ class MooncakeKVManager(CommonKVManager): raise RuntimeError( f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet." ) - if len(prefill_state_indices) < len(req.dst_state_indices): + effective_dst_state_indices = ( + np.asarray(dst_state_indices, dtype=np.int32) + if dst_state_indices is not None + else np.asarray(req.dst_state_indices, dtype=np.int32) + ) + if len(prefill_state_indices) < len(effective_dst_state_indices): logger.warning( - f"len(prefill_state_indices) = {len(prefill_state_indices)}, len(dst_state_indices) = {len(req.dst_state_indices)}" + f"len(prefill_state_indices) = {len(prefill_state_indices)}, len(dst_state_indices) = {len(effective_dst_state_indices)}" ) prefill_state_indices = prefill_state_indices[ - : len(req.dst_state_indices) + : len(effective_dst_state_indices) ] # Reuse _send_kvcache_generic interface to send extra pool data prefill_state_indices = np.array(prefill_state_indices, dtype=np.int32) - dst_state_indices = np.array(req.dst_state_indices, dtype=np.int32) return self._send_kvcache_generic( mooncake_session_id=req.mooncake_session_id, src_data_ptrs=self.kv_args.state_data_ptrs, dst_data_ptrs=dst_state_data_ptrs, item_lens=self.kv_args.state_item_lens, prefill_data_indices=prefill_state_indices, - dst_data_indices=dst_state_indices, + dst_data_indices=effective_dst_state_indices, executor=executor, ) else: @@ -807,7 +838,27 @@ class MooncakeKVManager(CommonKVManager): ) break - chunked_dst_kv_indice = req.dst_kv_indices[kv_chunk.index_slice] + if kv_chunk.logical_page_positions is not None: + chunked_dst_kv_indice = req.dst_kv_indices[ + kv_chunk.logical_page_positions + ] + else: + assert kv_chunk.index_slice is not None + chunked_dst_kv_indice = req.dst_kv_indices[ + kv_chunk.index_slice + ] + if envs.SGLANG_DEBUG_CP_SHARED_KV.get(): + _cp_shared_debug_log( + "transfer_worker_kv", + "transfer worker cp_rank=%s room=%s prefill_pages=%s " + "logical_positions=%s dst_pages=%s is_last=%s", + self.attn_cp_rank, + kv_chunk.room, + _np_summary(kv_chunk.prefill_kv_indices), + _np_summary(kv_chunk.logical_page_positions), + _np_summary(chunked_dst_kv_indice), + kv_chunk.is_last_chunk, + ) # NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices # is mismatched with the dst_kv_indices when page size > 1, this should never happen. @@ -871,12 +922,36 @@ class MooncakeKVManager(CommonKVManager): if kv_chunk.is_last_chunk: if kv_chunk.state_indices is not None: + dst_state_indices = None + if kv_chunk.state_logical_page_positions is not None: + from sglang.srt.disaggregation.utils import ( + select_pages_by_request_positions, + ) + + dst_state_indices = select_pages_by_request_positions( + req.dst_state_indices, + kv_chunk.state_logical_page_positions, + ) + if envs.SGLANG_DEBUG_CP_SHARED_KV.get(): + _cp_shared_debug_log( + "transfer_worker_state", + "transfer worker state cp_rank=%s room=%s prefill_state_pages=%s " + "state_positions=%s dst_state_pages=%s", + self.attn_cp_rank, + kv_chunk.room, + _np_summary(kv_chunk.state_indices), + _np_summary( + kv_chunk.state_logical_page_positions + ), + _np_summary(dst_state_indices), + ) self.maybe_send_extra( req, kv_chunk.state_indices, target_rank_registration_info.dst_state_data_ptrs, executor, target_rank_registration_info, + dst_state_indices, ) # Only the last chunk we need to send the aux data @@ -1049,10 +1124,12 @@ class MooncakeKVManager(CommonKVManager): self, bootstrap_room: int, kv_indices: npt.NDArray[np.int32], - index_slice: slice, + index_slice: Optional[slice], is_last_chunk: bool, aux_index: Optional[int] = None, state_indices: Optional[List[int]] = None, + logical_page_positions: Optional[npt.NDArray[np.int32]] = None, + state_logical_page_positions: Optional[npt.NDArray[np.int32]] = None, ): assert self.disaggregation_mode == DisaggregationMode.PREFILL assert not is_last_chunk or (is_last_chunk and aux_index is not None) @@ -1084,6 +1161,8 @@ class MooncakeKVManager(CommonKVManager): room=bootstrap_room, prefill_kv_indices=kv_indices, index_slice=index_slice, + logical_page_positions=logical_page_positions, + state_logical_page_positions=state_logical_page_positions, is_last_chunk=is_last_chunk, prefill_aux_index=aux_index, state_indices=state_indices, @@ -1147,9 +1226,57 @@ class MooncakeKVSender(CommonKVSender): index_slice = slice(self.curr_idx, self.curr_idx + len(kv_indices)) self.curr_idx += len(kv_indices) is_last_chunk = self.curr_idx == self.num_kv_indices + logical_page_positions = None + state_logical_page_positions = None + orig_kv_indices = None + orig_state_indices = None + if self.kv_mgr.server_args.enable_nsa_prefill_cp_shared_kv: + from sglang.srt.disaggregation.utils import filter_kv_pages_for_cp_shared_kv + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + chunk_page_start = index_slice.start + orig_kv_indices = np.asarray(kv_indices, dtype=np.int32).copy() + if state_indices is not None: + orig_state_indices = np.asarray(state_indices, dtype=np.int32).copy() + layout = CpSharedKVLayout( + page_size=self.kv_mgr.kv_args.page_size, + cp_size=self.kv_mgr.attn_cp_size, + cp_rank=self.kv_mgr.attn_cp_rank, + ) + kv_indices, logical_page_positions = filter_kv_pages_for_cp_shared_kv( + layout=layout, + logical_pages=kv_indices, + chunk_page_start=index_slice.start, + ) + if state_indices is not None: + state_indices, state_logical_page_positions = ( + filter_kv_pages_for_cp_shared_kv( + layout=layout, + logical_pages=state_indices, + chunk_page_start=0, + ) + ) + index_slice = None + if envs.SGLANG_DEBUG_CP_SHARED_KV.get(): + _cp_shared_debug_log( + "sender_filter", + "sender filter cp_rank=%s room=%s page_start=%s orig_kv_pages=%s " + "filtered_kv_pages=%s kv_positions=%s orig_state_pages=%s " + "filtered_state_pages=%s state_positions=%s is_last=%s", + self.kv_mgr.attn_cp_rank, + self.bootstrap_room, + chunk_page_start, + _np_summary(orig_kv_indices), + _np_summary(kv_indices), + _np_summary(logical_page_positions), + _np_summary(orig_state_indices), + _np_summary(state_indices), + _np_summary(state_logical_page_positions), + is_last_chunk, + ) # Special handling for cp - if self.kv_mgr.enable_all_cp_ranks_for_transfer: + elif self.kv_mgr.enable_all_cp_ranks_for_transfer: kv_indices, index_slice = filter_kv_indices_for_cp_rank( self.kv_mgr, kv_indices, @@ -1168,6 +1295,8 @@ class MooncakeKVSender(CommonKVSender): kv_indices, index_slice, False, + logical_page_positions=logical_page_positions, + state_logical_page_positions=state_logical_page_positions, ) else: self.kv_mgr.add_transfer_request( @@ -1177,6 +1306,8 @@ class MooncakeKVSender(CommonKVSender): True, aux_index=self.aux_index, state_indices=state_indices, + logical_page_positions=logical_page_positions, + state_logical_page_positions=state_logical_page_positions, ) def poll(self) -> KVPoll: diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index b7b3b0238..9a2b545e4 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -505,6 +505,44 @@ def filter_kv_indices_for_cp_rank( return new_kv_indices, new_index_slice +def filter_kv_pages_for_cp_shared_kv( + layout, + logical_pages: np.ndarray, + chunk_page_start: int, +) -> Tuple[np.ndarray, np.ndarray]: + logical_pages = np.asarray(logical_pages, dtype=np.int32) + if logical_pages.size > 0 and int(logical_pages.min()) < 0: + bad_pages = logical_pages[logical_pages < 0] + raise ValueError( + "CP shared KV transfer got negative logical_pages. " + f"bad_min={int(bad_pages.min())} bad_max={int(bad_pages.max())} " + f"chunk_page_start={chunk_page_start} num_pages={len(logical_pages)}" + ) + request_positions = np.arange( + chunk_page_start, + chunk_page_start + len(logical_pages), + dtype=np.int32, + ) + return layout.filter_owned_pages_np(logical_pages, request_positions) + + +def select_pages_by_request_positions( + pages: np.ndarray, + request_positions: np.ndarray, +) -> np.ndarray: + """Select per-request page entries by absolute request page positions.""" + pages = np.asarray(pages, dtype=np.int32) + request_positions = np.asarray(request_positions, dtype=np.int32) + if request_positions.size == 0: + return pages[:0] + if int(request_positions.max()) >= len(pages): + raise ValueError( + "request_positions contains an entry outside pages: " + f"max_position={int(request_positions.max())} pages={len(pages)}" + ) + return pages[request_positions] + + ######################### # Misc ######################### diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 4dcf0613b..2d81be044 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -202,6 +202,7 @@ class Envs: SGLANG_RECORD_STEP_TIME = EnvBool(False) SGLANG_FORCE_SHUTDOWN = EnvBool(False) SGLANG_DEBUG_MEMORY_POOL = EnvBool(False) + SGLANG_DEBUG_CP_SHARED_KV = EnvBool(False) SGLANG_TEST_REQUEST_TIME_STATS = EnvBool(False) SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK = EnvBool(False) SGLANG_SIMULATE_ACC_LEN = EnvFloat(-1) 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 new file mode 100644 index 000000000..3776e7f45 --- /dev/null +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py @@ -0,0 +1,603 @@ +from __future__ import annotations + +import logging + +import torch + +from sglang.srt.environ import envs +from sglang.srt.layers.dp_attention import get_attention_cp_group +from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + +logger = logging.getLogger(__name__) + +_DEBUG_LOG_COUNTS: dict[str, int] = {} + + +def cp_shared_kv_debug_enabled() -> bool: + return envs.SGLANG_DEBUG_CP_SHARED_KV.get() + + +def cp_shared_kv_debug_log( + key: str, + message: str, + *args, + limit: int = 32, +) -> None: + if not cp_shared_kv_debug_enabled(): + return + count = _DEBUG_LOG_COUNTS.get(key, 0) + if count >= limit: + return + _DEBUG_LOG_COUNTS[key] = count + 1 + logger.info("[CP_SHARED_KV_DEBUG] " + message, *args) + + +def tensor_debug_summary(tensor: torch.Tensor | None, *, sample: int = 8) -> str: + if tensor is None: + return "None" + try: + if tensor.numel() == 0: + return f"shape={tuple(tensor.shape)} dtype={tensor.dtype} numel=0" + flat = tensor.detach().reshape(-1) + head = flat[: min(sample, flat.numel())].detach().cpu().tolist() + if tensor.dtype.is_floating_point: + min_val = float(flat.min().item()) + max_val = float(flat.max().item()) + else: + min_val = int(flat.min().item()) + max_val = int(flat.max().item()) + return ( + f"shape={tuple(tensor.shape)} dtype={tensor.dtype} " + f"numel={tensor.numel()} min={min_val} max={max_val} head={head}" + ) + except Exception as exc: + return f"shape={tuple(tensor.shape)} dtype={tensor.dtype} summary_error={exc}" + + +def _debug_byte_sum(byte_view: torch.Tensor, indices: torch.Tensor) -> tuple[int, int]: + if indices.numel() == 0: + return 0, 0 + sampled = byte_view[indices].to(torch.int64) + return int(sampled.numel()), int(sampled.sum().item()) + + +def tensor_debug_checksum(tensor: torch.Tensor | None, *, max_bytes: int = 4096) -> str: + if tensor is None: + return "None" + try: + if tensor.numel() == 0: + return "numel=0 sampled_bytes=0 head_sum=0 spread_sum=0 tail_sum=0" + byte_view = tensor.detach().contiguous().view(torch.uint8).reshape(-1) + total_bytes = int(byte_view.numel()) + sample_count = min(max_bytes, total_bytes) + head_idx = torch.arange(sample_count, device=byte_view.device, dtype=torch.long) + if sample_count == total_bytes: + spread_idx = head_idx + tail_idx = head_idx + else: + # Sample across the whole storage so page-0 dummy zeros do not mask + # live data on later pages. + spread_idx = torch.linspace( + 0, + total_bytes - 1, + steps=sample_count, + device=byte_view.device, + dtype=torch.float32, + ).to(torch.long) + tail_idx = torch.arange( + total_bytes - sample_count, + total_bytes, + device=byte_view.device, + dtype=torch.long, + ) + head_n, head_sum = _debug_byte_sum(byte_view, head_idx) + spread_n, spread_sum = _debug_byte_sum(byte_view, spread_idx) + tail_n, tail_sum = _debug_byte_sum(byte_view, tail_idx) + return ( + f"sampled_bytes={head_n}/{spread_n}/{tail_n} " + f"head_sum={head_sum} spread_sum={spread_sum} tail_sum={tail_sum} " + f"total_bytes={total_bytes}" + ) + except Exception as exc: + return f"checksum_error={exc}" + + +def tensor_debug_nonzero_slice_checksum( + tensor: torch.Tensor | None, + *, + start: int, + max_bytes: int = 4096, +) -> str: + if tensor is None: + return "None" + try: + if tensor.shape[0] <= start: + return f"empty_after_start={start}" + return tensor_debug_checksum(tensor[start:], max_bytes=max_bytes) + except Exception as exc: + return f"slice_checksum_error={exc}" + + +def _debug_assert_no_tensor_values_below( + tensor: torch.Tensor, + *, + min_allowed: int, + context: str, + tensor_name: str, +) -> None: + """Debug-only tensor value guard. + + CUDA tensor predicates synchronize when used for Python control flow. Keep + these source-finding assertions behind SGLANG_DEBUG_CP_SHARED_KV so the + normal shared-KV fast path pays only one Python branch. + """ + if not cp_shared_kv_debug_enabled() or tensor.numel() == 0: + return + + invalid_mask = tensor < min_allowed + if not torch.any(invalid_mask): + return + + invalid_values = tensor[invalid_mask] + raise RuntimeError( + f"{context} got {tensor_name} below {min_allowed}. " + f"bad_min={int(invalid_values.min().item())} " + f"bad_max={int(invalid_values.max().item())} " + f"{tensor_name}={tensor_debug_summary(tensor)}" + ) + + +def _debug_assert_no_negative_tensor_values( + tensor: torch.Tensor, + *, + context: str, + tensor_name: str, +) -> None: + if not cp_shared_kv_debug_enabled() or tensor.numel() == 0: + return + + invalid_mask = tensor < 0 + if not torch.any(invalid_mask): + return + + invalid_values = tensor[invalid_mask] + raise RuntimeError( + f"{context} got negative {tensor_name}. " + f"bad_min={int(invalid_values.min().item())} " + f"bad_max={int(invalid_values.max().item())} " + f"{tensor_name}={tensor_debug_summary(tensor)}" + ) + + +def build_dense_page_remap(logical_pages: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Build a dense per-call page remap for shared-KV runtime materialization.""" + dense_pages = logical_pages.clone() + positive_mask = logical_pages > 0 + if not torch.any(positive_mask): + empty = logical_pages.new_empty((0,)) + return empty, dense_pages + + unique_logical_pages = torch.unique(logical_pages[positive_mask], sorted=True) + dense_pages[positive_mask] = remap_logical_pages_to_dense_pages( + logical_pages[positive_mask], + unique_logical_pages=unique_logical_pages, + ) + return unique_logical_pages, dense_pages + + +def remap_logical_pages_to_dense_pages( + logical_pages: torch.Tensor, + unique_logical_pages: torch.Tensor, +) -> torch.Tensor: + dense_pages = logical_pages.clone() + positive_mask = logical_pages > 0 + if not torch.any(positive_mask): + return dense_pages + + if unique_logical_pages.numel() == 0: + raise ValueError("unique_logical_pages is empty but logical_pages contains data") + + positive_pages = logical_pages[positive_mask] + insert_positions = torch.searchsorted(unique_logical_pages, positive_pages) + if insert_positions.numel() > 0: + if torch.any(insert_positions >= unique_logical_pages.numel()): + raise ValueError("logical_pages contains entries outside unique_logical_pages") + if not torch.equal(unique_logical_pages[insert_positions], positive_pages): + raise ValueError("logical_pages contains entries outside unique_logical_pages") + + dense_pages[positive_mask] = insert_positions.to(dense_pages.dtype) + 1 + return dense_pages + + +def remap_logical_locs_to_dense_locs( + logical_locs: torch.Tensor, + unique_logical_pages: torch.Tensor, + page_size: int, +) -> torch.Tensor: + dense_locs = logical_locs.clone() + valid_mask = logical_locs >= 0 + if not torch.any(valid_mask): + return dense_locs + + valid_locs = logical_locs[valid_mask] + logical_pages = torch.div(valid_locs, page_size, rounding_mode="floor") + offsets = torch.remainder(valid_locs, page_size) + dense_pages = remap_logical_pages_to_dense_pages( + logical_pages, + unique_logical_pages=unique_logical_pages, + ) + dense_locs[valid_mask] = dense_pages * page_size + offsets + return dense_locs + + +def logical_pages_from_locs(logical_locs: torch.Tensor, page_size: int) -> torch.Tensor: + logical_pages = logical_locs.clone() + valid_mask = logical_locs >= 0 + if torch.any(valid_mask): + logical_pages[valid_mask] = torch.div( + logical_locs[valid_mask], + page_size, + rounding_mode="floor", + ) + return logical_pages + + +def filter_locs_mappable_to_physical_pool( + logical_locs: torch.Tensor, + layout: CpSharedKVLayout, + physical_token_capacity: int, +) -> torch.Tensor: + """Debug-validate locs before mapping them into the physical pool. + + The checks below use CUDA tensor predicates in Python control flow and can + synchronize. Keep them on the debug path only; production materialization + trusts the logical-to-physical contract for speed. + """ + if not cp_shared_kv_debug_enabled(): + return logical_locs + + valid_mask = logical_locs >= 0 + if not torch.any(valid_mask): + return logical_locs + + filtered_locs = logical_locs.clone() + valid_locs = logical_locs[valid_mask] + physical_locs = layout.logical_locs_to_physical(valid_locs) + mappable_mask = physical_locs < physical_token_capacity + if torch.all(mappable_mask): + return filtered_locs + + invalid_logical_locs = valid_locs[~mappable_mask] + invalid_physical_locs = physical_locs[~mappable_mask] + raise RuntimeError( + "CP shared KV materialize got logical token locs outside the physical pool. " + f"invalid_logical_min={int(invalid_logical_locs.min().item())} " + f"invalid_logical_max={int(invalid_logical_locs.max().item())} " + f"invalid_physical_min={int(invalid_physical_locs.min().item())} " + f"invalid_physical_max={int(invalid_physical_locs.max().item())} " + f"physical_token_capacity={physical_token_capacity} " + f"page_size={layout.page_size} cp_rank={layout.cp_rank} cp_size={layout.cp_size}" + ) + + +def filter_pages_mappable_to_physical_pool( + logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + physical_page_capacity: int, +) -> torch.Tensor: + """Debug-validate page ids before mapping them into the physical pool.""" + if not cp_shared_kv_debug_enabled(): + return logical_pages + + positive_mask = logical_pages > 0 + if not torch.any(positive_mask): + return logical_pages + + filtered_pages = logical_pages.clone() + positive_pages = logical_pages[positive_mask] + physical_pages = layout.logical_pages_to_physical(positive_pages) + mappable_mask = physical_pages < physical_page_capacity + if torch.all(mappable_mask): + return filtered_pages + + invalid_logical_pages = positive_pages[~mappable_mask] + invalid_physical_pages = physical_pages[~mappable_mask] + raise RuntimeError( + "CP shared KV materialize got logical pages outside the physical page buffer. " + f"invalid_logical_page_min={int(invalid_logical_pages.min().item())} " + f"invalid_logical_page_max={int(invalid_logical_pages.max().item())} " + f"invalid_physical_page_min={int(invalid_physical_pages.min().item())} " + f"invalid_physical_page_max={int(invalid_physical_pages.max().item())} " + f"physical_page_capacity={physical_page_capacity} " + f"page_size={layout.page_size} cp_rank={layout.cp_rank} cp_size={layout.cp_size}" + ) + + +def filter_owned_logical_locs( + logical_locs: torch.Tensor, + layout: CpSharedKVLayout, +) -> tuple[torch.Tensor, torch.Tensor]: + """Return the tokens owned by this CP rank and their physical KV locs.""" + _debug_assert_no_negative_tensor_values( + logical_locs, + context="CP shared KV persistent write", + tensor_name="logical_locs", + ) + owned_mask = layout.owned_by_this_rank(logical_locs) + if not torch.any(owned_mask): + return owned_mask, logical_locs.new_empty((0,)) + + physical_locs = layout.logical_locs_to_physical(logical_locs[owned_mask]) + return owned_mask, physical_locs.contiguous() + + +def materialize_local_token_kv_pages( + kv_cache: torch.Tensor, + unique_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + page_size: int, +) -> torch.Tensor: + dense_num_pages = int(unique_logical_pages.numel()) + 1 + dense_kv_cache = kv_cache.new_zeros((dense_num_pages * page_size, *kv_cache.shape[1:])) + if unique_logical_pages.numel() == 0: + return dense_kv_cache + + owned_mask = layout.owned_pages_mask(unique_logical_pages) + if not torch.any(owned_mask): + return dense_kv_cache + + owned_logical_pages = unique_logical_pages[owned_mask].to(torch.int64) + owned_physical_pages = layout.logical_pages_to_physical(owned_logical_pages).to( + torch.long + ) + dense_page_ids = torch.nonzero(owned_mask, as_tuple=False).flatten().to(torch.long) + 1 + page_offsets = torch.arange(page_size, device=kv_cache.device, dtype=torch.long) + src_tokens = (owned_physical_pages[:, None] * page_size + page_offsets).reshape(-1) + dst_tokens = (dense_page_ids[:, None] * page_size + page_offsets).reshape(-1) + dense_kv_cache[dst_tokens] = kv_cache[src_tokens] + return dense_kv_cache + + +def token_page_copy_debug_checksum( + kv_cache: torch.Tensor, + dense_kv_cache: torch.Tensor, + unique_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, + page_size: int, +) -> str: + try: + owned_mask = layout.owned_pages_mask(unique_logical_pages) + if not torch.any(owned_mask): + return "owned_pages=0" + owned_logical_pages = unique_logical_pages[owned_mask].to(torch.int64) + owned_physical_pages = layout.logical_pages_to_physical( + owned_logical_pages + ).to(torch.long) + dense_page_ids = ( + torch.nonzero(owned_mask, as_tuple=False).flatten().to(torch.long) + 1 + ) + page_offsets = torch.arange(page_size, device=kv_cache.device, dtype=torch.long) + src_tokens = (owned_physical_pages[:, None] * page_size + page_offsets).reshape( + -1 + ) + dst_tokens = (dense_page_ids[:, None] * page_size + page_offsets).reshape(-1) + return ( + f"src_ck={tensor_debug_checksum(kv_cache[src_tokens])} " + f"dst_ck={tensor_debug_checksum(dense_kv_cache[dst_tokens])}" + ) + except Exception as exc: + return f"token_page_copy_ck_error={exc}" + + +def materialize_local_paged_buffer( + page_buffer: torch.Tensor, + unique_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, +) -> torch.Tensor: + dense_num_pages = int(unique_logical_pages.numel()) + 1 + dense_page_buffer = page_buffer.new_zeros((dense_num_pages, *page_buffer.shape[1:])) + if unique_logical_pages.numel() == 0: + return dense_page_buffer + + owned_mask = layout.owned_pages_mask(unique_logical_pages) + if not torch.any(owned_mask): + return dense_page_buffer + + owned_logical_pages = unique_logical_pages[owned_mask].to(torch.int64) + owned_physical_pages = layout.logical_pages_to_physical(owned_logical_pages).to( + torch.long + ) + dense_page_ids = torch.nonzero(owned_mask, as_tuple=False).flatten().to(torch.long) + 1 + dense_page_buffer[dense_page_ids] = page_buffer[owned_physical_pages] + return dense_page_buffer + + +def paged_copy_debug_checksum( + page_buffer: torch.Tensor, + dense_page_buffer: torch.Tensor, + unique_logical_pages: torch.Tensor, + layout: CpSharedKVLayout, +) -> str: + try: + owned_mask = layout.owned_pages_mask(unique_logical_pages) + if not torch.any(owned_mask): + return "owned_pages=0" + owned_logical_pages = unique_logical_pages[owned_mask].to(torch.int64) + owned_physical_pages = layout.logical_pages_to_physical( + owned_logical_pages + ).to(torch.long) + dense_page_ids = ( + torch.nonzero(owned_mask, as_tuple=False).flatten().to(torch.long) + 1 + ) + return ( + f"src_ck={tensor_debug_checksum(page_buffer[owned_physical_pages])} " + f"dst_ck={tensor_debug_checksum(dense_page_buffer[dense_page_ids])}" + ) + except Exception as exc: + return f"paged_copy_ck_error={exc}" + + +def _comm_view(buffer: torch.Tensor) -> torch.Tensor: + if buffer.dtype in (torch.float8_e4m3fn, torch.float8_e5m2): + return buffer.view(torch.uint8) + return buffer + + +def _all_reduce_materialized_buffer(buffer: torch.Tensor, cp_size: int) -> torch.Tensor: + if cp_size > 1: + comm_buffer = _comm_view(buffer) + cp_group = get_attention_cp_group() + if comm_buffer.dtype in (torch.float32, torch.float16, torch.bfloat16): + reduced = cp_group.all_reduce(comm_buffer) + if ( + reduced is not None + and reduced.untyped_storage().data_ptr() + != comm_buffer.untyped_storage().data_ptr() + ): + comm_buffer.copy_(reduced) + else: + torch.distributed.all_reduce(comm_buffer, group=cp_group.device_group) + return buffer + + +def materialize_shared_token_kv_buffer( + kv_cache: torch.Tensor, + logical_locs: torch.Tensor, + layout: CpSharedKVLayout, + page_size: int, + remap_logical_locs: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + _debug_assert_no_tensor_values_below( + logical_locs, + min_allowed=-1, + context="CP shared KV token materialize", + tensor_name="logical_locs", + ) + logical_locs = filter_locs_mappable_to_physical_pool( + logical_locs=logical_locs, + layout=layout, + physical_token_capacity=kv_cache.shape[0], + ) + + if remap_logical_locs is None: + remap_logical_locs = logical_locs + else: + _debug_assert_no_negative_tensor_values( + remap_logical_locs, + context="CP shared KV token materialize remap", + tensor_name="remap_logical_locs", + ) + remap_logical_locs = filter_locs_mappable_to_physical_pool( + logical_locs=remap_logical_locs, + layout=layout, + physical_token_capacity=kv_cache.shape[0], + ) + + remap_logical_pages = logical_pages_from_locs(remap_logical_locs, page_size) + unique_logical_pages, _ = build_dense_page_remap(remap_logical_pages) + dense_locs = remap_logical_locs_to_dense_locs( + logical_locs, + unique_logical_pages=unique_logical_pages, + page_size=page_size, + ) + dense_kv_cache = materialize_local_token_kv_pages( + kv_cache=kv_cache, + unique_logical_pages=unique_logical_pages, + layout=layout, + page_size=page_size, + ) + if cp_shared_kv_debug_enabled(): + owned_pages = unique_logical_pages[layout.owned_pages_mask(unique_logical_pages)] + physical_pages = layout.logical_pages_to_physical(owned_pages) + cp_shared_kv_debug_log( + "materialize_token_pre", + "token materialize pre cp_rank=%s logical_locs=%s remap_locs=%s " + "unique_pages=%s owned_pages=%s physical_pages=%s dense_locs=%s " + "local_ck=%s active_local_ck=%s owned_copy_ck=%s", + layout.cp_rank, + tensor_debug_summary(logical_locs), + tensor_debug_summary(remap_logical_locs), + tensor_debug_summary(unique_logical_pages), + tensor_debug_summary(owned_pages), + tensor_debug_summary(physical_pages), + tensor_debug_summary(dense_locs), + tensor_debug_checksum(dense_kv_cache), + tensor_debug_nonzero_slice_checksum(dense_kv_cache, start=page_size), + token_page_copy_debug_checksum( + kv_cache, + dense_kv_cache, + unique_logical_pages, + layout, + page_size, + ), + ) + dense_kv_cache = _all_reduce_materialized_buffer(dense_kv_cache, layout.cp_size) + if cp_shared_kv_debug_enabled(): + cp_shared_kv_debug_log( + "materialize_token_post", + "token materialize post cp_rank=%s dense_shape=%s reduced_ck=%s " + "active_reduced_ck=%s", + layout.cp_rank, + tuple(dense_kv_cache.shape), + tensor_debug_checksum(dense_kv_cache), + tensor_debug_nonzero_slice_checksum(dense_kv_cache, start=page_size), + ) + return dense_kv_cache, dense_locs + + +def materialize_shared_paged_buffer( + page_buffer: torch.Tensor, + logical_pages: torch.Tensor, + layout: CpSharedKVLayout, +) -> tuple[torch.Tensor, torch.Tensor]: + _debug_assert_no_negative_tensor_values( + logical_pages, + context="CP shared KV paged materialize", + tensor_name="logical_pages", + ) + logical_pages = filter_pages_mappable_to_physical_pool( + logical_pages=logical_pages, + layout=layout, + physical_page_capacity=page_buffer.shape[0], + ) + unique_logical_pages, dense_pages = build_dense_page_remap(logical_pages) + dense_page_buffer = materialize_local_paged_buffer( + page_buffer=page_buffer, + unique_logical_pages=unique_logical_pages, + layout=layout, + ) + if cp_shared_kv_debug_enabled(): + owned_pages = unique_logical_pages[layout.owned_pages_mask(unique_logical_pages)] + physical_pages = layout.logical_pages_to_physical(owned_pages) + cp_shared_kv_debug_log( + "materialize_paged_pre", + "paged materialize pre cp_rank=%s logical_pages=%s unique_pages=%s " + "owned_pages=%s physical_pages=%s dense_pages=%s local_ck=%s " + "active_local_ck=%s owned_copy_ck=%s", + layout.cp_rank, + tensor_debug_summary(logical_pages), + tensor_debug_summary(unique_logical_pages), + tensor_debug_summary(owned_pages), + tensor_debug_summary(physical_pages), + tensor_debug_summary(dense_pages), + tensor_debug_checksum(dense_page_buffer), + tensor_debug_nonzero_slice_checksum(dense_page_buffer, start=1), + paged_copy_debug_checksum( + page_buffer, + dense_page_buffer, + unique_logical_pages, + layout, + ), + ) + dense_page_buffer = _all_reduce_materialized_buffer(dense_page_buffer, layout.cp_size) + if cp_shared_kv_debug_enabled(): + cp_shared_kv_debug_log( + "materialize_paged_post", + "paged materialize post cp_rank=%s dense_shape=%s reduced_ck=%s " + "active_reduced_ck=%s", + layout.cp_rank, + tuple(dense_page_buffer.shape), + tensor_debug_checksum(dense_page_buffer), + tensor_debug_nonzero_slice_checksum(dense_page_buffer, start=1), + ) + return dense_page_buffer, dense_pages diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index bff848c1b..b2e3e540b 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -12,6 +12,15 @@ from sglang.jit_kernel.fused_store_index_cache import ( fused_store_index_k_cache, ) from sglang.srt.environ import envs +from sglang.srt.layers.attention.nsa import index_buf_accessor +from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( + cp_shared_kv_debug_enabled, + cp_shared_kv_debug_log, + filter_owned_logical_locs, + materialize_shared_paged_buffer, + tensor_debug_checksum, + tensor_debug_summary, +) from sglang.srt.layers.dp_attention import attn_tp_all_gather_into_tensor from sglang.srt.layers.layernorm import LayerNorm from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz @@ -226,6 +235,90 @@ class Indexer(MultiPlatformOp): self.scale_fmt = scale_fmt self.softmax_scale = self.head_dim**-0.5 + def _maybe_materialize_shared_index_buffer( + self, + forward_batch: ForwardBatch, + layer_id: int, + logical_page_table: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + index_buffer = forward_batch.token_to_kv_pool.get_index_k_with_scale_buffer( + layer_id=layer_id + ) + if not forward_batch.uses_cp_shared_kv: + return index_buffer, logical_page_table + + assert forward_batch.cp_shared_kv_layout is not None + layout = forward_batch.cp_shared_kv_layout + if cp_shared_kv_debug_enabled(): + cp_shared_kv_debug_log( + "index_materialize_call", + "NSA index materialize call cp_rank=%s layer=%s logical_page_table=%s", + layout.cp_rank, + layer_id, + tensor_debug_summary(logical_page_table), + ) + materialized, dense_pages = materialize_shared_paged_buffer( + page_buffer=index_buffer, + logical_pages=logical_page_table, + layout=layout, + ) + if cp_shared_kv_debug_enabled(): + cp_shared_kv_debug_log( + "index_materialize_result", + "NSA index materialize result cp_rank=%s layer=%s dense_pages=%s buffer_ck=%s", + layout.cp_rank, + layer_id, + tensor_debug_summary(dense_pages), + tensor_debug_checksum(materialized), + ) + return materialized, dense_pages + + def _filter_shared_index_write( + self, + forward_batch: ForwardBatch, + key: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + out_loc = forward_batch.out_cache_loc + if not forward_batch.uses_cp_shared_kv: + return out_loc, key + + if key.shape[0] != out_loc.numel(): + raise RuntimeError( + "CP shared KV index write shape mismatch: " + f"key_tokens={key.shape[0]} out_cache_loc={out_loc.numel()} " + f"forward_mode={forward_batch.forward_mode}" + ) + + assert forward_batch.cp_shared_kv_layout is not None + layout = forward_batch.cp_shared_kv_layout + owned_mask, physical_out_loc = filter_owned_logical_locs( + out_loc, + layout, + ) + if cp_shared_kv_debug_enabled(): + valid_mask = out_loc > 0 + owned_count = int(owned_mask.sum().item()) + valid_count = int(valid_mask.sum().item()) + cp_shared_kv_debug_log( + "index_write_filter", + "NSA index shard write filter cp_rank=%s forward_mode=%s logical_out_loc=%s " + "valid=%s owned=%s physical_out_loc=%s key_ck=%s key_valid_ck=%s " + "key_owned_ck=%s", + layout.cp_rank, + forward_batch.forward_mode, + tensor_debug_summary(out_loc), + valid_count, + owned_count, + tensor_debug_summary(physical_out_loc), + tensor_debug_checksum(key), + tensor_debug_checksum(key[valid_mask].contiguous()), + tensor_debug_checksum(key[owned_mask].contiguous()), + ) + if physical_out_loc.numel() == out_loc.numel(): + return physical_out_loc, key + + return physical_out_loc, key[owned_mask].contiguous() + @contextlib.contextmanager def _with_real_sm_count(self): # When pipeline parallelism is enabled, each PP rank initiates a recv operation after the _pp_launch_batch @@ -384,8 +477,10 @@ class Indexer(MultiPlatformOp): block_tables = metadata.get_page_table_64() max_seq_len = block_tables.shape[1] * page_size - kv_cache_fp8 = forward_batch.token_to_kv_pool.get_index_k_with_scale_buffer( - layer_id=layer_id + kv_cache_fp8, block_tables = self._maybe_materialize_shared_index_buffer( + forward_batch, + layer_id, + block_tables, ) blocksize = page_size @@ -546,12 +641,18 @@ class Indexer(MultiPlatformOp): indexer_seq_lens_cpu = metadata.get_indexer_seq_len_cpu() seq_len_sum = torch.sum(indexer_seq_lens_cpu).item() max_seq_len = torch.max(indexer_seq_lens_cpu).item() - k_fp8, k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_buffer( + index_buffer, block_tables = self._maybe_materialize_shared_index_buffer( + forward_batch, layer_id, - metadata.get_indexer_seq_len(), block_tables, - seq_len_sum, - max_seq_len, + ) + k_fp8, k_scale = index_buf_accessor.GetKAndS.execute( + forward_batch.token_to_kv_pool, + index_buffer, + page_indices=block_tables, + seq_len_tensor=metadata.get_indexer_seq_len(), + seq_len_sum=seq_len_sum, + max_seq_len=max_seq_len, ) if _is_fp8_fnuz: k_fp8 = k_fp8.view(torch.float8_e4m3fnuz) @@ -738,6 +839,11 @@ class Indexer(MultiPlatformOp): batch_idx_list = [] block_tables = metadata.get_page_table_64() + index_buffer, block_tables = self._maybe_materialize_shared_index_buffer( + forward_batch, + layer_id, + block_tables, + ) assert ( forward_batch.seq_lens_cpu is not None @@ -754,15 +860,17 @@ class Indexer(MultiPlatformOp): end_seq_position += pre_chunk_offset if offset == 0 and batch_idx != 0: offset += forward_batch.extend_seq_lens_cpu[batch_idx - 1] - k_fp8 = forward_batch.token_to_kv_pool.get_index_k_continuous( - layer_id, - end_seq_position, - block_tables[batch_idx], + k_fp8 = index_buf_accessor.GetK.execute( + forward_batch.token_to_kv_pool, + index_buffer, + seq_len=end_seq_position, + page_indices=block_tables[batch_idx], ) - k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_continuous( - layer_id, - end_seq_position, - block_tables[batch_idx], + k_scale = index_buf_accessor.GetS.execute( + forward_batch.token_to_kv_pool, + index_buffer, + seq_len=end_seq_position, + page_indices=block_tables[batch_idx], ) extend_seq_len = end_seq_position - start_seq_position @@ -810,32 +918,54 @@ class Indexer(MultiPlatformOp): batch_idx_list=batch_idx_list, ) else: - kv_len = ( + cp_kv_end = ( forward_batch.seq_lens_cpu[0].item() - forward_batch.extend_seq_lens_cpu[0] + kv_len ) - k_fp8 = forward_batch.token_to_kv_pool.get_index_k_continuous( - layer_id, - kv_len, - block_tables[0], + page_table_1 = metadata.get_page_table_1() + logical_kv_limit = min( + int(forward_batch.seq_lens_cpu[0].item()), + int(page_table_1.shape[1]), ) - k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_continuous( - layer_id, - kv_len, - block_tables[0], + ke_offset = torch.arange( + (cp_kv_end - actual_seq_q) + 1, + cp_kv_end + 1, + dtype=torch.int32, + device="cuda", + ) + valid_q_mask = ke_offset <= logical_kv_limit + valid_q_count = int(valid_q_mask.sum().item()) + topk_result = torch.full( + (actual_seq_q, self.index_topk), + -1, + dtype=torch.int32, + device=q_fp8.device, + ) + if valid_q_count == 0: + return topk_result + + q_fp8 = q_fp8[valid_q_mask].contiguous() + weights = weights[valid_q_mask].contiguous() + ke_offset = ke_offset[valid_q_mask].contiguous() + kv_len = min(cp_kv_end, logical_kv_limit) + k_fp8 = index_buf_accessor.GetK.execute( + forward_batch.token_to_kv_pool, + index_buffer, + seq_len=kv_len, + page_indices=block_tables[0], + ) + k_scale = index_buf_accessor.GetS.execute( + forward_batch.token_to_kv_pool, + index_buffer, + seq_len=kv_len, + page_indices=block_tables[0], ) k_fp8 = k_fp8.view(torch.float8_e4m3fn) k_scale = k_scale.view(torch.float32).squeeze(-1) kv_fp8 = (k_fp8, k_scale) - ks = torch.full((actual_seq_q,), offset, dtype=torch.int32, device="cuda") - ke_offset = torch.arange( - (kv_len - actual_seq_q) + 1, - kv_len + 1, - dtype=torch.int32, - device="cuda", - ) + ks = torch.full((valid_q_count,), offset, dtype=torch.int32, device="cuda") ke = ks + ke_offset with self._with_real_sm_count(): @@ -847,16 +977,17 @@ class Indexer(MultiPlatformOp): ke, clean_logits=False, ) - actual_seq_q = torch.tensor([actual_seq_q], dtype=torch.int32).to( + actual_seq_q_tensor = torch.tensor([valid_q_count], dtype=torch.int32).to( device="cuda", non_blocking=True ) - topk_result = metadata.topk_transform( + valid_topk_result = metadata.topk_transform( logits, self.index_topk, ks=ks, - cu_seqlens_q=actual_seq_q, + cu_seqlens_q=actual_seq_q_tensor, ke_offset=ke_offset, ) + topk_result[valid_q_mask] = valid_topk_result return topk_result @@ -890,6 +1021,11 @@ class Indexer(MultiPlatformOp): 0, block_tables.shape[-1], page_size, device="cuda" ) block_tables = block_tables[:, strided_indices] // page_size + index_buffer, block_tables = self._maybe_materialize_shared_index_buffer( + forward_batch, + layer_id, + block_tables, + ) q_len_start = 0 @@ -908,15 +1044,17 @@ class Indexer(MultiPlatformOp): weights_partial = weights[q_len_start:q_len_end] weights_partial = weights_partial.squeeze(-1).unsqueeze(0).contiguous() - k_fp8 = forward_batch.token_to_kv_pool.get_index_k_continuous( - layer_id, - seq_len, - block_tables[i], + k_fp8 = index_buf_accessor.GetK.execute( + forward_batch.token_to_kv_pool, + index_buffer, + seq_len=seq_len, + page_indices=block_tables[i], ) - k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_continuous( - layer_id, - seq_len, - block_tables[i], + k_scale = index_buf_accessor.GetS.execute( + forward_batch.token_to_kv_pool, + index_buffer, + seq_len=seq_len, + page_indices=block_tables[i], ) k_fp8 = k_fp8.view(torch.float8_e4m3fn).unsqueeze(0).contiguous() @@ -957,6 +1095,11 @@ class Indexer(MultiPlatformOp): Preferred: fused_store_index_k_cache(key, cache, out_cache_loc, page_size) Fallback : act_quant(key) + token_to_kv_pool.set_index_k_scale_buffer(...) """ + out_loc, key = self._filter_shared_index_write(forward_batch, key) + if out_loc.numel() == 0: + return + if not out_loc.is_contiguous(): + out_loc = out_loc.contiguous() # Fast path: JIT fused store (CUDA, page_size=64, non-fnuz) if ( @@ -964,7 +1107,7 @@ class Indexer(MultiPlatformOp): and (not _is_fp8_fnuz) and can_use_nsa_fused_store( key.dtype, - forward_batch.out_cache_loc.dtype, + out_loc.dtype, forward_batch.token_to_kv_pool.page_size, ) ): @@ -975,7 +1118,7 @@ class Indexer(MultiPlatformOp): fused_store_index_k_cache( key, buf, - forward_batch.out_cache_loc, + out_loc, forward_batch.token_to_kv_pool.page_size, ) return @@ -984,10 +1127,6 @@ class Indexer(MultiPlatformOp): assert act_quant is not None k_fp8, k_scale = act_quant(key, self.block_size, self.scale_fmt) - out_loc = forward_batch.out_cache_loc - if not out_loc.is_contiguous(): - out_loc = out_loc.contiguous() - forward_batch.token_to_kv_pool.set_index_k_scale_buffer( layer_id=layer_id, loc=out_loc, diff --git a/python/sglang/srt/layers/attention/nsa/transform_index.py b/python/sglang/srt/layers/attention/nsa/transform_index.py index 10b1068f5..4764de480 100644 --- a/python/sglang/srt/layers/attention/nsa/transform_index.py +++ b/python/sglang/srt/layers/attention/nsa/transform_index.py @@ -29,7 +29,7 @@ def transform_index_page_table_decode_kernel( offset = tl.arange(0, TOPK) # topk should be 2048 loaded_topk_indices = tl.load(topk_indices_ptr + offset) - mask = loaded_topk_indices >= 0 + mask = (loaded_topk_indices >= 0) & (loaded_topk_indices < max_seqlen_k) loaded_kv_indices = tl.load(page_table_ptr + loaded_topk_indices, mask=mask) tl.store(result_ptr + offset, loaded_kv_indices, mask=mask) tl.store(result_ptr + offset, -1, mask=~mask) @@ -102,13 +102,16 @@ def transform_index_page_table_decode_ref( if result is None: result = torch.empty_like(topk_indices, dtype=torch.int32) assert result.shape == topk_indices.shape + max_seqlen_k = page_table.shape[1] + valid_mask = (topk_indices >= 0) & (topk_indices < max_seqlen_k) + safe_indices = torch.where(valid_mask, topk_indices, torch.zeros_like(topk_indices)) torch.gather( page_table.to(result.dtype), dim=1, - index=topk_indices.clamp(min=0), + index=safe_indices, out=result, ) - result[topk_indices < 0] = -1 + result[~valid_mask] = -1 return result diff --git a/python/sglang/srt/layers/attention/nsa/utils.py b/python/sglang/srt/layers/attention/nsa/utils.py index 6ebc78114..a37dac0ef 100644 --- a/python/sglang/srt/layers/attention/nsa/utils.py +++ b/python/sglang/srt/layers/attention/nsa/utils.py @@ -609,6 +609,61 @@ def _round_robin_collect_last_token( return gathered[owner * bs : owner * bs + bs] +def _get_in_seq_last_token_owner_and_offset( + split_list: List[int], + cp_size: int, + actual_token_count: int, +) -> Tuple[int, int]: + """Return the CP rank and local offset for the real last token. + + In in-seq split, each rank owns two logical segments: + rank r gets segment r followed by segment (2 * cp_size - r - 1). + `split_list` is built from the padded model input length, while + `actual_token_count` is the non-padded extend length. The real last token + can therefore sit in a middle segment instead of rank 0's trailing segment. + """ + if cp_size <= 0: + raise ValueError(f"cp_size must be positive, got {cp_size}") + cp_segment_num = cp_size * 2 + if len(split_list) != cp_segment_num: + raise ValueError( + "in-seq split_list length must equal 2 * cp_size, " + f"got len={len(split_list)} cp_size={cp_size}" + ) + if actual_token_count <= 0: + raise ValueError( + f"actual_token_count must be positive, got {actual_token_count}" + ) + + total_split_tokens = sum(split_list) + if actual_token_count > total_split_tokens: + raise ValueError( + "actual_token_count exceeds split capacity: " + f"actual={actual_token_count} split_total={total_split_tokens}" + ) + + token_idx = actual_token_count - 1 + segment_idx = 0 + offset_in_segment = token_idx + for idx, segment_len in enumerate(split_list): + if offset_in_segment < segment_len: + segment_idx = idx + break + offset_in_segment -= segment_len + else: + raise ValueError( + "failed to locate last token in split_list: " + f"actual={actual_token_count} split_list={split_list}" + ) + + if segment_idx < cp_size: + return segment_idx, offset_in_segment + + owner = cp_segment_num - segment_idx - 1 + local_offset = split_list[owner] + offset_in_segment + return owner, local_offset + + def _in_seq_collect_last_token( hidden_states: torch.Tensor, forward_batch: "ForwardBatch", @@ -617,10 +672,18 @@ def _in_seq_collect_last_token( cp_rank = get_attention_cp_rank() bs = len(forward_batch.extend_seq_lens_cpu) owner = 0 - cp_group = get_attention_cp_group() + local_offset = hidden_states.shape[0] - bs - if cp_rank == owner and hidden_states.shape[0] > 0: - local_last = hidden_states[-bs:].contiguous() + if bs == 1 and forward_batch.nsa_cp_metadata is not None: + actual_token_count = sum(int(x) for x in forward_batch.extend_seq_lens_cpu) + owner, local_offset = _get_in_seq_last_token_owner_and_offset( + split_list=forward_batch.nsa_cp_metadata.split_list, + cp_size=cp_size, + actual_token_count=actual_token_count, + ) + + if cp_rank == owner and hidden_states.shape[0] >= local_offset + bs: + local_last = hidden_states[local_offset : local_offset + bs].contiguous() else: local_last = hidden_states.new_zeros((bs, hidden_states.shape[1])) diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index ad0d5840c..0c981ba0d 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -9,6 +9,14 @@ import torch from sglang.srt.configs.model_config import get_nsa_index_topk, is_deepseek_nsa from sglang.srt.environ import envs from sglang.srt.layers.attention.base_attn_backend import AttentionBackend +from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( + cp_shared_kv_debug_enabled, + cp_shared_kv_debug_log, + filter_owned_logical_locs, + materialize_shared_token_kv_buffer, + tensor_debug_checksum, + tensor_debug_summary, +) from sglang.srt.layers.attention.nsa.dequant_k_cache import dequantize_k_cache_paged from sglang.srt.layers.attention.nsa.nsa_backend_mtp_precompute import ( NativeSparseAttnBackendMTPPrecomputeMixin, @@ -151,6 +159,127 @@ class TopkTransformMethod(IntEnum): RAGGED = auto() +def _is_cuda_stream_capturing() -> bool: + if not torch.cuda.is_available(): + return False + try: + return torch.cuda.is_current_stream_capturing() + except RuntimeError: + return False + + +def _validate_paged_topk_transform_inputs( + score: torch.Tensor, + lengths: torch.Tensor | int, + page_table_size_1: torch.Tensor, + cu_seqlens_q: Optional[torch.Tensor], + row_starts: Optional[torch.Tensor], +) -> None: + if score.ndim != 2: + raise RuntimeError(f"NSA PAGED topk expects 2D score, got {score.shape=}") + if page_table_size_1.ndim != 2: + raise RuntimeError( + "NSA PAGED topk expects 2D page_table_size_1, " + f"got {page_table_size_1.shape=}" + ) + if cu_seqlens_q is None: + raise RuntimeError("NSA PAGED fused topk requires cu_seqlens_q") + + score_rows, score_cols = score.shape + if isinstance(lengths, torch.Tensor): + flat_lengths = lengths.reshape(-1) + if flat_lengths.numel() != score_rows: + raise RuntimeError( + "NSA PAGED fused topk length/score row mismatch: " + f"lengths={flat_lengths.numel()} score_rows={score_rows} " + f"score_shape={tuple(score.shape)} page_table_shape={tuple(page_table_size_1.shape)}" + ) + invalid_len = (flat_lengths < 0) | (flat_lengths > page_table_size_1.shape[1]) + if torch.any(invalid_len): + bad = flat_lengths[invalid_len] + raise RuntimeError( + "NSA PAGED fused topk lengths exceed page_table width. " + f"bad_min={int(bad.min().item())} bad_max={int(bad.max().item())} " + f"page_table_width={page_table_size_1.shape[1]} " + f"score_shape={tuple(score.shape)} page_table_shape={tuple(page_table_size_1.shape)} " + f"cu_seqlens_q_shape={tuple(cu_seqlens_q.shape)}" + ) + if row_starts is not None: + flat_starts = row_starts.reshape(-1) + if flat_starts.numel() != score_rows: + raise RuntimeError( + "NSA PAGED fused topk row_starts/score row mismatch: " + f"row_starts={flat_starts.numel()} score_rows={score_rows}" + ) + invalid_score_range = (flat_starts < 0) | (flat_starts + flat_lengths > score_cols) + if torch.any(invalid_score_range): + bad_starts = flat_starts[invalid_score_range] + bad_lengths = flat_lengths[invalid_score_range] + raise RuntimeError( + "NSA PAGED fused topk score range exceeds logits width. " + f"row_start_min={int(bad_starts.min().item())} " + f"row_start_max={int(bad_starts.max().item())} " + f"length_min={int(bad_lengths.min().item())} " + f"length_max={int(bad_lengths.max().item())} " + f"score_cols={score_cols} score_shape={tuple(score.shape)}" + ) + else: + invalid_score_range = flat_lengths > score_cols + if torch.any(invalid_score_range): + bad_lengths = flat_lengths[invalid_score_range] + raise RuntimeError( + "NSA PAGED fused topk lengths exceed logits width. " + f"length_min={int(bad_lengths.min().item())} " + f"length_max={int(bad_lengths.max().item())} " + f"score_cols={score_cols} score_shape={tuple(score.shape)}" + ) + else: + if lengths < 0 or lengths > page_table_size_1.shape[1]: + raise RuntimeError( + "NSA PAGED fused topk scalar length exceeds page_table width. " + f"length={lengths} page_table_width={page_table_size_1.shape[1]}" + ) + + if cu_seqlens_q.numel() != page_table_size_1.shape[0] + 1: + raise RuntimeError( + "NSA PAGED fused topk cu_seqlens/page_table batch mismatch: " + f"cu_seqlens_q={cu_seqlens_q.numel()} " + f"page_table_rows={page_table_size_1.shape[0]}" + ) + if int(cu_seqlens_q[-1].item()) != score_rows: + raise RuntimeError( + "NSA PAGED fused topk cu_seqlens does not cover score rows: " + f"cu_seqlens_q_last={int(cu_seqlens_q[-1].item())} score_rows={score_rows} " + f"cu_seqlens_q_shape={tuple(cu_seqlens_q.shape)}" + ) + + +def _validate_paged_topk_transform_output( + topk_result: torch.Tensor, + page_table_size_1: torch.Tensor, +) -> None: + if topk_result.numel() == 0 or page_table_size_1.numel() == 0: + return + + max_page_value = page_table_size_1.max() + invalid_values = (topk_result < -1) | (topk_result > max_page_value) + if not torch.any(invalid_values): + return + + bad_values = topk_result[invalid_values] + raise RuntimeError( + "NSA PAGED fused topk_transform produced values outside page_table_1. " + "This means the fused kernel returned entries that were not loaded from " + "the provided page table; inspect lengths/row_starts/cu_seqlens_q contract. " + f"bad_min={int(bad_values.min().item())} " + f"bad_max={int(bad_values.max().item())} " + f"page_table_min={int(page_table_size_1.min().item())} " + f"page_table_max={int(max_page_value.item())} " + f"topk_result_shape={tuple(topk_result.shape)} " + f"page_table_shape={tuple(page_table_size_1.shape)}" + ) + + @torch.compile def _compiled_cat(tensors: list[torch.Tensor], dim: int = -1) -> torch.Tensor: return torch.cat(tensors, dim=dim) @@ -178,6 +307,7 @@ class NSAIndexerMetadata(BaseIndexerMetadata): topk_transform_method: TopkTransformMethod paged_mqa_schedule_metadata: Optional[torch.Tensor] = None force_unfused_topk: bool = False + validate_paged_topk: bool = False def get_seqlens_int32(self) -> torch.Tensor: return self.attn_metadata.cache_seqlens_int32 @@ -251,7 +381,20 @@ class NSAIndexerMetadata(BaseIndexerMetadata): return fast_topk_v2(logits, seq_lens_topk, topk, row_starts=ks) elif self.topk_transform_method == TopkTransformMethod.PAGED: # NOTE(dark): if fused, we return a transformed page table directly - return fast_topk_transform_fused( + validate_paged_topk = ( + self.validate_paged_topk + and cp_shared_kv_debug_enabled() + and not _is_cuda_stream_capturing() + ) + if validate_paged_topk: + _validate_paged_topk_transform_inputs( + score=logits, + lengths=seq_lens_topk, + page_table_size_1=page_table_size_1, + cu_seqlens_q=cu_seqlens_q_topk, + row_starts=ks, + ) + topk_result = fast_topk_transform_fused( score=logits, lengths=seq_lens_topk, page_table_size_1=page_table_size_1, @@ -259,6 +402,9 @@ class NSAIndexerMetadata(BaseIndexerMetadata): topk=topk, row_starts=ks, ) + if validate_paged_topk: + _validate_paged_topk_transform_output(topk_result, page_table_size_1) + return topk_result elif self.topk_transform_method == TopkTransformMethod.RAGGED: return fast_topk_transform_ragged_fused( score=logits, @@ -368,6 +514,60 @@ class NativeSparseAttnBackend( ) return self._arange_buf[:l] + def _maybe_filter_shared_mla_kv_write( + self, + forward_batch: ForwardBatch, + cache_loc: torch.Tensor, + k: torch.Tensor, + k_rope: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + if not forward_batch.uses_cp_shared_kv: + return cache_loc, k, k_rope + + if k.shape[0] != cache_loc.numel() or k_rope.shape[0] != cache_loc.numel(): + raise RuntimeError( + "CP shared KV MLA write shape mismatch: " + f"k_tokens={k.shape[0]} k_rope_tokens={k_rope.shape[0]} " + f"cache_loc={cache_loc.numel()} forward_mode={forward_batch.forward_mode}" + ) + + assert forward_batch.cp_shared_kv_layout is not None + layout = forward_batch.cp_shared_kv_layout + owned_mask, physical_cache_loc = filter_owned_logical_locs( + cache_loc, + layout, + ) + if cp_shared_kv_debug_enabled(): + valid_mask = cache_loc > 0 + owned_count = int(owned_mask.sum().item()) + valid_count = int(valid_mask.sum().item()) + cp_shared_kv_debug_log( + "mla_write_filter", + "MLA shard write filter cp_rank=%s forward_mode=%s logical_loc=%s " + "valid=%s owned=%s physical_loc=%s k_ck=%s k_valid_ck=%s k_owned_ck=%s " + "k_rope_ck=%s k_rope_valid_ck=%s k_rope_owned_ck=%s", + layout.cp_rank, + forward_batch.forward_mode, + tensor_debug_summary(cache_loc), + valid_count, + owned_count, + tensor_debug_summary(physical_cache_loc), + tensor_debug_checksum(k), + tensor_debug_checksum(k[valid_mask].contiguous()), + tensor_debug_checksum(k[owned_mask].contiguous()), + tensor_debug_checksum(k_rope), + tensor_debug_checksum(k_rope[valid_mask].contiguous()), + tensor_debug_checksum(k_rope[owned_mask].contiguous()), + ) + if physical_cache_loc.numel() == cache_loc.numel(): + return physical_cache_loc, k, k_rope + + return ( + physical_cache_loc, + k[owned_mask].contiguous(), + k_rope[owned_mask].contiguous(), + ) + def _transform_table_1_to_real(self, page_table: torch.Tensor) -> torch.Tensor: page_size = self.real_page_size if page_size == 1: @@ -561,16 +761,17 @@ class NativeSparseAttnBackend( # Validate indices when logical tokens exceed physical capacity # This is likely to be triggered by PP with high kv reuse & parallelism - kv_cache_capacity = ( - forward_batch.token_to_kv_pool.size - + forward_batch.token_to_kv_pool.page_size - ) - if forward_batch.seq_lens_sum > kv_cache_capacity: - max_idx = page_table_1_flattened.max().item() - assert max_idx < kv_cache_capacity, ( - f"Invalid page table index: max={max_idx}, " - f"kv_cache_capacity={kv_cache_capacity}" + if not forward_batch.uses_cp_shared_kv: + kv_cache_capacity = ( + forward_batch.token_to_kv_pool.size + + forward_batch.token_to_kv_pool.page_size ) + if forward_batch.seq_lens_sum > kv_cache_capacity: + max_idx = page_table_1_flattened.max().item() + assert max_idx < kv_cache_capacity, ( + f"Invalid page table index: max={max_idx}, " + f"kv_cache_capacity={kv_cache_capacity}" + ) if topk_transform_method == TopkTransformMethod.RAGGED: topk_indices_offset = torch.repeat_interleave( @@ -1300,12 +1501,21 @@ class NativeSparseAttnBackend( if not layer.is_cross_attention else forward_batch.encoder_out_cache_loc ) - forward_batch.token_to_kv_pool.set_mla_kv_buffer( # type: ignore - layer, - cache_loc, - k, - k_rope, + cache_loc, k_to_write, k_rope_to_write = ( + self._maybe_filter_shared_mla_kv_write( + forward_batch, + cache_loc, + k, + k_rope, + ) ) + if cache_loc.numel() > 0: + forward_batch.token_to_kv_pool.set_mla_kv_buffer( # type: ignore + layer, + cache_loc, + k_to_write, + k_rope_to_write, + ) # Use MHA kernel if in MHA_ONE_SHOT mode if self.use_mha: @@ -1376,6 +1586,19 @@ class NativeSparseAttnBackend( ) ) + if ( + forward_batch.uses_cp_shared_kv + and topk_transform_method == TopkTransformMethod.PAGED + ): + assert forward_batch.cp_shared_kv_layout is not None + kv_cache, page_table_1 = materialize_shared_token_kv_buffer( + kv_cache=kv_cache, + logical_locs=page_table_1, + remap_logical_locs=metadata.page_table_1, + layout=forward_batch.cp_shared_kv_layout, + page_size=forward_batch.token_to_kv_pool.page_size, + ) + if nsa_impl == "tilelang": if q_rope is not None: q_all = concat_mla_absorb_q_general(q_nope, q_rope) @@ -1396,6 +1619,16 @@ class NativeSparseAttnBackend( self.forward_metadata.page_table_1_flattened ) assert page_table_1_flattened is not None + if forward_batch.uses_cp_shared_kv: + assert forward_batch.cp_shared_kv_layout is not None + kv_cache, page_table_1_flattened = ( + materialize_shared_token_kv_buffer( + kv_cache=kv_cache, + logical_locs=page_table_1_flattened, + layout=forward_batch.cp_shared_kv_layout, + page_size=forward_batch.token_to_kv_pool.page_size, + ) + ) kv_cache = dequantize_k_cache_paged( kv_cache, page_table_1_flattened ) @@ -1498,12 +1731,21 @@ class NativeSparseAttnBackend( if not layer.is_cross_attention else forward_batch.encoder_out_cache_loc ) - forward_batch.token_to_kv_pool.set_mla_kv_buffer( # type: ignore - layer, - cache_loc, - k, - k_rope, + cache_loc, k_to_write, k_rope_to_write = ( + self._maybe_filter_shared_mla_kv_write( + forward_batch, + cache_loc, + k, + k_rope, + ) ) + if cache_loc.numel() > 0: + forward_batch.token_to_kv_pool.set_mla_kv_buffer( # type: ignore + layer, + cache_loc, + k_to_write, + k_rope_to_write, + ) # Do absorbed multi-latent attention kv_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) @@ -1958,9 +2200,18 @@ class NativeSparseAttnBackend( if not layer.is_cross_attention else forward_batch.encoder_out_cache_loc ) - forward_batch.token_to_kv_pool.set_mla_kv_buffer( - layer, cache_loc, k, k_rope + cache_loc, k_to_write, k_rope_to_write = ( + self._maybe_filter_shared_mla_kv_write( + forward_batch, + cache_loc, + k, + k_rope, + ) ) + if cache_loc.numel() > 0: + forward_batch.token_to_kv_pool.set_mla_kv_buffer( + layer, cache_loc, k_to_write, k_rope_to_write + ) k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) kv_cache = k_cache.view(-1, self.real_page_size, self.kv_cache_dim).unsqueeze(1) @@ -2122,6 +2373,7 @@ class NativeSparseAttnBackend( topk_transform_method=self.get_topk_transform_method(), paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata, force_unfused_topk=force_unfused, + validate_paged_topk=forward_batch.uses_cp_shared_kv, ) def _compute_flashmla_metadata(self, cache_seqlens: torch.Tensor, seq_len_q: int): diff --git a/python/sglang/srt/mem_cache/allocator.py b/python/sglang/srt/mem_cache/allocator.py index c0ce15e20..c7d7a1a9d 100644 --- a/python/sglang/srt/mem_cache/allocator.py +++ b/python/sglang/srt/mem_cache/allocator.py @@ -517,3 +517,41 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): def load_cpu_copy(self, kv_cache_cpu, indices): return self._kvcache.load_cpu_copy(kv_cache_cpu, indices) + + +class CPSharedPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator): + """Paged allocator that returns CP-group logical KV locations. + + The allocator tracks logical pages visible to scheduler/radix/req_to_token. + The physical KV pool is smaller; logical-to-physical translation happens at + actual KV buffer access time through CpSharedKVLayout. + """ + + def __init__( + self, + logical_size: int, + physical_size: int, + page_size: int, + dtype: torch.dtype, + device: str, + kvcache: KVCache, + need_sort: bool, + cp_size: int, + cp_rank: int, + ): + if logical_size % page_size != 0: + raise ValueError("logical_size must be page aligned") + if physical_size % page_size != 0: + raise ValueError("physical_size must be page aligned") + if logical_size != physical_size * cp_size: + raise ValueError( + "logical_size must equal physical_size * cp_size, got " + f"{logical_size=} {physical_size=} {cp_size=}" + ) + if not 0 <= cp_rank < cp_size: + raise ValueError(f"cp_rank must be in [0, {cp_size}), got {cp_rank}") + + super().__init__(logical_size, page_size, dtype, device, kvcache, need_sort) + self.physical_size = physical_size + self.cp_size = cp_size + self.cp_rank = cp_rank diff --git a/python/sglang/srt/mem_cache/cp_shared_kv_layout.py b/python/sglang/srt/mem_cache/cp_shared_kv_layout.py new file mode 100644 index 000000000..7a038fe84 --- /dev/null +++ b/python/sglang/srt/mem_cache/cp_shared_kv_layout.py @@ -0,0 +1,85 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np +import torch + + +@dataclass(frozen=True) +class CpSharedKVLayout: + """Map CP-group logical KV locations to per-rank physical KV locations. + + Page 0 is the existing dummy/padding page. It is preserved as physical page + 0 on every rank and is not owned by any real CP rank. Usable logical pages + start from page 1 and are interleaved across CP ranks. + """ + + page_size: int + cp_size: int + cp_rank: int + + def __post_init__(self): + if self.page_size <= 0: + raise ValueError(f"page_size must be positive, got {self.page_size}") + if self.cp_size <= 0: + raise ValueError(f"cp_size must be positive, got {self.cp_size}") + if not 0 <= self.cp_rank < self.cp_size: + raise ValueError( + f"cp_rank must be in [0, {self.cp_size}), got {self.cp_rank}" + ) + + def owner_for_logical_pages(self, logical_pages: torch.Tensor) -> torch.Tensor: + owners = torch.remainder(logical_pages - 1, self.cp_size) + return torch.where(logical_pages <= 0, torch.full_like(owners, -1), owners) + + def owned_pages_mask(self, logical_pages: torch.Tensor) -> torch.Tensor: + return self.owner_for_logical_pages(logical_pages) == self.cp_rank + + def owned_by_this_rank(self, logical_locs: torch.Tensor) -> torch.Tensor: + logical_pages = torch.div(logical_locs, self.page_size, rounding_mode="floor") + return self.owned_pages_mask(logical_pages) + + def logical_pages_to_physical(self, logical_pages: torch.Tensor) -> torch.Tensor: + physical_pages = ( + torch.div(logical_pages - 1, self.cp_size, rounding_mode="floor") + 1 + ) + return torch.where( + logical_pages == 0, torch.zeros_like(physical_pages), physical_pages + ) + + def logical_locs_to_physical(self, logical_locs: torch.Tensor) -> torch.Tensor: + logical_pages = torch.div(logical_locs, self.page_size, rounding_mode="floor") + offsets = torch.remainder(logical_locs, self.page_size) + return self.logical_pages_to_physical(logical_pages) * self.page_size + offsets + + def owner_for_logical_pages_np(self, logical_pages: np.ndarray) -> np.ndarray: + pages = np.asarray(logical_pages, dtype=np.int64) + owners = (pages - 1) % self.cp_size + owners = np.where(pages <= 0, -1, owners) + return owners.astype(np.int32, copy=False) + + def logical_pages_to_physical_np(self, logical_pages: np.ndarray) -> np.ndarray: + pages = np.asarray(logical_pages, dtype=np.int64) + physical = (pages - 1) // self.cp_size + 1 + physical = np.where(pages == 0, 0, physical) + return physical.astype(np.int32, copy=False) + + def filter_owned_pages_np( + self, + logical_pages: np.ndarray, + request_page_positions: np.ndarray, + ) -> tuple[np.ndarray, np.ndarray]: + pages = np.asarray(logical_pages, dtype=np.int64) + positions = np.asarray(request_page_positions, dtype=np.int64) + if pages.shape != positions.shape: + raise ValueError( + "logical_pages and request_page_positions must have the same shape, " + f"got {pages.shape} and {positions.shape}" + ) + + mask = self.owner_for_logical_pages_np(pages) == self.cp_rank + return ( + self.logical_pages_to_physical_np(pages[mask]), + positions[mask].astype(np.int32, copy=False), + ) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index bec834965..0ae3e6614 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -53,6 +53,7 @@ from sglang.srt.layers.dp_attention import ( set_is_extend_in_batch, ) from sglang.srt.layers.utils.cp_utils import ContextParallelMetadata +from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import ( ForwardBatchDeepSeekMHAMixin, ) @@ -419,6 +420,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): attn_cp_metadata: Optional[ContextParallelMetadata] = None # Record the split metadata of the sequence number of NSA context parallels. nsa_cp_metadata: Optional[NSAContextParallelMetadata] = None + uses_cp_shared_kv: bool = False + cp_shared_kv_layout: Optional[CpSharedKVLayout] = None # For hidden states before normal return_hidden_states_before_norm: bool = False @@ -479,6 +482,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): return_hidden_states_before_norm=batch.return_hidden_states_before_norm, rids=[req.rid for req in batch.reqs], ) + ret.uses_cp_shared_kv = model_runner.uses_cp_shared_kv + ret.cp_shared_kv_layout = model_runner.cp_shared_kv_layout device = model_runner.device if batch.extend_input_logprob_token_ids is not None: diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 40c204468..4ae426b3c 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -99,6 +99,7 @@ from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp from sglang.srt.layers.attention.tbo_backend import TboAttnBackend from sglang.srt.layers.dp_attention import ( DpPaddingMode, + get_attention_cp_rank, get_attention_tp_group, get_attention_tp_size, initialize_dp_attention, @@ -338,6 +339,9 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.page_size = server_args.page_size self.req_to_token_pool = req_to_token_pool self.token_to_kv_pool_allocator = token_to_kv_pool_allocator + self.physical_max_total_num_tokens = None + self.uses_cp_shared_kv = server_args.enable_nsa_prefill_cp_shared_kv + self.cp_shared_kv_layout = None self.is_hybrid_swa = model_config.is_hybrid_swa self.is_hybrid_swa_compress = model_config.is_hybrid_swa_compress self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA @@ -918,6 +922,14 @@ class ModelRunner(ModelRunnerKVCacheMixin): server_args=self.server_args, model_config=self.model_config, ) + if self.uses_cp_shared_kv: + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + self.cp_shared_kv_layout = CpSharedKVLayout( + page_size=self.page_size, + cp_size=self.server_args.attn_cp_size, + cp_rank=get_attention_cp_rank(), + ) if is_npu(): register_sgl_tp_rank(self.gpu_id) @@ -2223,6 +2235,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): global_forward_mode=capture_forward_mode, lora_ids=lora_ids, ) + forward_batch.uses_cp_shared_kv = self.uses_cp_shared_kv + forward_batch.cp_shared_kv_layout = self.cp_shared_kv_layout if lora_ids is not None: self.lora_manager.prepare_lora_batch(forward_batch) diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index a021c6561..c5bf33336 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -8,8 +8,9 @@ import torch from sglang.srt.configs.model_config import get_nsa_index_head_dim, is_deepseek_nsa from sglang.srt.distributed.parallel_state import get_world_group -from sglang.srt.layers.dp_attention import get_attention_tp_size +from sglang.srt.layers.dp_attention import get_attention_cp_rank, get_attention_tp_size from sglang.srt.mem_cache.allocator import ( + CPSharedPagedTokenToKVPoolAllocator, PagedTokenToKVPoolAllocator, TokenToKVPoolAllocator, ) @@ -45,13 +46,16 @@ class MemoryPoolConfig: max_total_num_tokens: int max_running_requests: int + physical_max_total_num_tokens: Optional[int] = None full_max_total_num_tokens: Optional[int] = None swa_max_total_num_tokens: Optional[int] = None mem_fraction_static: Optional[float] = None def __post_init__(self): - if self.max_total_num_tokens <= 0: + if self.physical_max_total_num_tokens is None: + self.physical_max_total_num_tokens = self.max_total_num_tokens + if self.max_total_num_tokens <= 0 or self.physical_max_total_num_tokens <= 0: msg = "Not enough memory. Please try to increase --mem-fraction-static." if self.mem_fraction_static is not None: msg += f" Current value: mem_fraction_static={self.mem_fraction_static}" @@ -485,8 +489,13 @@ class ModelRunnerKVCacheMixin: end_layer=self.end_layer, ) elif self.use_mla_backend and is_nsa_model: + physical_kv_pool_size = ( + self.physical_max_total_num_tokens + if self.server_args.enable_nsa_prefill_cp_shared_kv + else self.max_total_num_tokens + ) nsa_pool_kwargs = dict( - size=self.max_total_num_tokens, + size=physical_kv_pool_size, page_size=self.page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, @@ -708,6 +717,20 @@ class ModelRunnerKVCacheMixin: kvcache=self.token_to_kv_pool, need_sort=need_sort, ) + elif self.server_args.enable_nsa_prefill_cp_shared_kv: + self.token_to_kv_pool_allocator = ( + CPSharedPagedTokenToKVPoolAllocator( + logical_size=self.max_total_num_tokens, + physical_size=self.physical_max_total_num_tokens, + page_size=self.page_size, + dtype=self.kv_cache_dtype, + device=self.device, + kvcache=self.token_to_kv_pool, + need_sort=need_sort, + cp_size=self.server_args.attn_cp_size, + cp_rank=get_attention_cp_rank(), + ) + ) else: self.token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator( self.max_total_num_tokens, @@ -785,6 +808,7 @@ class ModelRunnerKVCacheMixin: def _apply_memory_pool_config(self: ModelRunner, config: MemoryPoolConfig): """Apply a resolved MemoryPoolConfig and initialize pools.""" self.max_total_num_tokens = config.max_total_num_tokens + self.physical_max_total_num_tokens = config.physical_max_total_num_tokens self.max_running_requests = config.max_running_requests if self.is_hybrid_swa: self.full_max_total_num_tokens = config.full_max_total_num_tokens @@ -798,6 +822,13 @@ class ModelRunnerKVCacheMixin: """Profile GPU memory and resolve all pool parameters into a config.""" profiled_tokens = self.profile_max_num_token(pre_model_load_memory) token_capacity = self._resolve_token_capacity(profiled_tokens) + physical_token_capacity = token_capacity + if self.server_args.enable_nsa_prefill_cp_shared_kv: + assert self.server_args.page_size == self.page_size + physical_pages = physical_token_capacity // self.page_size + token_capacity = ( + physical_pages * self.server_args.attn_cp_size * self.page_size + ) full_tokens = None swa_tokens = None @@ -808,6 +839,7 @@ class ModelRunnerKVCacheMixin: return MemoryPoolConfig( max_total_num_tokens=token_capacity, + physical_max_total_num_tokens=physical_token_capacity, max_running_requests=self._resolve_max_num_reqs(token_capacity), full_max_total_num_tokens=full_tokens, swa_max_total_num_tokens=swa_tokens, @@ -826,6 +858,14 @@ class ModelRunnerKVCacheMixin: self._apply_memory_pool_config(self.memory_pool_config) + if self.server_args.enable_nsa_prefill_cp_shared_kv: + logger.info( + "CP shared KV enabled. physical_tokens_per_rank=%s, logical_tokens=%s, cp_size=%s, shard_policy=page_interleaved", + self.physical_max_total_num_tokens, + self.max_total_num_tokens, + self.server_args.attn_cp_size, + ) + logger.info( f"Memory pool end. " f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB" diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index cc067aa37..95aade3ce 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -667,6 +667,7 @@ class ServerArgs: # Context parallelism used in the long sequence prefill phase of DeepSeek v3.2 enable_nsa_prefill_context_parallel: bool = False nsa_prefill_cp_mode: str = "round-robin-split" + enable_nsa_prefill_cp_shared_kv: bool = False enable_fused_qk_norm_rope: bool = False enable_precise_embedding_interpolation: bool = False enable_fused_moe_sum_all_reduce: bool = False @@ -1565,7 +1566,32 @@ class ServerArgs: assert ( self.disaggregation_mode != "decode" ), "CP is only supported for prefill when PD disaggregation, please remove --enable-nsa-prefill-context-parallel." - + if self.enable_nsa_prefill_cp_shared_kv: + assert self.enable_nsa_prefill_context_parallel, ( + "--enable-nsa-prefill-cp-shared-kv requires " + "--enable-nsa-prefill-context-parallel." + ) + assert self.nsa_prefill_cp_mode == "in-seq-split", ( + "Phase 2 shared KV is initially validated only with " + "--nsa-prefill-cp-mode in-seq-split. The layout is " + "mode-neutral, but round-robin runtime wiring is not enabled yet." + ) + assert self.disaggregation_mode != "decode", ( + "Phase 2 shared KV supports prefill CP only; decode CP/shared KV is not supported." + ) + assert ( + self.page_size == 64 + ), "Phase 2 shared KV requires page_size=64 for NSA." + assert not self.enable_hisparse, ( + "Phase 2 shared KV initially supports NSATokenToKVPool only; " + "HiSparseNSATokenToKVPool is not wired yet." + ) + if self.disaggregation_mode == "prefill": + assert envs.SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER.get(), ( + "Phase 2 shared KV with PD disaggregation requires " + "SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER=1 so all " + "prefill CP ranks transfer shards." + ) else: # DeepSeek V3/R1/V3.1 if not self.disable_piecewise_cuda_graph: @@ -5556,6 +5582,14 @@ class ServerArgs: action="store_true", help="Enable context parallelism used in the long sequence prefill phase of DeepSeek v3.2.", ) + parser.add_argument( + "--enable-nsa-prefill-cp-shared-kv", + action="store_true", + help=( + "Enable Phase 2 shared/sharded persistent KV pool for NSA prefill CP. " + "Only prefill CP with NSA+MLA is supported; decode CP remains disabled." + ), + ) parser.add_argument( "--nsa-prefill-cp-mode", type=str, diff --git a/test/registered/unit/disaggregation/test_cp_shared_kv_transfer_mapping.py b/test/registered/unit/disaggregation/test_cp_shared_kv_transfer_mapping.py new file mode 100644 index 000000000..9379230f2 --- /dev/null +++ b/test/registered/unit/disaggregation/test_cp_shared_kv_transfer_mapping.py @@ -0,0 +1,61 @@ +import unittest + +import numpy as np + +from sglang.srt.disaggregation.utils import ( + filter_kv_pages_for_cp_shared_kv, + select_pages_by_request_positions, +) +from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-test-cpu") + + +class TestCPSharedKVTransferMapping(unittest.TestCase): + def test_filter_kv_pages_for_cp_shared_kv(self): + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=2) + logical_pages = [5, 6, 7, 8, 9, 10, 11] + src_pages, positions = filter_kv_pages_for_cp_shared_kv( + layout=layout, + logical_pages=logical_pages, + chunk_page_start=100, + ) + self.assertEqual(src_pages.tolist(), [2, 3]) + self.assertEqual(positions.tolist(), [102, 106]) + + def test_filter_kv_pages_for_cp_shared_kv_rejects_negative_pages(self): + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=2) + + with self.assertRaisesRegex( + ValueError, + "CP shared KV transfer got negative logical_pages", + ): + filter_kv_pages_for_cp_shared_kv( + layout=layout, + logical_pages=np.array([5, -1, 7], dtype=np.int32), + chunk_page_start=100, + ) + + def test_filter_nsa_state_pages_uses_physical_source_and_logical_destination(self): + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=2) + logical_state_pages = np.array([1, 2, 3, 4, 5, 6, 7, 8, 9], dtype=np.int32) + decode_state_pages = np.array([41, 42, 43, 44, 45, 46, 47, 48, 49], dtype=np.int32) + + prefill_physical_pages, logical_positions = filter_kv_pages_for_cp_shared_kv( + layout=layout, + logical_pages=logical_state_pages, + chunk_page_start=0, + ) + selected_decode_pages = select_pages_by_request_positions( + decode_state_pages, + logical_positions, + ) + + self.assertEqual(prefill_physical_pages.tolist(), [1, 2]) + self.assertEqual(logical_positions.tolist(), [2, 6]) + self.assertEqual(selected_decode_pages.tolist(), [43, 47]) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/layers/test_nsa_cp_utils.py b/test/registered/unit/layers/test_nsa_cp_utils.py new file mode 100644 index 000000000..e7e486fdc --- /dev/null +++ b/test/registered/unit/layers/test_nsa_cp_utils.py @@ -0,0 +1,41 @@ +import unittest + +from sglang.srt.layers.attention.nsa.utils import ( + _get_in_seq_last_token_owner_and_offset, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-test-cpu") + + +class TestNSAInSeqCPUtils(unittest.TestCase): + def test_last_token_owner_uses_actual_token_count_when_batch_is_padded(self): + # Padded prefill can have 64 model tokens while the real prompt has only + # 11 tokens. In in-seq split with cp_size=8, the real last token is in + # segment 2, not in rank 0's trailing padded segment. + split_list = [4] * 16 + + owner, local_offset = _get_in_seq_last_token_owner_and_offset( + split_list=split_list, + cp_size=8, + actual_token_count=11, + ) + + self.assertEqual(owner, 2) + self.assertEqual(local_offset, 2) + + def test_last_token_owner_keeps_existing_unpadded_fast_path_location(self): + split_list = [4] * 16 + + owner, local_offset = _get_in_seq_last_token_owner_and_offset( + split_list=split_list, + cp_size=8, + actual_token_count=64, + ) + + self.assertEqual(owner, 0) + self.assertEqual(local_offset, 7) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/layers/test_nsa_topk_transform.py b/test/registered/unit/layers/test_nsa_topk_transform.py new file mode 100644 index 000000000..33ebef59c --- /dev/null +++ b/test/registered/unit/layers/test_nsa_topk_transform.py @@ -0,0 +1,210 @@ +import unittest +from unittest.mock import patch + +import torch + +from sglang.srt.layers.attention.nsa_backend import ( + NSAIndexerMetadata, + NSAMetadata, + TopkTransformMethod, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-test-cpu") + + +class TestNSATopkTransform(unittest.TestCase): + def test_paged_topk_transform_raises_when_fused_output_is_not_from_page_table(self): + page_table = torch.tensor( + [ + [10, 11, 12, 13, 14, 15, 16], + [20, 21, 22, 23, 24, 25, 26], + ], + dtype=torch.int32, + ) + lengths_seen = {} + + def fake_fast_topk_transform_fused(**kwargs): + lengths_seen["value"] = kwargs["lengths"].clone() + return torch.tensor( + [ + [10, 1_039_799_618, 11, 12], + [20, 21, 22, 23], + ], + dtype=torch.int32, + ) + + metadata = NSAMetadata( + page_size=1, + cache_seqlens_int32=torch.tensor([7], dtype=torch.int32), + max_seq_len_q=1, + max_seq_len_k=7, + cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + cu_seqlens_k=torch.tensor([0, 7], dtype=torch.int32), + page_table_1=page_table, + real_page_table=page_table, + nsa_cache_seqlens_int32=torch.tensor([4, 7], dtype=torch.int32), + nsa_cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32), + nsa_cu_seqlens_k=torch.tensor([0, 4, 11], dtype=torch.int32), + nsa_extend_seq_lens_list=[2], + nsa_seqlens_expanded=torch.tensor([4, 7], dtype=torch.int32), + ) + indexer_metadata = NSAIndexerMetadata( + attn_metadata=metadata, + topk_transform_method=TopkTransformMethod.PAGED, + validate_paged_topk=True, + ) + + with patch( + "sgl_kernel.fast_topk_transform_fused", + side_effect=fake_fast_topk_transform_fused, + ), patch( + "sglang.srt.environ.envs.SGLANG_NSA_FUSE_TOPK.get", return_value=True + ), patch( + "sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", return_value=True + ): + with self.assertRaisesRegex( + RuntimeError, + "NSA PAGED fused topk_transform produced values outside page_table_1", + ): + indexer_metadata.topk_transform( + logits=torch.zeros((2, 7), dtype=torch.float32), + topk=4, + cu_seqlens_q=torch.tensor([1, 1], dtype=torch.int32), + ) + + self.assertEqual(lengths_seen["value"].tolist(), [4, 7]) + + def test_paged_topk_transform_rejects_lengths_exceeding_page_table_width(self): + page_table = torch.tensor([[10, 11, 12]], dtype=torch.int32) + metadata = NSAMetadata( + page_size=1, + cache_seqlens_int32=torch.tensor([3], dtype=torch.int32), + max_seq_len_q=1, + max_seq_len_k=3, + cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32), + page_table_1=page_table, + real_page_table=page_table, + nsa_cache_seqlens_int32=torch.tensor([4], dtype=torch.int32), + nsa_cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + nsa_cu_seqlens_k=torch.tensor([0, 4], dtype=torch.int32), + nsa_extend_seq_lens_list=[1], + nsa_seqlens_expanded=torch.tensor([4], dtype=torch.int32), + ) + indexer_metadata = NSAIndexerMetadata( + attn_metadata=metadata, + topk_transform_method=TopkTransformMethod.PAGED, + validate_paged_topk=True, + ) + + with patch( + "sgl_kernel.fast_topk_transform_fused", + side_effect=AssertionError("fused kernel should not be called"), + ), patch( + "sglang.srt.environ.envs.SGLANG_NSA_FUSE_TOPK.get", return_value=True + ), patch( + "sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", return_value=True + ): + with self.assertRaisesRegex( + RuntimeError, + "NSA PAGED fused topk lengths exceed page_table width", + ): + indexer_metadata.topk_transform( + logits=torch.zeros((1, 4), dtype=torch.float32), + topk=4, + ) + + def test_paged_topk_transform_skips_validation_during_cuda_graph_capture(self): + page_table = torch.tensor([[10, 11, 12]], dtype=torch.int32) + + def fake_fast_topk_transform_fused(**kwargs): + return torch.tensor([[1_039_799_618]], dtype=torch.int32) + + metadata = NSAMetadata( + page_size=1, + cache_seqlens_int32=torch.tensor([3], dtype=torch.int32), + max_seq_len_q=1, + max_seq_len_k=3, + cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32), + page_table_1=page_table, + real_page_table=page_table, + nsa_cache_seqlens_int32=torch.tensor([3], dtype=torch.int32), + nsa_cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + nsa_cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32), + nsa_extend_seq_lens_list=[1], + nsa_seqlens_expanded=torch.tensor([3], dtype=torch.int32), + ) + indexer_metadata = NSAIndexerMetadata( + attn_metadata=metadata, + topk_transform_method=TopkTransformMethod.PAGED, + validate_paged_topk=True, + ) + + with patch( + "sgl_kernel.fast_topk_transform_fused", + side_effect=fake_fast_topk_transform_fused, + ), patch( + "sglang.srt.layers.attention.nsa_backend._is_cuda_stream_capturing", + return_value=True, + ), patch( + "sglang.srt.environ.envs.SGLANG_NSA_FUSE_TOPK.get", return_value=True + ), patch( + "sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", return_value=True + ): + out = indexer_metadata.topk_transform( + logits=torch.zeros((1, 3), dtype=torch.float32), + topk=1, + ) + + self.assertEqual(out.tolist(), [[1_039_799_618]]) + + def test_paged_topk_transform_skips_validation_when_cp_shared_debug_disabled(self): + page_table = torch.tensor([[10, 11, 12]], dtype=torch.int32) + + def fake_fast_topk_transform_fused(**kwargs): + return torch.tensor([[1_039_799_618]], dtype=torch.int32) + + metadata = NSAMetadata( + page_size=1, + cache_seqlens_int32=torch.tensor([3], dtype=torch.int32), + max_seq_len_q=1, + max_seq_len_k=3, + cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32), + page_table_1=page_table, + real_page_table=page_table, + nsa_cache_seqlens_int32=torch.tensor([3], dtype=torch.int32), + nsa_cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32), + nsa_cu_seqlens_k=torch.tensor([0, 3], dtype=torch.int32), + nsa_extend_seq_lens_list=[1], + nsa_seqlens_expanded=torch.tensor([3], dtype=torch.int32), + ) + indexer_metadata = NSAIndexerMetadata( + attn_metadata=metadata, + topk_transform_method=TopkTransformMethod.PAGED, + validate_paged_topk=True, + ) + + with patch( + "sgl_kernel.fast_topk_transform_fused", + side_effect=fake_fast_topk_transform_fused, + ), patch( + "sglang.srt.layers.attention.nsa_backend._is_cuda_stream_capturing", + return_value=False, + ), patch( + "sglang.srt.environ.envs.SGLANG_NSA_FUSE_TOPK.get", return_value=True + ), patch( + "sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", return_value=False + ): + out = indexer_metadata.topk_transform( + logits=torch.zeros((1, 3), dtype=torch.float32), + topk=1, + ) + + self.assertEqual(out.tolist(), [[1_039_799_618]]) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/layers/test_nsa_transform_index.py b/test/registered/unit/layers/test_nsa_transform_index.py new file mode 100644 index 000000000..7384a7bff --- /dev/null +++ b/test/registered/unit/layers/test_nsa_transform_index.py @@ -0,0 +1,39 @@ +import unittest + +import torch + +from sglang.srt.layers.attention.nsa.transform_index import ( + transform_index_page_table_decode_ref, + transform_index_page_table_prefill_ref, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-test-cpu") + + +class TestNSATransformIndex(unittest.TestCase): + def test_decode_ref_masks_high_out_of_bounds_topk_indices(self): + page_table = torch.tensor([[10, 11, 12], [20, 21, 22]], dtype=torch.int32) + topk_indices = torch.tensor([[0, 2, 3, -1], [1, 4, 0, -1]], dtype=torch.int64) + + result = transform_index_page_table_decode_ref(page_table, topk_indices) + + self.assertEqual(result.tolist(), [[10, 12, -1, -1], [21, -1, 20, -1]]) + + def test_prefill_ref_masks_high_out_of_bounds_topk_indices(self): + page_table = torch.tensor([[10, 11, 12], [20, 21, 22]], dtype=torch.int32) + topk_indices = torch.tensor( + [[0, 3, -1], [2, 1, 9], [1, 0, -1]], dtype=torch.int64 + ) + + result = transform_index_page_table_prefill_ref( + page_table=page_table, + topk_indices=topk_indices, + extend_lens_cpu=[2, 1], + ) + + self.assertEqual(result.tolist(), [[10, -1, -1], [12, 11, -1], [21, 20, -1]]) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_cp_shared_kv_layout.py b/test/registered/unit/mem_cache/test_cp_shared_kv_layout.py new file mode 100644 index 000000000..b4f00efe0 --- /dev/null +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_layout.py @@ -0,0 +1,78 @@ +import unittest + +import numpy as np +import torch + +from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-test-cpu") + + +class TestCpSharedKVLayout(unittest.TestCase): + def test_page_owner_skips_dummy_page(self): + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=0) + pages = torch.tensor([-2, -1, 0, 1, 2, 3, 4, 5, 8, 9], dtype=torch.int64) + owners = layout.owner_for_logical_pages(pages) + self.assertEqual(owners.tolist(), [-1, -1, -1, 0, 1, 2, 3, 0, 3, 0]) + + def test_logical_to_physical_pages_keeps_dummy_zero(self): + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=0) + pages = torch.tensor([0, 1, 2, 3, 4, 5, 8, 9], dtype=torch.int64) + physical = layout.logical_pages_to_physical(pages) + self.assertEqual(physical.tolist(), [0, 1, 1, 1, 1, 2, 2, 3]) + + def test_owned_mask_and_loc_translation(self): + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=2) + locs = torch.tensor([0, 64, 128, 192, 256, 320, 384], dtype=torch.int64) + mask = layout.owned_by_this_rank(locs) + self.assertEqual( + mask.tolist(), [False, False, False, True, False, False, False] + ) + physical = layout.logical_locs_to_physical(locs[mask]) + self.assertEqual(physical.tolist(), [64]) + + def test_negative_and_dummy_locs_are_never_owned(self): + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=2) + locs = torch.tensor([-1, 0, 1, 63, 64, 128, 192], dtype=torch.int64) + mask = layout.owned_by_this_rank(locs) + self.assertEqual( + mask.tolist(), [False, False, False, False, False, False, True] + ) + + def test_numpy_filter_returns_request_absolute_positions(self): + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=1) + logical_pages = np.array([1, 2, 3, 4, 5, 6, 7, 8, 9], dtype=np.int32) + request_positions = np.arange(10, 19, dtype=np.int32) + src_physical, positions = layout.filter_owned_pages_np( + logical_pages, request_positions + ) + self.assertEqual(src_physical.tolist(), [1, 2]) + self.assertEqual(positions.tolist(), [11, 15]) + + +class TestCPSharedPagedAllocator(unittest.TestCase): + def test_shared_allocator_exposes_logical_capacity(self): + from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator + + allocator = CPSharedPagedTokenToKVPoolAllocator( + logical_size=64 * 8, + physical_size=64 * 2, + page_size=64, + dtype=torch.bfloat16, + device="cpu", + kvcache=None, + need_sort=False, + cp_size=4, + cp_rank=0, + ) + self.assertEqual(allocator.available_size(), 64 * 8) + locs = allocator.alloc(64 * 2) + self.assertEqual(locs.numel(), 64 * 2) + self.assertEqual(allocator.available_size(), 64 * 6) + allocator.free(locs) + self.assertEqual(allocator.available_size(), 64 * 8) + + +if __name__ == "__main__": + unittest.main() 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 new file mode 100644 index 000000000..a08385837 --- /dev/null +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py @@ -0,0 +1,435 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="stage-a-test-cpu") + + +class TestCpSharedKVRuntimeHelpers(unittest.TestCase): + def test_all_reduce_uses_group_fast_path_for_float_buffers(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime + + class DummyGroup: + device_group = object() + + def __init__(self): + self.called = False + + def all_reduce(self, tensor): + self.called = True + + dummy_group = DummyGroup() + buffer = torch.ones((2, 3), dtype=torch.float32) + + with patch.object( + runtime, "get_attention_cp_group", return_value=dummy_group + ), patch("torch.distributed.all_reduce") as dist_all_reduce: + out = runtime._all_reduce_materialized_buffer(buffer, cp_size=2) + + self.assertIs(out, buffer) + self.assertTrue(dummy_group.called) + dist_all_reduce.assert_not_called() + + def test_all_reduce_copies_out_of_place_group_result_for_float_buffers(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime + + class DummyGroup: + device_group = object() + + def all_reduce(self, tensor): + return tensor + 7 + + dummy_group = DummyGroup() + buffer = torch.ones((2, 3), dtype=torch.float32) + + with patch.object( + runtime, "get_attention_cp_group", return_value=dummy_group + ), patch("torch.distributed.all_reduce") as dist_all_reduce: + out = runtime._all_reduce_materialized_buffer(buffer, cp_size=2) + + self.assertIs(out, buffer) + self.assertTrue(torch.equal(buffer, torch.full((2, 3), 8.0))) + dist_all_reduce.assert_not_called() + + def test_all_reduce_bypasses_custom_group_for_byte_buffers(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime + + class DummyGroup: + device_group = object() + + def all_reduce(self, tensor): + raise AssertionError("byte buffers must bypass custom all-reduce") + + dummy_group = DummyGroup() + buffer = torch.ones((2, 3), dtype=torch.uint8) + + with patch.object( + runtime, "get_attention_cp_group", return_value=dummy_group + ), patch("torch.distributed.all_reduce") as dist_all_reduce: + out = runtime._all_reduce_materialized_buffer(buffer, cp_size=2) + + self.assertIs(out, buffer) + dist_all_reduce.assert_called_once() + self.assertIs(dist_all_reduce.call_args.args[0], buffer) + self.assertIs(dist_all_reduce.call_args.kwargs["group"], dummy_group.device_group) + + def test_build_dense_page_remap_preserves_sentinels(self): + from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( + build_dense_page_remap, + ) + + logical_pages = torch.tensor([0, 5, 2, -1, 5, 9, 2], dtype=torch.int32) + unique_pages, dense_pages = build_dense_page_remap(logical_pages) + + self.assertEqual(unique_pages.tolist(), [2, 5, 9]) + self.assertEqual(dense_pages.tolist(), [0, 2, 1, -1, 2, 3, 1]) + + def test_remap_logical_locs_to_dense_locs(self): + from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( + remap_logical_locs_to_dense_locs, + ) + + logical_locs = torch.tensor([0, 1, 8, 9, -1, 20], dtype=torch.int32) + unique_pages = torch.tensor([2, 5], dtype=torch.int32) + + dense_locs = remap_logical_locs_to_dense_locs( + logical_locs, + unique_logical_pages=unique_pages, + page_size=4, + ) + + self.assertEqual(dense_locs.tolist(), [0, 1, 4, 5, -1, 8]) + + def test_materialize_local_token_kv_pages(self): + from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( + build_dense_page_remap, + materialize_local_token_kv_pages, + ) + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + layout = CpSharedKVLayout(page_size=4, cp_size=4, cp_rank=1) + kv_cache = torch.arange(0, 16 * 2, dtype=torch.float32).view(16, 1, 2) + logical_pages = torch.tensor([1, 2, 5, 7, 10], dtype=torch.int32) + unique_pages, _ = build_dense_page_remap(logical_pages) + + dense_kv = materialize_local_token_kv_pages( + kv_cache=kv_cache, + unique_logical_pages=unique_pages, + layout=layout, + page_size=4, + ) + + self.assertEqual(list(dense_kv.shape), [24, 1, 2]) + self.assertTrue(torch.equal(dense_kv[8:12], kv_cache[4:8])) + self.assertTrue(torch.equal(dense_kv[20:24], kv_cache[12:16])) + self.assertEqual(float(dense_kv[4:8].abs().sum().item()), 0.0) + + def test_materialize_local_paged_index_buffer(self): + from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( + build_dense_page_remap, + materialize_local_paged_buffer, + ) + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + layout = CpSharedKVLayout(page_size=4, cp_size=4, cp_rank=1) + page_buffer = torch.arange(0, 5 * 3, dtype=torch.uint8).view(5, 3) + logical_pages = torch.tensor([1, 2, 5, 7, 10], dtype=torch.int32) + unique_pages, dense_pages = build_dense_page_remap(logical_pages) + + dense_page_buffer = materialize_local_paged_buffer( + page_buffer=page_buffer, + unique_logical_pages=unique_pages, + layout=layout, + ) + + self.assertEqual(dense_pages.tolist(), [1, 2, 3, 4, 5]) + self.assertEqual(list(dense_page_buffer.shape), [6, 3]) + self.assertTrue(torch.equal(dense_page_buffer[2], page_buffer[1])) + self.assertTrue(torch.equal(dense_page_buffer[5], page_buffer[3])) + self.assertEqual(int(dense_page_buffer[1].sum().item()), 0) + + def test_filter_owned_logical_locs_for_persistent_writes(self): + from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( + filter_owned_logical_locs, + ) + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=2) + logical_locs = torch.tensor( + [0, 64, 128, 192, 256, 320, 384, 448, 449], + dtype=torch.int64, + ) + + owned_mask, physical_locs = filter_owned_logical_locs(logical_locs, layout) + + self.assertEqual( + owned_mask.tolist(), + [False, False, False, True, False, False, False, True, True], + ) + self.assertEqual(physical_locs.tolist(), [64, 128, 129]) + + def test_materialize_token_kv_raises_on_locs_outside_physical_pool(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 + + layout = CpSharedKVLayout(page_size=4, cp_size=4, cp_rank=1) + kv_cache = torch.arange(0, 12 * 2, dtype=torch.float32).view(12, 1, 2) + logical_locs = torch.tensor([4, 8, 20, 40, -1], dtype=torch.int64) + + with patch.object( + runtime, "cp_shared_kv_debug_enabled", return_value=True + ), patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x): + with self.assertRaisesRegex( + RuntimeError, + "CP shared KV materialize got logical token locs outside the physical pool", + ): + runtime.materialize_shared_token_kv_buffer( + kv_cache=kv_cache, + logical_locs=logical_locs, + layout=layout, + page_size=4, + ) + + def test_materialize_token_kv_skips_physical_pool_validation_when_debug_disabled(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 + + layout = CpSharedKVLayout(page_size=4, cp_size=4, cp_rank=1) + kv_cache = torch.arange(0, 12 * 2, dtype=torch.float32).view(12, 1, 2) + # Logical page 9 maps outside the local physical pool, but is owned by + # rank 0. With debug disabled, this source-finding validation should + # not synchronize or raise on rank 1. + logical_locs = torch.tensor([4, 36, -1], dtype=torch.int64) + + with patch.object( + runtime, "cp_shared_kv_debug_enabled", return_value=False + ), patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x): + _, dense_locs = runtime.materialize_shared_token_kv_buffer( + kv_cache=kv_cache, + logical_locs=logical_locs, + layout=layout, + page_size=4, + ) + + self.assertEqual(dense_locs.tolist(), [4, 8, -1]) + + def test_debug_materialize_token_kv_allows_minus_one_sentinel(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 + + layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0) + kv_cache = torch.arange(0, 16, dtype=torch.float32).view(16, 1, 1) + logical_locs = torch.tensor([4, -1, 8], dtype=torch.int64) + + with patch.object(runtime, "cp_shared_kv_debug_enabled", return_value=True): + with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x): + _, dense_locs = runtime.materialize_shared_token_kv_buffer( + kv_cache=kv_cache, + logical_locs=logical_locs, + layout=layout, + page_size=4, + ) + + self.assertEqual(dense_locs.tolist(), [4, -1, 8]) + + def test_debug_materialize_token_kv_rejects_below_minus_one_locs(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 + + layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0) + kv_cache = torch.arange(0, 16, dtype=torch.float32).view(16, 1, 1) + logical_locs = torch.tensor([4, -2, 8], dtype=torch.int64) + + with patch.object(runtime, "cp_shared_kv_debug_enabled", return_value=True): + with self.assertRaisesRegex( + RuntimeError, + "CP shared KV token materialize got logical_locs below -1", + ): + runtime.materialize_shared_token_kv_buffer( + kv_cache=kv_cache, + logical_locs=logical_locs, + layout=layout, + page_size=4, + ) + + def test_debug_materialize_paged_buffer_rejects_negative_pages(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 + + layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0) + page_buffer = torch.arange(0, 4 * 3, dtype=torch.uint8).view(4, 3) + logical_pages = torch.tensor([1, -1, 2], dtype=torch.int32) + + with patch.object(runtime, "cp_shared_kv_debug_enabled", return_value=True): + with self.assertRaisesRegex( + RuntimeError, + "CP shared KV paged materialize got negative logical_pages", + ): + runtime.materialize_shared_paged_buffer( + page_buffer=page_buffer, + logical_pages=logical_pages, + layout=layout, + ) + + def test_debug_filter_owned_logical_locs_rejects_negative_write_locs(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 + + layout = CpSharedKVLayout(page_size=64, cp_size=4, cp_rank=2) + logical_locs = torch.tensor([192, -1], dtype=torch.int64) + + with patch.object(runtime, "cp_shared_kv_debug_enabled", return_value=True): + with self.assertRaisesRegex( + RuntimeError, + "CP shared KV persistent write got negative logical_locs", + ): + runtime.filter_owned_logical_locs(logical_locs, layout) + + def test_materialize_token_kv_uses_remap_source_for_target_subset(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 + + layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0) + kv_cache = torch.arange(0, 32, dtype=torch.float32).view(32, 1, 1) + logical_locs = torch.tensor([8, 20, -1], dtype=torch.int64) + remap_source_locs = torch.tensor([4, 8, 12, 16, 20], dtype=torch.int64) + + with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x): + dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer( + kv_cache=kv_cache, + logical_locs=logical_locs, + remap_logical_locs=remap_source_locs, + layout=layout, + page_size=4, + ) + + self.assertEqual(dense_locs.tolist(), [8, 20, -1]) + self.assertEqual(list(dense_kv.shape), [24, 1, 1]) + self.assertTrue(torch.equal(dense_kv[8:12], kv_cache[8:12])) + self.assertTrue(torch.equal(dense_kv[20:24], kv_cache[20:24])) + + def test_materialize_token_kv_keeps_dense_shape_for_shared_remap_source(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 + + layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0) + kv_cache = torch.arange(0, 32, dtype=torch.float32).view(32, 1, 1) + remap_source_locs = torch.tensor([4, 8, 12, 16, 20], dtype=torch.int64) + + with patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x): + dense_kv_a, dense_locs_a = runtime.materialize_shared_token_kv_buffer( + kv_cache=kv_cache, + logical_locs=torch.tensor([4, 20], dtype=torch.int64), + remap_logical_locs=remap_source_locs, + layout=layout, + page_size=4, + ) + dense_kv_b, dense_locs_b = runtime.materialize_shared_token_kv_buffer( + kv_cache=kv_cache, + logical_locs=torch.tensor([8, 16], dtype=torch.int64), + remap_logical_locs=remap_source_locs, + layout=layout, + page_size=4, + ) + + self.assertEqual(list(dense_kv_a.shape), [24, 1, 1]) + self.assertEqual(list(dense_kv_b.shape), [24, 1, 1]) + self.assertEqual(dense_locs_a.tolist(), [4, 20]) + self.assertEqual(dense_locs_b.tolist(), [8, 16]) + + def test_materialize_paged_buffer_raises_on_pages_outside_physical_pool(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 + + layout = CpSharedKVLayout(page_size=4, cp_size=4, cp_rank=1) + page_buffer = torch.arange(0, 3 * 4, dtype=torch.uint8).view(3, 4) + logical_pages = torch.tensor([1, 2, 6, 10], dtype=torch.int32) + + with patch.object( + runtime, "cp_shared_kv_debug_enabled", return_value=True + ), patch.object(runtime, "_all_reduce_materialized_buffer", lambda x, _: x): + with self.assertRaisesRegex( + RuntimeError, + "CP shared KV materialize got logical pages outside the physical page buffer", + ): + runtime.materialize_shared_paged_buffer( + page_buffer=page_buffer, + logical_pages=logical_pages, + layout=layout, + ) + + +class TestCpSharedKVLazyDebugLogging(unittest.TestCase): + def test_mla_write_filter_does_not_build_debug_summaries_when_debug_disabled(self): + from sglang.srt.layers.attention import nsa_backend + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1) + forward_batch = SimpleNamespace( + uses_cp_shared_kv=True, + cp_shared_kv_layout=layout, + forward_mode="test", + ) + cache_loc = torch.tensor([4, 8, 12, 16], dtype=torch.int64) + k = torch.arange(0, 8, dtype=torch.float32).view(4, 1, 2) + k_rope = torch.arange(0, 4, dtype=torch.float32).view(4, 1, 1) + + with patch.object( + nsa_backend, "tensor_debug_summary", side_effect=AssertionError + ), patch.object( + nsa_backend, "tensor_debug_checksum", side_effect=AssertionError + ), patch( + "sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", + return_value=False, + ): + physical_locs, k_to_write, k_rope_to_write = ( + nsa_backend.NativeSparseAttnBackend._maybe_filter_shared_mla_kv_write( + None, + forward_batch, + cache_loc, + k, + k_rope, + ) + ) + + self.assertEqual(physical_locs.tolist(), [4, 8]) + self.assertEqual(k_to_write.shape[0], 2) + self.assertEqual(k_rope_to_write.shape[0], 2) + + def test_index_write_filter_does_not_build_debug_summaries_when_debug_disabled(self): + from sglang.srt.layers.attention.nsa import nsa_indexer + from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout + + layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1) + forward_batch = SimpleNamespace( + uses_cp_shared_kv=True, + cp_shared_kv_layout=layout, + forward_mode="test", + out_cache_loc=torch.tensor([4, 8, 12, 16], dtype=torch.int64), + ) + key = torch.arange(0, 8, dtype=torch.float32).view(4, 2) + + with patch.object( + nsa_indexer, "tensor_debug_summary", side_effect=AssertionError + ), patch.object( + nsa_indexer, "tensor_debug_checksum", side_effect=AssertionError + ), patch( + "sglang.srt.environ.envs.SGLANG_DEBUG_CP_SHARED_KV.get", + return_value=False, + ): + physical_locs, key_to_write = nsa_indexer.Indexer._filter_shared_index_write( + None, + forward_batch, + key, + ) + + self.assertEqual(physical_locs.tolist(), [4, 8]) + self.assertEqual(key_to_write.shape[0], 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 1276493d3..e167d83c8 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -34,6 +34,25 @@ class TestPrepareServerArgs(CustomTestCase): ) +def test_enable_nsa_prefill_cp_shared_kv_parser_flag(): + import argparse + + parser = argparse.ArgumentParser() + ServerArgs.add_cli_args(parser) + raw_args = parser.parse_args( + [ + "--model-path", + "dummy", + "--enable-nsa-prefill-context-parallel", + "--enable-nsa-prefill-cp-shared-kv", + "--nsa-prefill-cp-mode", + "in-seq-split", + ] + ) + args = ServerArgs.from_cli_args(raw_args) + assert args.enable_nsa_prefill_cp_shared_kv is True + + class TestLoadBalanceMethod(unittest.TestCase): def test_non_pd_defaults_to_round_robin(self): server_args = ServerArgs(model_path="dummy", disaggregation_mode="null")