Files
sglang/python/sglang/srt/disaggregation/utils.py
T
62a4a6ea0e [NPU] Add NPU arch35 support and enhance DSV4 processing in DeepSeek-V4 (#37373)
Co-authored-by: AndyLi429 <AndyLi429@noreply.gitcode.com>
Co-authored-by: Kailong Lu <kelonlu@163.com>
Co-authored-by: cx <chengxin65@huawei.com>
Co-authored-by: ranjiewen <ranjiewen@huawei.com>
Co-authored-by: HEX1A0A <1a0ahex@gmail.com>
Co-authored-by: vstone-w <374330057@qq.com>
Co-authored-by: Even Zhou <even.y.zhou@outlook.com>
Co-authored-by: ClownBin <chaobin1993@126.com>
Co-authored-by: sglang-npu-bot <sglangnpu@163.com>
2026-09-07 21:08:06 +08:00

1640 lines
63 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
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.environ import envs
from sglang.srt.runtime_context import (
get_disagg,
)
from sglang.srt.utils import is_hip, 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"
_IS_HIP = is_hip()
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
return get_dsa_mtp_topk_width(hf_config)
def is_dsv4_c128_online_enabled() -> bool:
"""Return whether DSV4 C128 uses request-scoped online state."""
return not _IS_HIP and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get()
def get_dsv4_c4_state_indices(
req_pool_idx: int,
seq_len: int,
*,
ring_size: int,
) -> np.ndarray:
"""Return physical rows for the live C4 compressor history.
Prefill and decode may use different C4 ring sizes (8 without speculative
decoding and 16 with EAGLE/MTP). State transfer must therefore pair rows
by logical token position instead of copying a whole request-local bank.
The C4 overlap compressor keeps ``seq_len % 4 + 4`` live rows.
"""
if ring_size < 8 or ring_size % 4 != 0:
raise ValueError(
f"C4 ring_size must be a multiple of 4 and at least 8, got {ring_size}"
)
seq_len = max(0, int(seq_len))
state_len = seq_len % 4 + 4
positions = np.arange(max(0, seq_len - state_len), seq_len, dtype=np.int64)
rows = int(req_pool_idx) * int(ring_size) + positions % int(ring_size)
return rows.astype(np.int32)
def get_dsv4_c128_state_indices(
req_pool_idx: int,
seq_len: int,
*,
online: bool,
ring_size: int,
) -> np.ndarray:
"""Return the PD transfer row/page indices for DSV4 C128 state."""
if seq_len == 0 or seq_len % 128 == 0:
return np.empty((0,), dtype=np.int32)
if online:
return np.array([int(req_pool_idx)], dtype=np.int32)
assert ring_size % 128 == 0, f"C128 ring_size must be 128-aligned, got {ring_size}"
pages_per_req = ring_size // 128
page = int(req_pool_idx) * pages_per_req + ((seq_len - 1) % ring_size) // 128
return np.array([page], 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."""
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_top_logprobs_num: int = 128,
max_sampling_mask_tokens: Optional[int] = None,
custom_mem_pool: torch.cuda.MemPool = None,
output_dsa_topk_indices_dim: int = 0,
):
self.custom_mem_pool = custom_mem_pool
self.output_dsa_topk_indices_dim = output_dsa_topk_indices_dim
if max_sampling_mask_tokens is None:
max_sampling_mask_tokens = (
envs.SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS.get()
)
self.enable_sampling_mask = max_sampling_mask_tokens > 0
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
)
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,
]
if self.enable_sampling_mask:
bufs.extend(
[
self.output_token_sampling_mask_len,
self.output_token_sampling_mask_idx,
self.output_token_sampling_logprobs,
]
)
bufs.extend(
[
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)
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):
sampling_mask_len = None
sampling_mask_idx = None
sampling_logprobs = None
if self.enable_sampling_mask:
sampling_mask_len = self.output_token_sampling_mask_len[idx].clone()
sampling_mask_idx = self.output_token_sampling_mask_idx[idx].clone()
sampling_logprobs = self.output_token_sampling_logprobs[idx].clone()
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(),
sampling_mask_len,
sampling_mask_idx,
sampling_logprobs,
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:
if not self.enable_sampling_mask:
raise RuntimeError(
"return_sampling_mask with disaggregation requires "
"SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS > 0."
)
# 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 "
"SGLANG_DISAGGREGATION_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 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.
"""
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:
"""Populate ``kv_args`` state-buffer fields from the given pool.
Shared by prefill and decode bootstrap paths so the state_type dispatch
lives in one place.
"""
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.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 = []
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_c128_state_buf_infos"):
c128_ptrs, c128_lens, c128_item_lens = (
token_to_kv_pool.get_c128_state_buf_infos()
)
if c128_ptrs:
append_state_component(
kv_args,
StateType.C128_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)
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.total_kv_layers = total_kv_layers
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 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_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
)