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")