Files
sglang/python/sglang/srt/disaggregation/utils.py
T
2026-09-18 22:41:50 -07:00

1754 lines
68 KiB
Python

from __future__ import annotations
import random
from collections import deque
from contextlib import nullcontext
from enum import Enum
from typing import (
TYPE_CHECKING,
Iterable,
List,
Literal,
Optional,
Tuple,
Type,
overload,
)
import numpy as np
import torch
import torch.distributed as dist
from sglang.srt.configs.model_config import get_dsa_mtp_topk_width, is_deepseek_dsa
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.environ import envs
from sglang.srt.runtime_context import (
get_disagg,
get_spec,
)
from sglang.srt.utils import is_npu
if TYPE_CHECKING:
from sglang.srt.disaggregation.base.conn import KVArgs, StateType
from sglang.srt.disaggregation.common.conn import (
CommonKVBootstrapServer,
CommonKVManager,
CommonKVReceiver,
CommonKVSender,
)
from sglang.srt.managers.schedule_batch import Req
if is_npu():
from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import (
DSV4NPUTokenToKVPool,
)
#########################
# Constants & Enums
#########################
FAKE_BOOTSTRAP_HOST = "2.2.2.2"
def poll_and_all_reduce_pp(
rids: Iterable[str],
ready_poll: int,
pp_good_rids: Optional[List[str]] = None,
pp_bad_rids: Optional[List[str]] = None,
) -> List[Optional[int]]:
"""Map authoritative PP consensus to poll states without polling again."""
if pp_good_rids is None or pp_bad_rids is None:
raise ValueError("PP consensus is required")
good_rids = set(pp_good_rids)
bad_rids = set(pp_bad_rids)
return [
KVPoll.Failed if rid in bad_rids else ready_poll if rid in good_rids else None
for rid in rids
]
def get_dsa_seed_metadata_dim(hf_config) -> int:
"""Return the model-defined PD seed width, independent of local spec mode."""
if not getattr(hf_config, "index_share_for_mtp_iteration", False):
return 0
# QSA models reuse the same flag for their draft-side index sharing but
# carry no DSA seed metadata over PD.
if not is_deepseek_dsa(hf_config):
return 0
return get_dsa_mtp_topk_width(hf_config)
def get_qsa_pending_state_indices(req: Req) -> np.ndarray:
"""Return the request-pool row that owns a QSA pending-state ring."""
req_pool_idx = req.kv.req_pool_idx
if req_pool_idx is None:
raise ValueError("QSA pending-state transfer requires an allocated request row")
return np.array([int(req_pool_idx)], dtype=np.int32)
class DisaggregationMode(Enum):
NULL = "null"
PREFILL = "prefill"
DECODE = "decode"
@staticmethod
def to_engine_type(mode: str) -> str:
if mode == DisaggregationMode.PREFILL.value:
return "prefill"
elif mode == DisaggregationMode.DECODE.value:
return "decode"
return "unified"
def unified_memory_disagg_move_gate(scheduler):
"""Compaction move gate for a PD node running the unified memory pool.
Returns a predicate that is True only when no transfer can be in flight, so
compaction never relocates a page the RDMA engine is reading or writing.
Safe to read this state from here: every mover runs on the scheduler thread.
A page is exposed from the moment its address reaches the peer until the
transfer concludes, and for part of that lifetime the request is in NEITHER
end's queue -- so queue emptiness alone is not enough:
- PREFILL: scheduling the final chunk clears `chunked_req` while earlier
chunks may still be draining, and the request only reaches the inflight
queue later, in the result path.
- DECODE: `pop_preallocated` publishes one request's destinations and keeps
allocating for the next, whose allocation can urgently flush the peer
sub-allocator; the batch reaches the transfer queue only after the loop.
"""
if scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
def prefill_gate() -> bool:
return not (
scheduler.disagg_prefill_inflight_queue
or scheduler.disagg_prefill_pending_chunk_rids
)
return prefill_gate
if scheduler.disaggregation_mode == DisaggregationMode.DECODE:
def decode_gate() -> bool:
return not (
scheduler.disagg_decode_transfer_queue.queue
or scheduler.disagg_decode_prealloc_queue.has_published_destinations
)
return decode_gate
raise ValueError(
"unified_memory_disagg_move_gate: scheduler is not a PD node "
f"(mode={scheduler.disaggregation_mode})"
)
#########################
# Synchronization
#########################
def _poll_with_failure_injection(pollers) -> List[int]:
if (failure_prob := envs.SGLANG_TEST_DISAGG_FAILURE_PROB.get()) > 0:
return [
int(KVPoll.Failed) if random.random() < failure_prob else int(poller.poll())
for poller in pollers
]
return [int(poller.poll()) for poller in pollers]
def _is_fake_transfer(req: Req) -> bool:
return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
req.bootstrap_host is None
and get_disagg().disaggregation_transfer_backend == "fake"
)
def _apply_metadata_gate(polls, decode_reqs, metadata_buffers) -> None:
"""Downgrade Success → Transferring for requests whose metadata hasn't landed.
Mutates `polls` in-place. Called before all-reduce so that MIN across TP
ranks naturally prevents any rank from committing before all ranks are ready.
"""
for i, poll_val in enumerate(polls):
if poll_val == int(KVPoll.Success):
decode_req = decode_reqs[i]
if _is_fake_transfer(decode_req.req):
continue
actual_room = metadata_buffers.bootstrap_room[
decode_req.metadata_buffer_index, 0
].item()
if actual_room == 0:
polls[i] = int(KVPoll.Transferring)
def _all_reduce_polls(polls: List[int], group: dist.ProcessGroup) -> List[int]:
"""MIN-reduce poll states so no rank commits ahead of its peers."""
if dist.get_world_size(group) == 1:
return polls
tensor_to_reduce = torch.tensor(polls, dtype=torch.uint8, device="cpu")
dist.all_reduce(tensor_to_reduce, op=dist.ReduceOp.MIN, group=group)
return tensor_to_reduce.tolist()
def poll_and_all_reduce(
pollers,
gloo_group: dist.ProcessGroup,
decode_reqs=None,
metadata_buffers: Optional[MetadataBuffers] = None,
):
# at a certain prob, the poll is failed to simulate failure
polls = _poll_with_failure_injection(pollers)
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
if decode_reqs is not None and metadata_buffers is not None:
_apply_metadata_gate(polls, decode_reqs, metadata_buffers)
return _all_reduce_polls(polls, gloo_group)
def poll_and_all_reduce_attn_cp_tp_group(
pollers,
attn_cp_cpu_group: dist.ProcessGroup,
attn_tp_cpu_group: dist.ProcessGroup,
):
# First sync across attn-tp ranks so all TP participants for a given (dp, cp)
# shard observe the same status transitions.
polls = poll_and_all_reduce(pollers, attn_tp_cpu_group)
# Then sync across attn-cp ranks, so all TPxCP participants in one DP shard
# converge to the same global status.
return _all_reduce_polls(polls, attn_cp_cpu_group)
def poll_and_all_reduce_with_staging(
decode_reqs,
staging_handler,
gloo_group: dist.ProcessGroup,
metadata_buffers: Optional[MetadataBuffers] = None,
):
"""Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce."""
for decode_req in decode_reqs:
if decode_req.kv_receiver.require_staging and not staging_handler.is_done(
decode_req
):
staging_handler.advance_scatter(decode_req)
# allow test injection of failure probability at runtime
receivers = [dr.kv_receiver for dr in decode_reqs]
raw_polls = _poll_with_failure_injection(receivers)
for i, decode_req in enumerate(decode_reqs):
if decode_req.kv_receiver.require_staging and staging_handler.is_failed(
decode_req
):
# Staging completion timed out; KVPoll.Failed == 0 propagates
# through the MIN all_reduce.
raw_polls[i] = int(KVPoll.Failed)
continue
if raw_polls[i] == int(KVPoll.Success):
if decode_req.kv_receiver.require_staging and not staging_handler.is_done(
decode_req
):
raw_polls[i] = int(KVPoll.Transferring)
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
if metadata_buffers is not None:
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers)
return _all_reduce_polls(raw_polls, gloo_group)
#########################
# Metadata Buffers
#########################
class ReqToMetadataIdxAllocator:
"""A memory pool that maps a request to its first output token location."""
def __init__(
self,
size: int,
):
self.size = size
self.free_slots = deque(list(range(size)))
def available_size(self):
return len(self.free_slots)
def alloc(self) -> Optional[int]:
if len(self.free_slots) == 0:
return None
return self.free_slots.popleft()
def free(self, free_index: int):
self.free_slots.append(free_index)
class MetadataBuffers:
def __init__(
self,
size: int,
hidden_size: int,
hidden_states_dtype: torch.dtype,
max_sampling_mask_tokens: int,
max_top_logprobs_num: int = 128,
custom_mem_pool: torch.cuda.MemPool = None,
output_dsa_topk_indices_dim: int = 0,
*,
kv_checksum_enabled: bool = False,
):
self.custom_mem_pool = custom_mem_pool
self.output_dsa_topk_indices_dim = output_dsa_topk_indices_dim
self.enable_sampling_mask = envs.SGLANG_ENABLE_DISAGG_SAMPLING_MASK.get()
bootstrap_room_dtype = torch.uint64
device = "cpu"
if is_npu():
# For ascend backend, output tokens are placed in the NPU and will be transferred by D2D channel.
device = "npu"
# TODO: Fix me when npu backend supports torch.uint64
bootstrap_room_dtype = torch.int64
elif self.custom_mem_pool:
# TODO(shangming): Fix me (use 'cuda') when nvlink_transport of Mooncake is bug-free
device = "cpu"
elif envs.SGLANG_MOONCAKE_CUSTOM_MEM_POOL.get() == "INTRA_NODE_NVLINK":
device = "cuda"
with (
torch.cuda.use_mem_pool(self.custom_mem_pool)
if self.custom_mem_pool
else nullcontext()
):
# TODO: abort top_logprobs_num > 128 in PD
# We transfer the metadata of first output token to decode
# The minimal size for RDMA is 64Bytes, so we pad it to > 64Bytes
self.output_ids = torch.zeros((size, 16), dtype=torch.int32, device=device)
self.cached_tokens = torch.zeros(
(size, 16), dtype=torch.int32, device=device
)
self.output_token_logprobs_val = torch.zeros(
(size, 16), dtype=torch.float32, device=device
)
self.output_token_logprobs_idx = torch.zeros(
(size, 16), dtype=torch.int32, device=device
)
self.output_top_logprobs_val = torch.zeros(
(size, max_top_logprobs_num), dtype=torch.float32, device=device
)
self.output_top_logprobs_idx = torch.zeros(
(size, max_top_logprobs_num), dtype=torch.int32, device=device
)
self.output_token_sampling_mask_len = None
self.output_token_sampling_mask_idx = None
self.output_token_sampling_logprobs = None
if self.enable_sampling_mask:
self.output_token_sampling_mask_len = torch.zeros(
(size, 16), dtype=torch.int32, device=device
)
self.output_token_sampling_mask_idx = torch.zeros(
(size, max_sampling_mask_tokens), dtype=torch.int32, device=device
)
self.output_token_sampling_logprobs = torch.zeros(
(size, 16), dtype=torch.float32, device=device
)
# For PD + spec decode
self.output_topk_p = torch.zeros(
(size, 16), dtype=torch.float32, device=device
)
self.output_topk_index = torch.zeros(
(size, 16), dtype=torch.int64, device=device
)
self.output_hidden_states = torch.zeros(
(size, hidden_size), dtype=hidden_states_dtype, device=device
)
if self.output_dsa_topk_indices_dim > 0:
self.output_dsa_topk_indices = torch.full(
(size, self.output_dsa_topk_indices_dim),
-1,
dtype=torch.int32,
device=device,
)
else:
self.output_dsa_topk_indices = None
# Request validation: store bootstrap_room to detect metadata corruption
self.bootstrap_room = torch.zeros(
(size, 8), dtype=bootstrap_room_dtype, device=device
)
self.kv_checksum: torch.Tensor | None = None
if kv_checksum_enabled:
with (
torch.cuda.use_mem_pool(self.custom_mem_pool)
if self.custom_mem_pool
else nullcontext()
):
# Width 8 (uint64) keeps the per-row size at the 64 B RDMA minimum.
self.kv_checksum = torch.zeros(
(self.output_ids.shape[0], 8),
dtype=self.bootstrap_room.dtype,
device=self.bootstrap_room.device,
)
def set_kv_checksum(self, req: Req, value: int) -> None:
self.kv_checksum[req.metadata_buffer_index, 0] = value
def get_kv_checksum(self, idx: int) -> int:
return int(self.kv_checksum[idx, 0].item())
def get_buf_infos(self):
bufs = [
self.output_ids,
self.cached_tokens,
self.output_token_logprobs_val,
self.output_token_logprobs_idx,
self.output_top_logprobs_val,
self.output_top_logprobs_idx,
self.output_token_sampling_mask_len,
self.output_token_sampling_mask_idx,
self.output_token_sampling_logprobs,
self.output_topk_p,
self.output_topk_index,
self.output_hidden_states,
]
if self.output_dsa_topk_indices is not None:
bufs.append(self.output_dsa_topk_indices)
bufs.append(self.bootstrap_room)
if self.kv_checksum is not None:
bufs.append(self.kv_checksum)
bufs = [buf for buf in bufs if buf is not None]
ptrs = [buf.data_ptr() for buf in bufs]
data_lens = [buf.nbytes for buf in bufs]
item_lens = [buf[0].nbytes for buf in bufs]
return ptrs, data_lens, item_lens
def get_buf(self, idx: int):
return (
self.output_ids[idx].clone(),
self.cached_tokens[idx].clone(),
self.output_token_logprobs_val[idx].clone(),
self.output_token_logprobs_idx[idx].clone(),
self.output_top_logprobs_val[idx].clone(),
self.output_top_logprobs_idx[idx].clone(),
(
self.output_token_sampling_mask_len[idx].clone()
if self.enable_sampling_mask
else None
),
(
self.output_token_sampling_mask_idx[idx].clone()
if self.enable_sampling_mask
else None
),
(
self.output_token_sampling_logprobs[idx].clone()
if self.enable_sampling_mask
else None
),
self.output_topk_p[idx].clone(),
self.output_topk_index[idx].clone(),
self.output_hidden_states[idx].clone(),
(
self.output_dsa_topk_indices[idx].clone()
if self.output_dsa_topk_indices is not None
else None
),
self.bootstrap_room[idx].clone(),
)
def set_buf(self, req: Req):
self.output_ids[req.metadata_buffer_index][0] = req.output_ids[0]
# The cached_tokens buffer is (size, 16); slots 0-3 hold cached token
# counts and slots 4-6 are reused for multimodal prompt token counts
# (slots 7-15 remain spare). This avoids adding new RDMA buffers.
# Slot map: 0=cached 1=device 2=host 3=storage 4=image 5=audio 6=video.
self.cached_tokens[req.metadata_buffer_index][0] = req.cached_tokens
self.cached_tokens[req.metadata_buffer_index][1] = req.cached_tokens_device
self.cached_tokens[req.metadata_buffer_index][2] = req.cached_tokens_host
self.cached_tokens[req.metadata_buffer_index][3] = req.cached_tokens_storage
# Compute multimodal prompt token counts on the prefill node so decode
# can report them in usage.
if req.multimodal_inputs:
image_t, audio_t, video_t = req.multimodal_inputs.compute_mm_token_counts()
else:
image_t = audio_t = video_t = 0
self.cached_tokens[req.metadata_buffer_index][4] = image_t
self.cached_tokens[req.metadata_buffer_index][5] = audio_t
self.cached_tokens[req.metadata_buffer_index][6] = video_t
if req.return_logprob:
if req.logprob.output_token_logprobs_val: # not none or empty list
self.output_token_logprobs_val[req.metadata_buffer_index][0] = (
req.logprob.output_token_logprobs_val[0]
)
if req.logprob.output_token_logprobs_idx: # not none or empty list
self.output_token_logprobs_idx[req.metadata_buffer_index][0] = (
req.logprob.output_token_logprobs_idx[0]
)
if req.logprob.output_top_logprobs_val: # not none or empty list
top_logprobs_len = len(req.logprob.output_top_logprobs_val[0])
max_top_logprobs_len = self.output_top_logprobs_val.shape[1]
if top_logprobs_len > max_top_logprobs_len:
raise RuntimeError(
f"top_logprobs_num {top_logprobs_len} exceeds "
f"disaggregation metadata capacity {max_top_logprobs_len}. "
"Lower top_logprobs_num or increase the metadata buffer."
)
self.output_top_logprobs_val[req.metadata_buffer_index][
: len(req.logprob.output_top_logprobs_val[0])
] = torch.tensor(
req.logprob.output_top_logprobs_val[0],
dtype=torch.float32,
device="cpu",
)
if req.logprob.output_top_logprobs_idx: # not none or empty list
self.output_top_logprobs_idx[req.metadata_buffer_index][
: len(req.logprob.output_top_logprobs_idx[0])
] = torch.tensor(
req.logprob.output_top_logprobs_idx[0],
dtype=torch.int32,
device="cpu",
)
if req.return_sampling_mask:
# Sentinel -1: the decode side records None for this handoff token.
self.output_token_sampling_mask_len[req.metadata_buffer_index][0] = -1
sampling_masks = req.output_token_sampling_mask
sampling_logprobs = req.output_token_sampling_logprobs
if sampling_masks:
sampling_mask = sampling_masks[0]
sampling_logprob = sampling_logprobs[0] if sampling_logprobs else None
if sampling_mask is not None and sampling_logprob is not None:
mask_len = len(sampling_mask)
max_mask_len = self.output_token_sampling_mask_idx.shape[1]
if mask_len > max_mask_len:
raise RuntimeError(
f"Sampling mask length {mask_len} exceeds disaggregation "
f"metadata capacity {max_mask_len}. Increase "
"--sampling-mask-max-tokens."
)
self.output_token_sampling_mask_len[req.metadata_buffer_index][
0
] = mask_len
if mask_len:
self.output_token_sampling_mask_idx[
req.metadata_buffer_index, :mask_len
].copy_(
torch.tensor(
sampling_mask,
dtype=torch.int32,
device=self.output_token_sampling_mask_idx.device,
)
)
self.output_token_sampling_logprobs[req.metadata_buffer_index][
0
] = float(sampling_logprob)
# For PD + spec decode
if req.hidden_states_tensor is not None:
# speculative_eagle_topk should not be greater than 16 currently
topk = req.output_topk_p.size(0)
self.output_topk_p[req.metadata_buffer_index, :topk].copy_(
req.output_topk_p
)
self.output_topk_index[req.metadata_buffer_index, :topk].copy_(
req.output_topk_index
)
self.output_hidden_states[req.metadata_buffer_index].copy_(
req.hidden_states_tensor
)
if self.output_dsa_topk_indices is not None:
dsa_topk_indices = req.output_dsa_topk_indices
if dsa_topk_indices is not None:
self.output_dsa_topk_indices[req.metadata_buffer_index].copy_(
dsa_topk_indices
)
else:
self.output_dsa_topk_indices[req.metadata_buffer_index].fill_(-1)
# Store bootstrap_room for validation on decode side
self.bootstrap_room[req.metadata_buffer_index, 0] = (
req.bootstrap_room if req.bootstrap_room is not None else 0
)
#########################
# Transfer Backend
#########################
class TransferBackend(Enum):
MOONCAKE = "mooncake"
MORI = "mori"
NIXL = "nixl"
ASCEND = "ascend"
FAKE = "fake"
class KVClassType(Enum):
KVARGS = "kvargs"
MANAGER = "manager"
SENDER = "sender"
RECEIVER = "receiver"
BOOTSTRAP_SERVER = "bootstrap_server"
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.KVARGS]
) -> Type[KVArgs]: ...
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.MANAGER]
) -> Type[CommonKVManager]: ...
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.SENDER]
) -> Type[CommonKVSender]: ...
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.RECEIVER]
) -> Type[CommonKVReceiver]: ...
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.BOOTSTRAP_SERVER]
) -> Type[CommonKVBootstrapServer]: ...
def get_kv_class(
transfer_backend: TransferBackend, class_type: KVClassType
) -> Optional[Type]:
from sglang.srt.disaggregation.base import KVArgs
# Every backend shares the same KVArgs container.
if class_type == KVClassType.KVARGS:
return KVArgs
if transfer_backend == TransferBackend.MOONCAKE:
from sglang.srt.disaggregation.mooncake import (
MooncakeKVBootstrapServer,
MooncakeKVManager,
MooncakeKVReceiver,
MooncakeKVSender,
)
class_mapping = {
KVClassType.MANAGER: MooncakeKVManager,
KVClassType.SENDER: MooncakeKVSender,
KVClassType.RECEIVER: MooncakeKVReceiver,
KVClassType.BOOTSTRAP_SERVER: MooncakeKVBootstrapServer,
}
elif transfer_backend == TransferBackend.MORI:
from sglang.srt.disaggregation.mori import (
MoriKVBootstrapServer,
MoriKVManager,
MoriKVReceiver,
MoriKVSender,
)
class_mapping = {
KVClassType.MANAGER: MoriKVManager,
KVClassType.SENDER: MoriKVSender,
KVClassType.RECEIVER: MoriKVReceiver,
KVClassType.BOOTSTRAP_SERVER: MoriKVBootstrapServer,
}
elif transfer_backend == TransferBackend.ASCEND:
from sglang.srt.disaggregation.ascend import (
AscendKVBootstrapServer,
AscendKVManager,
AscendKVReceiver,
AscendKVSender,
)
class_mapping = {
KVClassType.MANAGER: AscendKVManager,
KVClassType.SENDER: AscendKVSender,
KVClassType.RECEIVER: AscendKVReceiver,
KVClassType.BOOTSTRAP_SERVER: AscendKVBootstrapServer,
}
elif transfer_backend == TransferBackend.NIXL:
from sglang.srt.disaggregation.nixl import (
NixlKVBootstrapServer,
NixlKVManager,
NixlKVReceiver,
NixlKVSender,
)
class_mapping = {
KVClassType.MANAGER: NixlKVManager,
KVClassType.SENDER: NixlKVSender,
KVClassType.RECEIVER: NixlKVReceiver,
KVClassType.BOOTSTRAP_SERVER: NixlKVBootstrapServer,
}
elif transfer_backend == TransferBackend.FAKE:
from sglang.srt.disaggregation.fake import (
FakeKVManager,
FakeKVReceiver,
FakeKVSender,
)
# No bootstrap server: the fake backend never registers one.
class_mapping = {
KVClassType.MANAGER: FakeKVManager,
KVClassType.SENDER: FakeKVSender,
KVClassType.RECEIVER: FakeKVReceiver,
}
else:
raise ValueError(f"Unsupported transfer backend: {transfer_backend}")
return class_mapping.get(class_type)
def _get_cp_rank_page_bounds(
total_pages: int, cp_rank: int, cp_size: int
) -> Tuple[int, int]:
base = total_pages // cp_size
rem = total_pages % cp_size
local_start = cp_rank * base + min(cp_rank, rem)
n_pages = base + (1 if cp_rank < rem else 0)
return local_start, local_start + n_pages
def filter_kv_indices_for_cp_rank(
kv_mgr: CommonKVManager,
kv_indices: np.ndarray,
index_slice: slice,
total_pages: Optional[int] = None,
) -> Tuple[np.ndarray, slice]:
"""Filters kv_indices and index_slice for the current CP rank."""
if total_pages is None:
total_pages = len(kv_indices)
cp_rank = kv_mgr.attn_cp_rank
cp_size = kv_mgr.attn_cp_size
if cp_size <= 1:
return kv_indices, index_slice
rank_start, rank_end = _get_cp_rank_page_bounds(total_pages, cp_rank, cp_size)
chunk_start = index_slice.start if index_slice.start is not None else 0
chunk_end = index_slice.stop if index_slice.stop is not None else total_pages
first_pos = max(rank_start, chunk_start) - chunk_start
last_pos = min(rank_end, chunk_end) - chunk_start
if last_pos <= first_pos:
new_kv_indices = kv_indices[:0]
new_index_slice = slice(chunk_start, chunk_start)
else:
new_kv_indices = kv_indices[first_pos:last_pos]
new_index_slice = slice(
chunk_start + first_pos,
chunk_start + last_pos,
)
return new_kv_indices, new_index_slice
#########################
# Misc
#########################
def is_mla_backend(target_kv_pool) -> bool:
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
return isinstance(target_kv_pool, (MLATokenToKVPool, DeepSeekV4TokenToKVPool))
def should_send_replicated_state(
*,
src_attn_tp_size: int,
dst_attn_tp_size: int,
local_tp_rank_in_group: int,
) -> bool:
"""Elect writers for state replicated within an attention-TP group.
Scatter (one source rank to several destination ranks) is a broadcast, so
the source sends to every destination registration. Aggregation has several
equivalent source copies targeting one destination; only the first source
in each aggregation group writes it.
"""
if src_attn_tp_size <= 0 or dst_attn_tp_size <= 0:
raise ValueError(
"Attention TP sizes must be positive for replicated-state transfer"
)
larger_tp_size = max(src_attn_tp_size, dst_attn_tp_size)
smaller_tp_size = min(src_attn_tp_size, dst_attn_tp_size)
if larger_tp_size % smaller_tp_size != 0:
raise ValueError(
"One attention TP size must divide the other for replicated-state "
f"transfer: src={src_attn_tp_size}, dst={dst_attn_tp_size}"
)
if not 0 <= local_tp_rank_in_group < src_attn_tp_size:
raise ValueError(
"Source attention TP rank is out of range for replicated-state "
f"transfer: rank={local_tp_rank_in_group}, size={src_attn_tp_size}"
)
if src_attn_tp_size <= dst_attn_tp_size:
return True
writers_per_decode = src_attn_tp_size // dst_attn_tp_size
return local_tp_rank_in_group % writers_per_decode == 0
def compute_mamba_state_slice_blocks(
src_dim: int,
dst_dim: int,
src_attn_tp_size: int,
dst_attn_tp_size: int,
dst_tp_rank_in_group: int,
local_tp_rank_in_group: int,
conv_shard_groups: Optional[List[int]] = None,
) -> List[Tuple[int, int, int]]:
"""Blocks to copy one mamba state item across differing attn-TP sizes.
Returns ``(src_dim_start, dst_dim_start, num_dims)`` triples in units of the
sliceable (3rd) dimension. Single-axis states (temporal_state, or when
``conv_shard_groups`` is None) return one contiguous block -- byte-identical to
the legacy behavior.
GDN conv_state is ``cat([query | key | value])`` where each sub-block (full
dims == ``conv_shard_groups``, e.g. ``[key_dim, key_dim, value_dim]``) is
head-sharded INDEPENDENTLY across attn-TP. In the SCATTER direction
(1 prefill rank -> several decode ranks) a single contiguous slice straddles
the q/k/v boundaries and delivers wrong channels. The AGGREGATION direction
(several prefill ranks -> 1 decode rank) has the symmetric problem: a single
contiguous write interleaves the sub-blocks by writer. Both directions emit one
block per sub-block for conv_state; temporal_state and non-GDN states (when
``conv_shard_groups`` is None) keep the single contiguous slice.
"""
use_subdims = (
conv_shard_groups is not None
and sum(conv_shard_groups) == src_dim * src_attn_tp_size
)
if src_attn_tp_size > dst_attn_tp_size:
# Aggregation: several prefill ranks each write their shard into one decode slot.
writers_per_decode = src_attn_tp_size // dst_attn_tp_size
local_writer_idx = local_tp_rank_in_group % writers_per_decode
if not use_subdims:
return [(0, local_writer_idx * src_dim, src_dim)]
# conv_state: a plain contiguous write would interleave the sub-blocks by
# writer ([q0,k0,v0,q1,k1,v1,...]); place this writer's shard of each
# independently head-sharded sub-block at its grouped offset so the decode
# buffer is [q0,q1,...,k0,k1,...,v0,v1,...].
blocks: List[Tuple[int, int, int]] = []
src_off = 0
dst_off = 0
for full_sd in conv_shard_groups:
src_sub = full_sd // src_attn_tp_size
dst_sub = full_sd // dst_attn_tp_size
blocks.append((src_off, dst_off + local_writer_idx * src_sub, src_sub))
src_off += src_sub
dst_off += dst_sub
return blocks
# Scatter: 1 prefill rank feeds several decode ranks.
if not use_subdims:
src_dim_start = (dst_tp_rank_in_group * dst_dim) % src_dim
return [(src_dim_start, 0, dst_dim)]
# conv_state: gather the decode rank's [q | k | v] shard from the three
# independently head-sharded sub-blocks of the src tensor. dst is contiguous.
blocks: List[Tuple[int, int, int]] = []
src_off = 0
dst_off = 0
for full_sd in conv_shard_groups:
src_sub = full_sd // src_attn_tp_size # this prefill rank's shard of sub-block
dst_sub = full_sd // dst_attn_tp_size # this decode rank's shard of sub-block
src_start = src_off + (dst_tp_rank_in_group * dst_sub) % src_sub
blocks.append((src_start, dst_off, dst_sub))
src_off += src_sub
dst_off += dst_sub
return blocks
def compute_mamba_state_slice_byte_blocks(
*,
src_item_len: int,
dst_item_len: int,
src_dim: int,
dst_dim: int,
outer_count: int,
src_attn_tp_size: int,
dst_attn_tp_size: int,
dst_tp_rank_in_group: int,
local_tp_rank_in_group: int,
conv_shard_groups: Optional[List[int]] = None,
) -> List[Tuple[int, int, int]]:
"""Convert logical TP slices into physical byte blocks for one state slot.
``outer_count`` is one for the usual ``[slice_dim, ...]`` layout. Kimi
conv state is ``[K - 1, slice_dim]``, so each logical channel slice expands
into one byte block per convolution row. A zero src/dst dim marks an item
replicated across attention TP and copies the whole item from an elected
source rank.
"""
if (src_dim == 0) != (dst_dim == 0):
raise ValueError(
"Mamba state replication metadata differs between prefill and decode"
)
if src_dim == 0:
if src_item_len != dst_item_len:
raise ValueError(
"Replicated Mamba state item lengths differ between prefill and "
f"decode: {src_item_len} != {dst_item_len}"
)
if not should_send_replicated_state(
src_attn_tp_size=src_attn_tp_size,
dst_attn_tp_size=dst_attn_tp_size,
local_tp_rank_in_group=local_tp_rank_in_group,
):
return []
return [(0, 0, src_item_len)]
src_bytes_per_dim = src_item_len // (src_dim * outer_count)
dst_bytes_per_dim = dst_item_len // (dst_dim * outer_count)
logical_blocks = compute_mamba_state_slice_blocks(
src_dim=src_dim,
dst_dim=dst_dim,
src_attn_tp_size=src_attn_tp_size,
dst_attn_tp_size=dst_attn_tp_size,
dst_tp_rank_in_group=dst_tp_rank_in_group,
local_tp_rank_in_group=local_tp_rank_in_group,
conv_shard_groups=conv_shard_groups,
)
blocks = []
for outer_idx in range(outer_count):
src_row_offset = outer_idx * src_dim * src_bytes_per_dim
dst_row_offset = outer_idx * dst_dim * dst_bytes_per_dim
for src_dim_start, dst_dim_start, num_dims in logical_blocks:
blocks.append(
(
src_row_offset + src_dim_start * src_bytes_per_dim,
dst_row_offset + dst_dim_start * dst_bytes_per_dim,
num_dims * src_bytes_per_dim,
)
)
return blocks
def build_transfer_entry_pairs(
src_layer_ids: List[int],
dst_layer_ids: List[int],
n_src: int,
n_dst: int,
allow_positional_fallback: bool = False,
) -> List[Tuple[int, int]]:
"""Pair prefill-local transfer entries with decode entries by layer id."""
if n_src == 0:
return []
if bool(src_layer_ids) != bool(dst_layer_ids):
if not allow_positional_fallback:
raise RuntimeError(
"Layer metadata must be provided by both PD peers or neither"
)
src_layer_ids = []
dst_layer_ids = []
if src_layer_ids:
if len(src_layer_ids) != n_src or len(dst_layer_ids) != n_dst:
raise RuntimeError(
"Layer metadata length must match transfer entries: "
f"src metadata={len(src_layer_ids)} entries={n_src}, "
f"dst metadata={len(dst_layer_ids)} entries={n_dst}"
)
# Layer ids can repeat across tensor groups (for example K/V or multiple
# state tensors), so pair occurrences in order rather than by plain lookup.
dst_pos = {}
for j, lid in enumerate(dst_layer_ids):
dst_pos.setdefault(lid, deque()).append(j)
pairs = []
for i, lid in enumerate(src_layer_ids):
if not dst_pos.get(lid):
raise RuntimeError(
f"Decode peer is missing a transfer entry for model layer {lid}"
)
pairs.append((i, dst_pos[lid].popleft()))
return pairs
if n_dst < n_src or (n_src != n_dst and not allow_positional_fallback):
# Without layer ids a positional pairing would silently transfer the
# wrong layers (e.g. PP prefill peered with a stale decode server).
raise RuntimeError(
"PP-heterogeneous transfer requires layer ids on "
f"both peers; got src={n_src} dst={n_dst} entries"
)
return [(i, i) for i in range(n_src)]
def build_kv_layer_ids(
*,
token_to_kv_pool,
draft_token_to_kv_pool,
num_draft_entries: int,
num_hidden_layers: int,
) -> List[int]:
"""Global layer id for every entry in ``kv_args.kv_data_ptrs``.
Draft KV buffers are appended after the target's, so they need ids of their
own: build_transfer_entry_pairs requires the id list to cover every entry,
and a target-only list would be rejected. The draft numbers its layers from
zero, which would collide with the target's, so its entries are remapped
into a reserved band above the target's layer range. Both PD peers run this
against the same draft config and so agree on the band.
Returns [] for pools that cannot report ids, leaving the peers on positional
pairing.
"""
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
if not isinstance(token_to_kv_pool, HybridLinearKVPool):
return []
layer_ids = token_to_kv_pool.get_kv_layer_ids()
if draft_token_to_kv_pool is None:
return layer_ids
draft_ids = _draft_entry_layer_ids(
pool=draft_token_to_kv_pool, num_entries=num_draft_entries
)
# Rank the draft's own ids by first appearance, so the band stays dense and
# contiguous whatever the draft config numbers its layers.
band_index = {lid: i for i, lid in enumerate(dict.fromkeys(draft_ids))}
return layer_ids + [num_hidden_layers + band_index[lid] for lid in draft_ids]
def _draft_entry_layer_ids(*, pool, num_entries: int) -> List[int]:
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
if isinstance(pool, HybridLinearKVPool):
ids = pool.get_kv_layer_ids()
else:
# Pools register k0..k(L-1) then v0..v(L-1), so ids repeat once per
# group; derive the group count rather than assuming MHA vs MLA.
if pool.layer_num <= 0 or num_entries % pool.layer_num != 0:
raise RuntimeError(
"Draft KV buffers must register a whole number of per-layer "
f"groups: entries={num_entries}, layers={pool.layer_num}"
)
ids = list(range(pool.layer_num)) * (num_entries // pool.layer_num)
if len(ids) != num_entries:
raise RuntimeError(
"Draft KV layer ids must cover every registered entry: "
f"ids={len(ids)}, entries={num_entries}"
)
return ids
def resolve_dcp_dst_entry_indices(
src_layer_ids: List[int],
dst_layer_ids: List[int],
n_src: int,
n_dst: int,
) -> List[int]:
"""Destination entry index for each local KV entry, for a DCP relayout.
DCP re-splits the KV by context while PP re-splits it by layer, so the two
index spaces only line up when neither peer is pipelined. Both backends
need the same resolution, hence the shared helper.
"""
if not src_layer_ids and not dst_layer_ids:
# Legacy/non-PP layout. n_dst may exceed n_src when the decode side
# runs speculative decoding and the prefill side does not.
return list(range(n_src))
# A one-sided mapping is rejected by build_transfer_entry_pairs itself.
return [
j
for _, j in build_transfer_entry_pairs(
src_layer_ids, dst_layer_ids, n_src, n_dst
)
]
def build_staging_slot_metadata(
*,
kv_layer_ids: List[int],
num_draft_entries: int,
kv_pool,
draft_kv_pool,
):
"""Buffers and per-slot layer ids for the staging gather.
The gather writes every k_buffer and then every v_buffer, while kv_layer_ids
follows kv_data_ptrs ([K target, V target, K draft, V draft]), so the two
orders diverge as soon as a draft pool is registered.
Returns (k_buffers, v_buffers, slot_layer_ids), or None for a pool that has
no contiguous K/V tensors to stage.
"""
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, MHATokenToKVPool
# A hybrid pool keeps its contiguous K/V tensors on the inner full-attention
# pool, and the draft pool is wrapped the same way.
if isinstance(kv_pool, HybridLinearKVPool):
kv_pool = kv_pool.full_kv_pool
if isinstance(draft_kv_pool, HybridLinearKVPool):
draft_kv_pool = draft_kv_pool.full_kv_pool
if not isinstance(kv_pool, MHATokenToKVPool):
return None
ids = list(kv_layer_ids or [])
num_target = len(ids) - num_draft_entries
half = num_target // 2
k_buffers, k_ids = list(kv_pool.k_buffer), ids[:half]
v_buffers, v_ids = list(kv_pool.v_buffer), ids[half:num_target]
draft_half = num_draft_entries // 2
if draft_half:
if not isinstance(draft_kv_pool, MHATokenToKVPool):
# An empty id list puts the sender back on kv_data_ptrs order, which
# is what staging did before draft KV existed.
return k_buffers, v_buffers, []
k_buffers += list(draft_kv_pool.k_buffer)
v_buffers += list(draft_kv_pool.v_buffer)
k_ids += ids[num_target : num_target + draft_half]
v_ids += ids[num_target + draft_half :]
return k_buffers, v_buffers, k_ids + v_ids
def append_state_component(
kv_args: KVArgs,
state_type: StateType,
data_ptrs: List[int],
data_lens: List[int],
item_lens: List[int],
dim_per_tensor: Optional[List[int]] = None,
conv_shard_groups: Optional[List[Optional[List[int]]]] = None,
slice_outer_counts: Optional[List[int]] = None,
layer_ids: Optional[List[int]] = None,
) -> None:
"""Append one state component. Caller orders state_types consistently
on prefill and decode sides."""
kv_args.state_types.append(state_type)
kv_args.state_data_ptrs.append(data_ptrs)
kv_args.state_data_lens.append(data_lens)
kv_args.state_item_lens.append(item_lens)
kv_args.state_dim_per_tensor.append(dim_per_tensor or [])
kv_args.state_conv_shard_groups.append(conv_shard_groups or [])
kv_args.state_slice_outer_counts.append(slice_outer_counts or [])
kv_args.state_layer_ids.append(layer_ids or [])
def get_dsa_tail_state_indices(pool, req_pool_idx: int, seq_len: int) -> List[int]:
if getattr(pool, "use_dsa", False):
pool = pool.full_kv_pool
if not pool.kpool_use_compress:
return []
pool_size = int(pool.index_kpool)
tail_size = pool_size + int(getattr(pool, "tail_extra_slots", 0))
if pool_size <= 1 or tail_size < pool_size:
raise ValueError(
"DSA kpool-compress requires pool_size > 1 and "
f"tail_size >= pool_size, got pool_size={pool_size}, "
f"tail_size={tail_size}"
)
n_valid = int(seq_len) % pool_size
if n_valid == 0:
return []
start_phys = (int(seq_len) - n_valid) % tail_size
first_n = min(n_valid, tail_size - start_phys)
second_n = n_valid - first_n
return [
int(req_pool_idx),
start_phys,
first_n,
0,
second_n,
tail_size,
]
def slice_dsa_tail_dst_ptrs_for_pp(
src_ptrs: List[int],
dst_ptrs: List[int],
start_layer: int,
end_layer: Optional[int],
) -> List[int]:
if len(src_ptrs) == len(dst_ptrs):
return list(dst_ptrs)
if len(src_ptrs) % 2 != 0 or len(dst_ptrs) % 2 != 0:
raise ValueError(
"DSA tail pointer lists must contain equal key/score halves, got "
f"src={len(src_ptrs)}, dst={len(dst_ptrs)}"
)
src_layers = len(src_ptrs) // 2
dst_layers = len(dst_ptrs) // 2
expected_end = start_layer + src_layers
if end_layer is not None and end_layer - start_layer == src_layers:
expected_end = end_layer
if start_layer < 0 or expected_end > dst_layers:
raise ValueError(
"DSA tail pointer count mismatch: "
f"src={len(src_ptrs)}, dst={len(dst_ptrs)}, "
f"prefill_layers=[{start_layer}, {expected_end})"
)
return list(dst_ptrs[start_layer:expected_end]) + list(
dst_ptrs[dst_layers + start_layer : dst_layers + expected_end]
)
def build_dsa_tail_transfer_blocks(
src_ptrs: List[int],
src_item_lens: List[int],
dst_ptrs: List[int],
src_indices: List[int],
dst_indices: List[int],
dst_item_lens: Optional[List[int]] = None,
) -> List[Tuple[int, int, int]]:
"""Remap live DSA tail tokens between rings with different speculative-slot counts."""
if not src_indices and not dst_indices:
return []
if not src_indices or not dst_indices:
raise ValueError(
f"DSA tail slot index missing: src={src_indices}, dst={dst_indices}"
)
if len(src_indices) != 6 or len(dst_indices) != 6:
raise ValueError(
"DSA tail slot indices must be 6-tuples, "
f"got src={src_indices}, dst={dst_indices}"
)
if dst_item_lens is None:
dst_item_lens = src_item_lens
if not (len(src_ptrs) == len(dst_ptrs) == len(src_item_lens) == len(dst_item_lens)):
raise ValueError(
"DSA tail pointer metadata mismatch: "
f"src_ptrs={len(src_ptrs)}, dst_ptrs={len(dst_ptrs)}, "
f"src_item_lens={len(src_item_lens)}, "
f"dst_item_lens={len(dst_item_lens)}"
)
src_tail_size = int(src_indices[5])
dst_tail_size = int(dst_indices[5])
if src_tail_size <= 0 or dst_tail_size <= 0:
raise ValueError(
"DSA tail ring sizes must be positive: "
f"src={src_tail_size}, dst={dst_tail_size}"
)
def parse_segments(indices: List[int], tail_size: int, side: str):
segments = []
for seg in (1, 2):
off = int(indices[seg * 2 - 1])
n = int(indices[seg * 2])
if min(off, n) < 0:
raise ValueError(
f"DSA tail {side} offsets and lengths must be non-negative"
)
if off + n > tail_size:
raise ValueError(
f"DSA tail {side} segment {seg} exceeds ring size "
f"{tail_size}: ({off}, {n})"
)
if n:
segments.append((off, n))
return segments
src_segments = parse_segments(src_indices, src_tail_size, "source")
dst_segments = parse_segments(dst_indices, dst_tail_size, "destination")
src_count = sum(n for _, n in src_segments)
dst_count = sum(n for _, n in dst_segments)
if src_count != dst_count:
raise ValueError(
f"DSA tail live-token count mismatch: src={src_count}, dst={dst_count}"
)
src_idx = int(src_indices[0])
dst_idx = int(dst_indices[0])
if src_idx < 0 or dst_idx < 0:
raise ValueError("DSA tail request row indices must be non-negative")
transfer_blocks = []
for src_ptr, src_row_bytes, dst_ptr, dst_row_bytes in zip(
src_ptrs, src_item_lens, dst_ptrs, dst_item_lens
):
src_row_bytes = int(src_row_bytes)
dst_row_bytes = int(dst_row_bytes)
if src_row_bytes == 0 and dst_row_bytes == 0:
continue
if src_row_bytes <= 0 or src_row_bytes % src_tail_size != 0:
raise ValueError(
f"DSA source tail row size {src_row_bytes} is not divisible by "
f"{src_tail_size}"
)
if dst_row_bytes <= 0 or dst_row_bytes % dst_tail_size != 0:
raise ValueError(
f"DSA destination tail row size {dst_row_bytes} is not "
f"divisible by {dst_tail_size}"
)
src_slot_bytes = src_row_bytes // src_tail_size
dst_slot_bytes = dst_row_bytes // dst_tail_size
if src_slot_bytes != dst_slot_bytes:
raise ValueError(
"DSA tail slot-size mismatch: "
f"src={src_slot_bytes}, dst={dst_slot_bytes}"
)
slot_bytes = src_slot_bytes
src_row_base = int(src_ptr) + src_row_bytes * src_idx
dst_row_base = int(dst_ptr) + dst_row_bytes * dst_idx
src_seg_idx = dst_seg_idx = 0
src_consumed = dst_consumed = 0
while src_seg_idx < len(src_segments):
src_off, src_n = src_segments[src_seg_idx]
dst_off, dst_n = dst_segments[dst_seg_idx]
n = min(src_n - src_consumed, dst_n - dst_consumed)
transfer_blocks.append(
(
src_row_base + (src_off + src_consumed) * slot_bytes,
dst_row_base + (dst_off + dst_consumed) * slot_bytes,
n * slot_bytes,
)
)
src_consumed += n
dst_consumed += n
if src_consumed == src_n:
src_seg_idx += 1
src_consumed = 0
if dst_consumed == dst_n:
dst_seg_idx += 1
dst_consumed = 0
return transfer_blocks
def setup_state_kv_args(
kv_args: KVArgs,
token_to_kv_pool,
draft_token_to_kv_pool=None,
total_kv_layers: int = None,
req_to_token_pool=None,
) -> None:
from sglang.srt.disaggregation.base.conn import StateType
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMLATokenToKVPool
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.mem_cache.memory_pool import (
DSATokenToKVPool,
HybridLinearKVPool,
MHATokenToKVPoolMXFP8,
MiniMaxSparseKVPool,
)
from sglang.srt.mem_cache.qsa_kv_pool import QSATokenToKVPool
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
kv_args.state_types = []
kv_args.state_data_ptrs = []
kv_args.state_data_lens = []
kv_args.state_item_lens = []
kv_args.state_dim_per_tensor = []
kv_args.state_slice_outer_counts = []
kv_args.state_layer_ids = []
kv_args.is_hybrid_mla_backend = False
kv_args.state_conv_shard_groups = []
# V4's KVCache is organized by compression-ratio buckets rather than by layer.
kv_args.mla_compression_ratios = (
list(token_to_kv_pool.compression_ratios)
if isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
else None
)
def append_dsa_tail(pool) -> None:
if not pool.kpool_use_compress:
return
tail_ptrs, tail_lens, tail_item_lens = pool.get_compress_tail_buf_infos()
if tail_ptrs:
append_state_component(
kv_args,
StateType.DSA_TAIL,
tail_ptrs,
tail_lens,
tail_item_lens,
)
if isinstance(token_to_kv_pool, MHATokenToKVPoolMXFP8):
append_state_component(
kv_args,
StateType.BLOCK_SCALE,
*token_to_kv_pool.get_kv_scale_buf_infos(),
)
if isinstance(token_to_kv_pool, MiniMaxSparseKVPool):
if token_to_kv_pool.index_kv_pool is not None:
raise NotImplementedError(
"PD disaggregation for MiniMax sparse layers with index value "
"(index_kv_pool) is not yet supported; only K-only sparse layers are."
)
if token_to_kv_pool.index_k_pool is not None:
dp, dl, il = token_to_kv_pool.get_index_k_state_buf_infos()
append_state_component(kv_args, StateType.MINIMAX_INDEX_K, dp, dl, il)
elif hasattr(token_to_kv_pool, "get_state_buf_infos"):
data_ptrs, data_lens, item_lens = token_to_kv_pool.get_state_buf_infos()
# DeepSeekV4TokenToKVPool inherits BaseSWAKVPool; its heterogeneous
# state list is described per-entry via get_state_buf_infos.
if isinstance(token_to_kv_pool, BaseSWAKVPool):
append_state_component(
kv_args, StateType.SWA, data_ptrs, data_lens, item_lens
)
# MXFP8 KV: each sub-pool's block scales ride as their own component
# so they inherit the index payload of the KV they describe.
# Only the concrete SWAKVPool owns a full sub-pool; other
# BaseSWAKVPool implementations describe their state per entry.
if isinstance(token_to_kv_pool, SWAKVPool) and isinstance(
token_to_kv_pool.full_kv_pool, MHATokenToKVPoolMXFP8
):
append_state_component(
kv_args,
StateType.BLOCK_SCALE,
*token_to_kv_pool.get_kv_scale_buf_infos(),
)
append_state_component(
kv_args,
StateType.BLOCK_SCALE_SWA,
*token_to_kv_pool.get_swa_kv_scale_buf_infos(),
)
# unified_kv: the SWA ring lives in the unified buffers (no separate
# swa_kv_pool) and is addressed per-row, so ship it as SWA_RING.
if getattr(token_to_kv_pool, "_unified_kv", False) and hasattr(
token_to_kv_pool, "get_unified_swa_ring_buf_infos"
):
ring_ptrs, ring_lens, ring_item_lens = (
token_to_kv_pool.get_unified_swa_ring_buf_infos()
)
if ring_ptrs:
append_state_component(
kv_args,
StateType.SWA_RING,
ring_ptrs,
ring_lens,
ring_item_lens,
)
if hasattr(token_to_kv_pool, "get_request_state_buf_infos"):
c128_ptrs, c128_lens, c128_item_lens = (
token_to_kv_pool.get_request_state_buf_infos()
)
if c128_ptrs:
append_state_component(
kv_args,
StateType.DSV4_REQUEST_STATE,
c128_ptrs,
c128_lens,
c128_item_lens,
)
elif isinstance(token_to_kv_pool, HybridLinearKVPool):
dim = (
token_to_kv_pool.get_state_dim_per_tensor()
if hasattr(token_to_kv_pool, "get_state_dim_per_tensor")
else None
)
kv_args.is_hybrid_mla_backend = is_mla_backend(
token_to_kv_pool.full_kv_pool
)
conv_shard_groups = (
token_to_kv_pool.get_state_conv_shard_groups()
if hasattr(token_to_kv_pool, "get_state_conv_shard_groups")
else None
)
slice_outer_counts = (
token_to_kv_pool.get_state_slice_outer_counts()
if hasattr(token_to_kv_pool, "get_state_slice_outer_counts")
else None
)
# Global layer ids let the sender pair src/dst entries when the
# prefill PP stage registers only its own subset of mamba layers.
layer_ids = token_to_kv_pool.get_state_layer_ids()
append_state_component(
kv_args,
StateType.MAMBA,
data_ptrs,
data_lens,
item_lens,
dim,
conv_shard_groups,
slice_outer_counts,
layer_ids,
)
# Hybrid DSA pools keep their index cache and kpool tail in the
# full-attention sub-pool rather than in the Mamba state above.
if getattr(token_to_kv_pool, "use_dsa", False):
dsa_pool = token_to_kv_pool.full_kv_pool
dsa_ptrs, dsa_lens, dsa_item_lens = dsa_pool.get_state_buf_infos()
append_state_component(
kv_args,
StateType.DSA,
dsa_ptrs,
dsa_lens,
dsa_item_lens,
)
append_dsa_tail(dsa_pool)
if isinstance(token_to_kv_pool, QSATokenToKVPool):
qsa_ptrs, qsa_lens, qsa_item_lens = (
token_to_kv_pool.get_qsa_pending_state_buf_infos()
)
append_state_component(
kv_args,
StateType.QSA_PENDING,
qsa_ptrs,
qsa_lens,
qsa_item_lens,
layer_ids=token_to_kv_pool.get_qsa_pending_state_layer_ids(),
)
compressed_ptrs, compressed_lens, compressed_item_lens = (
token_to_kv_pool.get_qsa_compressed_state_buf_infos()
)
append_state_component(
kv_args,
StateType.QSA_COMPRESSED,
compressed_ptrs,
compressed_lens,
compressed_item_lens,
layer_ids=token_to_kv_pool.get_qsa_compressed_state_layer_ids(),
)
elif isinstance(token_to_kv_pool, (DSATokenToKVPool, NPUMLATokenToKVPool)):
tail_ptrs, tail_lens, tail_item_lens = [], [], []
if isinstance(token_to_kv_pool, DSATokenToKVPool):
tail_ptrs, tail_lens, tail_item_lens = (
token_to_kv_pool.get_compress_tail_buf_infos()
)
if draft_token_to_kv_pool is not None and isinstance(
draft_token_to_kv_pool, DSATokenToKVPool
):
(
draft_data_ptrs,
draft_data_lens,
draft_item_lens,
) = draft_token_to_kv_pool.get_state_buf_infos()
data_ptrs = data_ptrs + draft_data_ptrs
data_lens = data_lens + draft_data_lens
item_lens = item_lens + draft_item_lens
draft_tail_ptrs, draft_tail_lens, draft_tail_item_lens = (
draft_token_to_kv_pool.get_compress_tail_buf_infos()
)
tail_ptrs = tail_ptrs + draft_tail_ptrs
tail_lens = tail_lens + draft_tail_lens
tail_item_lens = tail_item_lens + draft_tail_item_lens
if isinstance(token_to_kv_pool, NPUMLATokenToKVPool):
kv_args.kv_buf_groups = (
len(kv_args.kv_data_ptrs) // token_to_kv_pool.layer_num
)
kv_args.hidden_kv_layers = total_kv_layers
kv_args.draft_kv_layers = (
draft_token_to_kv_pool.layer_num if draft_token_to_kv_pool else 0
)
else:
append_state_component(
kv_args, StateType.DSA, data_ptrs, data_lens, item_lens
)
if tail_ptrs:
append_state_component(
kv_args,
StateType.DSA_TAIL,
tail_ptrs,
tail_lens,
tail_item_lens,
)
if is_npu() and isinstance(token_to_kv_pool, DSV4NPUTokenToKVPool):
from sglang.srt.disaggregation.ascend.conn import AscendStateType
c128_ptrs, c128_lens, c128_item_lens = token_to_kv_pool.get_c128_kv_buf_infos()
if c128_ptrs:
append_state_component(
kv_args,
AscendStateType.DSV4_C128,
c128_ptrs,
c128_lens,
c128_item_lens,
)
# On A5 (CYCLE cache_mode), C4 state uses request-local ring rows rather
# than SWA pages. Register it separately so P and D can independently
# map logical positions when their local ring sizes differ.
from sglang.srt.hardware_backend.npu.utils import is_npu_arch35
if is_npu_arch35():
c4_ptrs, c4_lens, c4_item_lens = token_to_kv_pool.get_c4_state_buf_infos()
if c4_ptrs:
append_state_component(
kv_args,
AscendStateType.DSV4_C4_STATE,
c4_ptrs,
c4_lens,
c4_item_lens,
)
# DSV4 NextN shares the target allocator, so target and draft use the same
# local SWA indices. Keep draft buffers in a separate positional component
# to avoid mixing them into the target's heterogeneous state layout, while
# reusing the existing SWA transport dispatch on both GPU and NPU.
if isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) and isinstance(
draft_token_to_kv_pool, DeepSeekV4TokenToKVPool
):
if not draft_token_to_kv_pool.compression_ratios or not all(
ratio == 0 for ratio in draft_token_to_kv_pool.compression_ratios
):
raise RuntimeError(
"DSV4 draft state transfer expects SWA-only NextN layers"
)
if token_to_kv_pool._unified_kv != draft_token_to_kv_pool._unified_kv:
raise RuntimeError(
"DSV4 target and draft pools must use the same unified-KV mode"
)
if token_to_kv_pool._unified_kv:
target_geometry = (
token_to_kv_pool.unified_swa_window,
token_to_kv_pool.unified_swa_ring_size,
token_to_kv_pool.unified_swa_pages,
)
draft_geometry = (
draft_token_to_kv_pool.unified_swa_window,
draft_token_to_kv_pool.unified_swa_ring_size,
draft_token_to_kv_pool.unified_swa_pages,
)
if target_geometry != draft_geometry:
raise RuntimeError(
"DSV4 target and draft pools must share SWA ring geometry: "
f"target={target_geometry}, draft={draft_geometry}"
)
draft_ptrs, draft_lens, draft_item_lens = (
draft_token_to_kv_pool.get_unified_swa_ring_buf_infos()
)
draft_state_type = StateType.SWA_RING
else:
if (
token_to_kv_pool.full_to_swa_index_mapping
is not draft_token_to_kv_pool.full_to_swa_index_mapping
):
raise RuntimeError(
"DSV4 target and draft pools must share the SWA index mapping"
)
target_geometry = (
token_to_kv_pool.page_size,
token_to_kv_pool.sliding_window,
)
draft_geometry = (
draft_token_to_kv_pool.page_size,
draft_token_to_kv_pool.sliding_window,
)
if target_geometry != draft_geometry:
raise RuntimeError(
"DSV4 target and draft pools must share paged SWA geometry: "
f"target={target_geometry}, draft={draft_geometry}"
)
draft_ptrs, draft_lens, draft_item_lens = (
draft_token_to_kv_pool.get_state_buf_infos()
)
draft_state_type = StateType.SWA
if draft_ptrs:
append_state_component(
kv_args,
draft_state_type,
draft_ptrs,
draft_lens,
draft_item_lens,
)
if (
StateType.MAMBA not in kv_args.state_types
and req_to_token_pool is not None
and hasattr(req_to_token_pool, "get_state_buf_infos")
):
data_ptrs, data_lens, item_lens = req_to_token_pool.get_state_buf_infos()
if data_ptrs:
dim = (
req_to_token_pool.get_state_dim_per_tensor()
if hasattr(req_to_token_pool, "get_state_dim_per_tensor")
else None
)
conv_shard_groups = (
req_to_token_pool.get_state_conv_shard_groups()
if hasattr(req_to_token_pool, "get_state_conv_shard_groups")
else None
)
slice_outer_counts = (
req_to_token_pool.get_state_slice_outer_counts()
if hasattr(req_to_token_pool, "get_state_slice_outer_counts")
else None
)
append_state_component(
kv_args,
StateType.MAMBA,
data_ptrs,
data_lens,
item_lens,
dim,
conv_shard_groups,
slice_outer_counts,
)
def get_dsv41_spec_layout(kv_args: KVArgs) -> Optional[dict]:
"""Describe the positional transfer layout without pool capacities or pointers."""
ratios = getattr(kv_args, "mla_compression_ratios", None) or []
if 2 not in ratios or str(get_spec().speculative_algorithm).upper() != "DSPARK":
return None
from sglang.srt.disaggregation.base.conn import StateType
if kv_args.state_types.count(StateType.SWA) != 2:
raise RuntimeError(
"DeepSeek-V4.1 DSpark PD requires target and draft SWA state"
)
return {
"num_draft_tokens": get_spec().speculative_num_draft_tokens,
"compression_ratios": list(ratios),
"kv_layer_ids": list(kv_args.kv_layer_ids),
"kv_item_lens": list(kv_args.kv_item_lens),
"state_types": [state_type.value for state_type in kv_args.state_types],
"state_item_lens": [list(items) for items in kv_args.state_item_lens],
}
def prepare_abort(req: Req, error_message: str, status_code=None):
from sglang.srt.managers.schedule_batch import FINISH_ABORT
# populate finish metadata and stream output
req.finished_reason = FINISH_ABORT(error_message, status_code)
if req.return_logprob:
req.logprob.input_token_logprobs_val = []
req.logprob.input_token_logprobs_idx = []
req.logprob.input_top_logprobs_val = []
req.logprob.input_top_logprobs_idx = []
req.logprob.input_token_ids_logprobs_val = []
req.logprob.input_token_ids_logprobs_idx = []
def is_unadmitted_reject(req: Req) -> bool:
"""A request rejected at intake, before it acquired anything.
A preempted or resumed request can also carry a pending abort -- "Abort
method 3" marks a *running* request and `filter_batch` does not drop it,
since `finished()` is still False -- and its queue owns the release of
whatever it still holds.
`req.is_retracted` catches the two re-entries that declare nothing:
priority preemption and the pause/retract-all path both requeue through a
bare `_add_request_to_queue`. `release_req` always calls
`reset_for_retract`, which sets it, and its clear sites all run downstream
of these doors. The resource markers stay as a second line of defence --
on their own they miss a `seqlen <= 1` preemption, whose KV is already
freed and whose `retraction_backup` was never taken.
`DecodePreallocQueue.add` still gates on its own `is_retracted` /
`is_rebootstrap` parameters as well, since they state the caller's intent
rather than inferring it.
"""
return is_aborted(req) and not (
req.is_retracted
or req.kv.holds_kv
or req.kv.holds_mamba
or req.metadata_buffer_index >= 0
or req.kv.retraction_backup is not None
)
def is_aborted(req: Req) -> bool:
from sglang.srt.managers.schedule_batch import FINISH_ABORT
return isinstance(req.to_finish, FINISH_ABORT) or isinstance(
req.finished_reason, FINISH_ABORT
)