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:
@@ -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)
|
||||
|
||||
@@ -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[
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
81
python/sglang/srt/mem_cache/cp_shared_kv_compute_owner.py
Normal file
81
python/sglang/srt/mem_cache/cp_shared_kv_compute_owner.py
Normal 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
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user