Expand prefill CP KV capacity by sharding persistent NSA KV

Prefill CP previously replicated NSA/MLA persistent KV on every CP rank, so CP8 consumed eight copies of KV memory while exposing only one rank of logical cache capacity. This change splits logical KV locs from per-rank physical storage, shards MLA latent KV and NSA index K/scale by deterministic page ownership, and keeps existing NSA attention kernels working through a full-view runtime materialization layer.

Mooncake PD transfer now sends each prefill CP rank's owned physical pages with explicit logical page positions so non-CP decode can reconstruct full-layout KV. The implementation is guarded by an explicit server flag and startup checks, and the design documentation records the implemented scope, debug environment, and Phase 3 boundary.

Constraint: Phase 2 must preserve existing NSA attention/index kernels via runtime full-view materialization
Constraint: Decode side remains non-CP and receives full KV through Mooncake
Rejected: Shard-aware NSA attention in this change | belongs to Phase 3 because it requires distributed topk/softmax/output contracts
Rejected: Request-contiguous CP ownership | unstable under chunked prefill and tied to attention split mode
Confidence: medium
Scope-risk: broad
Directive: Do not enable round-robin CP shared KV without wiring runtime materialization/PD transfer contracts for that split mode
Directive: Keep SGLANG_DEBUG_CP_SHARED_KV disabled for perf measurements; it intentionally enables CUDA-syncing diagnostics
Tested: Remote py_compile for shared-KV touched Python files in g0034 container
Tested: Remote pytest selected cp_shared/shared_kv/nsa suite: 37 passed, 34 deselected
Not-tested: Full GLM5 multi-node throughput/regression run after final doc update
Not-tested: Phase 3 shard-aware runtime, round-robin CP mode, and non-Mooncake PD backends
This commit is contained in:
laoyao0822
2026-04-26 04:11:17 +08:00
parent 76c830dda8
commit f8fca72635
22 changed files with 2489 additions and 90 deletions

View File

@@ -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 不开启 CPprefill 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. 验证计划

View File

@@ -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:

View File

@@ -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
#########################

View File

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

View File

@@ -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

View File

@@ -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,

View File

@@ -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

View File

@@ -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]))

View File

@@ -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):

View File

@@ -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

View File

@@ -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),
)

View File

@@ -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:

View File

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

View File

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

View File

@@ -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,

View File

@@ -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()

View File

@@ -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()

View File

@@ -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()

View File

@@ -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()

View File

@@ -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()

View File

@@ -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()

View File

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