Reduce CP shared KV materialize and direct-write overhead

Shared KV now relies on page-aligned CP metadata and compute-owner page allocation so persistent MLA KV and NSA index shards can be written by the rank that computed them. The compatibility read path keeps the dense full-view contract for existing topk and attention kernels, but removes duplicated prev/next index materialize, adds optional tai materialize integration, and tightens tests/docs around the fallback boundaries.

Constraint: Decode remains non-CP while prefill CP owns the shared-KV changes

Constraint: Existing attention/topk kernels still expect dense full-view KV/index inputs

Rejected: Change attention kernels to read owner-sharded KV directly | larger semantic change reserved for later phases

Rejected: Merge index K/scale storage with MLA KV storage | would couple topk and attention cache lifecycles before materialize overhead is isolated

Confidence: medium

Scope-risk: broad

Directive: Do not remove fallback logging or debug-gated assertions without reproducing long-context chunked/radix-hit paths

Tested: git diff --check --cached

Not-tested: Local pytest/runtime server verification not run in this commit step per current workflow constraints
This commit is contained in:
laoyao0822
2026-05-02 07:07:28 +08:00
parent 2317952a01
commit 5769b63082
17 changed files with 2524 additions and 96 deletions

View File

@@ -204,6 +204,7 @@ class Envs:
SGLANG_DEBUG_MEMORY_POOL = EnvBool(False)
SGLANG_DEBUG_CP_SHARED_KV = EnvBool(False)
SGLANG_CP_SHARED_KV_CURRENT_REUSE = EnvBool(False)
SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE = EnvBool(False)
SGLANG_TEST_REQUEST_TIME_STATS = EnvBool(False)
SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK = EnvBool(False)
SGLANG_SIMULATE_ACC_LEN = EnvFloat(-1)

View File

@@ -1,6 +1,7 @@
from __future__ import annotations
import logging
from functools import lru_cache
import torch
@@ -11,6 +12,7 @@ from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
logger = logging.getLogger(__name__)
_DEBUG_LOG_COUNTS: dict[str, int] = {}
_TAI_MATERIALIZE_FALLBACK_LOG_COUNTS: dict[str, int] = {}
def cp_shared_kv_debug_enabled() -> bool:
@@ -21,6 +23,125 @@ def cp_shared_kv_current_reuse_enabled() -> bool:
return envs.SGLANG_CP_SHARED_KV_CURRENT_REUSE.get()
def cp_shared_kv_tai_materialize_enabled() -> bool:
return envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get()
@lru_cache(maxsize=1)
def _load_tai_materialize_kernels():
try:
from tai_kernel.nsa_prefill import cp_shared_kv_materialize
return cp_shared_kv_materialize
except Exception as exc:
_log_tai_materialize_fallback(
"import_failed",
"CP shared KV tai materialize kernels are unavailable; "
"falling back to torch materialize. error=%s",
exc,
limit=1,
)
return None
def _tai_materialize_runtime_enabled() -> bool:
# Keep the debug path on the existing torch implementation. The debug path
# intentionally preserves tensor summaries and value assertions used for
# diagnosing shared-KV correctness; the Triton kernels optimize the normal
# runtime path only.
return cp_shared_kv_tai_materialize_enabled() and not cp_shared_kv_debug_enabled()
def _log_tai_materialize_fallback(
key: str,
message: str,
*args,
limit: int = 8,
) -> None:
count = _TAI_MATERIALIZE_FALLBACK_LOG_COUNTS.get(key, 0)
if count >= limit:
return
_TAI_MATERIALIZE_FALLBACK_LOG_COUNTS[key] = count + 1
logger.warning(message, *args)
def _contiguous_for_tai(tensor: torch.Tensor) -> torch.Tensor:
return tensor if tensor.is_contiguous() else tensor.contiguous()
def _try_tai_materialize_shared_pages(
page_buffer: torch.Tensor,
logical_pages: torch.Tensor,
layout: CpSharedKVLayout,
) -> tuple[torch.Tensor, torch.Tensor] | None:
if not _tai_materialize_runtime_enabled():
return None
kernels = _load_tai_materialize_kernels()
if kernels is None:
return None
try:
return kernels.materialize_shared_pages(
page_buffer,
_contiguous_for_tai(logical_pages),
cp_rank=layout.cp_rank,
cp_size=layout.cp_size,
)
except Exception as exc:
_log_tai_materialize_fallback(
"paged_failed",
"CP shared KV tai paged materialize failed; falling back to torch "
"materialize. error=%s",
exc,
)
return None
def _try_tai_materialize_token_kv_pages_and_locs(
kv_cache: torch.Tensor,
logical_locs: torch.Tensor,
slot_logical_pages: torch.Tensor,
logical_page_capacity: int,
layout: CpSharedKVLayout,
page_size: int,
) -> tuple[torch.Tensor, torch.Tensor] | None:
if not _tai_materialize_runtime_enabled():
return None
kernels = _load_tai_materialize_kernels()
if kernels is None:
return None
try:
tai_slot_logical_pages = _contiguous_for_tai(slot_logical_pages.reshape(-1))
page_inverse = kernels.build_slot_page_inverse(
tai_slot_logical_pages,
logical_page_capacity,
)
dense_locs = kernels.remap_logical_locs_to_slot_dense_locs(
_contiguous_for_tai(logical_locs),
page_inverse,
page_size=page_size,
)
dense_kv_cache = kernels.materialize_shared_token_kv_pages(
kv_cache,
tai_slot_logical_pages,
page_size=page_size,
cp_rank=layout.cp_rank,
cp_size=layout.cp_size,
)
return dense_kv_cache, dense_locs
except Exception as exc:
_log_tai_materialize_fallback(
"token_failed",
"CP shared KV tai token materialize failed; falling back to torch "
"materialize. error=%s",
exc,
)
return None
def is_current_only_extend_batch(forward_batch) -> bool:
"""Return whether an extend batch has no cached/history tokens.
@@ -737,6 +858,7 @@ def materialize_shared_token_kv_buffer(
physical_token_capacity=kv_cache.shape[0],
)
dense_kv_cache = None
if remap_logical_pages is None:
remap_pages_from_locs = logical_pages_from_locs(remap_logical_locs, page_size)
materialized_logical_pages, _ = build_dense_page_remap(remap_pages_from_locs)
@@ -757,21 +879,36 @@ def materialize_shared_token_kv_buffer(
layout=layout,
physical_page_capacity=kv_cache.shape[0] // page_size,
)
materialized_logical_pages, _ = build_slot_page_remap(remap_logical_pages)
logical_page_capacity = _logical_page_capacity_from_physical_page_capacity(
kv_cache.shape[0] // page_size,
layout,
)
page_inverse = build_slot_page_inverse(
materialized_logical_pages,
logical_page_capacity=logical_page_capacity,
)
dense_locs = remap_logical_locs_to_slot_dense_locs(
logical_locs,
page_inverse=page_inverse,
page_size=page_size,
)
use_slot_materialize = True
tai_result = None
if _tai_materialize_runtime_enabled():
materialized_logical_pages = remap_logical_pages.reshape(-1)
tai_result = _try_tai_materialize_token_kv_pages_and_locs(
kv_cache=kv_cache,
logical_locs=logical_locs,
slot_logical_pages=materialized_logical_pages,
logical_page_capacity=logical_page_capacity,
layout=layout,
page_size=page_size,
)
if tai_result is None:
materialized_logical_pages, _ = build_slot_page_remap(remap_logical_pages)
page_inverse = build_slot_page_inverse(
materialized_logical_pages,
logical_page_capacity=logical_page_capacity,
)
dense_locs = remap_logical_locs_to_slot_dense_locs(
logical_locs,
page_inverse=page_inverse,
page_size=page_size,
)
use_slot_materialize = True
else:
dense_kv_cache, dense_locs = tai_result
use_slot_materialize = False
if use_slot_materialize:
dense_kv_cache = materialize_local_token_kv_page_slots(
@@ -780,7 +917,7 @@ def materialize_shared_token_kv_buffer(
layout=layout,
page_size=page_size,
)
else:
elif dense_kv_cache is None:
dense_kv_cache = materialize_local_token_kv_pages(
kv_cache=kv_cache,
unique_logical_pages=materialized_logical_pages,
@@ -830,6 +967,7 @@ def materialize_shared_token_kv_buffer(
)
return dense_kv_cache, dense_locs
def materialize_shared_paged_buffer(
page_buffer: torch.Tensor,
logical_pages: torch.Tensor,
@@ -845,12 +983,21 @@ def materialize_shared_paged_buffer(
layout=layout,
physical_page_capacity=page_buffer.shape[0],
)
materialized_logical_pages, dense_pages = build_slot_page_remap(logical_pages)
dense_page_buffer = materialize_local_paged_buffer_page_slots(
tai_result = _try_tai_materialize_shared_pages(
page_buffer=page_buffer,
slot_logical_pages=materialized_logical_pages,
logical_pages=logical_pages,
layout=layout,
)
if tai_result is None:
materialized_logical_pages, dense_pages = build_slot_page_remap(logical_pages)
dense_page_buffer = materialize_local_paged_buffer_page_slots(
page_buffer=page_buffer,
slot_logical_pages=materialized_logical_pages,
layout=layout,
)
else:
dense_page_buffer, dense_pages = tai_result
materialized_logical_pages = logical_pages.reshape(-1)
if cp_shared_kv_debug_enabled():
owned_pages = materialized_logical_pages[

View File

@@ -52,8 +52,11 @@ from sglang.srt.distributed.parallel_state import get_pp_group
from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.attention.nsa.utils import (
cp_all_gather_rerange_output,
get_cp_shared_kv_local_out_cache_loc,
is_nsa_enable_prefill_cp,
is_nsa_prefill_cp_in_seq_split,
log_cp_shared_kv_direct_write_fallback,
nsa_use_prefill_cp,
split_in_seq_cp_local_pair,
)
from sglang.srt.layers.communicator import ScatterMode
@@ -442,6 +445,7 @@ class Indexer(MultiPlatformOp):
query = rotate_activation(query)
key = rotate_activation(key)
local_key = key
# allgather+rerrange
if forward_batch.nsa_cp_metadata is not None and self.nsa_enable_prefill_cp:
key = cp_all_gather_rerange_output(
@@ -450,7 +454,7 @@ class Indexer(MultiPlatformOp):
forward_batch,
torch.cuda.current_stream(),
)
return query, key
return query, key, local_key
def _get_k_bf16(
self,
@@ -839,6 +843,8 @@ class Indexer(MultiPlatformOp):
actual_seq_q: int,
cp_index: List[Tuple[int, int, int]] = None,
current_index_kv: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
shared_index_buffer: Optional[torch.Tensor] = None,
shared_block_tables: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if TYPE_CHECKING:
assert isinstance(forward_batch.token_to_kv_pool, NSATokenToKVPool)
@@ -855,15 +861,23 @@ class Indexer(MultiPlatformOp):
actual_seq_q_list = []
batch_idx_list = []
block_tables = metadata.get_page_table_64()
if current_index_kv is not None and cp_index is not None:
current_index_kv = None
if current_index_kv is None:
index_buffer, block_tables = self._maybe_materialize_shared_index_buffer(
forward_batch,
layer_id,
block_tables,
)
if shared_index_buffer is not None or shared_block_tables is not None:
if shared_index_buffer is None or shared_block_tables is None:
raise RuntimeError(
"shared index buffer and block tables must be provided together"
)
index_buffer = shared_index_buffer
block_tables = shared_block_tables
else:
block_tables = metadata.get_page_table_64()
index_buffer, block_tables = self._maybe_materialize_shared_index_buffer(
forward_batch,
layer_id,
block_tables,
)
else:
index_buffer = None
if cp_shared_kv_debug_enabled():
@@ -1032,6 +1046,76 @@ class Indexer(MultiPlatformOp):
return topk_result
def _get_topk_in_seq_cp_pair(
self,
forward_batch: ForwardBatch,
layer_id: int,
q_fp8: torch.Tensor,
weights: torch.Tensor,
metadata: BaseIndexerMetadata,
current_index_kv: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
) -> torch.Tensor:
assert forward_batch.nsa_cp_metadata is not None
kv_len_prev = forward_batch.nsa_cp_metadata.kv_len_prev
kv_len_next = forward_batch.nsa_cp_metadata.kv_len_next
actual_seq_q_prev = forward_batch.nsa_cp_metadata.actual_seq_q_prev
actual_seq_q_next = forward_batch.nsa_cp_metadata.actual_seq_q_next
# TODO support mutil-batch
# cp_batch_seq_index_prev = forward_batch.nsa_cp_metadata["cp_batch_seq_index_prev"]
# cp_batch_seq_index_next = forward_batch.nsa_cp_metadata["cp_batch_seq_index_next"]
q_fp8_prev, q_fp8_next = split_in_seq_cp_local_pair(
q_fp8,
actual_seq_q_prev,
actual_seq_q_next,
name="q_fp8",
)
weights_prev, weights_next = split_in_seq_cp_local_pair(
weights,
actual_seq_q_prev,
actual_seq_q_next,
name="weights",
)
shared_index_buffer = None
shared_block_tables = None
if current_index_kv is None:
shared_block_tables = metadata.get_page_table_64()
shared_index_buffer, shared_block_tables = (
self._maybe_materialize_shared_index_buffer(
forward_batch,
layer_id,
shared_block_tables,
)
)
topk_result_prev = self._get_topk_ragged_with_cp(
forward_batch,
layer_id,
q_fp8_prev,
weights_prev,
metadata,
kv_len_prev,
actual_seq_q_prev,
current_index_kv=current_index_kv,
shared_index_buffer=shared_index_buffer,
shared_block_tables=shared_block_tables,
)
topk_result_next = self._get_topk_ragged_with_cp(
forward_batch,
layer_id,
q_fp8_next,
weights_next,
metadata,
kv_len_next,
actual_seq_q_next,
current_index_kv=current_index_kv,
shared_index_buffer=shared_index_buffer,
shared_block_tables=shared_block_tables,
)
return torch.cat([topk_result_prev, topk_result_next], dim=0)
def forward_indexer(
self,
q_fp8: torch.Tensor,
@@ -1129,6 +1213,7 @@ class Indexer(MultiPlatformOp):
key: torch.Tensor,
*,
act_quant=None, # fallback only
out_loc_override: Optional[torch.Tensor] = None,
) -> None:
"""
Store NSA indexer K cache for current step.
@@ -1136,7 +1221,10 @@ 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_override is None:
out_loc, key = self._filter_shared_index_write(forward_batch, key)
else:
out_loc = out_loc_override
if out_loc.numel() == 0:
return
if not out_loc.is_contiguous():
@@ -1175,6 +1263,48 @@ class Indexer(MultiPlatformOp):
index_k_scale=k_scale,
)
def _store_cp_shared_local_index_k_cache(
self,
forward_batch: ForwardBatch,
layer_id: int,
local_key: torch.Tensor,
*,
act_quant,
) -> bool:
if not nsa_use_prefill_cp(forward_batch, self.nsa_enable_prefill_cp):
return False
local_out_loc = get_cp_shared_kv_local_out_cache_loc(forward_batch)
if local_out_loc is None:
return False
if local_key.shape[0] != local_out_loc.numel():
log_cp_shared_kv_direct_write_fallback(
"index_local_shape_mismatch",
"NSA index local key token count does not match local out_cache_loc: "
"local_key=%s local_out_cache_loc=%s layer_id=%s",
local_key.shape[0],
local_out_loc.numel(),
layer_id,
)
return False
if local_out_loc.numel() == 0:
return True
assert forward_batch.cp_shared_kv_layout is not None
physical_out_loc = (
forward_batch.cp_shared_kv_layout.logical_locs_to_physical(
local_out_loc
).contiguous()
)
self._store_index_k_cache(
forward_batch=forward_batch,
layer_id=layer_id,
key=local_key,
act_quant=act_quant,
out_loc_override=physical_out_loc,
)
return True
def forward_cuda(
self,
x: torch.Tensor,
@@ -1236,21 +1366,27 @@ class Indexer(MultiPlatformOp):
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
weights = self._project_and_scale_head_gates(x)
query, key = self._get_q_k_bf16(
query, key, local_key = self._get_q_k_bf16(
q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch
)
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
with torch.cuda.stream(self.alt_stream):
self._store_index_k_cache(
if not self._store_cp_shared_local_index_k_cache(
forward_batch=forward_batch,
layer_id=layer_id,
key=key,
local_key=local_key,
act_quant=act_quant,
)
):
self._store_index_k_cache(
forward_batch=forward_batch,
layer_id=layer_id,
key=key,
act_quant=act_quant,
)
current_stream.wait_stream(self.alt_stream)
weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale
else:
query, key = self._get_q_k_bf16(
query, key, local_key = self._get_q_k_bf16(
q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch
)
@@ -1260,21 +1396,33 @@ class Indexer(MultiPlatformOp):
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
with torch.cuda.stream(self.alt_stream):
if not self._store_cp_shared_local_index_k_cache(
forward_batch=forward_batch,
layer_id=layer_id,
local_key=local_key,
act_quant=act_quant,
):
self._store_index_k_cache(
forward_batch=forward_batch,
layer_id=layer_id,
key=key,
act_quant=act_quant,
)
current_stream.wait_stream(self.alt_stream)
else:
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
if not self._store_cp_shared_local_index_k_cache(
forward_batch=forward_batch,
layer_id=layer_id,
local_key=local_key,
act_quant=act_quant,
):
self._store_index_k_cache(
forward_batch=forward_batch,
layer_id=layer_id,
key=key,
act_quant=act_quant,
)
current_stream.wait_stream(self.alt_stream)
else:
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
self._store_index_k_cache(
forward_batch=forward_batch,
layer_id=layer_id,
key=key,
act_quant=act_quant,
)
# `_get_logits_head_gate` expects a Tensor. For tuple activations, dequantize
# to a float tensor here (callsite), keeping `_get_logits_head_gate` backend-agnostic.
@@ -1366,49 +1514,14 @@ class Indexer(MultiPlatformOp):
forward_batch.nsa_cp_metadata is not None
and is_nsa_prefill_cp_in_seq_split()
):
kv_len_prev = forward_batch.nsa_cp_metadata.kv_len_prev
kv_len_next = forward_batch.nsa_cp_metadata.kv_len_next
actual_seq_q_prev = forward_batch.nsa_cp_metadata.actual_seq_q_prev
actual_seq_q_next = forward_batch.nsa_cp_metadata.actual_seq_q_next
# TODO support mutil-batch
# cp_batch_seq_index_prev = forward_batch.nsa_cp_metadata["cp_batch_seq_index_prev"]
# cp_batch_seq_index_next = forward_batch.nsa_cp_metadata["cp_batch_seq_index_next"]
# TODO prev, next, combined into a single call
q_fp8_prev, q_fp8_next = split_in_seq_cp_local_pair(
return self._get_topk_in_seq_cp_pair(
forward_batch,
layer_id,
q_fp8,
actual_seq_q_prev,
actual_seq_q_next,
name="q_fp8",
)
weights_prev, weights_next = split_in_seq_cp_local_pair(
weights,
actual_seq_q_prev,
actual_seq_q_next,
name="weights",
)
topk_result_prev = self._get_topk_ragged_with_cp(
forward_batch,
layer_id,
q_fp8_prev,
weights_prev,
metadata,
kv_len_prev,
actual_seq_q_prev,
current_index_kv=current_index_kv,
)
topk_result_next = self._get_topk_ragged_with_cp(
forward_batch,
layer_id,
q_fp8_next,
weights_next,
metadata,
kv_len_next,
actual_seq_q_next,
current_index_kv=current_index_kv,
)
return torch.cat([topk_result_prev, topk_result_next], dim=0)
else:
topk_result = self._get_topk_ragged(
enable_dual_stream,

View File

@@ -1,4 +1,5 @@
# temp NSA debugging environ
import logging
from dataclasses import dataclass
from itertools import accumulate
from typing import TYPE_CHECKING, List, Tuple, Union
@@ -26,6 +27,25 @@ from sglang.srt.utils.common import ceil_align, ceil_div
if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
logger = logging.getLogger(__name__)
def log_cp_shared_kv_direct_write_fallback(
reason: str,
message: str,
*args,
) -> None:
"""Log every direct-write fallback event.
Warmup can hit the same fallback reason as a later real request, so
de-duplicating by reason hides correctness/performance issues after startup.
"""
logger.info(
"CP shared KV direct-write fallback (%s): " + message,
reason,
*args,
)
def compute_nsa_seqlens(original_seq_lens, nsa_index_topk: int):
return original_seq_lens.clamp(max=nsa_index_topk)
@@ -204,9 +224,12 @@ def build_page_aligned_in_seq_split_list(
) -> Tuple[List[int], PageAlignedInSeqSplitInfo]:
"""Build an in-seq split list whose real-token boundaries do not cut pages.
Phase 4 deliberately uses a conservative gate: at least `2 * cp_size` page
units are required so every zigzag segment has at least one page unit. When
the gate does not hold, this helper falls back to the existing token-balanced
Phase 4 deliberately uses a conservative gate for cache-miss chunks: at
least `2 * cp_size` page units are required so every zigzag segment has at
least one page unit. For radix-hit suffixes with a page-aligned prefix, the
gate is relaxed to `cp_size` page units so every CP rank still receives at
least one local page while second zigzag segments may be empty. When the
gate does not hold, this helper falls back to the existing token-balanced
split and marks the result as not page-aligned.
"""
@@ -230,7 +253,9 @@ def build_page_aligned_in_seq_split_list(
tail_tokens = extend_len % page_size
num_page_units = full_pages + (1 if tail_tokens > 0 else 0)
cp_segment_num = cp_size * 2
if num_page_units < cp_segment_num:
if num_page_units < cp_size or (
num_page_units < cp_segment_num and extend_prefix_len == 0
):
return fallback_split, fallback_info
base_units = num_page_units // cp_segment_num
@@ -303,6 +328,58 @@ def _build_in_seq_split_for_forward_batch(
)
def should_use_replicated_compute_for_short_radix_hit(
forward_batch: "ForwardBatch",
cp_size: int,
) -> bool:
"""Return whether a short radix-hit suffix should avoid CP splitting.
With CP shared KV, radix-hit suffixes can be page-aligned but shorter than
one page per CP rank. A page-aligned CP split would give some ranks zero
local tokens, which is unsafe for parts of the current CP collective/kernel
path. Instead, keep the original non-CP behavior: every rank computes the
short suffix, while shared-KV write filters persist only pages owned by the
local rank.
"""
if (
forward_batch is None
or cp_size <= 0
or not getattr(forward_batch, "uses_cp_shared_kv", False)
):
return False
extend_seq_lens_cpu = getattr(forward_batch, "extend_seq_lens_cpu", None)
extend_prefix_lens_cpu = getattr(forward_batch, "extend_prefix_lens_cpu", None)
if (
extend_seq_lens_cpu is None
or extend_prefix_lens_cpu is None
or len(extend_seq_lens_cpu) != 1
or len(extend_prefix_lens_cpu) != 1
):
return False
token_to_kv_pool = getattr(forward_batch, "token_to_kv_pool", None)
page_size = getattr(token_to_kv_pool, "page_size", None)
if page_size is None:
return False
page_size = int(page_size)
if page_size <= 1:
return False
extend_len = int(extend_seq_lens_cpu[0])
extend_prefix_len = int(extend_prefix_lens_cpu[0])
if (
extend_len <= 0
or extend_prefix_len <= 0
or extend_prefix_len % page_size != 0
):
return False
num_page_units = ceil_div(extend_len, page_size)
return 0 < num_page_units < cp_size
def can_cp_split(seq_len: int, cp_size: int, use_nsa: bool, forward_batch):
if is_nsa_prefill_cp_round_robin_split():
cur_cp_seq_len = seq_len // cp_size
@@ -313,6 +390,8 @@ def can_cp_split(seq_len: int, cp_size: int, use_nsa: bool, forward_batch):
# TODO current just support prefill batch=1 and len(input_ids) > self.cp_size * 2
# Note: (self.cp_size * 2) To achieve load balancing for seq computation,
# the seq data needs to be divided and recombined at twice the size of cp_size.
if should_use_replicated_compute_for_short_radix_hit(forward_batch, cp_size):
return False
cur_cp_seq_len = seq_len // (cp_size * 2)
if (
cur_cp_seq_len != 0
@@ -344,6 +423,110 @@ def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor):
return result
def cp_split_and_rebuild_1d(forward_batch, input_: torch.Tensor):
try:
round_robin_split = is_nsa_prefill_cp_round_robin_split()
except ValueError:
round_robin_split = False
if round_robin_split:
return nsa_cp_round_robin_split_data(input_)
input_list = list(
torch.split(input_, forward_batch.nsa_cp_metadata.split_list, dim=0)
)
return torch.cat(
[input_list[i] for i in forward_batch.nsa_cp_metadata.zigzag_index], dim=0
).view(-1)
def get_cp_shared_kv_local_out_cache_loc(forward_batch: "ForwardBatch"):
"""Return this CP rank's local logical out_cache_loc for direct writes.
`None` means the batch should keep using the compatibility path. This path
is intentionally conservative: it only enables direct writes after Phase 4
page-aligned split and after the logical page ids prove they are owned by
this CP rank.
"""
cached = getattr(forward_batch, "cp_local_out_cache_loc", None)
if cached is not None:
return cached
if not getattr(forward_batch, "uses_cp_shared_kv", False):
return None
metadata = getattr(forward_batch, "nsa_cp_metadata", None)
layout = getattr(forward_batch, "cp_shared_kv_layout", None)
out_cache_loc = getattr(forward_batch, "out_cache_loc", None)
if metadata is None:
log_cp_shared_kv_direct_write_fallback(
"missing_metadata",
"nsa_cp_metadata is missing",
)
return None
if layout is None:
log_cp_shared_kv_direct_write_fallback(
"missing_layout",
"cp_shared_kv_layout is missing",
)
return None
if out_cache_loc is None:
log_cp_shared_kv_direct_write_fallback(
"missing_out_cache_loc",
"out_cache_loc is missing",
)
return None
if not getattr(metadata, "page_aligned", False):
log_cp_shared_kv_direct_write_fallback(
"not_page_aligned",
"metadata is not page-aligned: page_size=%s extend_prefix_len=%s",
getattr(metadata, "page_size", None),
getattr(metadata, "extend_prefix_len", None),
)
return None
try:
in_seq_split = is_nsa_prefill_cp_in_seq_split()
except ValueError:
in_seq_split = True
if not in_seq_split:
log_cp_shared_kv_direct_write_fallback(
"not_in_seq_split",
"nsa_prefill_cp_mode is not in-seq-split",
)
return None
split_tokens = sum(int(x) for x in metadata.split_list)
out_cache_tokens = int(out_cache_loc.numel())
if split_tokens != out_cache_tokens:
log_cp_shared_kv_direct_write_fallback(
"split_out_cache_len_mismatch",
"split_list tokens=%s out_cache_loc tokens=%s",
split_tokens,
out_cache_tokens,
)
return None
local_out_cache_loc = cp_split_and_rebuild_1d(
forward_batch,
out_cache_loc.contiguous(),
)
if local_out_cache_loc.numel() == 0:
forward_batch.cp_local_out_cache_loc = local_out_cache_loc
return local_out_cache_loc
valid_locs = local_out_cache_loc[local_out_cache_loc > 0]
if valid_locs.numel() > 0 and not torch.all(layout.owned_by_this_rank(valid_locs)):
log_cp_shared_kv_direct_write_fallback(
"local_loc_owner_mismatch",
"local out_cache_loc contains pages not owned by this rank: cp_rank=%s cp_size=%s page_size=%s",
layout.cp_rank,
layout.cp_size,
layout.page_size,
)
return None
forward_batch.cp_local_out_cache_loc = local_out_cache_loc
return local_out_cache_loc
def cp_split_and_rebuild_position(forward_batch, positions: torch.Tensor):
if is_nsa_prefill_cp_round_robin_split():
cp_size = get_attention_cp_size()

View File

@@ -20,7 +20,7 @@ Page-aligned memory pool.
"""
import abc
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, List, Optional
import torch
import triton
@@ -555,3 +555,130 @@ class CPSharedPagedTokenToKVPoolAllocator(PagedTokenToKVPoolAllocator):
self.physical_size = physical_size
self.cp_size = cp_size
self.cp_rank = cp_rank
def compute_owner_lane_stats(
self,
page_compute_owners: List[int],
) -> tuple[List[int], List[int], List[int]]:
required = [0 for _ in range(self.cp_size)]
for owner in page_compute_owners:
if owner < 0 or owner >= self.cp_size:
raise ValueError(
f"compute owner must be in [0, {self.cp_size}), got {owner}"
)
required[owner] += 1
free_pages = self.free_pages
if len(self.release_pages) > 0:
free_pages = torch.cat((free_pages, self.release_pages))
available = [
int(
(
torch.remainder(free_pages - 1, self.cp_size) == owner
).sum().item()
)
for owner in range(self.cp_size)
]
deficits = [
max(0, required_count - available_count)
for required_count, available_count in zip(required, available)
]
return required, available, deficits
def _select_compute_owner_pages(
self,
page_compute_owners: List[int],
) -> Optional[torch.Tensor]:
selected_pages = []
lane_offsets = [0 for _ in range(self.cp_size)]
lane_pages = [
self.free_pages[
torch.remainder(self.free_pages - 1, self.cp_size) == owner
]
for owner in range(self.cp_size)
]
for owner in page_compute_owners:
if owner < 0 or owner >= self.cp_size:
raise ValueError(
f"compute owner must be in [0, {self.cp_size}), got {owner}"
)
lane_offset = lane_offsets[owner]
if lane_offset >= lane_pages[owner].numel():
return None
selected_pages.append(lane_pages[owner][lane_offset])
lane_offsets[owner] = lane_offset + 1
if not selected_pages:
return torch.empty((0,), dtype=torch.int64, device=self.device)
return torch.stack(selected_pages).to(torch.int64)
def alloc_extend_compute_owner(
self,
prefix_lens: torch.Tensor,
prefix_lens_cpu: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_cpu: torch.Tensor,
last_loc: torch.Tensor,
extend_num_tokens: int,
page_compute_owners: List[int],
):
"""Allocate extend KV locs so logical page owner matches CP compute rank.
The returned logical `out_cache_loc` is still full-order and identical on
every CP rank. Only the chosen logical page ids change: each newly
allocated request page comes from the modulo-owner lane that will compute
and directly persist that page.
"""
if len(prefix_lens_cpu) != 1 or len(seq_lens_cpu) != 1:
raise ValueError("compute-owner allocation supports batch size 1 only")
num_new_pages = get_num_new_pages(
seq_lens=seq_lens_cpu,
page_size=self.page_size,
prefix_lens=prefix_lens_cpu,
)
if num_new_pages != len(page_compute_owners):
raise ValueError(
"compute-owner page count mismatch: "
f"{num_new_pages=} page_compute_owners={len(page_compute_owners)}"
)
if self.need_sort and num_new_pages > len(self.free_pages):
self.merge_and_sort_free()
selected_pages = self._select_compute_owner_pages(page_compute_owners)
if selected_pages is None and self.need_sort and len(self.release_pages) > 0:
self.merge_and_sort_free()
selected_pages = self._select_compute_owner_pages(page_compute_owners)
if selected_pages is None:
return None
out_indices = torch.empty(
(extend_num_tokens,), dtype=torch.int64, device=self.device
)
alloc_extend_naive(
prefix_lens_cpu,
seq_lens_cpu,
last_loc,
selected_pages,
out_indices,
self.page_size,
self.device,
)
selected_mask = torch.isin(self.free_pages, selected_pages)
self.free_pages = self.free_pages[~selected_mask]
if self.debug_mode:
assert len(torch.unique(out_indices)) == len(out_indices)
selected_owners = torch.remainder(selected_pages - 1, self.cp_size)
expected_owners = torch.tensor(
page_compute_owners,
dtype=selected_owners.dtype,
device=selected_owners.device,
)
assert torch.equal(selected_owners, expected_owners)
return out_indices

View File

@@ -8,6 +8,10 @@ import triton
import triton.language as tl
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_in_seq_page_compute_owners,
get_in_seq_page_compute_owner_unavailable_reason,
)
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
from sglang.srt.server_args import get_global_server_args
@@ -24,6 +28,18 @@ MAMBA_STATE_PER_REQ_NO_CACHE = 1
logger = logging.getLogger(__name__)
def _log_cp_shared_kv_alloc_fallback(
reason: str,
message: str,
*args,
) -> None:
logger.info(
"CP shared KV compute-owner allocation fallback (%s): " + message,
reason,
*args,
)
@triton.jit
def write_req_to_token_pool_triton(
req_to_token_ptr, # [max_batch, max_context_len]
@@ -252,6 +268,47 @@ def evict_from_tree_cache(tree_cache: BasePrefixCache | None, num_tokens: int):
tree_cache.evict(EvictParams(num_tokens=num_tokens))
def _evict_for_compute_owner_lanes(
*,
tree_cache: BasePrefixCache | None,
allocator,
page_compute_owners: list[int],
) -> None:
if tree_cache is None or tree_cache.is_chunk_cache():
return
compute_owner_lane_stats = getattr(allocator, "compute_owner_lane_stats", None)
if compute_owner_lane_stats is None:
return
max_attempts = max(2, min(8, int(getattr(allocator, "cp_size", 1))))
for _ in range(max_attempts):
_required, _available, deficits = compute_owner_lane_stats(page_compute_owners)
deficit_pages = sum(deficits)
if deficit_pages <= 0:
return
try:
evictable_size = tree_cache.evictable_size()
except Exception:
evictable_size = allocator.page_size
if isinstance(evictable_size, tuple):
evictable_size = evictable_size[0]
if evictable_size <= 0:
return
evict_tokens = max(
allocator.page_size,
deficit_pages * allocator.page_size * int(getattr(allocator, "cp_size", 1)),
)
before_available = allocator.available_size()
evict_result = tree_cache.evict(EvictParams(num_tokens=evict_tokens))
after_available = allocator.available_size()
evicted_tokens = getattr(evict_result, "num_tokens_evicted", 0)
if after_available <= before_available and evicted_tokens <= 0:
return
def alloc_paged_token_slots_extend(
tree_cache: BasePrefixCache,
prefix_lens: torch.Tensor,
@@ -267,18 +324,114 @@ def alloc_paged_token_slots_extend(
num_tokens = extend_num_tokens + len(seq_lens_cpu) * allocator.page_size
evict_from_tree_cache(tree_cache, num_tokens)
alloc_extend_compute_owner = getattr(
allocator, "alloc_extend_compute_owner", None
)
page_compute_owners = None
compute_owner_unavailable_reason = None
if alloc_extend_compute_owner is not None and len(prefix_lens_cpu) == 1:
try:
server_args = get_global_server_args()
except ValueError:
server_args = None
if (
server_args is not None
and server_args.enable_nsa_prefill_cp_shared_kv
and server_args.enable_nsa_prefill_context_parallel
and server_args.nsa_prefill_cp_mode == "in-seq-split"
):
extend_len = int(seq_lens_cpu[0].item() - prefix_lens_cpu[0].item())
page_compute_owners = build_in_seq_page_compute_owners(
extend_len=extend_len,
extend_prefix_len=int(prefix_lens_cpu[0].item()),
page_size=int(allocator.page_size),
cp_size=int(allocator.cp_size),
)
if page_compute_owners is None:
compute_owner_unavailable_reason = (
get_in_seq_page_compute_owner_unavailable_reason(
extend_len=extend_len,
extend_prefix_len=int(prefix_lens_cpu[0].item()),
page_size=int(allocator.page_size),
cp_size=int(allocator.cp_size),
)
or "unknown"
)
else:
compute_owner_unavailable_reason = "server_args_not_enabled"
elif alloc_extend_compute_owner is not None:
compute_owner_unavailable_reason = (
"multi_batch" if len(prefix_lens_cpu) != 1 else "unknown"
)
if page_compute_owners is not None:
_evict_for_compute_owner_lanes(
tree_cache=tree_cache,
allocator=allocator,
page_compute_owners=page_compute_owners,
)
state = None
if backup_state:
state = allocator.backup_state()
out_cache_loc = allocator.alloc_extend(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
)
if page_compute_owners is not None:
out_cache_loc = alloc_extend_compute_owner(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
page_compute_owners,
)
if out_cache_loc is None:
required = available = deficits = None
compute_owner_lane_stats = getattr(
allocator, "compute_owner_lane_stats", None
)
if compute_owner_lane_stats is not None:
required, available, deficits = compute_owner_lane_stats(
page_compute_owners
)
_log_cp_shared_kv_alloc_fallback(
"owner_lane_exhausted",
"failed to allocate pages from compute-owner lanes; "
"falling back to legacy page allocation. extend_num_tokens=%s page_size=%s "
"required_by_owner=%s available_by_owner=%s deficit_by_owner=%s",
extend_num_tokens,
allocator.page_size,
required,
available,
deficits,
)
out_cache_loc = allocator.alloc_extend(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
)
else:
if alloc_extend_compute_owner is not None:
_log_cp_shared_kv_alloc_fallback(
compute_owner_unavailable_reason or "compute_owner_not_available",
"page-aligned compute-owner page assignment is unavailable; "
"falling back to legacy page allocation. batch_size=%s extend_num_tokens=%s page_size=%s reason=%s",
len(prefix_lens_cpu),
extend_num_tokens,
allocator.page_size,
compute_owner_unavailable_reason or "compute_owner_not_available",
)
out_cache_loc = allocator.alloc_extend(
prefix_lens,
prefix_lens_cpu,
seq_lens,
seq_lens_cpu,
last_loc,
extend_num_tokens,
)
if out_cache_loc is None:
error_msg = (

View File

@@ -0,0 +1,81 @@
from __future__ import annotations
from typing import List, Optional
def get_in_seq_page_compute_owner_unavailable_reason(
*,
extend_len: int,
extend_prefix_len: int,
page_size: int,
cp_size: int,
) -> Optional[str]:
if cp_size <= 0:
raise ValueError(f"cp_size must be positive, got {cp_size}")
if extend_len < 0:
raise ValueError(f"extend_len must be non-negative, got {extend_len}")
if page_size <= 1:
return "page_size_le_one"
if extend_len <= 0:
return "empty_extend"
if extend_prefix_len % page_size != 0:
return "prefix_not_page_aligned"
full_pages = extend_len // page_size
tail_tokens = extend_len % page_size
num_page_units = full_pages + (1 if tail_tokens > 0 else 0)
if num_page_units < cp_size * 2 and extend_prefix_len == 0:
return "too_short_for_page_aligned"
return None
def build_in_seq_page_compute_owners(
*,
extend_len: int,
extend_prefix_len: int,
page_size: int,
cp_size: int,
) -> Optional[List[int]]:
"""Return compute-owner CP rank for each newly allocated current page.
This mirrors the Phase 4 page-aligned `in-seq-split` segmentation for
normal CP chunks, but it only returns page-unit owners for the real extend
chunk. Short radix-hit suffixes with fewer pages than CP ranks are also
allowed: runtime keeps replicated compute for those chunks and the shared
KV write filters persist only locally owned pages. `None` means the batch
must stay on the legacy allocation/write path.
"""
if cp_size <= 0:
raise ValueError(f"cp_size must be positive, got {cp_size}")
if extend_len < 0:
raise ValueError(f"extend_len must be non-negative, got {extend_len}")
if (
get_in_seq_page_compute_owner_unavailable_reason(
extend_len=extend_len,
extend_prefix_len=extend_prefix_len,
page_size=page_size,
cp_size=cp_size,
)
is not None
):
return None
full_pages = extend_len // page_size
tail_tokens = extend_len % page_size
num_page_units = full_pages + (1 if tail_tokens > 0 else 0)
cp_segment_num = cp_size * 2
base_units = num_page_units // cp_segment_num
remainder_units = num_page_units % cp_segment_num
owners: List[int] = []
for segment_idx in range(cp_segment_num):
unit_count = base_units + (1 if segment_idx < remainder_units else 0)
if segment_idx < cp_size:
owner = segment_idx
else:
owner = cp_segment_num - segment_idx - 1
owners.extend([owner] * unit_count)
return owners

View File

@@ -422,6 +422,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
nsa_cp_metadata: Optional[NSAContextParallelMetadata] = None
uses_cp_shared_kv: bool = False
cp_shared_kv_layout: Optional[CpSharedKVLayout] = None
cp_local_out_cache_loc: Optional[torch.Tensor] = None
cp_shared_mla_direct_write_done: bool = False
# For hidden states before normal
return_hidden_states_before_norm: bool = False

View File

@@ -860,7 +860,7 @@ class ModelRunnerKVCacheMixin:
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",
"CP shared KV enabled. physical_tokens_per_rank=%s, logical_tokens=%s, cp_size=%s, shard_policy=compute_owner_page_aligned_when_available",
self.physical_max_total_num_tokens,
self.max_total_num_tokens,
self.server_args.attn_cp_size,

View File

@@ -6,7 +6,11 @@ import torch
from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph
from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.attention.nsa.utils import nsa_use_prefill_cp
from sglang.srt.layers.attention.nsa.utils import (
get_cp_shared_kv_local_out_cache_loc,
log_cp_shared_kv_direct_write_fallback,
nsa_use_prefill_cp,
)
from sglang.srt.layers.communicator import get_attn_tp_context
from sglang.srt.layers.quantization.fp8_kernel import (
fp8_dtype,
@@ -299,7 +303,16 @@ class DeepseekMLAForwardMixin:
):
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
shared_mla_direct_write_done = False
if nsa_use_prefill_cp(forward_batch):
shared_mla_direct_write_done = self._maybe_write_cp_shared_local_mla_kv(
forward_batch,
k_nope,
k_pe,
)
forward_batch.cp_shared_mla_direct_write_done = (
shared_mla_direct_write_done
)
# support allgather+rerrange
k_nope, k_pe = self.rebuild_cp_kv_cache(
latent_cache, forward_batch, k_nope, k_pe
@@ -329,7 +342,9 @@ class DeepseekMLAForwardMixin:
topk_indices,
llama_4_scaling,
):
save_kv_cache = True
save_kv_cache = not getattr(
forward_batch, "cp_shared_mla_direct_write_done", False
)
if self.current_attention_backend in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS:
extra_args = {}
@@ -345,6 +360,7 @@ class DeepseekMLAForwardMixin:
k_nope,
k_nope,
forward_batch,
save_kv_cache=save_kv_cache,
q_rope=q_pe,
k_rope=k_pe,
**extra_args,
@@ -512,6 +528,46 @@ class DeepseekMLAForwardMixin:
return output
def _maybe_write_cp_shared_local_mla_kv(
self: DeepseekV2AttentionMLA,
forward_batch: ForwardBatch,
k_nope: torch.Tensor,
k_pe: torch.Tensor,
) -> bool:
local_out_cache_loc = get_cp_shared_kv_local_out_cache_loc(forward_batch)
if local_out_cache_loc is None:
return False
if (
k_nope.shape[0] != local_out_cache_loc.numel()
or k_pe.shape[0] != local_out_cache_loc.numel()
):
log_cp_shared_kv_direct_write_fallback(
"mla_local_shape_mismatch",
"MLA local KV token count does not match local out_cache_loc: "
"k_nope=%s k_pe=%s local_out_cache_loc=%s layer_id=%s",
k_nope.shape[0],
k_pe.shape[0],
local_out_cache_loc.numel(),
self.attn_mqa.layer_id,
)
return False
if local_out_cache_loc.numel() == 0:
return True
assert forward_batch.cp_shared_kv_layout is not None
physical_out_cache_loc = (
forward_batch.cp_shared_kv_layout.logical_locs_to_physical(
local_out_cache_loc
).contiguous()
)
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
self.attn_mqa,
physical_out_cache_loc,
k_nope,
k_pe,
)
return True
def _fuse_rope_for_trtllm_mla(
self: DeepseekV2AttentionMLA, forward_batch: ForwardBatch
) -> bool: