NIXL: use prep+make API to improve performance (#26406)

This commit is contained in:
Ilia Yastrebov
2026-06-02 18:36:35 +02:00
committed by GitHub
parent 28f9c1ff24
commit 6c69756fa8
+314 -135
View File
@@ -8,7 +8,7 @@ import threading
import time
import uuid
from collections import defaultdict
from typing import TYPE_CHECKING, Dict, List, Optional, Set
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple
import numpy as np
import numpy.typing as npt
@@ -118,6 +118,7 @@ class KVArgsRegisterInfo:
decode_tp_size: int
decode_tp_rank: int
dst_kv_item_len: int
dst_num_slots: Optional[int] = None
dst_state_item_lens: List[List[int]] = dataclasses.field(default_factory=list)
dst_state_dim_per_tensor: List[List[int]] = dataclasses.field(default_factory=list)
# Keep last: optional, parsed from a variable-length tail of the ZMQ
@@ -135,6 +136,9 @@ class KVArgsRegisterInfo:
dst_state_dim_per_tensor = (
unpack_int_lists(msg[13], "I") if len(msg) > 13 and len(msg[13]) > 0 else []
)
dst_num_slots = (
int(msg[16].decode("ascii")) if len(msg) > 16 and msg[16] != b"" else None
)
return cls(
room=str(msg[0].decode("ascii")),
@@ -149,12 +153,50 @@ class KVArgsRegisterInfo:
decode_tp_size=int(msg[9].decode("ascii")),
decode_tp_rank=int(msg[10].decode("ascii")),
dst_kv_item_len=int(msg[11].decode("ascii")),
dst_num_slots=dst_num_slots,
dst_state_item_lens=dst_state_item_lens,
dst_state_dim_per_tensor=dst_state_dim_per_tensor,
staging=StagingRegisterInfo.from_zmq_fields(msg, 14),
)
def expand_page_indices_for_slice(
page_indices: npt.NDArray[np.int32],
num_ptr_pairs: int,
num_slots: int,
page_size: int,
num_groups: int = 1,
head_group_idx: int = 0,
) -> npt.NDArray[np.int32]:
"""Map page slot indices to flat dlist indices for the slice prepped path.
Dlist layout: num_ptr_pairs blocks of (num_slots * page_size * num_groups),
with [slot, token, group] interleaving. head_group_idx selects one group (0 for dst).
"""
token_offsets = np.arange(page_size, dtype=np.int32)
pair_stride = num_slots * page_size * num_groups
within_pair = (
page_indices[:, None] * (page_size * num_groups)
+ token_offsets[None, :] * num_groups
+ head_group_idx
).ravel()
pair_offsets = np.arange(num_ptr_pairs, dtype=np.int64) * pair_stride
return (pair_offsets[:, None] + within_pair[None, :]).ravel().astype(np.int32)
def repeat_indices_over_layers(
indices: npt.NDArray[np.int32], num_layers: int, layer_length: int
) -> npt.NDArray[np.int32]:
"""Map per-slot token indices to flat indices in a pre-built descriptor list.
Each of ``num_layers`` blocks has ``layer_length`` slots; block i is offset by
``i * layer_length``. Works uniformly for both MLA (one ptr/layer) and MHA
(K+V ptrs, 2×N entries).
"""
offsets = np.arange(num_layers, dtype=np.int32) * layer_length
return (offsets[:, None] + indices[None, :]).ravel().astype(np.int32)
@dataclasses.dataclass
class TransferStatus:
"""Used by KV Receiver to know when a transfer is done."""
@@ -249,8 +291,18 @@ class NixlKVManager(CommonKVManager):
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
self.kv_buffer_tensors = None
self.prep_handles: Dict[str, Any] = {}
self.prep_handle_slice_src: Optional[Tuple[Any, int, int, int]] = (
None # (handle, num_groups, num_ptr_pairs, num_slots)
)
self.prep_handles_slice_dst: Dict[str, Tuple[Any, int, int]] = {}
# peer_name -> (handle, num_slots, head_group_idx)
self._num_slots_src: int = 0
if self.disaggregation_mode == DisaggregationMode.PREFILL:
self._num_slots_src = (
self.kv_args.kv_data_lens[0] // self.kv_args.kv_item_lens[0]
)
transfer_queue_size = envs.SGLANG_DISAGGREGATION_QUEUE_SIZE.get()
self.transfer_queues: List[FastQueue] = [
FastQueue() for _ in range(transfer_queue_size)
@@ -454,6 +506,187 @@ class NixlKVManager(CommonKVManager):
def check_status(self, bootstrap_room: int):
return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput)
def _init_equal_tp_prep_handle(
self,
peer_name: str,
kv_ptrs: list[int],
gpu_id: int,
num_slots: Optional[int] = None,
):
"""Pre-build NIXL dlist: all KV slots × all layers.
peer_name="" = src side; agent name = dst side. num_slots overrides the local
slot count — pass decode's count for the dst dlist (may differ from prefill).
Uses prefill's kv_item_lens as stride; requires equal per-slot byte size (equal-TP or MLA).
"""
arrays = []
for base_ptr, item_len, data_len in zip(
kv_ptrs, self.kv_args.kv_item_lens, self.kv_args.kv_data_lens
):
n = num_slots if num_slots is not None else (data_len // item_len)
addrs = np.arange(n, dtype=np.int64) * item_len + base_ptr
arrays.append(
np.column_stack(
[
addrs,
np.full(n, item_len, dtype=np.int64),
np.full(n, gpu_id, dtype=np.int64),
]
)
)
self.prep_handles[peer_name] = self.agent.prep_xfer_dlist(
peer_name, np.vstack(arrays), "VRAM"
)
assert (
self.prep_handles[peer_name] is not None
), f"prep_xfer_dlist returned None for peer '{peer_name}'"
def _init_hetero_tp_prep_handle(
self, peer_name: str, decode_kv_args: KVArgsRegisterInfo
):
"""Pre-build NIXL dlists for TP-heterogeneous slice transfers.
Src dlist shared across decode peers (same TP size). prefill_tp < decode_tp:
interleave num_groups per token, peers select via head_group_idx.
prefill_tp > decode_tp: num_groups=1. Dst dlist is per-peer.
"""
decode_tp_size = decode_kv_args.decode_tp_size
dst_kv_item_len = decode_kv_args.dst_kv_item_len
prefill_tp_size = self.attn_tp_size
page_size = self.kv_args.page_size
total_kv_heads = getattr(self.kv_args, "total_kv_head_num", 0)
if total_kv_heads <= 0:
total_kv_heads = self.kv_args.kv_head_num * prefill_tp_size
src_heads_per_rank = max(1, total_kv_heads // prefill_tp_size)
dst_heads_per_rank = max(1, total_kv_heads // decode_tp_size)
bytes_per_head_slice = dst_kv_item_len // page_size // dst_heads_per_rank
if prefill_tp_size > decode_tp_size:
# Multiple prefill ranks feed one decode rank: each prefill rank sends
# all its src heads to a specific head-range in the decode rank.
src_replication = max(1, prefill_tp_size // total_kv_heads)
local_tp_rank_in_group = self.kv_args.engine_rank % prefill_tp_size
num_groups = 1
num_heads_to_send = src_heads_per_rank
head_group_idx = 0
unique_head_idx = local_tp_rank_in_group // src_replication
dst_head_start = (unique_head_idx * src_heads_per_rank) % dst_heads_per_rank
dst_head_offset = dst_head_start * bytes_per_head_slice
else:
# One prefill rank feeds multiple decode ranks: interleave num_groups
# head-groups in the src dlist so each decode rank picks its slice.
dst_tp_rank_in_group = decode_kv_args.decode_tp_rank % decode_tp_size
num_groups = decode_tp_size // prefill_tp_size
num_heads_to_send = dst_heads_per_rank
src_head_start = (
dst_tp_rank_in_group * dst_heads_per_rank
) % src_heads_per_rank
head_group_idx = src_head_start // dst_heads_per_rank
dst_head_offset = 0
src_kv_item_len = self.kv_args.kv_item_lens[0]
bytes_per_token_to_send = num_heads_to_send * bytes_per_head_slice
bytes_per_token_src = src_kv_item_len // page_size
bytes_per_token_dst = dst_kv_item_len // page_size
src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_pp = (
self.get_mha_kv_ptrs_with_pp(
self.kv_args.kv_data_ptrs, decode_kv_args.dst_kv_ptrs
)
)
src_ptrs = list(src_k_ptrs[:layers_pp]) + list(src_v_ptrs[:layers_pp])
dst_ptrs = list(dst_k_ptrs[:layers_pp]) + list(dst_v_ptrs[:layers_pp])
num_ptr_pairs = len(src_ptrs)
num_slots = self.kv_args.kv_data_lens[0] // src_kv_item_len
slots = np.arange(num_slots, dtype=np.int64)
tokens = np.arange(page_size, dtype=np.int64) # reused in dst dlist below
groups = np.arange(num_groups, dtype=np.int64)
# Src dlist built once and shared.
if self.prep_handle_slice_src is None:
# (ptr, slot, token, group) → ravel; groups interleaved per token.
src_ptrs_arr = np.array(src_ptrs, dtype=np.int64)
addrs = (
src_ptrs_arr[:, None, None, None]
+ slots[None, :, None, None] * src_kv_item_len
+ tokens[None, None, :, None] * bytes_per_token_src
+ groups[None, None, None, :] * bytes_per_token_to_send
).ravel()
src_array = np.column_stack(
[
addrs,
np.full(len(addrs), bytes_per_token_to_send, dtype=np.int64),
np.full(len(addrs), self.kv_args.gpu_id, dtype=np.int64),
]
)
src_handle = self.agent.prep_xfer_dlist("", src_array, "VRAM")
assert (
src_handle is not None
), f"prep_xfer_dlist returned None for slice src (decode_tp_size={decode_tp_size})"
self.prep_handle_slice_src = (
src_handle,
num_groups,
num_ptr_pairs,
num_slots,
)
# Dst dlist per-peer; use decode's slot count (may exceed prefill's).
num_slots_dst = (
decode_kv_args.dst_num_slots
if decode_kv_args.dst_num_slots is not None
else num_slots
)
dst_slots = np.arange(num_slots_dst, dtype=np.int64)
# (ptr, slot, token) → ravel.
dst_ptrs_arr = np.array(dst_ptrs, dtype=np.int64)
addrs = (
dst_ptrs_arr[:, None, None]
+ dst_slots[None, :, None] * dst_kv_item_len
+ tokens[None, None, :] * bytes_per_token_dst
+ dst_head_offset
).ravel()
dst_array = np.column_stack(
[
addrs,
np.full(len(addrs), bytes_per_token_to_send, dtype=np.int64),
np.full(len(addrs), decode_kv_args.gpu_id, dtype=np.int64),
]
)
dst_handle = self.agent.prep_xfer_dlist(peer_name, dst_array, "VRAM")
assert (
dst_handle is not None
), f"prep_xfer_dlist returned None for slice dst for peer '{peer_name}'"
self.prep_handles_slice_dst[peer_name] = (
dst_handle,
num_slots_dst,
head_group_idx,
)
def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo):
if self.is_mla_backend or peer_info.decode_tp_size == self.attn_tp_size:
# Safe to use prefill's kv_item_lens for the dst dlist stride:
# equal_tp guarantees identical heads-per-rank (same item_len);
# MLA latent shape is TP-invariant.
# Build the shared src dlist on the first equal-TP/MLA peer; later
# peers reuse it. Skipped entirely on heterogeneous-TP-only setups.
if "" not in self.prep_handles:
self._init_equal_tp_prep_handle(
"", self.kv_args.kv_data_ptrs, self.kv_args.gpu_id
)
self._init_equal_tp_prep_handle(
peer_info.agent_name,
peer_info.dst_kv_ptrs,
peer_info.gpu_id,
num_slots=peer_info.dst_num_slots,
)
else:
self._init_hetero_tp_prep_handle(peer_info.agent_name, peer_info)
def transfer_worker(self, queue: FastQueue, staging_buffer=None):
# Per-worker staging strategy: lazy-created on first chunk so we
# see kv_buffer_tensors (set by ModelRunner after engine init).
@@ -463,6 +696,7 @@ class NixlKVManager(CommonKVManager):
while True:
kv_chunk: TransferKVChunk = queue.get()
room = kv_chunk.room
handles: List[Any] = []
try:
if self.check_status(room) == KVPoll.Failed:
continue
@@ -481,7 +715,6 @@ class NixlKVManager(CommonKVManager):
self.update_status(room, KVPoll.Transferring)
reqs_to_be_processed = list(self.transfer_infos[room].values())
handles: List = []
# Set when staging allocation/watermark is not yet ready and
# the chunk has been re-enqueued. We then break out of the
@@ -516,6 +749,11 @@ class NixlKVManager(CommonKVManager):
: len(chunked_dst_kv_indice)
]
notif = (
f"{req.room}_kv_{kv_chunk.chunk_id}"
f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}"
)
# Decide which kv send path to use:
# 1. Staging (heterogeneous TP, both sides have
# registered staging, watermark/alloc ready)
@@ -551,10 +789,6 @@ class NixlKVManager(CommonKVManager):
# the slice path below.
if kv_xfer_handle is None:
notif = (
f"{req.room}_kv_{kv_chunk.chunk_id}"
f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}"
)
if self.is_mla_backend or (
decode_tp_size == self.attn_tp_size
):
@@ -570,14 +804,8 @@ class NixlKVManager(CommonKVManager):
kv_xfer_handle = self.send_kvcache_slice(
req.agent_name,
kv_chunk.prefill_kv_indices,
dst_info.dst_kv_ptrs,
chunked_dst_kv_indice,
dst_info.gpu_id,
notif,
prefill_tp_size=self.attn_tp_size,
decode_tp_size=decode_tp_size,
decode_tp_rank=dst_info.decode_tp_rank,
dst_kv_item_len=dst_info.dst_kv_item_len,
)
handles.append(kv_xfer_handle)
@@ -707,6 +935,8 @@ class NixlKVManager(CommonKVManager):
return
self.decode_kv_args_table[agent_name] = decode_kv_args
self.agent.add_remote_agent(decode_kv_args.agent_metadata)
if self.disaggregation_mode == DisaggregationMode.PREFILL:
self._prepare_payload_xfer(decode_kv_args)
def _send_kvcache_generic(
self,
@@ -721,6 +951,43 @@ class NixlKVManager(CommonKVManager):
):
"""Generic KV cache transfer supporting both MHA and MLA architectures.
Used by both send_kvcache and maybe_send_extra."""
# Prepped path (KV only; state transfers use the non-prepped path below).
if (
src_data_ptrs is self.kv_args.kv_data_ptrs
and "" in self.prep_handles
and peer_name in self.prep_handles
):
src_prep = self.prep_handles[""]
dst_prep = self.prep_handles[peer_name]
info = self.decode_kv_args_table[peer_name]
num_slots_dst = (
info.dst_num_slots
if info.dst_num_slots is not None
else self._num_slots_src
)
num_layers = len(item_lens)
src_indices = repeat_indices_over_layers(
prefill_data_indices, num_layers, self._num_slots_src
)
dst_indices = repeat_indices_over_layers(
dst_data_indices, num_layers, num_slots_dst
)
xfer_handle = self.agent.make_prepped_xfer(
"WRITE",
src_prep,
src_indices,
dst_prep,
dst_indices,
notif.encode("ascii"),
)
if not xfer_handle:
raise Exception("KVSender failed to create prepped transfer")
state = self.agent.transfer(xfer_handle)
if state == "ERR":
raise Exception("KVSender failed to post prepped transfer")
return xfer_handle
# Non-prepped path: used for state transfers (SWA/NSA) via maybe_send_extra.
# Convert pointer lists to np.uint64 arrays up front.
# torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set, e.g.
# 0xffff81ab54e01000). Casting here prevents overflow when these values
@@ -850,139 +1117,46 @@ class NixlKVManager(CommonKVManager):
self,
peer_name: str,
prefill_kv_indices: npt.NDArray[np.int32],
dst_kv_ptrs: list[int],
dst_kv_indices: npt.NDArray[np.int32],
dst_gpu_id: int,
notif: str,
prefill_tp_size: int,
decode_tp_size: int,
decode_tp_rank: int,
dst_kv_item_len: int,
):
# Get configuration from kv_args
local_tp_rank_in_group = self.kv_args.engine_rank % prefill_tp_size
dst_tp_rank_in_group = decode_tp_rank % decode_tp_size
src_kv_item_len = self.kv_args.kv_item_lens[0]
page_size = self.kv_args.page_size
# Use total KV head count (not per-rank) for correct head distribution.
# Per-rank kv_head_num is max(1, total//tp) which loses info when total < tp.
total_kv_heads = getattr(self.kv_args, "total_kv_head_num", 0)
if total_kv_heads <= 0:
total_kv_heads = self.kv_args.kv_head_num * prefill_tp_size
src_heads_per_rank = max(1, total_kv_heads // prefill_tp_size)
dst_heads_per_rank = max(1, total_kv_heads // decode_tp_size)
bytes_per_head_slice_to_send = (
dst_kv_item_len // page_size // dst_heads_per_rank
# Prepped path: src dlist is shared per decode_tp_size; dst is per peer.
assert self.prep_handle_slice_src is not None
assert peer_name in self.prep_handles_slice_dst
src_handle, num_groups, num_ptr_pairs, num_slots_src = (
self.prep_handle_slice_src
)
# GQA replication: how many prefill ranks share the same KV head
src_replication = max(1, prefill_tp_size // total_kv_heads)
# Determine which heads to send
if prefill_tp_size > decode_tp_size:
# Multiple prefill ranks to one decode rank
src_head_start_offset = 0
num_heads_to_send = src_heads_per_rank
unique_head_idx = local_tp_rank_in_group // src_replication
dst_head_start_offset = (
unique_head_idx * src_heads_per_rank
) % dst_heads_per_rank
else:
# Send KVCache from 1 prefill instance to multiple decode instances
src_head_start_offset = (
dst_tp_rank_in_group * dst_heads_per_rank
) % src_heads_per_rank
num_heads_to_send = dst_heads_per_rank
dst_head_start_offset = 0
# torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set, e.g.
# 0xffff81ab54e01000). Use np.uint64 to prevent overflow on XPU.
kv_data_ptrs = np.array(self.kv_args.kv_data_ptrs, dtype=np.uint64)
dst_kv_ptrs = np.array(dst_kv_ptrs, dtype=np.uint64)
src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_current_pp_stage = (
self.get_mha_kv_ptrs_with_pp(kv_data_ptrs, dst_kv_ptrs)
)
# Calculate precise byte offset and length for the sub-slice within the token
src_head_slice_offset = src_head_start_offset * bytes_per_head_slice_to_send
dst_head_slice_offset = dst_head_start_offset * bytes_per_head_slice_to_send
heads_bytes_per_token_to_send = num_heads_to_send * bytes_per_head_slice_to_send
src_dst_ptr_pairs = [
(
src_k_ptrs[layer_id],
dst_k_ptrs[layer_id],
)
for layer_id in range(layers_current_pp_stage)
] + [
(
src_v_ptrs[layer_id],
dst_v_ptrs[layer_id],
)
for layer_id in range(layers_current_pp_stage)
dst_handle, num_slots_dst, head_group_idx = self.prep_handles_slice_dst[
peer_name
]
prefill_indices = np.asarray(prefill_kv_indices, dtype=np.uint64)
dst_indices = np.asarray(dst_kv_indices, dtype=np.uint64)
bytes_per_token_prefill = src_kv_item_len // page_size
bytes_per_token_decode = dst_kv_item_len // page_size
token_offsets = np.arange(page_size, dtype=np.uint64)
src_addrs = []
dst_addrs = []
for src_ptr, dst_ptr in src_dst_ptr_pairs:
src_page_bases = src_ptr + prefill_indices * src_kv_item_len
dst_page_bases = dst_ptr + dst_indices * dst_kv_item_len
src_all = (
src_page_bases[:, None]
+ token_offsets[None, :] * bytes_per_token_prefill
+ src_head_slice_offset
).ravel()
dst_all = (
dst_page_bases[:, None]
+ token_offsets[None, :] * bytes_per_token_decode
+ dst_head_slice_offset
).ravel()
src_addrs.append(src_all)
dst_addrs.append(dst_all)
def make_req_array(addr_chunks, size, gpu):
if not addr_chunks:
return np.empty((0, 3), dtype=np.uint64)
flat_addrs = np.concatenate(addr_chunks).astype(np.uint64, copy=False)
return np.column_stack(
(
flat_addrs,
np.full_like(flat_addrs, size, dtype=np.uint64),
np.full_like(flat_addrs, gpu, dtype=np.uint64),
)
)
src_reqs = make_req_array(
src_addrs, heads_bytes_per_token_to_send, self.kv_args.gpu_id
page_size = self.kv_args.page_size
src_indices = expand_page_indices_for_slice(
np.asarray(prefill_kv_indices, dtype=np.int32),
num_ptr_pairs,
num_slots_src,
page_size,
num_groups=num_groups,
head_group_idx=head_group_idx,
)
dst_reqs = make_req_array(dst_addrs, heads_bytes_per_token_to_send, dst_gpu_id)
# Use NIXL agent for transfer
src_descs = self.agent.get_xfer_descs(src_reqs, "VRAM")
dst_descs = self.agent.get_xfer_descs(dst_reqs, "VRAM")
xfer_handle = self.agent.initialize_xfer(
"WRITE", src_descs, dst_descs, peer_name, notif.encode("ascii")
dst_indices = expand_page_indices_for_slice(
np.asarray(dst_kv_indices, dtype=np.int32),
num_ptr_pairs,
num_slots_dst,
page_size,
)
xfer_handle = self.agent.make_prepped_xfer(
"WRITE",
src_handle,
src_indices,
dst_handle,
dst_indices,
notif.encode("ascii"),
)
if not xfer_handle:
raise Exception("Failed to create sliced KV transfer")
raise Exception("KVSender failed to create prepped slice transfer")
state = self.agent.transfer(xfer_handle)
if state == "ERR":
raise Exception("Failed to post sliced KV transfer")
raise Exception("KVSender failed to post prepped slice transfer")
return xfer_handle
def send_kvcache_staged(
@@ -1910,6 +2084,10 @@ class NixlKVReceiver(CommonKVReceiver):
else:
packed_staging_base_ptr = b""
staging_total_size_str = b""
dst_num_slots = (
self.kv_mgr.kv_args.kv_data_lens[0]
// self.kv_mgr.kv_args.kv_item_lens[0]
)
with lock:
sock.send_multipart(
@@ -1931,6 +2109,7 @@ class NixlKVReceiver(CommonKVReceiver):
packed_state_dim_per_tensor,
packed_staging_base_ptr,
staging_total_size_str,
str(dst_num_slots).encode("ascii"),
]
)