[Feat][GLM5.2] Add DSA Cache Layer Split under Prefill CP (#29421)
Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com>
This commit is contained in:
@@ -49,6 +49,7 @@ class KVArgs:
|
|||||||
state_item_lens: List[List[int]]
|
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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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,12 +2892,28 @@ 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(
|
||||||
self.kv_buffer[layer_id - self.start_layer],
|
dst_buffer,
|
||||||
loc,
|
loc,
|
||||||
cache_k_nope,
|
cache_k_nope,
|
||||||
cache_k_rope,
|
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(
|
def get_mla_kv_buffer(
|
||||||
self,
|
self,
|
||||||
layer: RadixAttention,
|
layer: RadixAttention,
|
||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user