From 8e54517f027606a00054fc25c12898bc82acc65b Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Thu, 9 Jul 2026 18:03:56 +0800 Subject: [PATCH] [Feat][GLM5.2] Add DSA Cache Layer Split under Prefill CP (#29421) Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- python/sglang/srt/disaggregation/base/conn.py | 1 + .../sglang/srt/disaggregation/common/conn.py | 38 +- .../srt/disaggregation/mooncake/conn.py | 40 +- python/sglang/srt/disaggregation/prefill.py | 28 +- python/sglang/srt/disaggregation/utils.py | 4 + .../srt/layers/attention/dsa/dsa_indexer.py | 39 +- .../sglang/srt/layers/communicator_dsa_cp.py | 20 + python/sglang/srt/layers/cp/utils.py | 93 ++- .../srt/mem_cache/dsa_cache_layer_split.py | 581 ++++++++++++++++++ python/sglang/srt/mem_cache/memory_pool.py | 61 +- .../sglang/srt/mem_cache/memory_pool_host.py | 146 ++++- python/sglang/srt/mem_cache/pool_host/base.py | 35 ++ .../model_runner_kv_cache_mixin.py | 21 +- .../srt/model_executor/pool_configurator.py | 19 +- python/sglang/srt/models/deepseek_v2.py | 18 +- python/sglang/srt/server_args.py | 57 ++ .../test_dsa_glm52_cache_layer_split.py | 83 +++ .../mem_cache/test_dsa_layer_shard_utils.py | 123 ++++ .../test_dsa_layer_split_broadcast.py | 149 +++++ ...test_hicache_staged_write_back_dispatch.py | 22 +- .../model_executor/test_pool_configurator.py | 1 + 21 files changed, 1507 insertions(+), 72 deletions(-) create mode 100644 python/sglang/srt/mem_cache/dsa_cache_layer_split.py create mode 100644 test/registered/models_e2e/test_dsa_glm52_cache_layer_split.py create mode 100644 test/registered/unit/mem_cache/test_dsa_layer_shard_utils.py create mode 100644 test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 1cbe4e7f7..7412ec9f0 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -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 diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index a609acf8f..b51ee6290 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -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) diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 3e8153e96..970c6b7e9 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -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, diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 1e390518c..94dcf0930 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -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, ) diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 1ab8ec16b..fe0f011d6 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -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 ) diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 97ed54de6..427058691 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -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, diff --git a/python/sglang/srt/layers/communicator_dsa_cp.py b/python/sglang/srt/layers/communicator_dsa_cp.py index 397722eeb..c55d64b3b 100644 --- a/python/sglang/srt/layers/communicator_dsa_cp.py +++ b/python/sglang/srt/layers/communicator_dsa_cp.py @@ -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 diff --git a/python/sglang/srt/layers/cp/utils.py b/python/sglang/srt/layers/cp/utils.py index 21f980987..9ac1dc3ad 100644 --- a/python/sglang/srt/layers/cp/utils.py +++ b/python/sglang/srt/layers/cp/utils.py @@ -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", ] diff --git a/python/sglang/srt/mem_cache/dsa_cache_layer_split.py b/python/sglang/srt/mem_cache/dsa_cache_layer_split.py new file mode 100644 index 000000000..9f44ce0be --- /dev/null +++ b/python/sglang/srt/mem_cache/dsa_cache_layer_split.py @@ -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() diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 613492668..260cdef32 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -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): diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 43b6af776..e1bd76477 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -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 ) diff --git a/python/sglang/srt/mem_cache/pool_host/base.py b/python/sglang/srt/mem_cache/pool_host/base.py index 8c1b98e23..5dba1aa78 100644 --- a/python/sglang/srt/mem_cache/pool_host/base.py +++ b/python/sglang/srt/mem_cache/pool_host/base.py @@ -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() diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index 141804561..9ae58cede 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -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, diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index bf02c6693..4bc650ba5 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -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 diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 670d2adcb..ce6ecdc6d 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -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: diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 7c8763cc0..c3dbc3cc6 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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 diff --git a/test/registered/models_e2e/test_dsa_glm52_cache_layer_split.py b/test/registered/models_e2e/test_dsa_glm52_cache_layer_split.py new file mode 100644 index 000000000..217ed8ae8 --- /dev/null +++ b/test/registered/models_e2e/test_dsa_glm52_cache_layer_split.py @@ -0,0 +1,83 @@ +"""End-to-end GSM8K accuracy test for DSA cache layer split (GLM-5.2). + +Layer split shards the DSA GPU KV/indexer cache layers across prefill CP ranks +(``--enable-dsa-cache-layer-split``); non-owner ranks read a layer via an +owner-broadcast into a small remote scratch buffer. It only applies to PD +prefill workers running DSA prefill-CP (a unified server would decode on the +same worker, where non-owner ranks lack the full cache), so this test drives a +PD-disaggregated GLM-5.2 deployment: a layer-split prefill worker running +interleave prefill-CP + layer split, and an ordinary decode worker that receives +full cache shards via PD transfer. + +Sized for the 4-GPU B200 runner (prefill TP=2 + decode TP=2) rather than an +8-GPU deployment, since the 8-gpu-b200 runner is nightly-only. +""" + +import unittest + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.eval_accuracy_kit import GSM8KMixin +from sglang.test.server_fixtures.disaggregation_fixture import ( + PDDisaggregationServerBase, +) + +register_cuda_ci( + est_time=1200, + stage="extra-b", + runner_config="4-gpu-b200", + disabled="Temporarily disabled", +) + + +class TestGLM52DSACacheLayerSplit(PDDisaggregationServerBase, GSM8KMixin): + model = "nvidia/GLM-5.2-NVFP4" + + # Full GSM8K test set (1319 questions) with a tight accuracy floor. + gsm8k_accuracy_thres = 0.935 + gsm8k_num_questions = 1319 + gsm8k_num_threads = 200 + gsm8k_num_shots = 0 + + # Prefill worker: interleave prefill-CP + DSA cache layer split on 2 GPUs + # (TP=2 -> attn_cp_size=2, so KV/indexer layers shard 2-way across CP ranks). + extra_prefill_args = [ + "--tp", + "2", + "--dsa-prefill-backend", + "trtllm", + "--kv-cache-dtype", + "fp8_e4m3", + "--enable-dsa-cache-layer-split", + "--enable-prefill-cp", + "--cp-strategy", + "interleave", + "--mem-fraction-static", + "0.85", + "--chunked-prefill-size", + "4096", + "--max-prefill-tokens", + "4096", + ] + # Decode worker: ordinary local decode cache on the other 2 GPUs, receives + # full shards via PD transfer. + extra_decode_args = [ + "--tp", + "2", + "--dsa-decode-backend", + "trtllm", + "--kv-cache-dtype", + "fp8_e4m3", + "--mem-fraction-static", + "0.85", + "--base-gpu-id", + "2", + ] + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.launch_all() + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_dsa_layer_shard_utils.py b/test/registered/unit/mem_cache/test_dsa_layer_shard_utils.py new file mode 100644 index 000000000..b6ddb86eb --- /dev/null +++ b/test/registered/unit/mem_cache/test_dsa_layer_shard_utils.py @@ -0,0 +1,123 @@ +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.layers.cp.utils import get_layer_owner, get_layer_shard_range +from sglang.srt.mem_cache.dsa_cache_layer_split import LayerSplitDSATokenToKVPool +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class TestDSALayerShardUtils(CustomTestCase): + def test_balanced_layer_ranges_cover_all_layers_once(self): + ranges = [get_layer_shard_range(rank, 4, 10) for rank in range(4)] + self.assertEqual(ranges, [(0, 3), (3, 6), (6, 8), (8, 10)]) + + covered = [layer_id for start, end in ranges for layer_id in range(start, end)] + self.assertEqual(covered, list(range(10))) + + def test_owner_matches_uneven_layer_ranges(self): + self.assertEqual( + [get_layer_owner(i, 4, 10) for i in range(10)], + [0, 0, 0, 1, 1, 1, 2, 2, 3, 3], + ) + + def test_empty_tail_shards_have_empty_ranges(self): + ranges = [get_layer_shard_range(rank, 4, 2) for rank in range(4)] + self.assertEqual(ranges, [(0, 1), (1, 2), (2, 2), (2, 2)]) + + def test_prefetch_uses_sync_fallback_without_dedicated_communicator(self): + broadcasts = [] + counter = SimpleNamespace(wait_until=lambda _: self.fail("unexpected wait")) + pool = SimpleNamespace( + remote_kv_layer_id=None, + pending_remote_kv_broadcast=False, + pending_remote_kv_layer_id=None, + layer_broadcast_comm=None, + remote_kv_buffer=object(), + kv_buffer=[object()], + start_layer=0, + _local_layer_idx=lambda layer_id: layer_id, + _is_layer_owned=lambda _: True, + ) + + def broadcast(tensor, layer_id, *, src_tensor, use_layer_broadcast_comm): + broadcasts.append((tensor, layer_id, src_tensor, use_layer_broadcast_comm)) + + pool._broadcast_tensor_from_owner = broadcast + # Bind the real method against a lightweight stand-in so the sync + # (no dedicated NCCL comm) fallback path can be exercised on CPU. + LayerSplitDSATokenToKVPool.prefetch_kv_buffer( + pool, + layer_id=0, + layer_transfer_counter=counter, + layer_transfer_idx=3, + ) + + self.assertEqual(len(broadcasts), 1) + self.assertEqual(pool.remote_kv_layer_id, 0) + + def test_finalize_pending_broadcast_promotes_layer_id(self): + # After an async prefetch, finalizing must promote pending -> remote so a + # subsequent read of the same layer reuses the broadcast result. + pool = SimpleNamespace( + pending_remote_kv_broadcast=True, + pending_remote_kv_layer_id=7, + remote_kv_layer_id=None, + device_module=SimpleNamespace( + current_stream=lambda: SimpleNamespace(wait_stream=lambda _stream: None) + ), + kv_broadcast_stream=object(), + ) + LayerSplitDSATokenToKVPool._finalize_pending_kv_broadcast( + pool, set_remote_layer_id=True + ) + self.assertFalse(pool.pending_remote_kv_broadcast) + self.assertEqual(pool.remote_kv_layer_id, 7) + self.assertIsNone(pool.pending_remote_kv_layer_id) + + def test_get_broadcastable_kv_buffer_returns_owner_contents(self): + # A non-owner read must return the *owner's* KV bytes, copied into the + # remote scratch buffer by the broadcast. This checks prefetch_kv_buffer + # + _get_broadcastable_kv_buffer surface the correct contents. + layer_num = 4 + shard_size = 2 + owner_kv = { + layer_id: torch.full((3, 1, 8), float(layer_id + 1)) + for layer_id in range(layer_num) + } + remote = torch.zeros((3, 1, 8)) + + pool = SimpleNamespace( + layer_num=layer_num, + layer_shard_size=shard_size, + start_layer=0, + remote_kv_layer_id=None, + pending_remote_kv_broadcast=False, + pending_remote_kv_layer_id=None, + remote_kv_buffer=remote, + ) + pool._local_layer_idx = lambda layer_id: layer_id - pool.start_layer + pool._is_layer_owned = lambda layer_id: True + # kv_buffer holds this rank's owned layers; broadcast copies owner->remote. + pool.kv_buffer = [owner_kv[i] for i in range(layer_num)] + + def broadcast(tensor, layer_id, *, src_tensor, use_layer_broadcast_comm=False): + # Simulate the owner writing its layer into the remote scratch buffer. + tensor.copy_(owner_kv[layer_id]) + + pool._broadcast_tensor_from_owner = broadcast + + for layer_id in range(layer_num): + buf = LayerSplitDSATokenToKVPool._get_broadcastable_kv_buffer( + pool, layer_id + ) + self.assertTrue(torch.equal(buf, owner_kv[layer_id])) + self.assertEqual(pool.remote_kv_layer_id, layer_id) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py b/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py new file mode 100644 index 000000000..a56e01e4a --- /dev/null +++ b/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py @@ -0,0 +1,149 @@ +"""Multi-GPU integration test for LayerSplitDSATokenToKVPool owner-broadcast. + +Spawns ``world`` processes forming a single attention-CP group, builds a tiny +``LayerSplitDSATokenToKVPool`` on each rank, writes a rank-distinct value into +every owned layer, then verifies that reading ANY layer (owned or not) returns +the *owning* rank's bytes -- i.e. the owner-broadcast in +``_get_broadcastable_kv_buffer`` / ``prefetch_kv_buffer`` surfaces correct +contents. Also exercises the DSA indexer broadcast and the async prefetch path. + +Registered as a base-c 4-gpu-b200 unit test; uses up to 4 GPUs and skips when +fewer than 2 are visible. Run directly on 2+ GPUs: + CUDA_VISIBLE_DEVICES=0,1 python -m pytest \ + test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py +""" + +import os +import unittest + +import torch +import torch.multiprocessing as mp + +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=120, stage="base-c", runner_config="4-gpu-b200") + +LAYER_NUM = 4 +PAGE_SIZE = 64 +KV_LORA_RANK = 512 +QK_ROPE = 64 +INDEX_HEAD_DIM = 128 +SIZE = PAGE_SIZE * 3 # a few pages +PORT = 29711 + + +def _run(rank: int, world: int, port: int): + os.environ["MASTER_ADDR"] = "127.0.0.1" + os.environ["MASTER_PORT"] = str(port) + os.environ["RANK"] = str(rank) + os.environ["WORLD_SIZE"] = str(world) + os.environ.setdefault("no_proxy", "127.0.0.1,localhost") + torch.cuda.set_device(rank) + + from sglang.srt.distributed.parallel_state import ( + init_distributed_environment, + initialize_model_parallel, + ) + from sglang.srt.layers.dp_attention import ( + get_attention_cp_rank, + get_attention_cp_size, + ) + + init_distributed_environment( + world_size=world, + rank=rank, + local_rank=rank, + distributed_init_method=f"tcp://127.0.0.1:{port}", + backend="nccl", + ) + initialize_model_parallel( + tensor_model_parallel_size=world, + attention_context_model_parallel_size=world, + ) + + from sglang.srt.mem_cache.dsa_cache_layer_split import ( + LayerSplitDSATokenToKVPool, + ) + + cp_rank = get_attention_cp_rank() + cp_size = get_attention_cp_size() + assert cp_size == world + + pool = LayerSplitDSATokenToKVPool( + SIZE, + page_size=PAGE_SIZE, + kv_lora_rank=KV_LORA_RANK, + dtype=torch.bfloat16, + qk_rope_head_dim=QK_ROPE, + layer_num=LAYER_NUM, + device=f"cuda:{rank}", + index_head_dim=INDEX_HEAD_DIM, + enable_memory_saver=False, + kv_cache_dim=KV_LORA_RANK + QK_ROPE, + layer_shard_rank=cp_rank, + layer_shard_size=cp_size, + ) + + # Owner writes a layer-distinct constant into each owned kv_buffer layer. + for layer_id in range(LAYER_NUM): + if pool._is_layer_owned(layer_id): + pool.kv_buffer[layer_id].fill_(float(layer_id + 1)) + + torch.cuda.synchronize() + torch.distributed.barrier() + + # Every rank reads every layer; broadcast must surface the owner's value. + ok = True + for layer_id in range(LAYER_NUM): + buf = pool._get_broadcastable_kv_buffer(layer_id) + expected = float(layer_id + 1) + got = buf.float().mean().item() + if abs(got - expected) > 1e-3: + print(f"[rank {rank}] layer {layer_id}: expected {expected}, got {got}") + ok = False + assert ok, f"rank {rank} read stale/incorrect broadcast contents" + + # Indexer buffer owner-broadcast: owner writes a layer-distinct value, then + # every rank must read it back for every layer. + for layer_id in range(LAYER_NUM): + if pool._is_layer_owned(layer_id): + pool.index_k_with_scale_buffer[layer_id].fill_(layer_id + 10) + torch.cuda.synchronize() + torch.distributed.barrier() + for layer_id in range(LAYER_NUM): + # invalidate any cached remote copy so the read forces a fresh broadcast + pool.invalidate_index_buffer_for_layer(layer_id) + buf = pool._get_broadcastable_index_buffer(layer_id) + expected = layer_id + 10 + got = buf.float().mean().item() + if abs(got - expected) > 1e-3: + print(f"[rank {rank}] index layer {layer_id}: exp {expected}, got {got}") + ok = False + assert ok, f"rank {rank} read stale/incorrect index broadcast contents" + + # Async prefetch path: prefetch layer, then read must return owner value. + for layer_id in range(LAYER_NUM): + pool.remote_kv_layer_id = None # force a fresh broadcast + pool.prefetch_kv_buffer(layer_id) + buf = pool._get_broadcastable_kv_buffer(layer_id) + got = buf.float().mean().item() + if abs(got - float(layer_id + 1)) > 1e-3: + print(f"[rank {rank}] prefetch layer {layer_id}: got {got}") + ok = False + assert ok, f"rank {rank} prefetch path returned incorrect contents" + + print(f"[rank {rank}] OK: all {LAYER_NUM} layers read correct owner contents") + torch.distributed.barrier() + + +class TestLayerSplitDSABroadcast(CustomTestCase): + def test_owner_broadcast(self): + world = min(4, torch.cuda.device_count()) + if world < 2: + self.skipTest("LayerSplitDSATokenToKVPool broadcast test needs >= 2 GPUs") + mp.spawn(_run, args=(world, PORT), nprocs=world, join=True) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py index 7df5244c9..3a4bc7a55 100644 --- a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py +++ b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py @@ -49,6 +49,15 @@ def _ptr_key_from_tensor(ptrs: torch.Tensor) -> tuple[int, ...]: return tuple(int(ptr) for ptr in ptrs.cpu().tolist()) +def _device_pool_stub(*, layer_num: int, **fields) -> SimpleNamespace: + """Minimal device-pool stand-in with layer-split fields real pools expose.""" + return SimpleNamespace( + layer_num=layer_num, + layer_shard_enabled=False, + **fields, + ) + + def _cpu_staged_lf_pf_copy( src_registry, *, @@ -192,7 +201,8 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): ] expected_k = [layer[device_indices].clone() for layer in k_layers] expected_v = [layer[device_indices].clone() for layer in v_layers] - device_pool = SimpleNamespace( + device_pool = _device_pool_stub( + layer_num=layer_num, k_buffer=k_layers, v_buffer=v_layers, k_data_ptrs=torch.tensor( @@ -293,7 +303,8 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): for layer_id in range(layer_num) ] expected = [layer[device_indices].clone() for layer in device_layers] - device_pool = SimpleNamespace( + device_pool = _device_pool_stub( + layer_num=layer_num, kv_buffer=device_layers, data_ptrs=torch.tensor( [layer.data_ptr() for layer in device_layers], dtype=torch.uint64 @@ -301,6 +312,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): ) host = MLATokenToKVPoolHost.__new__(MLATokenToKVPoolHost) + host.device_pool = device_pool host.layout = "page_first" host.page_size = 1 host.layer_num = layer_num @@ -582,9 +594,13 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): for layer_id in range(layer_num) ] expected = [buffer[device_page_indices].clone() for buffer in device_layers] - device_pool = SimpleNamespace(index_k_with_scale_buffer=device_layers) + device_pool = _device_pool_stub( + layer_num=layer_num, + index_k_with_scale_buffer=device_layers, + ) host = DSAIndexerPoolHost.__new__(DSAIndexerPoolHost) + host.device_pool = device_pool host.layout = "page_first" host.page_size = page_size host.layer_num = layer_num diff --git a/test/registered/unit/model_executor/test_pool_configurator.py b/test/registered/unit/model_executor/test_pool_configurator.py index e774e086a..da8faa4f9 100644 --- a/test/registered/unit/model_executor/test_pool_configurator.py +++ b/test/registered/unit/model_executor/test_pool_configurator.py @@ -114,6 +114,7 @@ def _make_model_runner( sa.disaggregation_mode = disaggregation_mode sa.max_running_requests = max_running_requests sa.disaggregation_decode_extra_slots = disaggregation_decode_extra_slots + sa.enable_dsa_cache_layer_split = False mr.server_args = sa spec = MagicMock()