[Feat][GLM5.2] Add DSA Cache Layer Split under Prefill CP (#29421)
Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com>
This commit is contained in:
@@ -49,6 +49,7 @@ class KVArgs:
|
||||
state_item_lens: List[List[int]]
|
||||
# Per-tensor TP slice dim, used when prefill/decode attn_tp_size differ.
|
||||
state_dim_per_tensor: List[List[int]]
|
||||
is_hybrid_mla_backend: bool
|
||||
ib_device: str
|
||||
ib_traffic_class: str
|
||||
gpu_id: int
|
||||
|
||||
@@ -73,6 +73,7 @@ class PrefillServerInfo:
|
||||
page_size: Optional[int]
|
||||
kv_cache_dtype: Optional[str]
|
||||
follow_bootstrap_room: bool
|
||||
enable_dsa_cache_layer_split: bool = False
|
||||
|
||||
# PD true-retraction rebootstrap: the prefill's HTTP API port. The decode
|
||||
# already knows the prefill host (the bootstrap_addr host), so it can POST
|
||||
@@ -98,6 +99,7 @@ class PrefillServerInfo:
|
||||
str(self.kv_cache_dtype) if self.kv_cache_dtype is not None else None
|
||||
)
|
||||
self.follow_bootstrap_room = bool(self.follow_bootstrap_room)
|
||||
self.enable_dsa_cache_layer_split = bool(self.enable_dsa_cache_layer_split)
|
||||
self.prefill_http_port = (
|
||||
int(self.prefill_http_port) if self.prefill_http_port is not None else None
|
||||
)
|
||||
@@ -125,6 +127,7 @@ class CommonKVManager(BaseKVManager):
|
||||
self.kv_item_lens_sum = sum(args.kv_item_lens)
|
||||
self.state_item_lens_sum = sum(x for comp in args.state_item_lens for x in comp)
|
||||
self.is_mla_backend = is_mla_backend
|
||||
self.is_hybrid_mla_backend = getattr(args, "is_hybrid_mla_backend", False)
|
||||
self.disaggregation_mode = disaggregation_mode
|
||||
self.server_args = server_args
|
||||
# for p/d multi node infer
|
||||
@@ -146,8 +149,18 @@ class CommonKVManager(BaseKVManager):
|
||||
self.pp_size = server_args.pp_size
|
||||
self.pp_rank = self.kv_args.pp_rank
|
||||
self.local_ip = get_local_ip_auto()
|
||||
cp_sharded_prefill = self.attn_cp_size > 1 and (
|
||||
self.is_hybrid_mla_backend or server_args.enable_dsa_cache_layer_split
|
||||
)
|
||||
|
||||
hybrid_decode_pulls_all_ranks = (
|
||||
self.is_hybrid_mla_backend
|
||||
and disaggregation_mode == DisaggregationMode.DECODE
|
||||
)
|
||||
self.enable_all_cp_ranks_for_transfer = (
|
||||
envs.SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER.get()
|
||||
or cp_sharded_prefill
|
||||
or hybrid_decode_pulls_all_ranks
|
||||
)
|
||||
|
||||
# bind zmq socket
|
||||
@@ -450,7 +463,7 @@ class CommonKVManager(BaseKVManager):
|
||||
required_prefill_response_num = 1
|
||||
target_tp_ranks = [target_tp_rank]
|
||||
elif self.attn_tp_size > info.attn_tp_size:
|
||||
if not self.is_mla_backend:
|
||||
if not self.is_mla_backend and not self.is_hybrid_mla_backend:
|
||||
logger.warning_once(
|
||||
"Performance is NOT guaranteed when using different TP sizes for non-MLA models. "
|
||||
)
|
||||
@@ -461,7 +474,7 @@ class CommonKVManager(BaseKVManager):
|
||||
required_prefill_response_num = 1
|
||||
target_tp_ranks = [target_tp_rank]
|
||||
else:
|
||||
if not self.is_mla_backend:
|
||||
if not self.is_mla_backend and not self.is_hybrid_mla_backend:
|
||||
logger.warning_once(
|
||||
"Performance is NOT guaranteed when using different TP sizes for non-MLA models. "
|
||||
)
|
||||
@@ -495,7 +508,11 @@ class CommonKVManager(BaseKVManager):
|
||||
target_cp_ranks = [self.attn_cp_rank]
|
||||
else:
|
||||
target_cp_ranks = list(range(info.attn_cp_size))
|
||||
if not self.enable_all_cp_ranks_for_transfer:
|
||||
pull_from_all_cp_ranks = (
|
||||
self.enable_all_cp_ranks_for_transfer
|
||||
or info.enable_dsa_cache_layer_split
|
||||
)
|
||||
if not pull_from_all_cp_ranks:
|
||||
# Only retrieve from prefill CP rank 0 when not using all ranks
|
||||
target_cp_ranks = target_cp_ranks[:1]
|
||||
required_prefill_response_num *= 1
|
||||
@@ -582,6 +599,9 @@ class CommonKVManager(BaseKVManager):
|
||||
"page_size": self.kv_args.page_size,
|
||||
"kv_cache_dtype": self.server_args.kv_cache_dtype,
|
||||
"load_balance_method": self.server_args.load_balance_method,
|
||||
"enable_dsa_cache_layer_split": getattr(
|
||||
self.server_args, "enable_dsa_cache_layer_split", False
|
||||
),
|
||||
# Self-register the HTTP API port so the decode can derive the PD
|
||||
# retract rebootstrap /generate URL from bootstrap info instead of a
|
||||
# router-injected pd_rebootstrap_prefill_url.
|
||||
@@ -1041,7 +1061,10 @@ class CommonKVSender(BaseKVSender):
|
||||
self.curr_idx += len(kv_indices)
|
||||
is_last_chunk = self.curr_idx == self.num_kv_indices
|
||||
|
||||
if self.kv_mgr.enable_all_cp_ranks_for_transfer:
|
||||
if (
|
||||
self.kv_mgr.enable_all_cp_ranks_for_transfer
|
||||
and not self.kv_mgr.server_args.enable_dsa_cache_layer_split
|
||||
):
|
||||
kv_indices, index_slice = filter_kv_indices_for_cp_rank(
|
||||
self.kv_mgr,
|
||||
kv_indices,
|
||||
@@ -1396,6 +1419,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
||||
self.page_size = None
|
||||
self.kv_cache_dtype: Optional[str] = None
|
||||
self.follow_bootstrap_room: Optional[bool] = None
|
||||
self.enable_dsa_cache_layer_split: Optional[bool] = None
|
||||
self.prefill_http_port: Optional[int] = None
|
||||
self.prefill_port_table: Dict[
|
||||
int, Dict[int, Dict[int, Dict[int, PrefillRankInfo]]]
|
||||
@@ -1492,6 +1516,11 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
||||
)
|
||||
self.follow_bootstrap_room = load_balance_method == "follow_bootstrap_room"
|
||||
|
||||
if self.enable_dsa_cache_layer_split is None:
|
||||
self.enable_dsa_cache_layer_split = bool(
|
||||
data.get("enable_dsa_cache_layer_split", False)
|
||||
)
|
||||
|
||||
if system_dp_size == 1:
|
||||
dp_group = attn_dp_rank
|
||||
else:
|
||||
@@ -1555,6 +1584,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
||||
if self.follow_bootstrap_room is not None
|
||||
else True
|
||||
),
|
||||
enable_dsa_cache_layer_split=bool(self.enable_dsa_cache_layer_split),
|
||||
prefill_http_port=self.prefill_http_port,
|
||||
)
|
||||
return web.json_response(dataclasses.asdict(info), status=200)
|
||||
|
||||
@@ -605,7 +605,7 @@ class MooncakeKVManager(CommonKVManager):
|
||||
layers_params = None
|
||||
|
||||
# Decode pp size should be equal to prefill pp size or 1
|
||||
if self.is_mla_backend or force_flat:
|
||||
if self.is_mla_backend or self.is_hybrid_mla_backend or force_flat:
|
||||
src_kv_ptrs, dst_kv_ptrs, layers_current_pp_stage = (
|
||||
self.get_mla_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs, state_type)
|
||||
)
|
||||
@@ -924,6 +924,31 @@ class MooncakeKVManager(CommonKVManager):
|
||||
f"Received AUX_DATA for bootstrap_room {room} with length:{len(data)}"
|
||||
)
|
||||
|
||||
def _get_dsa_cache_transfer_skip_flags(
|
||||
self, info: Optional[KVArgsRegisterInfo]
|
||||
) -> Tuple[bool, bool]:
|
||||
skip_kv = False
|
||||
skip_state = False
|
||||
if not self.is_hybrid_mla_backend:
|
||||
return skip_kv, skip_state
|
||||
|
||||
if info is not None and self.attn_tp_size > info.dst_attn_tp_size:
|
||||
sub_rank = (self.kv_args.engine_rank % self.attn_tp_size) % (
|
||||
self.attn_tp_size // info.dst_attn_tp_size
|
||||
)
|
||||
if sub_rank != 0:
|
||||
skip_kv = True
|
||||
skip_state = True
|
||||
|
||||
if (
|
||||
self.attn_cp_size > 1
|
||||
and self.attn_cp_rank != 0
|
||||
and not self.server_args.enable_dsa_cache_layer_split
|
||||
):
|
||||
skip_state = True
|
||||
|
||||
return skip_kv, skip_state
|
||||
|
||||
def maybe_send_extra(
|
||||
self,
|
||||
req: TransferInfo,
|
||||
@@ -1308,10 +1333,15 @@ class MooncakeKVManager(CommonKVManager):
|
||||
target_rank_registration_info: KVArgsRegisterInfo = (
|
||||
self.decode_kv_args_table[req.mooncake_session_id]
|
||||
)
|
||||
if len(kv_chunk.prefill_kv_indices) == 0:
|
||||
skip_kv, skip_state = self._get_dsa_cache_transfer_skip_flags(
|
||||
target_rank_registration_info
|
||||
)
|
||||
if len(kv_chunk.prefill_kv_indices) == 0 or skip_kv:
|
||||
ret = 0
|
||||
elif self.is_mla_backend or (
|
||||
self.attn_tp_size
|
||||
elif (
|
||||
self.is_mla_backend
|
||||
or self.is_hybrid_mla_backend
|
||||
or self.attn_tp_size
|
||||
== target_rank_registration_info.dst_attn_tp_size
|
||||
):
|
||||
ret = self.send_kvcache(
|
||||
@@ -1376,7 +1406,7 @@ class MooncakeKVManager(CommonKVManager):
|
||||
break
|
||||
|
||||
if kv_chunk.is_last_chunk:
|
||||
if kv_chunk.state_indices:
|
||||
if kv_chunk.state_indices and not skip_state:
|
||||
self.maybe_send_extra(
|
||||
req,
|
||||
kv_chunk.state_indices,
|
||||
|
||||
@@ -154,14 +154,34 @@ class PrefillBootstrapQueue:
|
||||
kv_args.engine_rank = self.tp_rank
|
||||
kv_args.pp_rank = self.pp_rank
|
||||
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
|
||||
kv_args.prefill_start_layer = self.token_to_kv_pool.start_layer
|
||||
kv_args.prefill_end_layer = getattr(self.token_to_kv_pool, "end_layer", None)
|
||||
layer_shard_enabled = getattr(
|
||||
self.token_to_kv_pool, "layer_shard_enabled", False
|
||||
)
|
||||
layer_shard_rank = getattr(self.token_to_kv_pool, "layer_shard_rank", None)
|
||||
layer_shard_size = getattr(self.token_to_kv_pool, "layer_shard_size", 1)
|
||||
transfer_draft_cache = (
|
||||
not layer_shard_enabled or layer_shard_rank == layer_shard_size - 1
|
||||
)
|
||||
kv_args.prefill_start_layer = (
|
||||
getattr(
|
||||
self.token_to_kv_pool,
|
||||
"layer_shard_start",
|
||||
self.token_to_kv_pool.start_layer,
|
||||
)
|
||||
if layer_shard_enabled
|
||||
else self.token_to_kv_pool.start_layer
|
||||
)
|
||||
kv_args.mla_compression_ratios = None
|
||||
kv_data_ptrs, kv_data_lens, kv_item_lens = (
|
||||
self.token_to_kv_pool.get_contiguous_buf_infos()
|
||||
)
|
||||
kv_args.prefill_end_layer = (
|
||||
kv_args.prefill_start_layer + len(kv_data_ptrs)
|
||||
if layer_shard_enabled
|
||||
else getattr(self.token_to_kv_pool, "end_layer", None)
|
||||
)
|
||||
|
||||
if self.draft_token_to_kv_pool is not None:
|
||||
if self.draft_token_to_kv_pool is not None and transfer_draft_cache:
|
||||
# We should also transfer draft model kv cache. The indices are
|
||||
# always shared with a target model.
|
||||
draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = (
|
||||
@@ -191,7 +211,7 @@ class PrefillBootstrapQueue:
|
||||
setup_state_kv_args(
|
||||
kv_args,
|
||||
self.token_to_kv_pool,
|
||||
self.draft_token_to_kv_pool,
|
||||
self.draft_token_to_kv_pool if transfer_draft_cache else None,
|
||||
self.scheduler.model_config.num_hidden_layers,
|
||||
req_to_token_pool=req_to_token_pool,
|
||||
)
|
||||
|
||||
@@ -680,6 +680,7 @@ def setup_state_kv_args(
|
||||
kv_args.state_data_lens = []
|
||||
kv_args.state_item_lens = []
|
||||
kv_args.state_dim_per_tensor = []
|
||||
kv_args.is_hybrid_mla_backend = False
|
||||
|
||||
if isinstance(token_to_kv_pool, MiniMaxSparseKVPool):
|
||||
if token_to_kv_pool.index_kv_pool is not None:
|
||||
@@ -733,6 +734,9 @@ def setup_state_kv_args(
|
||||
if hasattr(token_to_kv_pool, "get_state_dim_per_tensor")
|
||||
else None
|
||||
)
|
||||
kv_args.is_hybrid_mla_backend = is_mla_backend(
|
||||
token_to_kv_pool.full_kv_pool
|
||||
)
|
||||
append_state_component(
|
||||
kv_args, StateType.MAMBA, data_ptrs, data_lens, item_lens, dim
|
||||
)
|
||||
|
||||
@@ -673,6 +673,10 @@ class Indexer(MultiPlatformOp):
|
||||
out_cache_loc = forward_batch.out_cache_loc
|
||||
pool = get_token_to_kv_pool()
|
||||
page_size = pool.page_size
|
||||
if hasattr(pool, "invalidate_index_buffer_for_layer"):
|
||||
pool.invalidate_index_buffer_for_layer(layer_id)
|
||||
if hasattr(pool, "_is_layer_owned") and not pool._is_layer_owned(layer_id):
|
||||
return
|
||||
if (
|
||||
not _is_fp8_fnuz
|
||||
and out_cache_loc is not None
|
||||
@@ -801,6 +805,15 @@ class Indexer(MultiPlatformOp):
|
||||
return
|
||||
dst.copy_(src)
|
||||
|
||||
@staticmethod
|
||||
def _get_index_k_read_buffer(pool, layer_id: int) -> torch.Tensor:
|
||||
# Read path: prefer the owner-broadcast scratch buffer under DSA cache
|
||||
# layer split; fall back to the owned buffer for plain pools. Stores go
|
||||
# through get_index_k_with_scale_buffer() (owned buffer) instead.
|
||||
if hasattr(pool, "get_broadcastable_index_k_with_scale_buffer"):
|
||||
return pool.get_broadcastable_index_k_with_scale_buffer(layer_id)
|
||||
return pool.get_index_k_with_scale_buffer(layer_id=layer_id)
|
||||
|
||||
@staticmethod
|
||||
def _pad_heads_for_deep_gemm(q_fp8, weights):
|
||||
"""Pad q and weights to 32 heads when num_heads < 32,
|
||||
@@ -878,9 +891,7 @@ class Indexer(MultiPlatformOp):
|
||||
block_tables = metadata.get_page_table_64()
|
||||
|
||||
max_seq_len = block_tables.shape[1] * page_size
|
||||
kv_cache_fp8 = get_token_to_kv_pool().get_index_k_with_scale_buffer(
|
||||
layer_id=layer_id
|
||||
)
|
||||
kv_cache_fp8 = self._get_index_k_read_buffer(get_token_to_kv_pool(), layer_id)
|
||||
|
||||
blocksize = page_size
|
||||
if (
|
||||
@@ -1627,24 +1638,28 @@ class Indexer(MultiPlatformOp):
|
||||
if out_cache_loc is None:
|
||||
out_cache_loc = forward_batch.out_cache_loc
|
||||
|
||||
pool = get_token_to_kv_pool()
|
||||
if hasattr(pool, "invalidate_index_buffer_for_layer"):
|
||||
pool.invalidate_index_buffer_for_layer(layer_id)
|
||||
if hasattr(pool, "_is_layer_owned") and not pool._is_layer_owned(layer_id):
|
||||
return
|
||||
|
||||
if (
|
||||
_is_cuda
|
||||
and (not _is_fp8_fnuz)
|
||||
and can_use_dsa_fused_store(
|
||||
key.dtype,
|
||||
out_cache_loc.dtype,
|
||||
get_token_to_kv_pool().page_size,
|
||||
pool.page_size,
|
||||
)
|
||||
):
|
||||
# NOTE: wrapper already normalizes shape/contiguity and asserts dtypes.
|
||||
buf = get_token_to_kv_pool().get_index_k_with_scale_buffer(
|
||||
layer_id=layer_id
|
||||
)
|
||||
buf = pool.get_index_k_with_scale_buffer(layer_id=layer_id)
|
||||
fused_store_index_k_cache(
|
||||
key,
|
||||
buf,
|
||||
out_cache_loc,
|
||||
get_token_to_kv_pool().page_size,
|
||||
pool.page_size,
|
||||
)
|
||||
return
|
||||
|
||||
@@ -1654,10 +1669,8 @@ class Indexer(MultiPlatformOp):
|
||||
# layout with page_size=1; the same kv_cache.view works for both cases
|
||||
# because page_size is 1 there.
|
||||
if _use_aiter:
|
||||
page_size = get_token_to_kv_pool().page_size
|
||||
buf = get_token_to_kv_pool().get_index_k_with_scale_buffer(
|
||||
layer_id=layer_id
|
||||
)
|
||||
page_size = pool.page_size
|
||||
buf = pool.get_index_k_with_scale_buffer(layer_id=layer_id)
|
||||
kv_cache = buf.view(-1, page_size, 132).view(fp8_dtype)
|
||||
out_loc = forward_batch.out_cache_loc
|
||||
if not out_loc.is_contiguous():
|
||||
@@ -1679,7 +1692,7 @@ class Indexer(MultiPlatformOp):
|
||||
if not out_cache_loc.is_contiguous():
|
||||
out_cache_loc = out_cache_loc.contiguous()
|
||||
|
||||
get_token_to_kv_pool().set_index_k_scale_buffer(
|
||||
pool.set_index_k_scale_buffer(
|
||||
layer_id=layer_id,
|
||||
loc=out_cache_loc,
|
||||
index_k=k_fp8,
|
||||
|
||||
@@ -38,6 +38,7 @@ from sglang.srt.layers.dp_attention import (
|
||||
)
|
||||
from sglang.srt.layers.utils.cp_utils import mla_use_prefill_cp
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
|
||||
@@ -48,6 +49,25 @@ def dsa_enable_prefill_cp():
|
||||
return is_dsa_enable_prefill_cp()
|
||||
|
||||
|
||||
def maybe_prefetch_next_full_attention_kv(
|
||||
forward_batch: ForwardBatch,
|
||||
next_full_attention_layer_id: Optional[int],
|
||||
) -> None:
|
||||
"""Prefetch (owner-broadcast) the next layer's DSA KV under layer split.
|
||||
|
||||
No-op unless the current batch runs DSA prefill-CP and the active KV pool is
|
||||
a layer-sharded pool exposing ``prefetch_kv_buffer`` (i.e.
|
||||
``LayerSplitDSATokenToKVPool``). Kicking the broadcast off one layer ahead
|
||||
overlaps it with the current layer's attention compute.
|
||||
"""
|
||||
if next_full_attention_layer_id is None or not dsa_use_prefill_cp(forward_batch):
|
||||
return
|
||||
|
||||
prefetch_kv_buffer = getattr(get_token_to_kv_pool(), "prefetch_kv_buffer", None)
|
||||
if prefetch_kv_buffer is not None:
|
||||
prefetch_kv_buffer(next_full_attention_layer_id)
|
||||
|
||||
|
||||
def dsa_cp_gather_hidden_states(hidden_states: torch.Tensor):
|
||||
attn_dp_size = get_parallel().attn_dp_size
|
||||
attn_tp_size = get_parallel().attn_tp_size
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
|
||||
"""Public import facade and runtime helpers for context parallel strategies."""
|
||||
|
||||
from typing import Any, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any, Optional, Tuple
|
||||
|
||||
from sglang.srt.layers.cp.base import (
|
||||
BaseContextParallelMetadata,
|
||||
@@ -33,6 +33,9 @@ from sglang.srt.layers.cp.zigzag import (
|
||||
ZigzagCPStrategy,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
|
||||
CP_V2_DEFAULT_MODEL_CLASSES = frozenset(
|
||||
{
|
||||
"Qwen3MoeForCausalLM",
|
||||
@@ -40,6 +43,89 @@ CP_V2_DEFAULT_MODEL_CLASSES = frozenset(
|
||||
)
|
||||
|
||||
|
||||
def is_glm_dsa_cache_layer_split_enabled(model_runner: "ModelRunner") -> bool:
|
||||
"""Whether DSA GPU KV/indexer cache layers are sharded across CP ranks.
|
||||
|
||||
Layer split is a prefill-CP-only optimization for DSA (DeepSeek Sparse
|
||||
Attention) MLA models (e.g. GLM-5.2). Draft workers keep the full cache.
|
||||
"""
|
||||
from sglang.srt.configs.model_config import is_deepseek_dsa
|
||||
|
||||
return (
|
||||
not model_runner.is_draft_worker
|
||||
and model_runner.server_args.enable_dsa_cache_layer_split
|
||||
and model_runner.use_mla_backend
|
||||
and is_deepseek_dsa(model_runner.model_config.hf_config)
|
||||
)
|
||||
|
||||
|
||||
def get_glm_dsa_cp_layer_shard_info(
|
||||
model_runner: "ModelRunner",
|
||||
) -> Tuple[Optional[int], int]:
|
||||
"""Return ``(layer_shard_rank, layer_shard_size)`` for the DSA KV pool.
|
||||
|
||||
``(None, 1)`` disables sharding (feature off or only one CP rank).
|
||||
"""
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
get_attention_cp_rank,
|
||||
get_attention_cp_size,
|
||||
)
|
||||
|
||||
if not is_glm_dsa_cache_layer_split_enabled(model_runner):
|
||||
return None, 1
|
||||
shard_size = get_attention_cp_size()
|
||||
if shard_size <= 1:
|
||||
return None, 1
|
||||
return get_attention_cp_rank(), shard_size
|
||||
|
||||
|
||||
def get_glm_dsa_layer_split_effective_num_layers(
|
||||
model_runner: "ModelRunner", num_layers: int
|
||||
) -> int:
|
||||
"""Per-rank owned layer count used when sizing the DSA KV cell.
|
||||
|
||||
Under layer split each CP rank only stores ``ceil(num_layers / shard_size)``
|
||||
layers, plus one extra layer for the remote scratch buffer used when reading
|
||||
a layer owned by another CP rank.
|
||||
"""
|
||||
from sglang.srt.layers.dp_attention import get_attention_cp_size
|
||||
|
||||
if not is_glm_dsa_cache_layer_split_enabled(model_runner):
|
||||
return num_layers
|
||||
shard_size = get_attention_cp_size()
|
||||
if shard_size <= 1:
|
||||
return num_layers
|
||||
owned_layers_upper_bound = (num_layers + shard_size - 1) // shard_size
|
||||
return max(1, owned_layers_upper_bound + 1)
|
||||
|
||||
|
||||
def get_layer_shard_range(
|
||||
rank: int, shard_size: int, total_layers: int
|
||||
) -> Tuple[int, int]:
|
||||
"""Contiguous ``[start, end)`` local-layer range owned by ``rank``.
|
||||
|
||||
Layers are split as evenly as possible; the first ``total_layers %
|
||||
shard_size`` ranks own one extra layer.
|
||||
"""
|
||||
base = total_layers // shard_size
|
||||
rem = total_layers % shard_size
|
||||
start = rank * base + min(rank, rem)
|
||||
end = start + base + (1 if rank < rem else 0)
|
||||
return start, end
|
||||
|
||||
|
||||
def get_layer_owner(local_layer_idx: int, shard_size: int, total_layers: int) -> int:
|
||||
"""CP rank that owns ``local_layer_idx`` under the contiguous split."""
|
||||
for rank in range(shard_size):
|
||||
start, end = get_layer_shard_range(rank, shard_size, total_layers)
|
||||
if start <= local_layer_idx < end:
|
||||
return rank
|
||||
raise ValueError(
|
||||
f"Invalid local_layer_idx={local_layer_idx} for "
|
||||
f"shard_size={shard_size}, total_layers={total_layers}"
|
||||
)
|
||||
|
||||
|
||||
def enable_cp_v2() -> bool:
|
||||
"""Return whether the CP-v2 path is enabled for this process."""
|
||||
from sglang.srt.environ import envs
|
||||
@@ -140,4 +226,9 @@ __all__ = [
|
||||
"cp_gather_after_forward",
|
||||
"cp_split_before_forward",
|
||||
"prepare_cp_forward",
|
||||
"is_glm_dsa_cache_layer_split_enabled",
|
||||
"get_glm_dsa_cp_layer_shard_info",
|
||||
"get_glm_dsa_layer_split_effective_num_layers",
|
||||
"get_layer_shard_range",
|
||||
"get_layer_owner",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,581 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
|
||||
"""Layer-sharded DSA KV cache pool for context-parallel prefill.
|
||||
|
||||
``LayerSplitDSATokenToKVPool`` splits the DSA (DeepSeek Sparse Attention) GPU
|
||||
KV/indexer cache layers across context-parallel (CP) ranks so that each rank
|
||||
only materializes the layers it owns, reducing per-rank KV memory. When a rank
|
||||
needs to read a layer it does not own, the owning rank broadcasts that layer's
|
||||
buffer into a small per-rank remote scratch buffer.
|
||||
|
||||
This subclass keeps the core ``KVCache`` / ``MLATokenToKVPool`` /
|
||||
``DSATokenToKVPool`` pools untouched: all sharding, broadcast, and remote-scratch
|
||||
bookkeeping lives here. Layer split is only ever enabled for DSA MLA models on
|
||||
PD prefill workers under prefill-CP (see
|
||||
``sglang.srt.layers.cp.utils.is_glm_dsa_cache_layer_split_enabled``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from contextlib import nullcontext
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.dsa import index_buf_accessor
|
||||
from sglang.srt.layers.cp.utils import get_layer_owner, get_layer_shard_range
|
||||
from sglang.srt.layers.dp_attention import get_attention_cp_group
|
||||
from sglang.srt.mem_cache.memory_pool import (
|
||||
GPU_MEMORY_TYPE_KV_CACHE,
|
||||
DSATokenToKVPool,
|
||||
RadixAttention,
|
||||
get_tensor_size_bytes,
|
||||
maybe_detect_oob,
|
||||
unwrap_write_loc,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.cache_controller import LayerDoneCounter
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class LayerSplitDSATokenToKVPool(DSATokenToKVPool):
|
||||
"""DSA KV pool that shards layers across CP ranks with owner-broadcast reads."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*args,
|
||||
layer_shard_rank: int,
|
||||
layer_shard_size: int,
|
||||
**kwargs,
|
||||
):
|
||||
assert (
|
||||
layer_shard_rank is not None and layer_shard_size > 1
|
||||
), "LayerSplitDSATokenToKVPool requires layer_shard_size > 1"
|
||||
self.layer_shard_rank = layer_shard_rank
|
||||
self.layer_shard_size = layer_shard_size
|
||||
self.layer_shard_enabled = True
|
||||
self.layer_broadcast_comm = None
|
||||
super().__init__(*args, **kwargs)
|
||||
# First global layer index owned by this rank (used by PD transfer to
|
||||
# label the contiguous owned-buffer range).
|
||||
my_start, _ = self._owned_local_layer_range()
|
||||
self.layer_shard_start = self.start_layer + my_start
|
||||
|
||||
# ---- layer ownership helpers ------------------------------------------
|
||||
|
||||
def _local_layer_idx(self, layer_id: int) -> int:
|
||||
return layer_id - self.start_layer
|
||||
|
||||
def _owned_local_layer_range(self) -> tuple[int, int]:
|
||||
return get_layer_shard_range(
|
||||
self.layer_shard_rank, self.layer_shard_size, self.layer_num
|
||||
)
|
||||
|
||||
def _is_layer_owned(self, layer_id: int) -> bool:
|
||||
local_idx = self._local_layer_idx(layer_id)
|
||||
owned_start, owned_end = self._owned_local_layer_range()
|
||||
return owned_start <= local_idx < owned_end
|
||||
|
||||
def _get_layer_owner_rank(self, layer_id: int) -> int:
|
||||
return get_layer_owner(
|
||||
self._local_layer_idx(layer_id), self.layer_shard_size, self.layer_num
|
||||
)
|
||||
|
||||
def _log_layer_shard_plan(self) -> None:
|
||||
partitions = []
|
||||
for rank in range(self.layer_shard_size):
|
||||
st, ed = get_layer_shard_range(rank, self.layer_shard_size, self.layer_num)
|
||||
partitions.append(f"r{rank}:[{st},{ed})")
|
||||
my_start, my_end = self._owned_local_layer_range()
|
||||
logger.info(
|
||||
"Layer shard plan (continuous): "
|
||||
f"layer_num={self.layer_num}, shard_size={self.layer_shard_size}, "
|
||||
f"rank={self.layer_shard_rank}, local=[{my_start},{my_end}), "
|
||||
f"global=[{self.start_layer + my_start},{self.start_layer + my_end}), "
|
||||
f"partitions={'; '.join(partitions)}"
|
||||
)
|
||||
|
||||
# ---- broadcast plumbing -----------------------------------------------
|
||||
|
||||
def _init_layer_broadcast_comm(self) -> None:
|
||||
cp_group = get_attention_cp_group()
|
||||
if cp_group.world_size <= 1 or cp_group.pynccl_comm is None:
|
||||
return
|
||||
|
||||
from sglang.srt.distributed.device_communicators.pynccl import (
|
||||
PyNcclCommunicator,
|
||||
)
|
||||
|
||||
self.layer_broadcast_comm = PyNcclCommunicator(
|
||||
group=cp_group.cpu_group,
|
||||
device=cp_group.device,
|
||||
)
|
||||
logger.info(
|
||||
"Initialized dedicated layer-shard broadcast NCCL communicator: "
|
||||
f"rank={cp_group.rank_in_group}, world_size={cp_group.world_size}"
|
||||
)
|
||||
|
||||
def _broadcast_tensor_from_owner(
|
||||
self,
|
||||
tensor: torch.Tensor,
|
||||
layer_id: int,
|
||||
src_tensor: Optional[torch.Tensor] = None,
|
||||
use_layer_broadcast_comm: bool = False,
|
||||
) -> torch.Tensor:
|
||||
owner_rank = self._get_layer_owner_rank(layer_id)
|
||||
if self.layer_shard_rank == owner_rank:
|
||||
assert src_tensor is not None
|
||||
if tensor.data_ptr() != src_tensor.data_ptr():
|
||||
tensor.copy_(src_tensor)
|
||||
|
||||
cp_group = get_attention_cp_group()
|
||||
comm = (
|
||||
self.layer_broadcast_comm
|
||||
if use_layer_broadcast_comm and self.layer_broadcast_comm is not None
|
||||
else cp_group.pynccl_comm
|
||||
)
|
||||
if comm is not None:
|
||||
# PyNcclCommunicator defaults to disabled=True (it is only enabled
|
||||
# inside CUDA-graph capture via change_state). Without re-enabling it
|
||||
# here, comm.broadcast() is a silent no-op and non-owner CP ranks read
|
||||
# stale remote buffers, corrupting layer-split attention. Mirror the
|
||||
# standard usage in parallel_state.py.
|
||||
with comm.change_state(enable=True):
|
||||
comm.broadcast(tensor, src=owner_rank)
|
||||
else:
|
||||
torch.distributed.broadcast(
|
||||
tensor, src=owner_rank, group=cp_group.cpu_group
|
||||
)
|
||||
return tensor
|
||||
|
||||
# ---- buffer allocation (owned-only + remote scratch) ------------------
|
||||
|
||||
def _create_buffers(self):
|
||||
self._log_layer_shard_plan()
|
||||
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||||
with (
|
||||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||||
if self.custom_mem_pool
|
||||
else nullcontext()
|
||||
):
|
||||
# Owned layers get the full buffer; non-owned layers allocate a
|
||||
# 0-row placeholder so ``kv_buffer`` stays index-aligned by layer.
|
||||
self.kv_buffer = [
|
||||
torch.zeros(
|
||||
(
|
||||
(
|
||||
(self.size + self.page_size)
|
||||
if self._is_layer_owned(self.start_layer + i)
|
||||
else 0
|
||||
),
|
||||
1,
|
||||
self.kv_cache_dim,
|
||||
),
|
||||
dtype=self.store_dtype,
|
||||
device=self.device,
|
||||
)
|
||||
for i in range(self.layer_num)
|
||||
]
|
||||
self.remote_kv_buffer = torch.empty(
|
||||
(self.size + self.page_size, 1, self.kv_cache_dim),
|
||||
dtype=self.store_dtype,
|
||||
device=self.device,
|
||||
)
|
||||
self.remote_kv_layer_id: Optional[int] = None
|
||||
self.device_module = torch.get_device_module(self.device)
|
||||
self.kv_broadcast_stream = self.device_module.Stream()
|
||||
self.pending_remote_kv_layer_id: Optional[int] = None
|
||||
self.pending_remote_kv_broadcast = False
|
||||
self._init_layer_broadcast_comm()
|
||||
|
||||
def _create_index_buffers(self):
|
||||
num_pages = (self.index_buf_size + self.page_size + 1) // self.page_size
|
||||
with (
|
||||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||||
if self.custom_mem_pool
|
||||
else nullcontext()
|
||||
):
|
||||
self.index_k_with_scale_buffer = [
|
||||
torch.zeros(
|
||||
self._index_buffer_shape(
|
||||
num_pages if self._is_layer_owned(self.start_layer + i) else 0
|
||||
),
|
||||
dtype=self.index_k_with_scale_buffer_dtype,
|
||||
device=self.device,
|
||||
)
|
||||
for i in range(self.layer_num)
|
||||
]
|
||||
self.remote_index_k_with_scale_buffer = torch.empty(
|
||||
self._index_buffer_shape(num_pages),
|
||||
dtype=self.index_k_with_scale_buffer_dtype,
|
||||
device=self.device,
|
||||
)
|
||||
self.remote_index_layer_id: Optional[int] = None
|
||||
|
||||
def _clear_buffers(self):
|
||||
del self.kv_buffer
|
||||
del self.remote_kv_buffer
|
||||
del self.remote_index_k_with_scale_buffer
|
||||
del self.index_k_with_scale_buffer
|
||||
|
||||
# ---- MLA latent KV: owned-only writes, owner-broadcast reads ----------
|
||||
|
||||
def get_kv_size_bytes(self):
|
||||
kv_size_bytes = 0
|
||||
for kv_cache in self.kv_buffer:
|
||||
kv_size_bytes += get_tensor_size_bytes(kv_cache)
|
||||
for index_k_cache in self.index_k_with_scale_buffer:
|
||||
kv_size_bytes += get_tensor_size_bytes(index_k_cache)
|
||||
return kv_size_bytes
|
||||
|
||||
def get_contiguous_buf_infos(self):
|
||||
# Only report buffers owned by the current CP rank; non-owned layers
|
||||
# are empty and are pulled from their owner via PD transfer.
|
||||
owned_layer_ids = [
|
||||
i
|
||||
for i in range(self.layer_num)
|
||||
if self._is_layer_owned(self.start_layer + i)
|
||||
]
|
||||
kv_data_ptrs = [self.kv_buffer[i].data_ptr() for i in owned_layer_ids]
|
||||
kv_data_lens = [self.kv_buffer[i].nbytes for i in owned_layer_ids]
|
||||
kv_item_lens = [
|
||||
self.kv_buffer[i][0].nbytes * self.page_size for i in owned_layer_ids
|
||||
]
|
||||
return kv_data_ptrs, kv_data_lens, kv_item_lens
|
||||
|
||||
def get_key_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
|
||||
kv_buffer = self._get_broadcastable_kv_buffer(layer_id)
|
||||
if self.store_dtype != self.dtype:
|
||||
return kv_buffer.view(self.dtype)
|
||||
return kv_buffer
|
||||
|
||||
def get_value_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
|
||||
kv_buffer = self._get_broadcastable_kv_buffer(layer_id)
|
||||
if self.store_dtype != self.dtype:
|
||||
return kv_buffer[..., : self.kv_lora_rank].view(self.dtype)
|
||||
return kv_buffer[..., : self.kv_lora_rank]
|
||||
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
loc_info,
|
||||
cache_k: torch.Tensor,
|
||||
cache_v: torch.Tensor,
|
||||
):
|
||||
loc, _, _ = unwrap_write_loc(loc_info)
|
||||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_kv_buffer (MLA)")
|
||||
layer_id = layer.layer_id
|
||||
assert not self.dsa_kv_cache_store_fp8
|
||||
# A write invalidates any cached remote copy for this layer.
|
||||
if self.pending_remote_kv_layer_id == layer_id:
|
||||
self._finalize_pending_kv_broadcast(set_remote_layer_id=False)
|
||||
if self.remote_kv_layer_id == layer_id:
|
||||
self.remote_kv_layer_id = None
|
||||
if not self._is_layer_owned(layer_id):
|
||||
return
|
||||
if cache_k.dtype != self.dtype:
|
||||
cache_k = cache_k.to(self.dtype)
|
||||
if self.store_dtype != self.dtype:
|
||||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k.view(
|
||||
self.store_dtype
|
||||
)
|
||||
else:
|
||||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
|
||||
|
||||
def set_mla_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
loc: torch.Tensor,
|
||||
cache_k_nope: torch.Tensor,
|
||||
cache_k_rope: torch.Tensor,
|
||||
):
|
||||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_mla_kv_buffer (MLA)")
|
||||
layer_id = layer.layer_id
|
||||
if self.pending_remote_kv_layer_id == layer_id:
|
||||
self._finalize_pending_kv_broadcast(set_remote_layer_id=True)
|
||||
remote_kv_updatable = self.remote_kv_layer_id == layer_id
|
||||
if remote_kv_updatable:
|
||||
self._write_mla_kv_buffer(
|
||||
self.remote_kv_buffer, loc, cache_k_nope, cache_k_rope
|
||||
)
|
||||
if not self._is_layer_owned(layer_id):
|
||||
return
|
||||
self._write_mla_kv_buffer(
|
||||
self.kv_buffer[layer_id - self.start_layer],
|
||||
loc,
|
||||
cache_k_nope,
|
||||
cache_k_rope,
|
||||
)
|
||||
if not remote_kv_updatable and self.remote_kv_layer_id == layer_id:
|
||||
self.remote_kv_layer_id = None
|
||||
|
||||
def _finalize_pending_kv_broadcast(
|
||||
self, *, set_remote_layer_id: bool = True
|
||||
) -> None:
|
||||
if not self.pending_remote_kv_broadcast:
|
||||
return
|
||||
self.device_module.current_stream().wait_stream(self.kv_broadcast_stream)
|
||||
self.pending_remote_kv_broadcast = False
|
||||
if set_remote_layer_id and self.pending_remote_kv_layer_id is not None:
|
||||
self.remote_kv_layer_id = self.pending_remote_kv_layer_id
|
||||
self.pending_remote_kv_layer_id = None
|
||||
|
||||
def prefetch_kv_buffer(
|
||||
self,
|
||||
layer_id: int,
|
||||
layer_transfer_counter: Optional[LayerDoneCounter] = None,
|
||||
layer_transfer_idx: Optional[int] = None,
|
||||
) -> None:
|
||||
"""Kick off an async owner-broadcast of ``layer_id``'s latent KV.
|
||||
|
||||
Called ahead of the layer's attention so the remote scratch buffer is
|
||||
ready by the time a non-owner rank reads it (see the prefetch wiring in
|
||||
``DeepseekV2DecoderLayer``).
|
||||
"""
|
||||
if self.remote_kv_layer_id == layer_id:
|
||||
return
|
||||
if self.pending_remote_kv_broadcast:
|
||||
if self.pending_remote_kv_layer_id == layer_id:
|
||||
return
|
||||
self._finalize_pending_kv_broadcast(set_remote_layer_id=False)
|
||||
|
||||
local_idx = self._local_layer_idx(layer_id)
|
||||
src_tensor = (
|
||||
self.kv_buffer[local_idx] if self._is_layer_owned(layer_id) else None
|
||||
)
|
||||
if self.layer_broadcast_comm is None:
|
||||
self._broadcast_tensor_from_owner(
|
||||
self.remote_kv_buffer,
|
||||
layer_id,
|
||||
src_tensor=src_tensor,
|
||||
use_layer_broadcast_comm=True,
|
||||
)
|
||||
self.remote_kv_layer_id = layer_id
|
||||
return
|
||||
|
||||
self.kv_broadcast_stream.wait_stream(self.device_module.current_stream())
|
||||
with self.device_module.stream(self.kv_broadcast_stream):
|
||||
if layer_transfer_counter is not None and layer_transfer_idx is not None:
|
||||
layer_transfer_counter.wait_until(layer_transfer_idx)
|
||||
self._broadcast_tensor_from_owner(
|
||||
self.remote_kv_buffer,
|
||||
layer_id,
|
||||
src_tensor=src_tensor,
|
||||
use_layer_broadcast_comm=True,
|
||||
)
|
||||
self.pending_remote_kv_layer_id = layer_id
|
||||
self.pending_remote_kv_broadcast = True
|
||||
|
||||
def _get_broadcastable_kv_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
if self.pending_remote_kv_broadcast:
|
||||
self._finalize_pending_kv_broadcast(
|
||||
set_remote_layer_id=self.pending_remote_kv_layer_id == layer_id
|
||||
)
|
||||
if self.remote_kv_layer_id != layer_id:
|
||||
local_idx = self._local_layer_idx(layer_id)
|
||||
src_tensor = (
|
||||
self.kv_buffer[local_idx] if self._is_layer_owned(layer_id) else None
|
||||
)
|
||||
self._broadcast_tensor_from_owner(
|
||||
self.remote_kv_buffer,
|
||||
layer_id,
|
||||
src_tensor=src_tensor,
|
||||
use_layer_broadcast_comm=True,
|
||||
)
|
||||
self.remote_kv_layer_id = layer_id
|
||||
return self.remote_kv_buffer
|
||||
|
||||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||||
size_limit = self.size + self.page_size
|
||||
maybe_detect_oob(tgt_loc, 0, size_limit, "move_kv_cache tgt_loc")
|
||||
maybe_detect_oob(src_loc, 0, size_limit, "move_kv_cache src_loc")
|
||||
if tgt_loc.numel() == 0:
|
||||
return
|
||||
tgt_loc_flat = tgt_loc.view(-1).long()
|
||||
src_loc_flat = src_loc.view(-1).long()
|
||||
for kv_cache in self.kv_buffer:
|
||||
if kv_cache.shape[0] == 0:
|
||||
continue
|
||||
kv_cache[tgt_loc_flat] = kv_cache[src_loc_flat]
|
||||
for index_k in self.index_k_with_scale_buffer:
|
||||
if index_k.shape[0] == 0:
|
||||
continue
|
||||
index_k[tgt_loc_flat] = index_k[src_loc_flat]
|
||||
|
||||
# ---- DSA indexer buffer: owned-only writes, owner-broadcast reads -----
|
||||
|
||||
def get_broadcastable_index_k_with_scale_buffer(
|
||||
self, layer_id: int
|
||||
) -> torch.Tensor:
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
return self._get_broadcastable_index_buffer(layer_id)
|
||||
|
||||
def get_index_k_continuous(self, layer_id, seq_len, page_indices):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
buf = self._get_broadcastable_index_buffer(layer_id)
|
||||
return index_buf_accessor.GetK.execute(
|
||||
self, buf, seq_len=seq_len, page_indices=page_indices
|
||||
)
|
||||
|
||||
def get_index_k_scale_continuous(self, layer_id, seq_len, page_indices):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
buf = self._get_broadcastable_index_buffer(layer_id)
|
||||
return index_buf_accessor.GetS.execute(
|
||||
self, buf, seq_len=seq_len, page_indices=page_indices
|
||||
)
|
||||
|
||||
def get_index_k_scale_buffer(
|
||||
self, layer_id, seq_len_tensor, page_indices, seq_len_sum, max_seq_len
|
||||
):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
buf = self._get_broadcastable_index_buffer(layer_id)
|
||||
# Overlap the latent-KV owner-broadcast with the indexer read.
|
||||
self.prefetch_kv_buffer(layer_id)
|
||||
return index_buf_accessor.GetKAndS.execute(
|
||||
self,
|
||||
buf,
|
||||
page_indices=page_indices,
|
||||
seq_len_tensor=seq_len_tensor,
|
||||
seq_len_sum=seq_len_sum,
|
||||
max_seq_len=max_seq_len,
|
||||
)
|
||||
|
||||
def set_index_k_scale_buffer(self, layer_id, loc, index_k, index_k_scale) -> None:
|
||||
self.invalidate_index_buffer_for_layer(layer_id)
|
||||
if not self._is_layer_owned(layer_id):
|
||||
return
|
||||
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer]
|
||||
index_buf_accessor.SetKAndS.execute(
|
||||
pool=self, buf=buf, loc=loc, index_k=index_k, index_k_scale=index_k_scale
|
||||
)
|
||||
|
||||
def invalidate_index_buffer_for_layer(self, layer_id: int) -> None:
|
||||
if self.remote_index_layer_id == layer_id:
|
||||
self.remote_index_layer_id = None
|
||||
|
||||
def _get_broadcastable_index_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
if self.remote_index_layer_id != layer_id:
|
||||
local_idx = self._local_layer_idx(layer_id)
|
||||
src_tensor = (
|
||||
self.index_k_with_scale_buffer[local_idx]
|
||||
if self._is_layer_owned(layer_id)
|
||||
else None
|
||||
)
|
||||
self._broadcast_tensor_from_owner(
|
||||
self.remote_index_k_with_scale_buffer,
|
||||
layer_id,
|
||||
src_tensor=src_tensor,
|
||||
)
|
||||
self.remote_index_layer_id = layer_id
|
||||
return self.remote_index_k_with_scale_buffer
|
||||
|
||||
def get_state_buf_infos(self):
|
||||
owned_layer_ids = [
|
||||
i
|
||||
for i in range(self.layer_num)
|
||||
if self._is_layer_owned(self.start_layer + i)
|
||||
]
|
||||
data_ptrs = [
|
||||
self.index_k_with_scale_buffer[i].data_ptr() for i in owned_layer_ids
|
||||
]
|
||||
data_lens = [self.index_k_with_scale_buffer[i].nbytes for i in owned_layer_ids]
|
||||
item_lens = [
|
||||
self.index_k_with_scale_buffer[i][0].nbytes for i in owned_layer_ids
|
||||
]
|
||||
return data_ptrs, data_lens, item_lens
|
||||
|
||||
# ---- HiCache CPU offload: skip empty (non-owned) layers ---------------
|
||||
|
||||
def get_cpu_copy(self, indices, mamba_indices=None):
|
||||
from sglang.srt.utils import current_platform
|
||||
|
||||
current_platform.synchronize()
|
||||
kv_cache_cpu = []
|
||||
chunk_size = self.cpu_offloading_chunk_size
|
||||
for layer_id in range(self.layer_num):
|
||||
kv_cache_cpu.append([])
|
||||
if self.kv_buffer[layer_id].shape[0] == 0:
|
||||
continue
|
||||
for i in range(0, len(indices), chunk_size):
|
||||
chunk_indices = indices[i : i + chunk_size]
|
||||
kv_cpu = self.kv_buffer[layer_id][chunk_indices].to(
|
||||
"cpu", non_blocking=True
|
||||
)
|
||||
kv_cache_cpu[-1].append(kv_cpu)
|
||||
current_platform.synchronize()
|
||||
|
||||
page_indices = indices[:: self.page_size] // self.page_size
|
||||
torch.cuda.synchronize()
|
||||
index_k_cpu = []
|
||||
page_chunk_size = max(1, chunk_size // self.page_size)
|
||||
for layer_id in range(self.layer_num):
|
||||
index_k_cpu.append([])
|
||||
if self.index_k_with_scale_buffer[layer_id].shape[0] == 0:
|
||||
continue
|
||||
for i in range(0, len(page_indices), page_chunk_size):
|
||||
chunk_page_indices = page_indices[i : i + page_chunk_size]
|
||||
idx_cpu = self.index_k_with_scale_buffer[layer_id][
|
||||
chunk_page_indices
|
||||
].to("cpu", non_blocking=True)
|
||||
index_k_cpu[-1].append(idx_cpu)
|
||||
torch.cuda.synchronize()
|
||||
return {"kv": kv_cache_cpu, "index_k": index_k_cpu}
|
||||
|
||||
def load_cpu_copy(self, kv_cache_cpu_dict, indices, mamba_indices=None):
|
||||
from sglang.srt.utils import current_platform
|
||||
|
||||
kv_cache_cpu = kv_cache_cpu_dict["kv"]
|
||||
current_platform.synchronize()
|
||||
chunk_size = self.cpu_offloading_chunk_size
|
||||
for layer_id in range(self.layer_num):
|
||||
if self.kv_buffer[layer_id].shape[0] == 0:
|
||||
continue
|
||||
for i in range(0, len(indices), chunk_size):
|
||||
chunk_indices = indices[i : i + chunk_size]
|
||||
kv_cpu = kv_cache_cpu[layer_id][i // chunk_size]
|
||||
assert kv_cpu.shape[0] == len(chunk_indices)
|
||||
kv_chunk = kv_cpu.to(self.kv_buffer[layer_id].device, non_blocking=True)
|
||||
self.kv_buffer[layer_id][chunk_indices] = kv_chunk
|
||||
current_platform.synchronize()
|
||||
|
||||
page_indices = indices[:: self.page_size] // self.page_size
|
||||
index_k_cpu = kv_cache_cpu_dict["index_k"]
|
||||
torch.cuda.synchronize()
|
||||
page_chunk_size = max(1, chunk_size // self.page_size)
|
||||
for layer_id in range(self.layer_num):
|
||||
if self.index_k_with_scale_buffer[layer_id].shape[0] == 0:
|
||||
continue
|
||||
for i in range(0, len(page_indices), page_chunk_size):
|
||||
chunk_page_indices = page_indices[i : i + page_chunk_size]
|
||||
idx_cpu = index_k_cpu[layer_id][i // page_chunk_size]
|
||||
assert idx_cpu.shape[0] == len(chunk_page_indices)
|
||||
idx_chunk = idx_cpu.to(
|
||||
self.index_k_with_scale_buffer[layer_id].device, non_blocking=True
|
||||
)
|
||||
self.index_k_with_scale_buffer[layer_id][chunk_page_indices] = idx_chunk
|
||||
torch.cuda.synchronize()
|
||||
@@ -1226,6 +1226,7 @@ class KvBufferDesc:
|
||||
|
||||
|
||||
class KVCache(abc.ABC):
|
||||
layer_shard_enabled: bool = False
|
||||
post_capture_active: bool = False
|
||||
|
||||
@abc.abstractmethod
|
||||
@@ -2838,7 +2839,6 @@ class MLATokenToKVPool(KVCache):
|
||||
if not valid_mask.all():
|
||||
loc = loc[valid_mask]
|
||||
cache_k = cache_k[valid_mask]
|
||||
|
||||
if cache_k.dtype != self.dtype:
|
||||
cache_k = cache_k.to(self.dtype)
|
||||
|
||||
@@ -2849,21 +2849,18 @@ class MLATokenToKVPool(KVCache):
|
||||
else:
|
||||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
|
||||
|
||||
def set_mla_kv_buffer(
|
||||
def _write_mla_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
dst_buffer: torch.Tensor,
|
||||
loc: torch.Tensor,
|
||||
cache_k_nope: torch.Tensor,
|
||||
cache_k_rope: torch.Tensor,
|
||||
):
|
||||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_mla_kv_buffer (MLA)")
|
||||
layer_id = layer.layer_id
|
||||
|
||||
) -> None:
|
||||
if _is_hip and self.use_dsa and self.dtype == fp8_dtype:
|
||||
# HIP FP8 path uses raw MLA KV layout (nope + rope) without per-block scales.
|
||||
# Fuse BF16/FP16 -> FP8 cast with paged KV write.
|
||||
set_mla_kv_buffer_triton_fp8_quant(
|
||||
self.kv_buffer[layer_id - self.start_layer],
|
||||
dst_buffer,
|
||||
loc,
|
||||
cache_k_nope,
|
||||
cache_k_rope,
|
||||
@@ -2881,7 +2878,7 @@ class MLATokenToKVPool(KVCache):
|
||||
# cache_k_nope_fp8: (num_tokens, 1, 528) uint8 [nope_fp8(512) | scales(16)]
|
||||
# cache_k_rope_fp8: (num_tokens, 1, 128) uint8 [rope_bf16_bytes(128)]
|
||||
set_mla_kv_buffer_triton(
|
||||
self.kv_buffer[layer_id - self.start_layer],
|
||||
dst_buffer,
|
||||
loc,
|
||||
cache_k_nope_fp8,
|
||||
cache_k_rope_fp8,
|
||||
@@ -2895,12 +2892,28 @@ class MLATokenToKVPool(KVCache):
|
||||
cache_k_rope = cache_k_rope.view(self.store_dtype)
|
||||
|
||||
set_mla_kv_buffer_triton(
|
||||
self.kv_buffer[layer_id - self.start_layer],
|
||||
dst_buffer,
|
||||
loc,
|
||||
cache_k_nope,
|
||||
cache_k_rope,
|
||||
)
|
||||
|
||||
def set_mla_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
loc: torch.Tensor,
|
||||
cache_k_nope: torch.Tensor,
|
||||
cache_k_rope: torch.Tensor,
|
||||
):
|
||||
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_mla_kv_buffer (MLA)")
|
||||
layer_id = layer.layer_id
|
||||
self._write_mla_kv_buffer(
|
||||
self.kv_buffer[layer_id - self.start_layer],
|
||||
loc,
|
||||
cache_k_nope,
|
||||
cache_k_rope,
|
||||
)
|
||||
|
||||
def get_mla_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
@@ -3150,6 +3163,7 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
self.index_head_dim = index_head_dim
|
||||
if index_buf_size is None:
|
||||
index_buf_size = size
|
||||
self.index_buf_size = index_buf_size
|
||||
# num head == 1 and head dim == 128 for index_k in DSA
|
||||
assert index_head_dim == 128
|
||||
|
||||
@@ -3164,6 +3178,18 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
), f"HIP legacy DSA path requires page_size == 1, got {self.page_size}"
|
||||
else:
|
||||
assert self.page_size == 64
|
||||
self._create_index_buffers()
|
||||
self._finalize_allocation_log(size)
|
||||
|
||||
def _index_buffer_shape(self, num_pages: int) -> tuple[int, int]:
|
||||
return (
|
||||
num_pages,
|
||||
self.page_size
|
||||
* (self.index_head_dim + self.index_head_dim // self.quant_block_size * 4),
|
||||
)
|
||||
|
||||
def _create_index_buffers(self):
|
||||
num_pages = (self.index_buf_size + self.page_size + 1) // self.page_size
|
||||
with (
|
||||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||||
if self.custom_mem_pool
|
||||
@@ -3177,22 +3203,15 @@ class DSATokenToKVPool(MLATokenToKVPool):
|
||||
# data: for page i,
|
||||
# * buf[i, :page_size * head_dim] for fp8 data
|
||||
# * buf[i, page_size * head_dim:].view(float32) for scale
|
||||
(
|
||||
(index_buf_size + page_size + 1) // self.page_size,
|
||||
self.page_size
|
||||
* (
|
||||
index_head_dim + index_head_dim // self.quant_block_size * 4
|
||||
),
|
||||
),
|
||||
self._index_buffer_shape(num_pages),
|
||||
dtype=self.index_k_with_scale_buffer_dtype,
|
||||
device=device,
|
||||
device=self.device,
|
||||
)
|
||||
for _ in range(layer_num)
|
||||
for _ in range(self.layer_num)
|
||||
]
|
||||
self._finalize_allocation_log(size)
|
||||
|
||||
def _clear_buffers(self):
|
||||
del self.kv_buffer
|
||||
super()._clear_buffers()
|
||||
del self.index_k_with_scale_buffer
|
||||
|
||||
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
|
||||
|
||||
@@ -127,7 +127,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
def get_size_per_token(self):
|
||||
self.kv_lora_rank = self.device_pool.kv_lora_rank
|
||||
self.qk_rope_head_dim = self.device_pool.qk_rope_head_dim
|
||||
self.layer_num = self.device_pool.layer_num
|
||||
self.layer_num = self._effective_host_layer_num()
|
||||
self.kv_cache_dim = self.override_kv_cache_dim or (
|
||||
self.kv_lora_rank + self.qk_rope_head_dim
|
||||
)
|
||||
@@ -244,19 +244,23 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
def load_to_device_per_layer(
|
||||
self, device_pool, host_indices, device_indices, layer_id, io_backend
|
||||
):
|
||||
if not self._is_device_layer_owned(device_pool, layer_id):
|
||||
return
|
||||
host_layer = self._host_layer_index(layer_id)
|
||||
|
||||
if io_backend == "kernel":
|
||||
if self.layout == "layer_first":
|
||||
if self.can_use_jit:
|
||||
jit_transfer_hicache_one_layer_mla(
|
||||
cache_dst=device_pool.kv_buffer[layer_id],
|
||||
cache_src=self.kv_buffer[layer_id],
|
||||
cache_src=self.kv_buffer[host_layer],
|
||||
indices_dst=device_indices,
|
||||
indices_src=host_indices,
|
||||
element_dim=self.kv_cache_dim,
|
||||
)
|
||||
else:
|
||||
transfer_kv_per_layer_mla(
|
||||
src=self.kv_buffer[layer_id],
|
||||
src=self.kv_buffer[host_layer],
|
||||
dst=device_pool.kv_buffer[layer_id],
|
||||
src_indices=host_indices,
|
||||
dst_indices=device_indices,
|
||||
@@ -266,7 +270,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
if self.can_use_jit:
|
||||
jit_transfer_hicache_one_layer_mla(
|
||||
cache_dst=device_pool.kv_buffer[layer_id],
|
||||
cache_src=self.data_refs[layer_id],
|
||||
cache_src=self.data_refs[host_layer],
|
||||
indices_dst=device_indices,
|
||||
indices_src=host_indices,
|
||||
element_dim=self.kv_cache_dim,
|
||||
@@ -277,7 +281,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
dst=device_pool.kv_buffer[layer_id],
|
||||
src_indices=host_indices,
|
||||
dst_indices=device_indices,
|
||||
layer_id=layer_id,
|
||||
layer_id=host_layer,
|
||||
item_size=self.token_stride_size,
|
||||
src_layout_dim=self.layout_dim,
|
||||
)
|
||||
@@ -286,7 +290,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
elif io_backend == "direct":
|
||||
if self.layout == "layer_first":
|
||||
transfer_kv_direct(
|
||||
src_layers=[self.kv_buffer[layer_id]],
|
||||
src_layers=[self.kv_buffer[host_layer]],
|
||||
dst_layers=[device_pool.kv_buffer[layer_id]],
|
||||
src_indices=host_indices,
|
||||
dst_indices=device_indices,
|
||||
@@ -298,7 +302,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
dst_ptrs=[device_pool.kv_buffer[layer_id]],
|
||||
src_indices=host_indices,
|
||||
dst_indices=device_indices,
|
||||
layer_id=layer_id,
|
||||
layer_id=host_layer,
|
||||
page_size=self.page_size,
|
||||
)
|
||||
else:
|
||||
@@ -324,9 +328,75 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
else:
|
||||
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
||||
|
||||
def _backup_from_device_per_layer(
|
||||
self, device_pool, host_indices, device_indices, layer_id, io_backend
|
||||
):
|
||||
host_layer = self._host_layer_index(layer_id)
|
||||
if io_backend == "kernel":
|
||||
if self.layout == "layer_first":
|
||||
if self.can_use_jit:
|
||||
jit_transfer_hicache_one_layer_mla(
|
||||
cache_dst=self.kv_buffer[host_layer],
|
||||
cache_src=device_pool.kv_buffer[layer_id],
|
||||
indices_dst=host_indices,
|
||||
indices_src=device_indices,
|
||||
element_dim=self.kv_cache_dim,
|
||||
)
|
||||
else:
|
||||
transfer_kv_per_layer_mla(
|
||||
src=device_pool.kv_buffer[layer_id],
|
||||
dst=self.kv_buffer[host_layer],
|
||||
src_indices=device_indices,
|
||||
dst_indices=host_indices,
|
||||
item_size=self.token_stride_size,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
if self.can_use_jit:
|
||||
jit_transfer_hicache_one_layer_mla(
|
||||
cache_dst=self.data_refs[host_layer],
|
||||
cache_src=device_pool.kv_buffer[layer_id],
|
||||
indices_dst=host_indices,
|
||||
indices_src=device_indices,
|
||||
element_dim=self.kv_cache_dim,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Layer-sharded MLA HiCache backup with page_first layout "
|
||||
"requires the JIT one-layer kernel."
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Layer-sharded HiCache backup does not support layout: {self.layout}"
|
||||
)
|
||||
elif io_backend == "direct":
|
||||
if self.layout == "layer_first":
|
||||
transfer_kv_direct(
|
||||
src_layers=[device_pool.kv_buffer[layer_id]],
|
||||
dst_layers=[self.kv_buffer[host_layer]],
|
||||
src_indices=device_indices,
|
||||
dst_indices=host_indices,
|
||||
page_size=self.page_size,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Layer-sharded direct HiCache backup only supports "
|
||||
f"layer_first layout, got {self.layout}"
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Layer-sharded HiCache backup does not support IO backend: {io_backend}"
|
||||
)
|
||||
|
||||
def backup_from_device_all_layer(
|
||||
self, device_pool, host_indices, device_indices, io_backend
|
||||
):
|
||||
if self._is_device_layer_sharded(device_pool):
|
||||
for layer_id in self._owned_device_layer_ids(device_pool):
|
||||
self._backup_from_device_per_layer(
|
||||
device_pool, host_indices, device_indices, layer_id, io_backend
|
||||
)
|
||||
return
|
||||
|
||||
if io_backend == "kernel":
|
||||
if self.layout == "layer_first":
|
||||
if self.can_use_jit:
|
||||
@@ -2109,7 +2179,7 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
self.dtype = device_pool.store_dtype
|
||||
self.start_layer = device_pool.start_layer
|
||||
self.end_layer = device_pool.end_layer
|
||||
self.layer_num = device_pool.layer_num
|
||||
self.layer_num = self._effective_host_layer_num()
|
||||
|
||||
self.index_head_dim = device_pool.index_head_dim
|
||||
self.indexer_quant_block_size = device_pool.quant_block_size
|
||||
@@ -2242,6 +2312,10 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
def load_to_device_per_layer(
|
||||
self, device_pool, host_indices, device_indices, layer_id, io_backend
|
||||
):
|
||||
if not self._is_device_layer_owned(device_pool, layer_id):
|
||||
return
|
||||
host_layer = self._host_layer_index(layer_id)
|
||||
|
||||
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||
host_indices, device_indices
|
||||
)
|
||||
@@ -2249,7 +2323,7 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
if use_kernel:
|
||||
if self.layout == "layer_first":
|
||||
transfer_kv_per_layer_mla(
|
||||
src=self.index_k_with_scale_buffer[layer_id],
|
||||
src=self.index_k_with_scale_buffer[host_layer],
|
||||
dst=device_pool.index_k_with_scale_buffer[layer_id],
|
||||
src_indices=host_page_indices,
|
||||
dst_indices=device_page_indices,
|
||||
@@ -2261,7 +2335,7 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
dst=device_pool.index_k_with_scale_buffer[layer_id],
|
||||
src_indices=host_page_indices,
|
||||
dst_indices=device_page_indices,
|
||||
layer_id=layer_id,
|
||||
layer_id=host_layer,
|
||||
item_size=self.indexer_page_stride_size,
|
||||
src_layout_dim=self.indexer_layout_dim,
|
||||
)
|
||||
@@ -2270,7 +2344,7 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
elif io_backend == "direct":
|
||||
if self.layout == "layer_first":
|
||||
transfer_kv_direct(
|
||||
src_layers=[self.index_k_with_scale_buffer[layer_id]],
|
||||
src_layers=[self.index_k_with_scale_buffer[host_layer]],
|
||||
dst_layers=[device_pool.index_k_with_scale_buffer[layer_id]],
|
||||
src_indices=host_page_indices,
|
||||
dst_indices=device_page_indices,
|
||||
@@ -2282,7 +2356,7 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
dst_ptrs=[device_pool.index_k_with_scale_buffer[layer_id]],
|
||||
src_indices=host_page_indices,
|
||||
dst_indices=device_page_indices,
|
||||
layer_id=layer_id,
|
||||
layer_id=host_layer,
|
||||
page_size=1,
|
||||
)
|
||||
else:
|
||||
@@ -2290,9 +2364,57 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
else:
|
||||
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
||||
|
||||
def _backup_from_device_per_layer(
|
||||
self, device_pool, host_indices, device_indices, layer_id, io_backend
|
||||
):
|
||||
host_layer = self._host_layer_index(layer_id)
|
||||
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||
host_indices, device_indices
|
||||
)
|
||||
use_kernel = io_backend == "kernel" and self.indexer_page_stride_size % 8 == 0
|
||||
if use_kernel:
|
||||
if self.layout == "layer_first":
|
||||
transfer_kv_per_layer_mla(
|
||||
src=device_pool.index_k_with_scale_buffer[layer_id],
|
||||
dst=self.index_k_with_scale_buffer[host_layer],
|
||||
src_indices=device_page_indices,
|
||||
dst_indices=host_page_indices,
|
||||
item_size=self.indexer_page_stride_size,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
raise ValueError(
|
||||
"Layer-sharded DSA indexer HiCache backup with page_first "
|
||||
"layout is not supported without a per-layer LF->PF kernel."
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||
elif io_backend == "direct":
|
||||
if self.layout == "layer_first":
|
||||
transfer_kv_direct(
|
||||
src_layers=[device_pool.index_k_with_scale_buffer[layer_id]],
|
||||
dst_layers=[self.index_k_with_scale_buffer[host_layer]],
|
||||
src_indices=device_page_indices,
|
||||
dst_indices=host_page_indices,
|
||||
page_size=1,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Layer-sharded direct DSA indexer backup only supports "
|
||||
f"layer_first layout, got {self.layout}"
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
||||
|
||||
def backup_from_device_all_layer(
|
||||
self, device_pool, host_indices, device_indices, io_backend
|
||||
):
|
||||
if self._is_device_layer_sharded(device_pool):
|
||||
for layer_id in self._owned_device_layer_ids(device_pool):
|
||||
self._backup_from_device_per_layer(
|
||||
device_pool, host_indices, device_indices, layer_id, io_backend
|
||||
)
|
||||
return
|
||||
|
||||
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||
host_indices, device_indices
|
||||
)
|
||||
|
||||
@@ -168,6 +168,41 @@ class HostKVCache(abc.ABC):
|
||||
def get_size_per_token(self):
|
||||
raise NotImplementedError()
|
||||
|
||||
def _is_device_layer_sharded(self, device_pool=None) -> bool:
|
||||
device_pool = device_pool or self.device_pool
|
||||
return bool(device_pool.layer_shard_enabled)
|
||||
|
||||
def _device_owned_layer_range(self, device_pool=None) -> tuple[int, int]:
|
||||
"""Contiguous ``[start, end)`` local device layers this rank stores.
|
||||
|
||||
``(0, layer_num)`` when the device pool is not layer-sharded.
|
||||
"""
|
||||
device_pool = device_pool or self.device_pool
|
||||
if not self._is_device_layer_sharded(device_pool):
|
||||
return 0, device_pool.layer_num
|
||||
return device_pool._owned_local_layer_range()
|
||||
|
||||
def _effective_host_layer_num(self, device_pool=None) -> int:
|
||||
"""Number of layers the host pool allocates for this rank."""
|
||||
device_pool = device_pool or self.device_pool
|
||||
if not self._is_device_layer_sharded(device_pool):
|
||||
return device_pool.layer_num
|
||||
shard_size = device_pool.layer_shard_size
|
||||
return (device_pool.layer_num + shard_size - 1) // shard_size
|
||||
|
||||
def _is_device_layer_owned(self, device_pool, layer_id: int) -> bool:
|
||||
start, end = self._device_owned_layer_range(device_pool)
|
||||
return start <= layer_id < end
|
||||
|
||||
def _host_layer_index(self, layer_id: int, device_pool=None) -> int:
|
||||
"""Map a full local device layer id to its compacted host-buffer slot."""
|
||||
start, _ = self._device_owned_layer_range(device_pool)
|
||||
return layer_id - start
|
||||
|
||||
def _owned_device_layer_ids(self, device_pool) -> list[int]:
|
||||
start, end = self._device_owned_layer_range(device_pool)
|
||||
return list(range(start, end))
|
||||
|
||||
@abc.abstractmethod
|
||||
def init_kv_buffer(self):
|
||||
raise NotImplementedError()
|
||||
|
||||
@@ -947,16 +947,31 @@ class ModelRunnerKVCacheMixin:
|
||||
end_layer=self.end_layer,
|
||||
)
|
||||
elif self.use_mla_backend and is_dsa_model:
|
||||
PoolCls = (
|
||||
HiSparseDSATokenToKVPool if self.enable_hisparse else DSATokenToKVPool
|
||||
)
|
||||
from sglang.srt.layers.cp.utils import get_glm_dsa_cp_layer_shard_info
|
||||
|
||||
(
|
||||
dsa_cp_layer_shard_rank,
|
||||
dsa_cp_layer_shard_size,
|
||||
) = get_glm_dsa_cp_layer_shard_info(self)
|
||||
pool_kwargs = {}
|
||||
if self.enable_hisparse:
|
||||
PoolCls = HiSparseDSATokenToKVPool
|
||||
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
||||
|
||||
pool_kwargs["host_to_device_ratio"] = parse_hisparse_config(
|
||||
self.server_args
|
||||
).host_to_device_ratio
|
||||
elif dsa_cp_layer_shard_rank is not None:
|
||||
# DSA cache layer split: shard KV/indexer layers across CP ranks.
|
||||
from sglang.srt.mem_cache.dsa_cache_layer_split import (
|
||||
LayerSplitDSATokenToKVPool,
|
||||
)
|
||||
|
||||
PoolCls = LayerSplitDSATokenToKVPool
|
||||
pool_kwargs["layer_shard_rank"] = dsa_cp_layer_shard_rank
|
||||
pool_kwargs["layer_shard_size"] = dsa_cp_layer_shard_size
|
||||
else:
|
||||
PoolCls = DSATokenToKVPool
|
||||
self.token_to_kv_pool = PoolCls(
|
||||
self.max_total_num_tokens,
|
||||
page_size=self.page_size,
|
||||
|
||||
@@ -177,6 +177,13 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
||||
# args to config cell size
|
||||
model_config = mr.model_config
|
||||
kv_cache_dtype = mr.kv_cache_dtype
|
||||
from sglang.srt.layers.cp.utils import (
|
||||
get_glm_dsa_layer_split_effective_num_layers,
|
||||
)
|
||||
|
||||
effective_num_layers = get_glm_dsa_layer_split_effective_num_layers(
|
||||
mr, num_layers
|
||||
)
|
||||
|
||||
kv_size = torch._utils._element_size(kv_cache_dtype)
|
||||
tp_size = get_parallel().attn_tp_size
|
||||
@@ -184,7 +191,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
||||
if mr.use_mla_backend:
|
||||
cell_size = (
|
||||
(model_config.kv_lora_rank + model_config.qk_rope_head_dim)
|
||||
* num_layers
|
||||
* effective_num_layers
|
||||
* kv_size
|
||||
)
|
||||
if is_float4_e2m1fn_x2(kv_cache_dtype):
|
||||
@@ -195,7 +202,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
||||
(model_config.kv_lora_rank + model_config.qk_rope_head_dim)
|
||||
// scale_block_size
|
||||
)
|
||||
* num_layers
|
||||
* effective_num_layers
|
||||
* kv_size
|
||||
)
|
||||
|
||||
@@ -209,7 +216,9 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
||||
element_size = torch._utils._element_size(
|
||||
DSATokenToKVPool.index_k_with_scale_buffer_dtype
|
||||
)
|
||||
cell_size += indexer_size_per_token * num_layers * element_size
|
||||
cell_size += (
|
||||
indexer_size_per_token * effective_num_layers * element_size
|
||||
)
|
||||
elif is_minimax_sparse(model_config.hf_config):
|
||||
# Mirrors MiniMaxSparseKVPool: main pool (K+V all layers) + indexer pool
|
||||
# (sparse-only, single-head; kv layers store K+V, k-only layers store K).
|
||||
@@ -252,7 +261,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
||||
cell_size = (
|
||||
model_config.get_num_kv_heads(tp_size)
|
||||
* (model_config.head_dim + model_config.v_head_dim)
|
||||
* num_layers
|
||||
* effective_num_layers
|
||||
* kv_size
|
||||
)
|
||||
|
||||
@@ -262,7 +271,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
||||
n = model_config.get_num_kv_heads(tp_size)
|
||||
k = model_config.head_dim
|
||||
cell_size = (cell_size // 2) + (
|
||||
(n * k * num_layers * 2 * kv_size) // scale_block_size
|
||||
(n * k * effective_num_layers * 2 * kv_size) // scale_block_size
|
||||
)
|
||||
|
||||
return cell_size
|
||||
|
||||
@@ -70,7 +70,10 @@ from sglang.srt.layers.communicator import (
|
||||
enable_moe_dense_fully_dp,
|
||||
get_attn_tp_context,
|
||||
)
|
||||
from sglang.srt.layers.communicator_dsa_cp import DSACPLayerCommunicator
|
||||
from sglang.srt.layers.communicator_dsa_cp import (
|
||||
DSACPLayerCommunicator,
|
||||
maybe_prefetch_next_full_attention_kv,
|
||||
)
|
||||
from sglang.srt.layers.dcp import dcp_enabled, get_attention_dcp_world_size
|
||||
from sglang.srt.layers.dcp.planner import (
|
||||
prepare_decode_context_parallel_metadata,
|
||||
@@ -2207,6 +2210,7 @@ class DeepseekV2DecoderLayer(nn.Module):
|
||||
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||
prev_topk_indices: Optional[torch.Tensor] = None,
|
||||
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
|
||||
next_full_attention_layer_id: Optional[int] = None,
|
||||
) -> torch.Tensor:
|
||||
hidden_states_orig = hidden_states
|
||||
hidden_states, residual = (
|
||||
@@ -2234,6 +2238,10 @@ class DeepseekV2DecoderLayer(nn.Module):
|
||||
topk_indices = None
|
||||
get_attn_tp_context().clear_attn_inputs()
|
||||
|
||||
maybe_prefetch_next_full_attention_kv(
|
||||
forward_batch, next_full_attention_layer_id
|
||||
)
|
||||
|
||||
hidden_states, residual = self.layer_communicator.prepare_mlp(
|
||||
hidden_states, residual, forward_batch
|
||||
)
|
||||
@@ -2429,6 +2437,11 @@ class DeepseekV2Model(nn.Module):
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
local_layer_ids = list(range(self.start_layer, self.end_layer))
|
||||
self.next_full_attention_layer_id = dict(
|
||||
zip(local_layer_ids, local_layer_ids[1:])
|
||||
)
|
||||
if self.pp_group.is_last_rank:
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
else:
|
||||
@@ -2597,6 +2610,9 @@ class DeepseekV2Model(nn.Module):
|
||||
captured_last_layer_outputs=(
|
||||
aux_hidden_states if i in self.layers_to_capture else None
|
||||
),
|
||||
next_full_attention_layer_id=self.next_full_attention_layer_id.get(
|
||||
i
|
||||
),
|
||||
)
|
||||
|
||||
if normal_end_layer != self.end_layer:
|
||||
|
||||
@@ -916,6 +916,11 @@ class ServerArgs:
|
||||
choices=("zigzag", "interleave"),
|
||||
),
|
||||
] = None
|
||||
# Split DSA GPU KV/indexer cache layers across CP ranks.
|
||||
enable_dsa_cache_layer_split: A[
|
||||
bool,
|
||||
"Split DSA (DeepSeek Sparse Attention) GPU KV/indexer cache layers across context-parallel ranks to reduce per-rank KV memory. Currently only supported with the mooncake transfer backend (mooncake / mooncake_tcp); mori/nixl support will be added later by the community.",
|
||||
] = False
|
||||
enable_dsa_prefill_context_parallel: A[bool, Arg(no_cli=True)] = False
|
||||
dsa_prefill_cp_mode: A[str, Arg(no_cli=True)] = "round-robin-split"
|
||||
enable_prefill_context_parallel: A[bool, Arg(no_cli=True)] = False
|
||||
@@ -4058,6 +4063,12 @@ class ServerArgs:
|
||||
hf_config = self.get_model_config().hf_config
|
||||
model_arch = hf_config.architectures[0]
|
||||
|
||||
if self.enable_dsa_cache_layer_split and not is_deepseek_dsa(hf_config):
|
||||
raise ValueError(
|
||||
"--enable-dsa-cache-layer-split is only supported for DSA "
|
||||
"(DeepSeek Sparse Attention) models."
|
||||
)
|
||||
|
||||
_hybrid_spec = get_linear_attn_spec_by_arch(model_arch)
|
||||
if _hybrid_spec is not None and _hybrid_spec.uses_mamba_radix_cache:
|
||||
self._handle_mamba_radix_cache(model_arch=model_arch)
|
||||
@@ -4158,6 +4169,52 @@ class ServerArgs:
|
||||
assert (
|
||||
self.disaggregation_mode != "decode"
|
||||
), "CP is only supported for prefill when PD disaggregation, please remove --enable-prefill-cp."
|
||||
if (
|
||||
self.enable_dsa_cache_layer_split
|
||||
and self.disaggregation_mode != "prefill"
|
||||
):
|
||||
if self.disaggregation_mode == "decode":
|
||||
raise ValueError(
|
||||
"--enable-dsa-cache-layer-split is not supported on "
|
||||
"decode workers. This flag is a prefill-CP "
|
||||
"optimization; decode receives full cache shards "
|
||||
"through PD transfer."
|
||||
)
|
||||
raise ValueError(
|
||||
"--enable-dsa-cache-layer-split is only supported on PD "
|
||||
"prefill workers. Non-PD workers also run decode and "
|
||||
"require ordinary local decode cache semantics."
|
||||
)
|
||||
if self.enable_dsa_cache_layer_split and (
|
||||
not self.enable_prefill_cp or self.cp_strategy != "interleave"
|
||||
):
|
||||
raise ValueError(
|
||||
"--enable-dsa-cache-layer-split requires "
|
||||
"--enable-prefill-cp and --cp-strategy interleave "
|
||||
"(or legacy --enable-nsa-prefill-context-parallel with "
|
||||
"--nsa-prefill-cp-mode round-robin-split)."
|
||||
)
|
||||
# Layer split relies on the mooncake all-CP-rank KV/indexer
|
||||
# transfer path. mori/nixl support is a temporary limitation
|
||||
# and will be added later by the community.
|
||||
if (
|
||||
self.enable_dsa_cache_layer_split
|
||||
and self.disaggregation_transfer_backend != "mooncake"
|
||||
):
|
||||
raise ValueError(
|
||||
"--enable-dsa-cache-layer-split currently only supports "
|
||||
"the mooncake transfer backend (mooncake / mooncake_tcp). "
|
||||
f"Got --disaggregation-transfer-backend "
|
||||
f"{self.disaggregation_transfer_backend!r}. mori/nixl "
|
||||
"support will be added later by the community."
|
||||
)
|
||||
if self.enable_dsa_cache_layer_split and self.pp_size > 1:
|
||||
raise ValueError(
|
||||
"--enable-dsa-cache-layer-split is not supported with "
|
||||
"pipeline parallelism (pp_size > 1) yet. It requires "
|
||||
"prefill context parallelism, and CP + PP has not been "
|
||||
"validated for this feature."
|
||||
)
|
||||
|
||||
else:
|
||||
# DeepSeek V3/R1/V3.1
|
||||
|
||||
Reference in New Issue
Block a user