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
@@ -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:
+38
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
#########################
+1
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)
@@ -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
@@ -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,
@@ -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
@@ -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]))
+274 -22
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):
+38
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
@@ -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),
)
@@ -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:
@@ -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)
@@ -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"
+35 -1
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,