hisparse: support NIXL DRAM KV destinations for HiSparse (#27563)
Co-authored-by: Zhangheng <hzh0425@apache.org> Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
Zhangheng
Shangming Cai
parent
c1b5c7e499
commit
b030b1a5f3
@@ -108,10 +108,15 @@ Pass as a JSON string via `--hisparse-config`:
|
||||
<td>int</td>
|
||||
<td>Ratio of logical pool size to device pool size, determining host memory capacity</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td><code>swap_in_block_size</code></td>
|
||||
<td>int / 960</td>
|
||||
<td>CUDA thread-block size for the HiSparse swap-in kernel</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
Example: `--hisparse-config='{"top_k": 2048, "device_buffer_size": 6144, "host_to_device_ratio": 10}'`
|
||||
Example: `--hisparse-config='{"top_k": 2048, "device_buffer_size": 6144, "host_to_device_ratio": 10, "swap_in_block_size": 960}'`
|
||||
|
||||
## Deployment
|
||||
|
||||
@@ -149,7 +154,7 @@ python3 -m sglang.launch_server \
|
||||
--dist-init-addr 127.0.0.1:5757 \
|
||||
--nnodes 1 --node-rank 0 \
|
||||
--enable-hisparse \
|
||||
--hisparse-config='{"top_k": 2048, "device_buffer_size": 6144, "host_to_device_ratio": 10}'
|
||||
--hisparse-config='{"top_k": 2048, "device_buffer_size": 6144, "host_to_device_ratio": 10, "swap_in_block_size": 960}'
|
||||
```
|
||||
|
||||
> **Note**: For DSA models, `--kv-cache-dtype` defaults to `auto`, which resolves to `fp8_e4m3` on SM100+ (Blackwell) and `bfloat16` on older architectures. The DSA decode backend is automatically selected based on KV dtype (`bfloat16` → `flashmla_sparse`, `fp8_e4m3` → `flashmla_kv`). DSA backend flags apply only to DSA models; DeepSeek V4 uses its own `dsv4` attention backend.
|
||||
|
||||
@@ -100,7 +100,6 @@ class BaseKVManager(ABC):
|
||||
|
||||
|
||||
class BaseKVSender(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def __init__(
|
||||
self,
|
||||
@@ -156,7 +155,6 @@ class BaseKVSender(ABC):
|
||||
|
||||
|
||||
class BaseKVReceiver(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def __init__(
|
||||
self,
|
||||
|
||||
@@ -403,6 +403,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
kv_data_ptrs, kv_data_lens, kv_item_lens = (
|
||||
transfer_kv_pool.get_contiguous_buf_infos()
|
||||
)
|
||||
kv_data_mem_kinds = (
|
||||
["DRAM"] * len(kv_data_ptrs)
|
||||
if self.scheduler.enable_hisparse
|
||||
else ["VRAM"] * len(kv_data_ptrs)
|
||||
)
|
||||
if self.scheduler.enable_hisparse and isinstance(
|
||||
self.token_to_kv_pool, DeepSeekV4TokenToKVPool
|
||||
):
|
||||
@@ -413,6 +418,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
kv_data_ptrs += device_kv_data_ptrs[c4_layer_num:]
|
||||
kv_data_lens += device_kv_data_lens[c4_layer_num:]
|
||||
kv_item_lens += device_kv_item_lens[c4_layer_num:]
|
||||
kv_data_mem_kinds += ["VRAM"] * len(device_kv_data_ptrs[c4_layer_num:])
|
||||
if self.draft_token_to_kv_pool is not None:
|
||||
# We should also transfer draft model kv cache. The indices are
|
||||
# always shared with a target model.
|
||||
@@ -422,10 +428,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
kv_data_ptrs += draft_kv_data_ptrs
|
||||
kv_data_lens += draft_kv_data_lens
|
||||
kv_item_lens += draft_kv_item_lens
|
||||
kv_data_mem_kinds += ["VRAM"] * len(draft_kv_data_ptrs)
|
||||
|
||||
kv_args.kv_data_ptrs = kv_data_ptrs
|
||||
kv_args.kv_data_lens = kv_data_lens
|
||||
kv_args.kv_item_lens = kv_item_lens
|
||||
if self.transfer_backend == TransferBackend.NIXL:
|
||||
kv_args.kv_data_mem_kinds = kv_data_mem_kinds
|
||||
kv_args.page_size = self.token_to_kv_pool.page_size
|
||||
|
||||
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
|
||||
|
||||
@@ -54,6 +54,87 @@ except ImportError:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
GUARD = "NixlMsgGuard".encode("ascii")
|
||||
KV_MEM_KINDS = {"VRAM", "DRAM"}
|
||||
|
||||
|
||||
def _normalize_kv_mem_kinds(kinds: Optional[List[str]], expected_len: int) -> List[str]:
|
||||
if kinds is None:
|
||||
return ["VRAM"] * expected_len
|
||||
kinds = [str(kind) for kind in kinds]
|
||||
if len(kinds) != expected_len:
|
||||
raise ValueError(
|
||||
f"kv_data_mem_kinds length mismatch: got {len(kinds)}, "
|
||||
f"expected {expected_len}"
|
||||
)
|
||||
invalid = sorted(set(kinds) - KV_MEM_KINDS)
|
||||
if invalid:
|
||||
raise ValueError(f"Unsupported NIXL KV memory kind(s): {invalid}")
|
||||
return kinds
|
||||
|
||||
|
||||
def _pack_kv_mem_kinds(kinds: List[str]) -> bytes:
|
||||
return ",".join(kinds).encode("ascii")
|
||||
|
||||
|
||||
def _unpack_kv_mem_kinds(buf: bytes, expected_len: int) -> List[str]:
|
||||
if not buf:
|
||||
return ["VRAM"] * expected_len
|
||||
return _normalize_kv_mem_kinds(buf.decode("ascii").split(","), expected_len)
|
||||
|
||||
|
||||
def _nixl_device_id(mem_kind: str, gpu_id: int) -> int:
|
||||
return gpu_id if mem_kind == "VRAM" else 0
|
||||
|
||||
|
||||
def _homogeneous_kv_mem_kind(kinds: List[str], context: str) -> str:
|
||||
unique = set(kinds)
|
||||
if len(unique) != 1:
|
||||
raise NotImplementedError(
|
||||
f"NIXL {context} mixed KV memory kinds are not implemented safely yet: "
|
||||
f"{sorted(unique)}"
|
||||
)
|
||||
return next(iter(unique))
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class _KVXferMemSegment:
|
||||
start: int
|
||||
end: int
|
||||
src_mem_kind: str
|
||||
dst_mem_kind: str
|
||||
|
||||
|
||||
def _kv_xfer_mem_segments(
|
||||
src_kinds: List[str], dst_kinds: List[str]
|
||||
) -> List[_KVXferMemSegment]:
|
||||
if len(src_kinds) != len(dst_kinds):
|
||||
raise ValueError(
|
||||
f"KV source/destination memory kind length mismatch: "
|
||||
f"src={len(src_kinds)}, dst={len(dst_kinds)}"
|
||||
)
|
||||
if not src_kinds:
|
||||
return []
|
||||
|
||||
segments = []
|
||||
start = 0
|
||||
cur = (src_kinds[0], dst_kinds[0])
|
||||
for i, pair in enumerate(zip(src_kinds, dst_kinds)):
|
||||
if pair == cur:
|
||||
continue
|
||||
segments.append(_KVXferMemSegment(start, i, cur[0], cur[1]))
|
||||
start = i
|
||||
cur = pair
|
||||
segments.append(_KVXferMemSegment(start, len(src_kinds), cur[0], cur[1]))
|
||||
return segments
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class _KVXferPreparedSegment:
|
||||
start: int
|
||||
end: int
|
||||
src_handle: Any
|
||||
dst_handle: Any
|
||||
dst_num_slots: int
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
@@ -113,21 +194,42 @@ class KVArgsRegisterInfo:
|
||||
agent_name: str
|
||||
agent_metadata: bytes
|
||||
dst_kv_ptrs: list[int]
|
||||
dst_kv_mem_kinds: list[str]
|
||||
dst_aux_ptrs: list[int]
|
||||
dst_state_data_ptrs: List[List[int]]
|
||||
gpu_id: int
|
||||
decode_tp_size: int
|
||||
decode_tp_rank: int
|
||||
dst_kv_item_len: int
|
||||
dst_kv_item_lens: list[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)
|
||||
dst_homogeneous_mem_kind: Optional[str] = None
|
||||
kv_xfer_segments: Optional[List[_KVXferPreparedSegment]] = None
|
||||
# Keep last: optional, parsed from a variable-length tail of the ZMQ
|
||||
# frame in from_zmq() below, so positional construction stays stable.
|
||||
staging: Optional[StagingRegisterInfo] = None
|
||||
|
||||
@classmethod
|
||||
def from_zmq(cls, msg: List[bytes]):
|
||||
dst_kv_ptrs = list(struct.unpack(f"{len(msg[5]) // 8}Q", msg[5]))
|
||||
dst_kv_mem_kinds = (
|
||||
_unpack_kv_mem_kinds(msg[17], len(dst_kv_ptrs))
|
||||
if len(msg) > 17
|
||||
else ["VRAM"] * len(dst_kv_ptrs)
|
||||
)
|
||||
dst_kv_item_len = int(msg[11].decode("ascii"))
|
||||
dst_kv_item_lens = (
|
||||
list(struct.unpack(f"{len(msg[18]) // 8}Q", msg[18]))
|
||||
if len(msg) > 18 and msg[18] != b""
|
||||
else [dst_kv_item_len] * len(dst_kv_ptrs)
|
||||
)
|
||||
if len(dst_kv_item_lens) != len(dst_kv_ptrs):
|
||||
raise ValueError(
|
||||
"dst_kv_item_lens length mismatch: "
|
||||
f"got {len(dst_kv_item_lens)}, expected {len(dst_kv_ptrs)}"
|
||||
)
|
||||
dst_state_data_ptrs = (
|
||||
unpack_int_lists(msg[7], "Q") if len(msg) > 7 and msg[7] != b"" else []
|
||||
)
|
||||
@@ -147,13 +249,15 @@ class KVArgsRegisterInfo:
|
||||
dst_port=int(msg[2].decode("ascii")),
|
||||
agent_name=msg[3].decode("ascii"),
|
||||
agent_metadata=msg[4],
|
||||
dst_kv_ptrs=list(struct.unpack(f"{len(msg[5]) // 8}Q", msg[5])),
|
||||
dst_kv_ptrs=dst_kv_ptrs,
|
||||
dst_kv_mem_kinds=dst_kv_mem_kinds,
|
||||
dst_aux_ptrs=list(struct.unpack(f"{len(msg[6]) // 8}Q", msg[6])),
|
||||
dst_state_data_ptrs=dst_state_data_ptrs,
|
||||
gpu_id=int(msg[8].decode("ascii")),
|
||||
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_kv_item_len=dst_kv_item_len,
|
||||
dst_kv_item_lens=dst_kv_item_lens,
|
||||
dst_num_slots=dst_num_slots,
|
||||
dst_state_item_lens=dst_state_item_lens,
|
||||
dst_state_dim_per_tensor=dst_state_dim_per_tensor,
|
||||
@@ -216,6 +320,10 @@ class TransferStatus:
|
||||
received_state_per_pp: Set[int] = dataclasses.field(default_factory=set)
|
||||
# Whether state data is expected (set based on state_type).
|
||||
expects_state: bool = False
|
||||
# KV part notifications for mixed-memory transfers. Keyed by
|
||||
# (pp_rank, chunk_id); normal homogeneous transfers bypass this.
|
||||
received_kv_parts_per_pp: Optional[Dict[Tuple[int, int], Set[int]]] = None
|
||||
expected_kv_parts_per_pp: Optional[Dict[Tuple[int, int], int]] = None
|
||||
|
||||
def is_done(self):
|
||||
if self.num_pp_ranks_expected is None or not self.received_aux:
|
||||
@@ -245,6 +353,15 @@ class NixlKVManager(CommonKVManager):
|
||||
is_mla_backend: Optional[bool] = False,
|
||||
):
|
||||
super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
|
||||
self.kv_args.kv_data_mem_kinds = _normalize_kv_mem_kinds(
|
||||
getattr(self.kv_args, "kv_data_mem_kinds", None),
|
||||
len(self.kv_args.kv_data_ptrs),
|
||||
)
|
||||
self.src_mem_kind = (
|
||||
_homogeneous_kv_mem_kind(self.kv_args.kv_data_mem_kinds, "source")
|
||||
if disaggregation_mode == DisaggregationMode.PREFILL
|
||||
else None
|
||||
)
|
||||
try:
|
||||
from nixl._api import nixl_agent, nixl_agent_config, nixl_thread_sync_t
|
||||
except ImportError as e:
|
||||
@@ -304,6 +421,7 @@ class NixlKVManager(CommonKVManager):
|
||||
)
|
||||
self.prep_handles_slice_dst: Dict[str, Tuple[Any, int, int]] = {}
|
||||
# peer_name -> (handle, num_slots, head_group_idx)
|
||||
self.prep_handles_segment_src: Dict[Tuple[int, int, str], Any] = {}
|
||||
self._num_slots_src: int = 0
|
||||
|
||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
@@ -513,47 +631,99 @@ 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(
|
||||
def _prep_equal_tp_dlist(
|
||||
self,
|
||||
peer_name: str,
|
||||
kv_ptrs: list[int],
|
||||
kv_item_lens: list[int],
|
||||
kv_data_lens: list[int],
|
||||
gpu_id: int,
|
||||
num_slots: Optional[int] = None,
|
||||
mem_kind: str = "VRAM",
|
||||
kv_xfer_lens: Optional[list[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).
|
||||
"""
|
||||
if kv_xfer_lens is None:
|
||||
kv_xfer_lens = kv_item_lens
|
||||
if not (
|
||||
len(kv_ptrs) == len(kv_item_lens) == len(kv_data_lens) == len(kv_xfer_lens)
|
||||
):
|
||||
raise ValueError(
|
||||
"NIXL prepared dlist geometry length mismatch: "
|
||||
f"ptrs={len(kv_ptrs)}, item_lens={len(kv_item_lens)}, "
|
||||
f"data_lens={len(kv_data_lens)}, xfer_lens={len(kv_xfer_lens)}"
|
||||
)
|
||||
device_id = _nixl_device_id(mem_kind, gpu_id)
|
||||
arrays = []
|
||||
# torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set).
|
||||
# Convert once at entry; all downstream arithmetic stays in uint64.
|
||||
kv_ptrs_u64 = np.array(kv_ptrs, dtype=np.uint64)
|
||||
for base_ptr, item_len, data_len in zip(
|
||||
kv_ptrs_u64, self.kv_args.kv_item_lens, self.kv_args.kv_data_lens
|
||||
for base_ptr, item_len, data_len, xfer_len in zip(
|
||||
kv_ptrs_u64, kv_item_lens, kv_data_lens, kv_xfer_lens
|
||||
):
|
||||
if xfer_len > item_len:
|
||||
raise ValueError(
|
||||
"NIXL prepared dlist transfer length exceeds item stride: "
|
||||
f"xfer_len={xfer_len}, item_len={item_len}, mem_kind={mem_kind}"
|
||||
)
|
||||
n = num_slots if num_slots is not None else (data_len // item_len)
|
||||
addrs = np.arange(n, dtype=np.uint64) * np.uint64(item_len) + base_ptr
|
||||
arrays.append(
|
||||
np.column_stack(
|
||||
[
|
||||
addrs,
|
||||
np.full(n, item_len, dtype=np.uint64),
|
||||
np.full(n, gpu_id, dtype=np.uint64),
|
||||
np.full(n, xfer_len, dtype=np.uint64),
|
||||
np.full(n, device_id, dtype=np.uint64),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
self.prep_handles[peer_name] = self.agent.prep_xfer_dlist(
|
||||
peer_name, np.vstack(arrays), "VRAM"
|
||||
)
|
||||
prep_handle = self.agent.prep_xfer_dlist(peer_name, np.vstack(arrays), mem_kind)
|
||||
assert (
|
||||
self.prep_handles[peer_name] is not None
|
||||
prep_handle is not None
|
||||
), f"prep_xfer_dlist returned None for peer '{peer_name}'"
|
||||
return prep_handle
|
||||
|
||||
def _init_equal_tp_prep_handle(
|
||||
self,
|
||||
peer_name: str,
|
||||
kv_ptrs: list[int],
|
||||
gpu_id: int,
|
||||
num_slots: Optional[int] = None,
|
||||
mem_kind: str = "VRAM",
|
||||
kv_item_lens: Optional[list[int]] = None,
|
||||
kv_data_lens: Optional[list[int]] = None,
|
||||
kv_xfer_lens: Optional[list[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).
|
||||
Source dlists use prefill geometry; destination dlists must use decode
|
||||
stride geometry but source transfer lengths, because HiSparse can transfer
|
||||
directly into a host pool whose slot stride differs from prefill.
|
||||
"""
|
||||
if kv_item_lens is None:
|
||||
kv_item_lens = self.kv_args.kv_item_lens
|
||||
if kv_data_lens is None:
|
||||
kv_data_lens = self.kv_args.kv_data_lens
|
||||
self.prep_handles[peer_name] = self._prep_equal_tp_dlist(
|
||||
peer_name,
|
||||
kv_ptrs,
|
||||
kv_item_lens,
|
||||
kv_data_lens,
|
||||
gpu_id,
|
||||
num_slots=num_slots,
|
||||
mem_kind=mem_kind,
|
||||
kv_xfer_lens=kv_xfer_lens,
|
||||
)
|
||||
|
||||
def _init_hetero_tp_prep_handle(
|
||||
self, peer_name: str, decode_kv_args: KVArgsRegisterInfo
|
||||
self,
|
||||
peer_name: str,
|
||||
decode_kv_args: KVArgsRegisterInfo,
|
||||
src_mem_kind: str = "VRAM",
|
||||
dst_mem_kind: str = "VRAM",
|
||||
):
|
||||
"""Pre-build NIXL dlists for TP-heterogeneous slice transfers.
|
||||
|
||||
@@ -630,10 +800,14 @@ class NixlKVManager(CommonKVManager):
|
||||
[
|
||||
addrs,
|
||||
np.full(len(addrs), bytes_per_token_to_send, dtype=np.uint64),
|
||||
np.full(len(addrs), self.kv_args.gpu_id, dtype=np.uint64),
|
||||
np.full(
|
||||
len(addrs),
|
||||
_nixl_device_id(src_mem_kind, self.kv_args.gpu_id),
|
||||
dtype=np.uint64,
|
||||
),
|
||||
]
|
||||
)
|
||||
src_handle = self.agent.prep_xfer_dlist("", src_array, "VRAM")
|
||||
src_handle = self.agent.prep_xfer_dlist("", src_array, src_mem_kind)
|
||||
assert (
|
||||
src_handle is not None
|
||||
), f"prep_xfer_dlist returned None for slice src (decode_tp_size={decode_tp_size})"
|
||||
@@ -663,10 +837,14 @@ class NixlKVManager(CommonKVManager):
|
||||
[
|
||||
addrs,
|
||||
np.full(len(addrs), bytes_per_token_to_send, dtype=np.uint64),
|
||||
np.full(len(addrs), decode_kv_args.gpu_id, dtype=np.uint64),
|
||||
np.full(
|
||||
len(addrs),
|
||||
_nixl_device_id(dst_mem_kind, decode_kv_args.gpu_id),
|
||||
dtype=np.uint64,
|
||||
),
|
||||
]
|
||||
)
|
||||
dst_handle = self.agent.prep_xfer_dlist(peer_name, dst_array, "VRAM")
|
||||
dst_handle = self.agent.prep_xfer_dlist(peer_name, dst_array, dst_mem_kind)
|
||||
assert (
|
||||
dst_handle is not None
|
||||
), f"prep_xfer_dlist returned None for slice dst for peer '{peer_name}'"
|
||||
@@ -676,25 +854,119 @@ class NixlKVManager(CommonKVManager):
|
||||
head_group_idx,
|
||||
)
|
||||
|
||||
def _init_mixed_equal_tp_prep_handles(
|
||||
self,
|
||||
peer_info: KVArgsRegisterInfo,
|
||||
mem_segments: List[_KVXferMemSegment],
|
||||
):
|
||||
prepared_segments = []
|
||||
for seg in mem_segments:
|
||||
src_key = (seg.start, seg.end, seg.src_mem_kind)
|
||||
src_handle = self.prep_handles_segment_src.get(src_key)
|
||||
if src_handle is None:
|
||||
src_handle = self._prep_equal_tp_dlist(
|
||||
"",
|
||||
self.kv_args.kv_data_ptrs[seg.start : seg.end],
|
||||
self.kv_args.kv_item_lens[seg.start : seg.end],
|
||||
self.kv_args.kv_data_lens[seg.start : seg.end],
|
||||
self.kv_args.gpu_id,
|
||||
mem_kind=seg.src_mem_kind,
|
||||
)
|
||||
self.prep_handles_segment_src[src_key] = src_handle
|
||||
|
||||
dst_num_slots = (
|
||||
peer_info.dst_num_slots
|
||||
if peer_info.dst_num_slots is not None
|
||||
else self._num_slots_src
|
||||
)
|
||||
dst_kv_item_lens = peer_info.dst_kv_item_lens[seg.start : seg.end]
|
||||
dst_kv_data_lens = [
|
||||
item_len * dst_num_slots for item_len in dst_kv_item_lens
|
||||
]
|
||||
dst_handle = self._prep_equal_tp_dlist(
|
||||
peer_info.agent_name,
|
||||
peer_info.dst_kv_ptrs[seg.start : seg.end],
|
||||
dst_kv_item_lens,
|
||||
dst_kv_data_lens,
|
||||
peer_info.gpu_id,
|
||||
num_slots=peer_info.dst_num_slots,
|
||||
mem_kind=seg.dst_mem_kind,
|
||||
kv_xfer_lens=self.kv_args.kv_item_lens[seg.start : seg.end],
|
||||
)
|
||||
prepared_segments.append(
|
||||
_KVXferPreparedSegment(
|
||||
start=seg.start,
|
||||
end=seg.end,
|
||||
src_handle=src_handle,
|
||||
dst_handle=dst_handle,
|
||||
dst_num_slots=dst_num_slots,
|
||||
)
|
||||
)
|
||||
peer_info.kv_xfer_segments = prepared_segments
|
||||
|
||||
def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo):
|
||||
assert self.src_mem_kind is not None
|
||||
src_mem_kind = self.src_mem_kind
|
||||
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.
|
||||
dst_mem_kind = None
|
||||
try:
|
||||
dst_mem_kind = _homogeneous_kv_mem_kind(
|
||||
peer_info.dst_kv_mem_kinds, "destination"
|
||||
)
|
||||
except NotImplementedError:
|
||||
mem_segments = _kv_xfer_mem_segments(
|
||||
self.kv_args.kv_data_mem_kinds, peer_info.dst_kv_mem_kinds
|
||||
)
|
||||
if not mem_segments:
|
||||
raise ValueError("NIXL KV transfer has no KV memory segments")
|
||||
self._init_mixed_equal_tp_prep_handles(peer_info, mem_segments)
|
||||
return
|
||||
|
||||
peer_info.dst_homogeneous_mem_kind = dst_mem_kind
|
||||
# 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.kv_args.kv_data_ptrs,
|
||||
self.kv_args.gpu_id,
|
||||
mem_kind=src_mem_kind,
|
||||
)
|
||||
dst_num_slots = (
|
||||
peer_info.dst_num_slots
|
||||
if peer_info.dst_num_slots is not None
|
||||
else self._num_slots_src
|
||||
)
|
||||
dst_kv_item_lens = peer_info.dst_kv_item_lens
|
||||
dst_kv_data_lens = [
|
||||
item_len * dst_num_slots for item_len in dst_kv_item_lens
|
||||
]
|
||||
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,
|
||||
mem_kind=dst_mem_kind,
|
||||
kv_item_lens=dst_kv_item_lens,
|
||||
kv_data_lens=dst_kv_data_lens,
|
||||
kv_xfer_lens=self.kv_args.kv_item_lens,
|
||||
)
|
||||
else:
|
||||
self._init_hetero_tp_prep_handle(peer_info.agent_name, peer_info)
|
||||
dst_mem_kind = _homogeneous_kv_mem_kind(
|
||||
peer_info.dst_kv_mem_kinds, "destination"
|
||||
)
|
||||
peer_info.dst_homogeneous_mem_kind = dst_mem_kind
|
||||
if dst_mem_kind != "VRAM":
|
||||
raise NotImplementedError(
|
||||
"NIXL heterogeneous-TP direct-to-host KV transfer is not "
|
||||
"implemented safely yet"
|
||||
)
|
||||
self._init_hetero_tp_prep_handle(
|
||||
peer_info.agent_name,
|
||||
peer_info,
|
||||
src_mem_kind=src_mem_kind,
|
||||
dst_mem_kind=dst_mem_kind,
|
||||
)
|
||||
|
||||
def transfer_worker(self, queue: FastQueue, staging_buffer=None):
|
||||
# Per-worker staging strategy: lazy-created on first chunk so we
|
||||
@@ -758,6 +1030,8 @@ class NixlKVManager(CommonKVManager):
|
||||
: len(chunked_dst_kv_indice)
|
||||
]
|
||||
|
||||
src_prefill_kv_indices = kv_chunk.prefill_kv_indices
|
||||
|
||||
notif = (
|
||||
f"{req.room}_kv_{kv_chunk.chunk_id}"
|
||||
f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}"
|
||||
@@ -782,6 +1056,7 @@ class NixlKVManager(CommonKVManager):
|
||||
kv_xfer_handle, deferred = self._do_staging_transfer(
|
||||
staging_strategy,
|
||||
kv_chunk,
|
||||
src_prefill_kv_indices,
|
||||
req,
|
||||
dst_info,
|
||||
queue,
|
||||
@@ -801,23 +1076,41 @@ class NixlKVManager(CommonKVManager):
|
||||
if self.is_mla_backend or (
|
||||
decode_tp_size == self.attn_tp_size
|
||||
):
|
||||
kv_xfer_handle = self.send_kvcache(
|
||||
req.agent_name,
|
||||
kv_chunk.prefill_kv_indices,
|
||||
dst_info.dst_kv_ptrs,
|
||||
chunked_dst_kv_indice,
|
||||
dst_info.gpu_id,
|
||||
notif,
|
||||
)
|
||||
if dst_info.kv_xfer_segments is None:
|
||||
if dst_info.dst_homogeneous_mem_kind is None:
|
||||
raise RuntimeError(
|
||||
"Missing NIXL destination KV memory kind"
|
||||
)
|
||||
kv_xfer_handle = self.send_kvcache(
|
||||
req.agent_name,
|
||||
src_prefill_kv_indices,
|
||||
dst_info.dst_kv_ptrs,
|
||||
chunked_dst_kv_indice,
|
||||
dst_info.gpu_id,
|
||||
notif,
|
||||
dst_mem_kind=(
|
||||
dst_info.dst_homogeneous_mem_kind
|
||||
),
|
||||
)
|
||||
else:
|
||||
handles.extend(
|
||||
self.send_kvcache_mixed(
|
||||
req.agent_name,
|
||||
src_prefill_kv_indices,
|
||||
chunked_dst_kv_indice,
|
||||
notif,
|
||||
)
|
||||
)
|
||||
else:
|
||||
kv_xfer_handle = self.send_kvcache_slice(
|
||||
req.agent_name,
|
||||
kv_chunk.prefill_kv_indices,
|
||||
src_prefill_kv_indices,
|
||||
chunked_dst_kv_indice,
|
||||
notif,
|
||||
)
|
||||
|
||||
handles.append(kv_xfer_handle)
|
||||
if kv_xfer_handle is not None:
|
||||
handles.append(kv_xfer_handle)
|
||||
|
||||
if kv_chunk.is_last_chunk:
|
||||
dst_info = self.decode_kv_args_table[req.agent_name]
|
||||
@@ -863,10 +1156,16 @@ class NixlKVManager(CommonKVManager):
|
||||
continue
|
||||
|
||||
while handles:
|
||||
states = [self.agent.check_xfer_state(h) for h in handles]
|
||||
if any(s == "ERR" for s in states):
|
||||
raise RuntimeError(f"NIXL transfer encountered ERR room={room}")
|
||||
if all(s == "DONE" for s in states):
|
||||
all_done = True
|
||||
for handle in handles:
|
||||
state = self.agent.check_xfer_state(handle)
|
||||
if state == "ERR":
|
||||
raise RuntimeError(
|
||||
f"NIXL transfer encountered ERR room={room}"
|
||||
)
|
||||
if state != "DONE":
|
||||
all_done = False
|
||||
if all_done:
|
||||
break
|
||||
time.sleep(0)
|
||||
|
||||
@@ -899,13 +1198,34 @@ class NixlKVManager(CommonKVManager):
|
||||
self.update_status(room, KVPoll.Failed)
|
||||
|
||||
def register_buffer_to_engine(self):
|
||||
kv_addrs = []
|
||||
for kv_data_ptr, kv_data_len in zip(
|
||||
self.kv_args.kv_data_ptrs, self.kv_args.kv_data_lens
|
||||
self.kv_descs = []
|
||||
kv_addrs_by_mem_kind = {"VRAM": [], "DRAM": []}
|
||||
for kv_data_ptr, kv_data_len, kv_mem_kind in zip(
|
||||
self.kv_args.kv_data_ptrs,
|
||||
self.kv_args.kv_data_lens,
|
||||
self.kv_args.kv_data_mem_kinds,
|
||||
):
|
||||
kv_addrs.append((kv_data_ptr, kv_data_len, self.kv_args.gpu_id, ""))
|
||||
self.kv_descs = self.agent.register_memory(kv_addrs, "VRAM")
|
||||
logger.debug(f"Register kv tensors, len(kv_addr)= {len(kv_addrs)}")
|
||||
kv_addrs_by_mem_kind[kv_mem_kind].append(
|
||||
(
|
||||
kv_data_ptr,
|
||||
kv_data_len,
|
||||
_nixl_device_id(kv_mem_kind, self.kv_args.gpu_id),
|
||||
"",
|
||||
)
|
||||
)
|
||||
for mem_kind in ("VRAM", "DRAM"):
|
||||
kv_addrs = kv_addrs_by_mem_kind[mem_kind]
|
||||
if not kv_addrs:
|
||||
continue
|
||||
kv_descs = self.agent.register_memory(kv_addrs, mem_kind)
|
||||
logger.debug(
|
||||
f"Register kv tensors, kind={mem_kind}, len(kv_addr)= {len(kv_addrs)}"
|
||||
)
|
||||
if not kv_descs:
|
||||
raise Exception(
|
||||
f"NIXL memory registration failed for {mem_kind} kv tensors"
|
||||
)
|
||||
self.kv_descs.append(kv_descs)
|
||||
if not self.kv_descs:
|
||||
raise Exception("NIXL memory registration failed for kv tensors")
|
||||
aux_addrs = []
|
||||
@@ -957,9 +1277,12 @@ class NixlKVManager(CommonKVManager):
|
||||
dst_data_indices: npt.NDArray[np.int32],
|
||||
dst_gpu_id: int,
|
||||
notif: str,
|
||||
src_mem_kind: str = "VRAM",
|
||||
dst_mem_kind: str = "VRAM",
|
||||
):
|
||||
"""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
|
||||
@@ -1079,14 +1402,18 @@ class NixlKVManager(CommonKVManager):
|
||||
)
|
||||
)
|
||||
|
||||
src_reqs = make_req_array(src_addrs, src_lens, self.kv_args.gpu_id)
|
||||
dst_reqs = make_req_array(dst_addrs, dst_lens, dst_gpu_id)
|
||||
src_reqs = make_req_array(
|
||||
src_addrs, src_lens, _nixl_device_id(src_mem_kind, self.kv_args.gpu_id)
|
||||
)
|
||||
dst_reqs = make_req_array(
|
||||
dst_addrs, dst_lens, _nixl_device_id(dst_mem_kind, dst_gpu_id)
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"len(src_addrs): before group: {len(prefill_data_indices)}, after group: {len(src_addrs)}"
|
||||
)
|
||||
src_descs = self.agent.get_xfer_descs(src_reqs, "VRAM")
|
||||
dst_descs = self.agent.get_xfer_descs(dst_reqs, "VRAM")
|
||||
src_descs = self.agent.get_xfer_descs(src_reqs, src_mem_kind)
|
||||
dst_descs = self.agent.get_xfer_descs(dst_reqs, dst_mem_kind)
|
||||
# Transfer data
|
||||
xfer_handle = self.agent.initialize_xfer(
|
||||
"WRITE",
|
||||
@@ -1110,7 +1437,9 @@ class NixlKVManager(CommonKVManager):
|
||||
dst_kv_indices: npt.NDArray[np.int32],
|
||||
dst_gpu_id: int,
|
||||
notif: str,
|
||||
dst_mem_kind: str = "VRAM",
|
||||
):
|
||||
assert self.src_mem_kind is not None
|
||||
return self._send_kvcache_generic(
|
||||
peer_name=peer_name,
|
||||
src_data_ptrs=self.kv_args.kv_data_ptrs,
|
||||
@@ -1120,8 +1449,50 @@ class NixlKVManager(CommonKVManager):
|
||||
dst_data_indices=dst_kv_indices,
|
||||
dst_gpu_id=dst_gpu_id,
|
||||
notif=notif,
|
||||
src_mem_kind=self.src_mem_kind,
|
||||
dst_mem_kind=dst_mem_kind,
|
||||
)
|
||||
|
||||
def send_kvcache_mixed(
|
||||
self,
|
||||
peer_name: str,
|
||||
prefill_kv_indices: npt.NDArray[np.int32],
|
||||
dst_kv_indices: npt.NDArray[np.int32],
|
||||
notif: str,
|
||||
):
|
||||
info = self.decode_kv_args_table[peer_name]
|
||||
segments = info.kv_xfer_segments
|
||||
assert segments is not None
|
||||
if not segments:
|
||||
raise RuntimeError(f"Missing NIXL mixed KV transfer plan for {peer_name}")
|
||||
|
||||
num_parts = len(segments)
|
||||
handles = []
|
||||
for part_idx, seg in enumerate(segments):
|
||||
num_layers = seg.end - seg.start
|
||||
src_indices = repeat_indices_over_layers(
|
||||
prefill_kv_indices, num_layers, self._num_slots_src
|
||||
)
|
||||
dst_indices = repeat_indices_over_layers(
|
||||
dst_kv_indices, num_layers, seg.dst_num_slots
|
||||
)
|
||||
part_notif = f"{notif}_part_{part_idx}_{num_parts}"
|
||||
xfer_handle = self.agent.make_prepped_xfer(
|
||||
"WRITE",
|
||||
seg.src_handle,
|
||||
src_indices,
|
||||
seg.dst_handle,
|
||||
dst_indices,
|
||||
part_notif.encode("ascii"),
|
||||
)
|
||||
if not xfer_handle:
|
||||
raise Exception("KVSender failed to create mixed prepped transfer")
|
||||
state = self.agent.transfer(xfer_handle)
|
||||
if state == "ERR":
|
||||
raise Exception("KVSender failed to post mixed prepped transfer")
|
||||
handles.append(xfer_handle)
|
||||
return handles
|
||||
|
||||
def send_kvcache_slice(
|
||||
self,
|
||||
peer_name: str,
|
||||
@@ -1300,6 +1671,7 @@ class NixlKVManager(CommonKVManager):
|
||||
self,
|
||||
staging_strategy,
|
||||
kv_chunk: TransferKVChunk,
|
||||
src_prefill_kv_indices: npt.NDArray[np.int32],
|
||||
req: TransferInfo,
|
||||
dst_info: KVArgsRegisterInfo,
|
||||
queue: FastQueue,
|
||||
@@ -1348,7 +1720,7 @@ class NixlKVManager(CommonKVManager):
|
||||
)
|
||||
handle = self.send_kvcache_staged(
|
||||
req.agent_name,
|
||||
kv_chunk.prefill_kv_indices,
|
||||
src_prefill_kv_indices,
|
||||
dst_info.staging.base_ptr + c_offset,
|
||||
dst_info.staging.total_size - c_offset,
|
||||
dst_info.gpu_id,
|
||||
@@ -1693,6 +2065,7 @@ class NixlKVManager(CommonKVManager):
|
||||
for msg in messages:
|
||||
# Notification tag layouts (underscore-separated):
|
||||
# kv: {room}_kv_{chunk_id}_{is_last}_{pp_rank} -> 5 fields
|
||||
# kvpart:{room}_kv_{chunk_id}_{is_last}_{pp_rank}_part_{i}_{n}-> 8 fields
|
||||
# stg: {room}_stg_{chunk_id}_{is_last}_{pp_rank}_{chunk_idx}
|
||||
# _{page_start}_{num_pages}_{agent_name} -> 9 fields
|
||||
# aux: {room}_aux -> 2 fields
|
||||
@@ -1707,7 +2080,17 @@ class NixlKVManager(CommonKVManager):
|
||||
chunk_id = int(components[2])
|
||||
is_last_chunk = bool(int(components[3]))
|
||||
pp_rank = int(components[4]) if len(components) > 4 else 0
|
||||
self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
|
||||
if len(components) > 7 and components[5] == "part":
|
||||
self._track_kv_part_arrival(
|
||||
room,
|
||||
chunk_id,
|
||||
is_last_chunk,
|
||||
pp_rank,
|
||||
int(components[6]),
|
||||
int(components[7]),
|
||||
)
|
||||
else:
|
||||
self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
|
||||
elif tag == "stg":
|
||||
self._handle_stg_notification(components, room)
|
||||
elif tag == "aux":
|
||||
@@ -1778,6 +2161,45 @@ class NixlKVManager(CommonKVManager):
|
||||
):
|
||||
self._maybe_submit_last_scatter(room)
|
||||
|
||||
def _track_kv_part_arrival(
|
||||
self,
|
||||
room: int,
|
||||
chunk_id: int,
|
||||
is_last_chunk: bool,
|
||||
pp_rank: int,
|
||||
part_idx: int,
|
||||
num_parts: int,
|
||||
):
|
||||
"""Track one segment of a mixed-memory KV transfer."""
|
||||
if num_parts <= 1:
|
||||
self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
|
||||
return
|
||||
if part_idx < 0 or part_idx >= num_parts:
|
||||
raise RuntimeError(
|
||||
f"NIXL KV part index out of range for room={room}, "
|
||||
f"chunk={chunk_id}, pp_rank={pp_rank}: part={part_idx}, "
|
||||
f"num_parts={num_parts}"
|
||||
)
|
||||
|
||||
key = (pp_rank, chunk_id)
|
||||
status = self.transfer_statuses[room]
|
||||
if status.received_kv_parts_per_pp is None:
|
||||
status.received_kv_parts_per_pp = defaultdict(set)
|
||||
if status.expected_kv_parts_per_pp is None:
|
||||
status.expected_kv_parts_per_pp = {}
|
||||
expected = status.expected_kv_parts_per_pp.setdefault(key, num_parts)
|
||||
if expected != num_parts:
|
||||
raise RuntimeError(
|
||||
f"NIXL KV part count mismatch for room={room}, chunk={chunk_id}, "
|
||||
f"pp_rank={pp_rank}: got {num_parts}, expected {expected}"
|
||||
)
|
||||
parts = status.received_kv_parts_per_pp[key]
|
||||
parts.add(part_idx)
|
||||
if len(parts) == num_parts:
|
||||
status.received_kv_parts_per_pp.pop(key, None)
|
||||
status.expected_kv_parts_per_pp.pop(key, None)
|
||||
self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
|
||||
|
||||
def _handle_staging_chunk_arrived(
|
||||
self,
|
||||
room: int,
|
||||
@@ -2095,6 +2517,13 @@ class NixlKVReceiver(CommonKVReceiver):
|
||||
packed_kv_data_ptrs = b"".join(
|
||||
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs
|
||||
)
|
||||
packed_kv_data_mem_kinds = _pack_kv_mem_kinds(
|
||||
self.kv_mgr.kv_args.kv_data_mem_kinds
|
||||
)
|
||||
packed_kv_item_lens = b"".join(
|
||||
struct.pack("Q", item_len)
|
||||
for item_len in self.kv_mgr.kv_args.kv_item_lens
|
||||
)
|
||||
packed_aux_data_ptrs = b"".join(
|
||||
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.aux_data_ptrs
|
||||
)
|
||||
@@ -2145,6 +2574,8 @@ class NixlKVReceiver(CommonKVReceiver):
|
||||
packed_staging_base_ptr,
|
||||
staging_total_size_str,
|
||||
str(dst_num_slots).encode("ascii"),
|
||||
packed_kv_data_mem_kinds,
|
||||
packed_kv_item_lens,
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@@ -264,7 +264,7 @@ class PrefillBootstrapQueue:
|
||||
def finalize_bootstrap(self, req: Req) -> bool:
|
||||
"""Initialize the sender after bootstrap completes.
|
||||
Returns False if no metadata buffer is available (non-terminal)."""
|
||||
assert req.pending_bootstrap, f"finalize_bootstrap is not idempotent"
|
||||
assert req.pending_bootstrap, "finalize_bootstrap is not idempotent"
|
||||
if not self.ensure_metadata_buffer(req):
|
||||
return False
|
||||
|
||||
@@ -737,7 +737,6 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
undone_reqs: List[Req] = []
|
||||
# Check .poll() for the reqs in disagg_prefill_inflight_queue. If Success, respond to the client and remove it from the queue
|
||||
for req, poll in zip(self.disagg_prefill_inflight_queue, polls):
|
||||
|
||||
if rids_to_check is not None:
|
||||
if req.rid not in rids_to_check:
|
||||
undone_reqs.append(req)
|
||||
|
||||
@@ -57,12 +57,14 @@ class HiSparseCoordinator:
|
||||
device: str,
|
||||
tp_group,
|
||||
host_to_device_ratio: int = 2,
|
||||
swap_in_block_size: int = 960,
|
||||
):
|
||||
self.req_to_token_pool = req_to_token_pool
|
||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
||||
self.top_k = top_k
|
||||
self.device_buffer_size = device_buffer_size
|
||||
self.device = device
|
||||
self.swap_in_block_size = swap_in_block_size
|
||||
self.compress_ratio = self.token_to_kv_pool_allocator.compress_ratio
|
||||
|
||||
self.is_dsv4_hisparse = isinstance(
|
||||
@@ -815,8 +817,6 @@ class HiSparseCoordinator:
|
||||
top_k_indices = self.top_k_device_locs_buffer[:num_reqs]
|
||||
top_k_indices.fill_(-1)
|
||||
|
||||
# todo, adjustable for performance
|
||||
block_size = 1024
|
||||
swap_in_fn = (
|
||||
load_cache_to_device_buffer_dsv4_mla
|
||||
if self.is_dsv4_hisparse
|
||||
@@ -837,7 +837,7 @@ class HiSparseCoordinator:
|
||||
num_top_k=self.top_k,
|
||||
hot_buffer_size=self.device_buffer_size,
|
||||
page_size=1,
|
||||
block_size=block_size,
|
||||
block_size=self.swap_in_block_size,
|
||||
num_real_reqs=self.num_real_reqs,
|
||||
)
|
||||
return top_k_indices
|
||||
|
||||
@@ -58,6 +58,7 @@ class SparseConfig:
|
||||
top_k: int = 2048
|
||||
device_buffer_size: int = 4096
|
||||
host_to_device_ratio: int = 2
|
||||
swap_in_block_size: int = 960
|
||||
algorithm: Optional[str] = None
|
||||
backend: Optional[str] = None
|
||||
page_size: Optional[int] = None
|
||||
|
||||
@@ -62,7 +62,7 @@ def _parse_sparse_config(server_args) -> SparseConfig:
|
||||
"""Parse hierarchical sparse config from JSON string.
|
||||
|
||||
Required fields with defaults: top_k (2048), device_buffer_size (2*top_k),
|
||||
host_to_device_ratio (2).
|
||||
host_to_device_ratio (2), swap_in_block_size (960).
|
||||
Optional fields (default None): algorithm, backend, min_sparse_prompt_len,
|
||||
page_size. All remaining fields go to sparse_extra_config.
|
||||
"""
|
||||
@@ -78,11 +78,20 @@ def _parse_sparse_config(server_args) -> SparseConfig:
|
||||
top_k = extra_config.pop("top_k", 2048)
|
||||
device_buffer_size = extra_config.pop("device_buffer_size", 2 * top_k)
|
||||
host_to_device_ratio = extra_config.pop("host_to_device_ratio", 2)
|
||||
swap_in_block_size = extra_config.pop("swap_in_block_size", 960)
|
||||
|
||||
if device_buffer_size < top_k:
|
||||
raise ValueError(
|
||||
f"device_buffer_size ({device_buffer_size}) must be no smaller than top_k ({top_k})"
|
||||
)
|
||||
if not isinstance(swap_in_block_size, int) or isinstance(swap_in_block_size, bool):
|
||||
raise ValueError(
|
||||
f"swap_in_block_size must be an integer, got {swap_in_block_size!r}"
|
||||
)
|
||||
if swap_in_block_size <= 0 or swap_in_block_size > 1024:
|
||||
raise ValueError(
|
||||
f"swap_in_block_size ({swap_in_block_size}) must be in the range [1, 1024]"
|
||||
)
|
||||
|
||||
algorithm = extra_config.pop("algorithm", None)
|
||||
backend = extra_config.pop("backend", None)
|
||||
@@ -93,6 +102,7 @@ def _parse_sparse_config(server_args) -> SparseConfig:
|
||||
top_k=top_k,
|
||||
device_buffer_size=device_buffer_size,
|
||||
host_to_device_ratio=host_to_device_ratio,
|
||||
swap_in_block_size=swap_in_block_size,
|
||||
algorithm=algorithm,
|
||||
backend=backend,
|
||||
page_size=page_size,
|
||||
|
||||
@@ -856,6 +856,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
else self.tp_group.cpu_group
|
||||
),
|
||||
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
||||
swap_in_block_size=hisparse_cfg.swap_in_block_size,
|
||||
)
|
||||
|
||||
self.init_routed_experts_capturer()
|
||||
|
||||
@@ -17,10 +17,6 @@ register_cuda_ci(est_time=500, stage="base-c", runner_config="deepep-8-gpu-h200"
|
||||
DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8"
|
||||
|
||||
DEEPEP_CONFIG = '{"normal_dispatch":{"num_sms":96},"normal_combine":{"num_sms":96}}'
|
||||
DSV4_FLASH_LOADER_CONFIG = '{"enable_multithread_load": true, "num_threads": 64}'
|
||||
DSV4_HISPARSE_CONFIG = (
|
||||
'{"top_k":512,"device_buffer_size":4096,"host_to_device_ratio":2}'
|
||||
)
|
||||
|
||||
DSV4_FLASH_ENV = {
|
||||
"SGLANG_DSV4_FP4_EXPERTS": "0",
|
||||
@@ -129,104 +125,5 @@ class TestDisaggregationDSV4(SpecDecodingMixin, PDDisaggregationServerBase, GSM8
|
||||
)
|
||||
|
||||
|
||||
class TestDisaggregationDSV4HiSparseMooncake(PDDisaggregationServerBase, GSM8KMixin):
|
||||
gsm8k_accuracy_thres = 0.93
|
||||
gsm8k_num_questions = 200
|
||||
gsm8k_num_shots = 20
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
|
||||
cls.model = try_cached_model(DSV4_FLASH_MODEL)
|
||||
|
||||
cls.start_prefill()
|
||||
cls.start_decode()
|
||||
|
||||
cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill)
|
||||
cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode)
|
||||
|
||||
cls.launch_lb()
|
||||
|
||||
@classmethod
|
||||
def start_prefill(cls):
|
||||
prefill_args = [
|
||||
"--trust-remote-code",
|
||||
"--disaggregation-mode",
|
||||
"prefill",
|
||||
"--disaggregation-bootstrap-port",
|
||||
cls.bootstrap_port,
|
||||
"--tp",
|
||||
4,
|
||||
"--page-size",
|
||||
256,
|
||||
"--chunked-prefill-size",
|
||||
8192,
|
||||
"--max-running-requests",
|
||||
16,
|
||||
"--mem-fraction-static",
|
||||
0.9,
|
||||
"--skip-server-warmup",
|
||||
"--reasoning-parser",
|
||||
"deepseek-v4",
|
||||
"--tool-call-parser",
|
||||
"deepseekv4",
|
||||
"--model-loader-extra-config",
|
||||
DSV4_FLASH_LOADER_CONFIG,
|
||||
"--watchdog-timeout",
|
||||
"900",
|
||||
]
|
||||
prefill_args += cls.transfer_backend + cls.rdma_devices
|
||||
cls.process_prefill = popen_launch_pd_server(
|
||||
cls.model,
|
||||
cls.prefill_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=prefill_args,
|
||||
env=DSV4_FLASH_ENV,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def start_decode(cls):
|
||||
decode_args = [
|
||||
"--trust-remote-code",
|
||||
"--disaggregation-mode",
|
||||
"decode",
|
||||
"--disaggregation-bootstrap-port",
|
||||
cls.bootstrap_port,
|
||||
"--tp",
|
||||
4,
|
||||
"--base-gpu-id",
|
||||
4,
|
||||
"--page-size",
|
||||
256,
|
||||
"--chunked-prefill-size",
|
||||
8192,
|
||||
"--max-running-requests",
|
||||
16,
|
||||
"--mem-fraction-static",
|
||||
0.9,
|
||||
"--skip-server-warmup",
|
||||
"--reasoning-parser",
|
||||
"deepseek-v4",
|
||||
"--tool-call-parser",
|
||||
"deepseekv4",
|
||||
"--model-loader-extra-config",
|
||||
DSV4_FLASH_LOADER_CONFIG,
|
||||
"--enable-hisparse",
|
||||
"--hisparse-config",
|
||||
DSV4_HISPARSE_CONFIG,
|
||||
"--watchdog-timeout",
|
||||
"900",
|
||||
]
|
||||
decode_args += cls.transfer_backend + cls.rdma_devices
|
||||
cls.process_decode = popen_launch_pd_server(
|
||||
cls.model,
|
||||
cls.decode_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=decode_args,
|
||||
env=DSV4_FLASH_ENV,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
import unittest
|
||||
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
|
||||
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||
PDDisaggregationServerBase,
|
||||
)
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
is_in_ci,
|
||||
popen_launch_pd_server,
|
||||
try_cached_model,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=1000, stage="extra-b", runner_config="deepep-8-gpu-h200")
|
||||
|
||||
DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8"
|
||||
DSV4_FLASH_LOADER_CONFIG = '{"enable_multithread_load": true, "num_threads": 64}'
|
||||
DSV4_HISPARSE_CONFIG = (
|
||||
'{"top_k":512,"device_buffer_size":4096,"host_to_device_ratio":2}'
|
||||
)
|
||||
|
||||
DSV4_FLASH_ENV = {
|
||||
"SGLANG_DSV4_FP4_EXPERTS": "0",
|
||||
"SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "256",
|
||||
}
|
||||
DSV4_NIXL_SERVER_LAUNCH_TIMEOUT = 1800
|
||||
|
||||
|
||||
def _has_nixl():
|
||||
try:
|
||||
import nixl._api # noqa: F401
|
||||
except Exception:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
class TestDisaggregationDSV4HiSparseBase(PDDisaggregationServerBase, GSM8KMixin):
|
||||
gsm8k_accuracy_thres = 0.93
|
||||
gsm8k_num_questions = 200
|
||||
gsm8k_num_shots = 20
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
|
||||
cls.model = try_cached_model(DSV4_FLASH_MODEL)
|
||||
cls.start_prefill()
|
||||
cls.start_decode()
|
||||
|
||||
cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill)
|
||||
cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode)
|
||||
cls.launch_lb()
|
||||
|
||||
@classmethod
|
||||
def start_prefill(cls):
|
||||
prefill_args = [
|
||||
"--trust-remote-code",
|
||||
"--disaggregation-mode",
|
||||
"prefill",
|
||||
"--disaggregation-bootstrap-port",
|
||||
cls.bootstrap_port,
|
||||
"--tp",
|
||||
4,
|
||||
"--page-size",
|
||||
256,
|
||||
"--chunked-prefill-size",
|
||||
8192,
|
||||
"--max-running-requests",
|
||||
16,
|
||||
"--mem-fraction-static",
|
||||
0.9,
|
||||
"--skip-server-warmup",
|
||||
"--reasoning-parser",
|
||||
"deepseek-v4",
|
||||
"--tool-call-parser",
|
||||
"deepseekv4",
|
||||
"--model-loader-extra-config",
|
||||
DSV4_FLASH_LOADER_CONFIG,
|
||||
"--watchdog-timeout",
|
||||
"900",
|
||||
]
|
||||
prefill_args += cls.transfer_backend + cls.rdma_devices
|
||||
cls.process_prefill = popen_launch_pd_server(
|
||||
cls.model,
|
||||
cls.prefill_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=prefill_args,
|
||||
env=DSV4_FLASH_ENV,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def start_decode(cls):
|
||||
decode_args = [
|
||||
"--trust-remote-code",
|
||||
"--disaggregation-mode",
|
||||
"decode",
|
||||
"--disaggregation-bootstrap-port",
|
||||
cls.bootstrap_port,
|
||||
"--tp",
|
||||
4,
|
||||
"--base-gpu-id",
|
||||
4,
|
||||
"--page-size",
|
||||
256,
|
||||
"--chunked-prefill-size",
|
||||
8192,
|
||||
"--max-running-requests",
|
||||
16,
|
||||
"--mem-fraction-static",
|
||||
0.9,
|
||||
"--skip-server-warmup",
|
||||
"--reasoning-parser",
|
||||
"deepseek-v4",
|
||||
"--tool-call-parser",
|
||||
"deepseekv4",
|
||||
"--model-loader-extra-config",
|
||||
DSV4_FLASH_LOADER_CONFIG,
|
||||
"--enable-hisparse",
|
||||
"--hisparse-config",
|
||||
DSV4_HISPARSE_CONFIG,
|
||||
"--watchdog-timeout",
|
||||
"900",
|
||||
]
|
||||
decode_args += cls.transfer_backend + cls.rdma_devices
|
||||
cls.process_decode = popen_launch_pd_server(
|
||||
cls.model,
|
||||
cls.decode_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=decode_args,
|
||||
env=DSV4_FLASH_ENV,
|
||||
)
|
||||
|
||||
|
||||
@unittest.skipUnless(
|
||||
is_in_ci() or _has_nixl(),
|
||||
"NIXL is required for DSV4 HiSparse disaggregation coverage.",
|
||||
)
|
||||
class TestDisaggregationDSV4HiSparseNixl(TestDisaggregationDSV4HiSparseBase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
PDDisaggregationServerBase.setUpClass.__func__(cls)
|
||||
|
||||
cls.transfer_backend = ["--disaggregation-transfer-backend", "nixl"]
|
||||
cls.model = try_cached_model(DSV4_FLASH_MODEL)
|
||||
|
||||
cls.start_prefill()
|
||||
cls.start_decode()
|
||||
|
||||
cls.wait_server_ready(
|
||||
cls.prefill_url + "/health",
|
||||
timeout=DSV4_NIXL_SERVER_LAUNCH_TIMEOUT,
|
||||
process=cls.process_prefill,
|
||||
)
|
||||
cls.wait_server_ready(
|
||||
cls.decode_url + "/health",
|
||||
timeout=DSV4_NIXL_SERVER_LAUNCH_TIMEOUT,
|
||||
process=cls.process_decode,
|
||||
)
|
||||
|
||||
cls.launch_lb()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -212,6 +212,9 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
|
||||
pack_int_lists(state_dims, "I"),
|
||||
struct.pack("Q", staging_ptr),
|
||||
b"1048576",
|
||||
b"64",
|
||||
b"DRAM,DRAM",
|
||||
b"".join(struct.pack("Q", item_len) for item_len in [1024, 2048]),
|
||||
]
|
||||
|
||||
info = KVArgsRegisterInfo.from_zmq(msg)
|
||||
@@ -228,6 +231,9 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
|
||||
self.assertEqual(info.decode_tp_size, 4)
|
||||
self.assertEqual(info.decode_tp_rank, 1)
|
||||
self.assertEqual(info.dst_kv_item_len, 1024)
|
||||
self.assertEqual(info.dst_kv_item_lens, [1024, 2048])
|
||||
self.assertEqual(info.dst_num_slots, 64)
|
||||
self.assertEqual(info.dst_kv_mem_kinds, ["DRAM", "DRAM"])
|
||||
self.assertEqual(info.dst_state_item_lens, state_item_lens)
|
||||
self.assertEqual(info.dst_state_dim_per_tensor, state_dims)
|
||||
self.assertIsNotNone(info.staging)
|
||||
@@ -255,6 +261,7 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
|
||||
self.assertEqual(info.dst_state_data_ptrs, [])
|
||||
self.assertEqual(info.dst_state_item_lens, [])
|
||||
self.assertEqual(info.dst_state_dim_per_tensor, [])
|
||||
self.assertEqual(info.dst_kv_item_lens, [256])
|
||||
self.assertIsNone(info.staging)
|
||||
|
||||
|
||||
@@ -523,6 +530,33 @@ class TestNixlStaging(CustomTestCase):
|
||||
mgr.server_args = SimpleNamespace(chunked_prefill_size=4)
|
||||
return mgr
|
||||
|
||||
def test_register_buffer_to_engine_groups_kv_memory_kinds_in_one_pass(self):
|
||||
agent = StagingFakeAgent(register_result=["desc"])
|
||||
mgr = self._make_manager(agent)
|
||||
mgr.kv_args.kv_data_ptrs = [0x1000, 0x2000, 0x3000]
|
||||
mgr.kv_args.kv_data_lens = [64, 128, 256]
|
||||
mgr.kv_args.kv_data_mem_kinds = ["VRAM", "DRAM", "VRAM"]
|
||||
mgr.kv_args.aux_data_ptrs = [0x4000]
|
||||
mgr.kv_args.aux_data_lens = [32]
|
||||
mgr.kv_args.state_data_ptrs = []
|
||||
mgr.kv_args.state_data_lens = []
|
||||
|
||||
mgr.register_buffer_to_engine()
|
||||
|
||||
self.assertEqual(
|
||||
agent.register_memory_calls,
|
||||
[
|
||||
(
|
||||
[(0x1000, 64, 1, ""), (0x3000, 256, 1, "")],
|
||||
"VRAM",
|
||||
),
|
||||
([(0x2000, 128, 0, "")], "DRAM"),
|
||||
([(0x4000, 32, 0, "")], "DRAM"),
|
||||
],
|
||||
)
|
||||
self.assertEqual(mgr.kv_descs, [["desc"], ["desc"]])
|
||||
self.assertEqual(mgr.aux_descs, ["desc"])
|
||||
|
||||
def test_register_staging_memory_uses_vram_and_fails_on_empty_descs(self):
|
||||
agent = StagingFakeAgent(register_result=["staging"])
|
||||
mgr = self._make_manager(agent)
|
||||
@@ -601,7 +635,12 @@ class TestNixlStaging(CustomTestCase):
|
||||
},
|
||||
):
|
||||
handle, deferred = mgr._do_staging_transfer(
|
||||
strategy, kv_chunk, req, SimpleNamespace(), queue
|
||||
strategy,
|
||||
kv_chunk,
|
||||
kv_chunk.prefill_kv_indices,
|
||||
req,
|
||||
SimpleNamespace(),
|
||||
queue,
|
||||
)
|
||||
|
||||
self.assertIsNone(handle)
|
||||
@@ -640,6 +679,7 @@ class TestNixlStaging(CustomTestCase):
|
||||
mgr._do_staging_transfer(
|
||||
strategy,
|
||||
kv_chunk,
|
||||
kv_chunk.prefill_kv_indices,
|
||||
SimpleNamespace(room=3, agent_name="decode_agent"),
|
||||
SimpleNamespace(),
|
||||
FakeQueue(),
|
||||
@@ -666,12 +706,14 @@ class TestNixlStaging(CustomTestCase):
|
||||
agent_name="decode_agent",
|
||||
agent_metadata=b"",
|
||||
dst_kv_ptrs=[],
|
||||
dst_kv_mem_kinds=[],
|
||||
dst_aux_ptrs=[],
|
||||
dst_state_data_ptrs=[],
|
||||
gpu_id=5,
|
||||
decode_tp_size=1,
|
||||
decode_tp_rank=0,
|
||||
dst_kv_item_len=128,
|
||||
dst_kv_item_lens=[],
|
||||
staging=SimpleNamespace(base_ptr=0x8000, total_size=4096),
|
||||
)
|
||||
calls = []
|
||||
@@ -682,6 +724,7 @@ class TestNixlStaging(CustomTestCase):
|
||||
handle, deferred = mgr._do_staging_transfer(
|
||||
strategy,
|
||||
kv_chunk,
|
||||
kv_chunk.prefill_kv_indices,
|
||||
SimpleNamespace(room=3, agent_name="decode_agent"),
|
||||
dst_info,
|
||||
FakeQueue(),
|
||||
|
||||
Reference in New Issue
Block a user