[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]] state_item_lens: List[List[int]]
# Per-tensor TP slice dim, used when prefill/decode attn_tp_size differ. # Per-tensor TP slice dim, used when prefill/decode attn_tp_size differ.
state_dim_per_tensor: List[List[int]] state_dim_per_tensor: List[List[int]]
is_hybrid_mla_backend: bool
ib_device: str ib_device: str
ib_traffic_class: str ib_traffic_class: str
gpu_id: int gpu_id: int
@@ -73,6 +73,7 @@ class PrefillServerInfo:
page_size: Optional[int] page_size: Optional[int]
kv_cache_dtype: Optional[str] kv_cache_dtype: Optional[str]
follow_bootstrap_room: bool follow_bootstrap_room: bool
enable_dsa_cache_layer_split: bool = False
# PD true-retraction rebootstrap: the prefill's HTTP API port. The decode # 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 # 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 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.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 = ( self.prefill_http_port = (
int(self.prefill_http_port) if self.prefill_http_port is not None else None 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.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.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_mla_backend = is_mla_backend
self.is_hybrid_mla_backend = getattr(args, "is_hybrid_mla_backend", False)
self.disaggregation_mode = disaggregation_mode self.disaggregation_mode = disaggregation_mode
self.server_args = server_args self.server_args = server_args
# for p/d multi node infer # for p/d multi node infer
@@ -146,8 +149,18 @@ class CommonKVManager(BaseKVManager):
self.pp_size = server_args.pp_size self.pp_size = server_args.pp_size
self.pp_rank = self.kv_args.pp_rank self.pp_rank = self.kv_args.pp_rank
self.local_ip = get_local_ip_auto() 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 = ( self.enable_all_cp_ranks_for_transfer = (
envs.SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER.get() envs.SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER.get()
or cp_sharded_prefill
or hybrid_decode_pulls_all_ranks
) )
# bind zmq socket # bind zmq socket
@@ -450,7 +463,7 @@ class CommonKVManager(BaseKVManager):
required_prefill_response_num = 1 required_prefill_response_num = 1
target_tp_ranks = [target_tp_rank] target_tp_ranks = [target_tp_rank]
elif self.attn_tp_size > info.attn_tp_size: 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( logger.warning_once(
"Performance is NOT guaranteed when using different TP sizes for non-MLA models. " "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 required_prefill_response_num = 1
target_tp_ranks = [target_tp_rank] target_tp_ranks = [target_tp_rank]
else: else:
if not self.is_mla_backend: if not self.is_mla_backend and not self.is_hybrid_mla_backend:
logger.warning_once( logger.warning_once(
"Performance is NOT guaranteed when using different TP sizes for non-MLA models. " "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] target_cp_ranks = [self.attn_cp_rank]
else: else:
target_cp_ranks = list(range(info.attn_cp_size)) 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 # Only retrieve from prefill CP rank 0 when not using all ranks
target_cp_ranks = target_cp_ranks[:1] target_cp_ranks = target_cp_ranks[:1]
required_prefill_response_num *= 1 required_prefill_response_num *= 1
@@ -582,6 +599,9 @@ class CommonKVManager(BaseKVManager):
"page_size": self.kv_args.page_size, "page_size": self.kv_args.page_size,
"kv_cache_dtype": self.server_args.kv_cache_dtype, "kv_cache_dtype": self.server_args.kv_cache_dtype,
"load_balance_method": self.server_args.load_balance_method, "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 # Self-register the HTTP API port so the decode can derive the PD
# retract rebootstrap /generate URL from bootstrap info instead of a # retract rebootstrap /generate URL from bootstrap info instead of a
# router-injected pd_rebootstrap_prefill_url. # router-injected pd_rebootstrap_prefill_url.
@@ -1041,7 +1061,10 @@ class CommonKVSender(BaseKVSender):
self.curr_idx += len(kv_indices) self.curr_idx += len(kv_indices)
is_last_chunk = self.curr_idx == self.num_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( kv_indices, index_slice = filter_kv_indices_for_cp_rank(
self.kv_mgr, self.kv_mgr,
kv_indices, kv_indices,
@@ -1396,6 +1419,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
self.page_size = None self.page_size = None
self.kv_cache_dtype: Optional[str] = None self.kv_cache_dtype: Optional[str] = None
self.follow_bootstrap_room: Optional[bool] = 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_http_port: Optional[int] = None
self.prefill_port_table: Dict[ self.prefill_port_table: Dict[
int, Dict[int, Dict[int, Dict[int, PrefillRankInfo]]] 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" 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: if system_dp_size == 1:
dp_group = attn_dp_rank dp_group = attn_dp_rank
else: else:
@@ -1555,6 +1584,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
if self.follow_bootstrap_room is not None if self.follow_bootstrap_room is not None
else True else True
), ),
enable_dsa_cache_layer_split=bool(self.enable_dsa_cache_layer_split),
prefill_http_port=self.prefill_http_port, prefill_http_port=self.prefill_http_port,
) )
return web.json_response(dataclasses.asdict(info), status=200) return web.json_response(dataclasses.asdict(info), status=200)
@@ -605,7 +605,7 @@ class MooncakeKVManager(CommonKVManager):
layers_params = None layers_params = None
# Decode pp size should be equal to prefill pp size or 1 # 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 = ( 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) 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)}" 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( def maybe_send_extra(
self, self,
req: TransferInfo, req: TransferInfo,
@@ -1308,10 +1333,15 @@ class MooncakeKVManager(CommonKVManager):
target_rank_registration_info: KVArgsRegisterInfo = ( target_rank_registration_info: KVArgsRegisterInfo = (
self.decode_kv_args_table[req.mooncake_session_id] 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 ret = 0
elif self.is_mla_backend or ( elif (
self.attn_tp_size self.is_mla_backend
or self.is_hybrid_mla_backend
or self.attn_tp_size
== target_rank_registration_info.dst_attn_tp_size == target_rank_registration_info.dst_attn_tp_size
): ):
ret = self.send_kvcache( ret = self.send_kvcache(
@@ -1376,7 +1406,7 @@ class MooncakeKVManager(CommonKVManager):
break break
if kv_chunk.is_last_chunk: if kv_chunk.is_last_chunk:
if kv_chunk.state_indices: if kv_chunk.state_indices and not skip_state:
self.maybe_send_extra( self.maybe_send_extra(
req, req,
kv_chunk.state_indices, kv_chunk.state_indices,
+24 -4
View File
@@ -154,14 +154,34 @@ class PrefillBootstrapQueue:
kv_args.engine_rank = self.tp_rank kv_args.engine_rank = self.tp_rank
kv_args.pp_rank = self.pp_rank kv_args.pp_rank = self.pp_rank
kv_args.system_dp_rank = self.scheduler.ps.dp_rank kv_args.system_dp_rank = self.scheduler.ps.dp_rank
kv_args.prefill_start_layer = self.token_to_kv_pool.start_layer layer_shard_enabled = getattr(
kv_args.prefill_end_layer = getattr(self.token_to_kv_pool, "end_layer", None) 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_args.mla_compression_ratios = None
kv_data_ptrs, kv_data_lens, kv_item_lens = ( kv_data_ptrs, kv_data_lens, kv_item_lens = (
self.token_to_kv_pool.get_contiguous_buf_infos() 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 # We should also transfer draft model kv cache. The indices are
# always shared with a target model. # always shared with a target model.
draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = ( draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = (
@@ -191,7 +211,7 @@ class PrefillBootstrapQueue:
setup_state_kv_args( setup_state_kv_args(
kv_args, kv_args,
self.token_to_kv_pool, 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, self.scheduler.model_config.num_hidden_layers,
req_to_token_pool=req_to_token_pool, 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_data_lens = []
kv_args.state_item_lens = [] kv_args.state_item_lens = []
kv_args.state_dim_per_tensor = [] kv_args.state_dim_per_tensor = []
kv_args.is_hybrid_mla_backend = False
if isinstance(token_to_kv_pool, MiniMaxSparseKVPool): if isinstance(token_to_kv_pool, MiniMaxSparseKVPool):
if token_to_kv_pool.index_kv_pool is not None: 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") if hasattr(token_to_kv_pool, "get_state_dim_per_tensor")
else None else None
) )
kv_args.is_hybrid_mla_backend = is_mla_backend(
token_to_kv_pool.full_kv_pool
)
append_state_component( append_state_component(
kv_args, StateType.MAMBA, data_ptrs, data_lens, item_lens, dim 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 out_cache_loc = forward_batch.out_cache_loc
pool = get_token_to_kv_pool() pool = get_token_to_kv_pool()
page_size = pool.page_size 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 ( if (
not _is_fp8_fnuz not _is_fp8_fnuz
and out_cache_loc is not None and out_cache_loc is not None
@@ -801,6 +805,15 @@ class Indexer(MultiPlatformOp):
return return
dst.copy_(src) 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 @staticmethod
def _pad_heads_for_deep_gemm(q_fp8, weights): def _pad_heads_for_deep_gemm(q_fp8, weights):
"""Pad q and weights to 32 heads when num_heads < 32, """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() block_tables = metadata.get_page_table_64()
max_seq_len = block_tables.shape[1] * page_size max_seq_len = block_tables.shape[1] * page_size
kv_cache_fp8 = get_token_to_kv_pool().get_index_k_with_scale_buffer( kv_cache_fp8 = self._get_index_k_read_buffer(get_token_to_kv_pool(), layer_id)
layer_id=layer_id
)
blocksize = page_size blocksize = page_size
if ( if (
@@ -1627,24 +1638,28 @@ class Indexer(MultiPlatformOp):
if out_cache_loc is None: if out_cache_loc is None:
out_cache_loc = forward_batch.out_cache_loc 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 ( if (
_is_cuda _is_cuda
and (not _is_fp8_fnuz) and (not _is_fp8_fnuz)
and can_use_dsa_fused_store( and can_use_dsa_fused_store(
key.dtype, key.dtype,
out_cache_loc.dtype, out_cache_loc.dtype,
get_token_to_kv_pool().page_size, pool.page_size,
) )
): ):
# NOTE: wrapper already normalizes shape/contiguity and asserts dtypes. # NOTE: wrapper already normalizes shape/contiguity and asserts dtypes.
buf = get_token_to_kv_pool().get_index_k_with_scale_buffer( buf = pool.get_index_k_with_scale_buffer(layer_id=layer_id)
layer_id=layer_id
)
fused_store_index_k_cache( fused_store_index_k_cache(
key, key,
buf, buf,
out_cache_loc, out_cache_loc,
get_token_to_kv_pool().page_size, pool.page_size,
) )
return return
@@ -1654,10 +1669,8 @@ class Indexer(MultiPlatformOp):
# layout with page_size=1; the same kv_cache.view works for both cases # layout with page_size=1; the same kv_cache.view works for both cases
# because page_size is 1 there. # because page_size is 1 there.
if _use_aiter: if _use_aiter:
page_size = get_token_to_kv_pool().page_size page_size = pool.page_size
buf = get_token_to_kv_pool().get_index_k_with_scale_buffer( buf = pool.get_index_k_with_scale_buffer(layer_id=layer_id)
layer_id=layer_id
)
kv_cache = buf.view(-1, page_size, 132).view(fp8_dtype) kv_cache = buf.view(-1, page_size, 132).view(fp8_dtype)
out_loc = forward_batch.out_cache_loc out_loc = forward_batch.out_cache_loc
if not out_loc.is_contiguous(): if not out_loc.is_contiguous():
@@ -1679,7 +1692,7 @@ class Indexer(MultiPlatformOp):
if not out_cache_loc.is_contiguous(): if not out_cache_loc.is_contiguous():
out_cache_loc = out_cache_loc.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, layer_id=layer_id,
loc=out_cache_loc, loc=out_cache_loc,
index_k=k_fp8, 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.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_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 from sglang.srt.runtime_context import get_parallel
@@ -48,6 +49,25 @@ def dsa_enable_prefill_cp():
return is_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): def dsa_cp_gather_hidden_states(hidden_states: torch.Tensor):
attn_dp_size = get_parallel().attn_dp_size attn_dp_size = get_parallel().attn_dp_size
attn_tp_size = get_parallel().attn_tp_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.""" """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 ( from sglang.srt.layers.cp.base import (
BaseContextParallelMetadata, BaseContextParallelMetadata,
@@ -33,6 +33,9 @@ from sglang.srt.layers.cp.zigzag import (
ZigzagCPStrategy, ZigzagCPStrategy,
) )
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
CP_V2_DEFAULT_MODEL_CLASSES = frozenset( CP_V2_DEFAULT_MODEL_CLASSES = frozenset(
{ {
"Qwen3MoeForCausalLM", "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: def enable_cp_v2() -> bool:
"""Return whether the CP-v2 path is enabled for this process.""" """Return whether the CP-v2 path is enabled for this process."""
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -140,4 +226,9 @@ __all__ = [
"cp_gather_after_forward", "cp_gather_after_forward",
"cp_split_before_forward", "cp_split_before_forward",
"prepare_cp_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()
+39 -20
View File
@@ -1226,6 +1226,7 @@ class KvBufferDesc:
class KVCache(abc.ABC): class KVCache(abc.ABC):
layer_shard_enabled: bool = False
post_capture_active: bool = False post_capture_active: bool = False
@abc.abstractmethod @abc.abstractmethod
@@ -2838,7 +2839,6 @@ class MLATokenToKVPool(KVCache):
if not valid_mask.all(): if not valid_mask.all():
loc = loc[valid_mask] loc = loc[valid_mask]
cache_k = cache_k[valid_mask] cache_k = cache_k[valid_mask]
if cache_k.dtype != self.dtype: if cache_k.dtype != self.dtype:
cache_k = cache_k.to(self.dtype) cache_k = cache_k.to(self.dtype)
@@ -2849,21 +2849,18 @@ class MLATokenToKVPool(KVCache):
else: else:
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
def set_mla_kv_buffer( def _write_mla_kv_buffer(
self, self,
layer: RadixAttention, dst_buffer: torch.Tensor,
loc: torch.Tensor, loc: torch.Tensor,
cache_k_nope: torch.Tensor, cache_k_nope: torch.Tensor,
cache_k_rope: torch.Tensor, cache_k_rope: torch.Tensor,
): ) -> None:
maybe_detect_oob(loc, 0, self.size + self.page_size, "set_mla_kv_buffer (MLA)")
layer_id = layer.layer_id
if _is_hip and self.use_dsa and self.dtype == fp8_dtype: 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. # HIP FP8 path uses raw MLA KV layout (nope + rope) without per-block scales.
# Fuse BF16/FP16 -> FP8 cast with paged KV write. # Fuse BF16/FP16 -> FP8 cast with paged KV write.
set_mla_kv_buffer_triton_fp8_quant( set_mla_kv_buffer_triton_fp8_quant(
self.kv_buffer[layer_id - self.start_layer], dst_buffer,
loc, loc,
cache_k_nope, cache_k_nope,
cache_k_rope, 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_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)] # cache_k_rope_fp8: (num_tokens, 1, 128) uint8 [rope_bf16_bytes(128)]
set_mla_kv_buffer_triton( set_mla_kv_buffer_triton(
self.kv_buffer[layer_id - self.start_layer], dst_buffer,
loc, loc,
cache_k_nope_fp8, cache_k_nope_fp8,
cache_k_rope_fp8, cache_k_rope_fp8,
@@ -2895,6 +2892,22 @@ class MLATokenToKVPool(KVCache):
cache_k_rope = cache_k_rope.view(self.store_dtype) cache_k_rope = cache_k_rope.view(self.store_dtype)
set_mla_kv_buffer_triton( set_mla_kv_buffer_triton(
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], self.kv_buffer[layer_id - self.start_layer],
loc, loc,
cache_k_nope, cache_k_nope,
@@ -3150,6 +3163,7 @@ class DSATokenToKVPool(MLATokenToKVPool):
self.index_head_dim = index_head_dim self.index_head_dim = index_head_dim
if index_buf_size is None: if index_buf_size is None:
index_buf_size = size index_buf_size = size
self.index_buf_size = index_buf_size
# num head == 1 and head dim == 128 for index_k in DSA # num head == 1 and head dim == 128 for index_k in DSA
assert index_head_dim == 128 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}" ), f"HIP legacy DSA path requires page_size == 1, got {self.page_size}"
else: else:
assert self.page_size == 64 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 ( with (
torch.cuda.use_mem_pool(self.custom_mem_pool) torch.cuda.use_mem_pool(self.custom_mem_pool)
if self.custom_mem_pool if self.custom_mem_pool
@@ -3177,22 +3203,15 @@ class DSATokenToKVPool(MLATokenToKVPool):
# data: for page i, # data: for page i,
# * buf[i, :page_size * head_dim] for fp8 data # * buf[i, :page_size * head_dim] for fp8 data
# * buf[i, page_size * head_dim:].view(float32) for scale # * buf[i, page_size * head_dim:].view(float32) for scale
( self._index_buffer_shape(num_pages),
(index_buf_size + page_size + 1) // self.page_size,
self.page_size
* (
index_head_dim + index_head_dim // self.quant_block_size * 4
),
),
dtype=self.index_k_with_scale_buffer_dtype, 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): def _clear_buffers(self):
del self.kv_buffer super()._clear_buffers()
del self.index_k_with_scale_buffer del self.index_k_with_scale_buffer
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor): 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): def get_size_per_token(self):
self.kv_lora_rank = self.device_pool.kv_lora_rank self.kv_lora_rank = self.device_pool.kv_lora_rank
self.qk_rope_head_dim = self.device_pool.qk_rope_head_dim 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_cache_dim = self.override_kv_cache_dim or (
self.kv_lora_rank + self.qk_rope_head_dim self.kv_lora_rank + self.qk_rope_head_dim
) )
@@ -244,19 +244,23 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
def load_to_device_per_layer( def load_to_device_per_layer(
self, device_pool, host_indices, device_indices, layer_id, io_backend 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 io_backend == "kernel":
if self.layout == "layer_first": if self.layout == "layer_first":
if self.can_use_jit: if self.can_use_jit:
jit_transfer_hicache_one_layer_mla( jit_transfer_hicache_one_layer_mla(
cache_dst=device_pool.kv_buffer[layer_id], 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_dst=device_indices,
indices_src=host_indices, indices_src=host_indices,
element_dim=self.kv_cache_dim, element_dim=self.kv_cache_dim,
) )
else: else:
transfer_kv_per_layer_mla( transfer_kv_per_layer_mla(
src=self.kv_buffer[layer_id], src=self.kv_buffer[host_layer],
dst=device_pool.kv_buffer[layer_id], dst=device_pool.kv_buffer[layer_id],
src_indices=host_indices, src_indices=host_indices,
dst_indices=device_indices, dst_indices=device_indices,
@@ -266,7 +270,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
if self.can_use_jit: if self.can_use_jit:
jit_transfer_hicache_one_layer_mla( jit_transfer_hicache_one_layer_mla(
cache_dst=device_pool.kv_buffer[layer_id], 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_dst=device_indices,
indices_src=host_indices, indices_src=host_indices,
element_dim=self.kv_cache_dim, element_dim=self.kv_cache_dim,
@@ -277,7 +281,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
dst=device_pool.kv_buffer[layer_id], dst=device_pool.kv_buffer[layer_id],
src_indices=host_indices, src_indices=host_indices,
dst_indices=device_indices, dst_indices=device_indices,
layer_id=layer_id, layer_id=host_layer,
item_size=self.token_stride_size, item_size=self.token_stride_size,
src_layout_dim=self.layout_dim, src_layout_dim=self.layout_dim,
) )
@@ -286,7 +290,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
elif io_backend == "direct": elif io_backend == "direct":
if self.layout == "layer_first": if self.layout == "layer_first":
transfer_kv_direct( 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]], dst_layers=[device_pool.kv_buffer[layer_id]],
src_indices=host_indices, src_indices=host_indices,
dst_indices=device_indices, dst_indices=device_indices,
@@ -298,7 +302,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
dst_ptrs=[device_pool.kv_buffer[layer_id]], dst_ptrs=[device_pool.kv_buffer[layer_id]],
src_indices=host_indices, src_indices=host_indices,
dst_indices=device_indices, dst_indices=device_indices,
layer_id=layer_id, layer_id=host_layer,
page_size=self.page_size, page_size=self.page_size,
) )
else: else:
@@ -324,9 +328,75 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
else: else:
raise ValueError(f"Unsupported IO backend: {io_backend}") 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( def backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend 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 io_backend == "kernel":
if self.layout == "layer_first": if self.layout == "layer_first":
if self.can_use_jit: if self.can_use_jit:
@@ -2109,7 +2179,7 @@ class DSAIndexerPoolHost(HostKVCache):
self.dtype = device_pool.store_dtype self.dtype = device_pool.store_dtype
self.start_layer = device_pool.start_layer self.start_layer = device_pool.start_layer
self.end_layer = device_pool.end_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.index_head_dim = device_pool.index_head_dim
self.indexer_quant_block_size = device_pool.quant_block_size self.indexer_quant_block_size = device_pool.quant_block_size
@@ -2242,6 +2312,10 @@ class DSAIndexerPoolHost(HostKVCache):
def load_to_device_per_layer( def load_to_device_per_layer(
self, device_pool, host_indices, device_indices, layer_id, io_backend 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_page_indices, device_page_indices = self._get_indexer_page_indices(
host_indices, device_indices host_indices, device_indices
) )
@@ -2249,7 +2323,7 @@ class DSAIndexerPoolHost(HostKVCache):
if use_kernel: if use_kernel:
if self.layout == "layer_first": if self.layout == "layer_first":
transfer_kv_per_layer_mla( 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], dst=device_pool.index_k_with_scale_buffer[layer_id],
src_indices=host_page_indices, src_indices=host_page_indices,
dst_indices=device_page_indices, dst_indices=device_page_indices,
@@ -2261,7 +2335,7 @@ class DSAIndexerPoolHost(HostKVCache):
dst=device_pool.index_k_with_scale_buffer[layer_id], dst=device_pool.index_k_with_scale_buffer[layer_id],
src_indices=host_page_indices, src_indices=host_page_indices,
dst_indices=device_page_indices, dst_indices=device_page_indices,
layer_id=layer_id, layer_id=host_layer,
item_size=self.indexer_page_stride_size, item_size=self.indexer_page_stride_size,
src_layout_dim=self.indexer_layout_dim, src_layout_dim=self.indexer_layout_dim,
) )
@@ -2270,7 +2344,7 @@ class DSAIndexerPoolHost(HostKVCache):
elif io_backend == "direct": elif io_backend == "direct":
if self.layout == "layer_first": if self.layout == "layer_first":
transfer_kv_direct( 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]], dst_layers=[device_pool.index_k_with_scale_buffer[layer_id]],
src_indices=host_page_indices, src_indices=host_page_indices,
dst_indices=device_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]], dst_ptrs=[device_pool.index_k_with_scale_buffer[layer_id]],
src_indices=host_page_indices, src_indices=host_page_indices,
dst_indices=device_page_indices, dst_indices=device_page_indices,
layer_id=layer_id, layer_id=host_layer,
page_size=1, page_size=1,
) )
else: else:
@@ -2290,9 +2364,57 @@ class DSAIndexerPoolHost(HostKVCache):
else: else:
raise ValueError(f"Unsupported IO backend: {io_backend}") 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( def backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend 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_page_indices, device_page_indices = self._get_indexer_page_indices(
host_indices, device_indices host_indices, device_indices
) )
@@ -168,6 +168,41 @@ class HostKVCache(abc.ABC):
def get_size_per_token(self): def get_size_per_token(self):
raise NotImplementedError() 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 @abc.abstractmethod
def init_kv_buffer(self): def init_kv_buffer(self):
raise NotImplementedError() raise NotImplementedError()
@@ -947,16 +947,31 @@ class ModelRunnerKVCacheMixin:
end_layer=self.end_layer, end_layer=self.end_layer,
) )
elif self.use_mla_backend and is_dsa_model: elif self.use_mla_backend and is_dsa_model:
PoolCls = ( from sglang.srt.layers.cp.utils import get_glm_dsa_cp_layer_shard_info
HiSparseDSATokenToKVPool if self.enable_hisparse else DSATokenToKVPool
) (
dsa_cp_layer_shard_rank,
dsa_cp_layer_shard_size,
) = get_glm_dsa_cp_layer_shard_info(self)
pool_kwargs = {} pool_kwargs = {}
if self.enable_hisparse: if self.enable_hisparse:
PoolCls = HiSparseDSATokenToKVPool
from sglang.srt.mem_cache.sparsity import parse_hisparse_config from sglang.srt.mem_cache.sparsity import parse_hisparse_config
pool_kwargs["host_to_device_ratio"] = parse_hisparse_config( pool_kwargs["host_to_device_ratio"] = parse_hisparse_config(
self.server_args self.server_args
).host_to_device_ratio ).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.token_to_kv_pool = PoolCls(
self.max_total_num_tokens, self.max_total_num_tokens,
page_size=self.page_size, page_size=self.page_size,
@@ -177,6 +177,13 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
# args to config cell size # args to config cell size
model_config = mr.model_config model_config = mr.model_config
kv_cache_dtype = mr.kv_cache_dtype 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) kv_size = torch._utils._element_size(kv_cache_dtype)
tp_size = get_parallel().attn_tp_size tp_size = get_parallel().attn_tp_size
@@ -184,7 +191,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
if mr.use_mla_backend: if mr.use_mla_backend:
cell_size = ( cell_size = (
(model_config.kv_lora_rank + model_config.qk_rope_head_dim) (model_config.kv_lora_rank + model_config.qk_rope_head_dim)
* num_layers * effective_num_layers
* kv_size * kv_size
) )
if is_float4_e2m1fn_x2(kv_cache_dtype): 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) (model_config.kv_lora_rank + model_config.qk_rope_head_dim)
// scale_block_size // scale_block_size
) )
* num_layers * effective_num_layers
* kv_size * kv_size
) )
@@ -209,7 +216,9 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
element_size = torch._utils._element_size( element_size = torch._utils._element_size(
DSATokenToKVPool.index_k_with_scale_buffer_dtype 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): elif is_minimax_sparse(model_config.hf_config):
# Mirrors MiniMaxSparseKVPool: main pool (K+V all layers) + indexer pool # Mirrors MiniMaxSparseKVPool: main pool (K+V all layers) + indexer pool
# (sparse-only, single-head; kv layers store K+V, k-only layers store K). # (sparse-only, single-head; kv layers store K+V, k-only layers store K).
@@ -252,7 +261,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
cell_size = ( cell_size = (
model_config.get_num_kv_heads(tp_size) model_config.get_num_kv_heads(tp_size)
* (model_config.head_dim + model_config.v_head_dim) * (model_config.head_dim + model_config.v_head_dim)
* num_layers * effective_num_layers
* kv_size * kv_size
) )
@@ -262,7 +271,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
n = model_config.get_num_kv_heads(tp_size) n = model_config.get_num_kv_heads(tp_size)
k = model_config.head_dim k = model_config.head_dim
cell_size = (cell_size // 2) + ( 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 return cell_size
+17 -1
View File
@@ -70,7 +70,10 @@ from sglang.srt.layers.communicator import (
enable_moe_dense_fully_dp, enable_moe_dense_fully_dp,
get_attn_tp_context, 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 import dcp_enabled, get_attention_dcp_world_size
from sglang.srt.layers.dcp.planner import ( from sglang.srt.layers.dcp.planner import (
prepare_decode_context_parallel_metadata, prepare_decode_context_parallel_metadata,
@@ -2207,6 +2210,7 @@ class DeepseekV2DecoderLayer(nn.Module):
llama_4_scaling: Optional[torch.Tensor] = None, llama_4_scaling: Optional[torch.Tensor] = None,
prev_topk_indices: Optional[torch.Tensor] = None, prev_topk_indices: Optional[torch.Tensor] = None,
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None, captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
next_full_attention_layer_id: Optional[int] = None,
) -> torch.Tensor: ) -> torch.Tensor:
hidden_states_orig = hidden_states hidden_states_orig = hidden_states
hidden_states, residual = ( hidden_states, residual = (
@@ -2234,6 +2238,10 @@ class DeepseekV2DecoderLayer(nn.Module):
topk_indices = None topk_indices = None
get_attn_tp_context().clear_attn_inputs() 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 = self.layer_communicator.prepare_mlp(
hidden_states, residual, forward_batch 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: if self.pp_group.is_last_rank:
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
else: else:
@@ -2597,6 +2610,9 @@ class DeepseekV2Model(nn.Module):
captured_last_layer_outputs=( captured_last_layer_outputs=(
aux_hidden_states if i in self.layers_to_capture else None 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: if normal_end_layer != self.end_layer:
+57
View File
@@ -916,6 +916,11 @@ class ServerArgs:
choices=("zigzag", "interleave"), choices=("zigzag", "interleave"),
), ),
] = None ] = 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 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" dsa_prefill_cp_mode: A[str, Arg(no_cli=True)] = "round-robin-split"
enable_prefill_context_parallel: A[bool, Arg(no_cli=True)] = False 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 hf_config = self.get_model_config().hf_config
model_arch = hf_config.architectures[0] 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) _hybrid_spec = get_linear_attn_spec_by_arch(model_arch)
if _hybrid_spec is not None and _hybrid_spec.uses_mamba_radix_cache: if _hybrid_spec is not None and _hybrid_spec.uses_mamba_radix_cache:
self._handle_mamba_radix_cache(model_arch=model_arch) self._handle_mamba_radix_cache(model_arch=model_arch)
@@ -4158,6 +4169,52 @@ class ServerArgs:
assert ( assert (
self.disaggregation_mode != "decode" self.disaggregation_mode != "decode"
), "CP is only supported for prefill when PD disaggregation, please remove --enable-prefill-cp." ), "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: else:
# DeepSeek V3/R1/V3.1 # DeepSeek V3/R1/V3.1
@@ -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()
@@ -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()
@@ -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()
@@ -49,6 +49,15 @@ def _ptr_key_from_tensor(ptrs: torch.Tensor) -> tuple[int, ...]:
return tuple(int(ptr) for ptr in ptrs.cpu().tolist()) 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( def _cpu_staged_lf_pf_copy(
src_registry, src_registry,
*, *,
@@ -192,7 +201,8 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
] ]
expected_k = [layer[device_indices].clone() for layer in k_layers] expected_k = [layer[device_indices].clone() for layer in k_layers]
expected_v = [layer[device_indices].clone() for layer in v_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, k_buffer=k_layers,
v_buffer=v_layers, v_buffer=v_layers,
k_data_ptrs=torch.tensor( k_data_ptrs=torch.tensor(
@@ -293,7 +303,8 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
for layer_id in range(layer_num) for layer_id in range(layer_num)
] ]
expected = [layer[device_indices].clone() for layer in device_layers] 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, kv_buffer=device_layers,
data_ptrs=torch.tensor( data_ptrs=torch.tensor(
[layer.data_ptr() for layer in device_layers], dtype=torch.uint64 [layer.data_ptr() for layer in device_layers], dtype=torch.uint64
@@ -301,6 +312,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
) )
host = MLATokenToKVPoolHost.__new__(MLATokenToKVPoolHost) host = MLATokenToKVPoolHost.__new__(MLATokenToKVPoolHost)
host.device_pool = device_pool
host.layout = "page_first" host.layout = "page_first"
host.page_size = 1 host.page_size = 1
host.layer_num = layer_num host.layer_num = layer_num
@@ -582,9 +594,13 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
for layer_id in range(layer_num) for layer_id in range(layer_num)
] ]
expected = [buffer[device_page_indices].clone() for buffer in device_layers] 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 = DSAIndexerPoolHost.__new__(DSAIndexerPoolHost)
host.device_pool = device_pool
host.layout = "page_first" host.layout = "page_first"
host.page_size = page_size host.page_size = page_size
host.layer_num = layer_num host.layer_num = layer_num
@@ -114,6 +114,7 @@ def _make_model_runner(
sa.disaggregation_mode = disaggregation_mode sa.disaggregation_mode = disaggregation_mode
sa.max_running_requests = max_running_requests sa.max_running_requests = max_running_requests
sa.disaggregation_decode_extra_slots = disaggregation_decode_extra_slots sa.disaggregation_decode_extra_slots = disaggregation_decode_extra_slots
sa.enable_dsa_cache_layer_split = False
mr.server_args = sa mr.server_args = sa
spec = MagicMock() spec = MagicMock()