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:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
#########################
|
||||
|
||||
@@ -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]))
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user