[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:
Shijin Zhang
2026-07-09 03:03:56 -07:00
committed by GitHub
parent 336b64ecce
commit 8e54517f02
21 changed files with 1507 additions and 72 deletions
@@ -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,
+24 -4
View File
@@ -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
+92 -1
View File
@@ -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()
+40 -21
View File
@@ -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):
+134 -12
View File
@@ -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
+17 -1
View File
@@ -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:
+57
View File
@@ -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