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>int</td>
|
||||||
<td>Ratio of logical pool size to device pool size, determining host memory capacity</td>
|
<td>Ratio of logical pool size to device pool size, determining host memory capacity</td>
|
||||||
</tr>
|
</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>
|
</tbody>
|
||||||
</table>
|
</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
|
## Deployment
|
||||||
|
|
||||||
@@ -149,7 +154,7 @@ python3 -m sglang.launch_server \
|
|||||||
--dist-init-addr 127.0.0.1:5757 \
|
--dist-init-addr 127.0.0.1:5757 \
|
||||||
--nnodes 1 --node-rank 0 \
|
--nnodes 1 --node-rank 0 \
|
||||||
--enable-hisparse \
|
--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.
|
> **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):
|
class BaseKVSender(ABC):
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -156,7 +155,6 @@ class BaseKVSender(ABC):
|
|||||||
|
|
||||||
|
|
||||||
class BaseKVReceiver(ABC):
|
class BaseKVReceiver(ABC):
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -403,6 +403,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
kv_data_ptrs, kv_data_lens, kv_item_lens = (
|
kv_data_ptrs, kv_data_lens, kv_item_lens = (
|
||||||
transfer_kv_pool.get_contiguous_buf_infos()
|
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(
|
if self.scheduler.enable_hisparse and isinstance(
|
||||||
self.token_to_kv_pool, DeepSeekV4TokenToKVPool
|
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_ptrs += device_kv_data_ptrs[c4_layer_num:]
|
||||||
kv_data_lens += device_kv_data_lens[c4_layer_num:]
|
kv_data_lens += device_kv_data_lens[c4_layer_num:]
|
||||||
kv_item_lens += device_kv_item_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:
|
if self.draft_token_to_kv_pool is not None:
|
||||||
# We should also transfer draft model kv cache. The indices are
|
# We should also transfer draft model kv cache. The indices are
|
||||||
# always shared with a target model.
|
# always shared with a target model.
|
||||||
@@ -422,10 +428,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
kv_data_ptrs += draft_kv_data_ptrs
|
kv_data_ptrs += draft_kv_data_ptrs
|
||||||
kv_data_lens += draft_kv_data_lens
|
kv_data_lens += draft_kv_data_lens
|
||||||
kv_item_lens += draft_kv_item_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_ptrs = kv_data_ptrs
|
||||||
kv_args.kv_data_lens = kv_data_lens
|
kv_args.kv_data_lens = kv_data_lens
|
||||||
kv_args.kv_item_lens = kv_item_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.page_size = self.token_to_kv_pool.page_size
|
||||||
|
|
||||||
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
GUARD = "NixlMsgGuard".encode("ascii")
|
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
|
@dataclasses.dataclass
|
||||||
@@ -113,21 +194,42 @@ class KVArgsRegisterInfo:
|
|||||||
agent_name: str
|
agent_name: str
|
||||||
agent_metadata: bytes
|
agent_metadata: bytes
|
||||||
dst_kv_ptrs: list[int]
|
dst_kv_ptrs: list[int]
|
||||||
|
dst_kv_mem_kinds: list[str]
|
||||||
dst_aux_ptrs: list[int]
|
dst_aux_ptrs: list[int]
|
||||||
dst_state_data_ptrs: List[List[int]]
|
dst_state_data_ptrs: List[List[int]]
|
||||||
gpu_id: int
|
gpu_id: int
|
||||||
decode_tp_size: int
|
decode_tp_size: int
|
||||||
decode_tp_rank: int
|
decode_tp_rank: int
|
||||||
dst_kv_item_len: int
|
dst_kv_item_len: int
|
||||||
|
dst_kv_item_lens: list[int]
|
||||||
dst_num_slots: Optional[int] = None
|
dst_num_slots: Optional[int] = None
|
||||||
dst_state_item_lens: List[List[int]] = dataclasses.field(default_factory=list)
|
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_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
|
# Keep last: optional, parsed from a variable-length tail of the ZMQ
|
||||||
# frame in from_zmq() below, so positional construction stays stable.
|
# frame in from_zmq() below, so positional construction stays stable.
|
||||||
staging: Optional[StagingRegisterInfo] = None
|
staging: Optional[StagingRegisterInfo] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_zmq(cls, msg: List[bytes]):
|
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 = (
|
dst_state_data_ptrs = (
|
||||||
unpack_int_lists(msg[7], "Q") if len(msg) > 7 and msg[7] != b"" else []
|
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")),
|
dst_port=int(msg[2].decode("ascii")),
|
||||||
agent_name=msg[3].decode("ascii"),
|
agent_name=msg[3].decode("ascii"),
|
||||||
agent_metadata=msg[4],
|
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_aux_ptrs=list(struct.unpack(f"{len(msg[6]) // 8}Q", msg[6])),
|
||||||
dst_state_data_ptrs=dst_state_data_ptrs,
|
dst_state_data_ptrs=dst_state_data_ptrs,
|
||||||
gpu_id=int(msg[8].decode("ascii")),
|
gpu_id=int(msg[8].decode("ascii")),
|
||||||
decode_tp_size=int(msg[9].decode("ascii")),
|
decode_tp_size=int(msg[9].decode("ascii")),
|
||||||
decode_tp_rank=int(msg[10].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_num_slots=dst_num_slots,
|
||||||
dst_state_item_lens=dst_state_item_lens,
|
dst_state_item_lens=dst_state_item_lens,
|
||||||
dst_state_dim_per_tensor=dst_state_dim_per_tensor,
|
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)
|
received_state_per_pp: Set[int] = dataclasses.field(default_factory=set)
|
||||||
# Whether state data is expected (set based on state_type).
|
# Whether state data is expected (set based on state_type).
|
||||||
expects_state: bool = False
|
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):
|
def is_done(self):
|
||||||
if self.num_pp_ranks_expected is None or not self.received_aux:
|
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,
|
is_mla_backend: Optional[bool] = False,
|
||||||
):
|
):
|
||||||
super().__init__(args, disaggregation_mode, server_args, is_mla_backend)
|
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:
|
try:
|
||||||
from nixl._api import nixl_agent, nixl_agent_config, nixl_thread_sync_t
|
from nixl._api import nixl_agent, nixl_agent_config, nixl_thread_sync_t
|
||||||
except ImportError as e:
|
except ImportError as e:
|
||||||
@@ -304,6 +421,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
)
|
)
|
||||||
self.prep_handles_slice_dst: Dict[str, Tuple[Any, int, int]] = {}
|
self.prep_handles_slice_dst: Dict[str, Tuple[Any, int, int]] = {}
|
||||||
# peer_name -> (handle, num_slots, head_group_idx)
|
# 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
|
self._num_slots_src: int = 0
|
||||||
|
|
||||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||||
@@ -513,47 +631,99 @@ class NixlKVManager(CommonKVManager):
|
|||||||
def check_status(self, bootstrap_room: int):
|
def check_status(self, bootstrap_room: int):
|
||||||
return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput)
|
return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput)
|
||||||
|
|
||||||
def _init_equal_tp_prep_handle(
|
def _prep_equal_tp_dlist(
|
||||||
self,
|
self,
|
||||||
peer_name: str,
|
peer_name: str,
|
||||||
kv_ptrs: list[int],
|
kv_ptrs: list[int],
|
||||||
|
kv_item_lens: list[int],
|
||||||
|
kv_data_lens: list[int],
|
||||||
gpu_id: int,
|
gpu_id: int,
|
||||||
num_slots: Optional[int] = None,
|
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.
|
if kv_xfer_lens is None:
|
||||||
|
kv_xfer_lens = kv_item_lens
|
||||||
peer_name="" = src side; agent name = dst side. num_slots overrides the local
|
if not (
|
||||||
slot count — pass decode's count for the dst dlist (may differ from prefill).
|
len(kv_ptrs) == len(kv_item_lens) == len(kv_data_lens) == len(kv_xfer_lens)
|
||||||
Uses prefill's kv_item_lens as stride; requires equal per-slot byte size (equal-TP or MLA).
|
):
|
||||||
"""
|
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 = []
|
arrays = []
|
||||||
# torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set).
|
# torch.int exceeds np.int64 range on Intel XPU (addresses have bit 63 set).
|
||||||
# Convert once at entry; all downstream arithmetic stays in uint64.
|
# Convert once at entry; all downstream arithmetic stays in uint64.
|
||||||
kv_ptrs_u64 = np.array(kv_ptrs, dtype=np.uint64)
|
kv_ptrs_u64 = np.array(kv_ptrs, dtype=np.uint64)
|
||||||
for base_ptr, item_len, data_len in zip(
|
for base_ptr, item_len, data_len, xfer_len in zip(
|
||||||
kv_ptrs_u64, self.kv_args.kv_item_lens, self.kv_args.kv_data_lens
|
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)
|
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
|
addrs = np.arange(n, dtype=np.uint64) * np.uint64(item_len) + base_ptr
|
||||||
arrays.append(
|
arrays.append(
|
||||||
np.column_stack(
|
np.column_stack(
|
||||||
[
|
[
|
||||||
addrs,
|
addrs,
|
||||||
np.full(n, item_len, dtype=np.uint64),
|
np.full(n, xfer_len, dtype=np.uint64),
|
||||||
np.full(n, gpu_id, dtype=np.uint64),
|
np.full(n, device_id, dtype=np.uint64),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
self.prep_handles[peer_name] = self.agent.prep_xfer_dlist(
|
prep_handle = self.agent.prep_xfer_dlist(peer_name, np.vstack(arrays), mem_kind)
|
||||||
peer_name, np.vstack(arrays), "VRAM"
|
|
||||||
)
|
|
||||||
assert (
|
assert (
|
||||||
self.prep_handles[peer_name] is not None
|
prep_handle is not None
|
||||||
), f"prep_xfer_dlist returned None for peer '{peer_name}'"
|
), 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(
|
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.
|
"""Pre-build NIXL dlists for TP-heterogeneous slice transfers.
|
||||||
|
|
||||||
@@ -630,10 +800,14 @@ class NixlKVManager(CommonKVManager):
|
|||||||
[
|
[
|
||||||
addrs,
|
addrs,
|
||||||
np.full(len(addrs), bytes_per_token_to_send, dtype=np.uint64),
|
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 (
|
assert (
|
||||||
src_handle is not None
|
src_handle is not None
|
||||||
), f"prep_xfer_dlist returned None for slice src (decode_tp_size={decode_tp_size})"
|
), f"prep_xfer_dlist returned None for slice src (decode_tp_size={decode_tp_size})"
|
||||||
@@ -663,10 +837,14 @@ class NixlKVManager(CommonKVManager):
|
|||||||
[
|
[
|
||||||
addrs,
|
addrs,
|
||||||
np.full(len(addrs), bytes_per_token_to_send, dtype=np.uint64),
|
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 (
|
assert (
|
||||||
dst_handle is not None
|
dst_handle is not None
|
||||||
), f"prep_xfer_dlist returned None for slice dst for peer '{peer_name}'"
|
), f"prep_xfer_dlist returned None for slice dst for peer '{peer_name}'"
|
||||||
@@ -676,25 +854,119 @@ class NixlKVManager(CommonKVManager):
|
|||||||
head_group_idx,
|
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):
|
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:
|
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:
|
dst_mem_kind = None
|
||||||
# equal_tp guarantees identical heads-per-rank (same item_len);
|
try:
|
||||||
# MLA latent shape is TP-invariant.
|
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
|
# Build the shared src dlist on the first equal-TP/MLA peer; later
|
||||||
# peers reuse it. Skipped entirely on heterogeneous-TP-only setups.
|
# peers reuse it. Skipped entirely on heterogeneous-TP-only setups.
|
||||||
if "" not in self.prep_handles:
|
if "" not in self.prep_handles:
|
||||||
self._init_equal_tp_prep_handle(
|
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(
|
self._init_equal_tp_prep_handle(
|
||||||
peer_info.agent_name,
|
peer_info.agent_name,
|
||||||
peer_info.dst_kv_ptrs,
|
peer_info.dst_kv_ptrs,
|
||||||
peer_info.gpu_id,
|
peer_info.gpu_id,
|
||||||
num_slots=peer_info.dst_num_slots,
|
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:
|
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):
|
def transfer_worker(self, queue: FastQueue, staging_buffer=None):
|
||||||
# Per-worker staging strategy: lazy-created on first chunk so we
|
# Per-worker staging strategy: lazy-created on first chunk so we
|
||||||
@@ -758,6 +1030,8 @@ class NixlKVManager(CommonKVManager):
|
|||||||
: len(chunked_dst_kv_indice)
|
: len(chunked_dst_kv_indice)
|
||||||
]
|
]
|
||||||
|
|
||||||
|
src_prefill_kv_indices = kv_chunk.prefill_kv_indices
|
||||||
|
|
||||||
notif = (
|
notif = (
|
||||||
f"{req.room}_kv_{kv_chunk.chunk_id}"
|
f"{req.room}_kv_{kv_chunk.chunk_id}"
|
||||||
f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}"
|
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(
|
kv_xfer_handle, deferred = self._do_staging_transfer(
|
||||||
staging_strategy,
|
staging_strategy,
|
||||||
kv_chunk,
|
kv_chunk,
|
||||||
|
src_prefill_kv_indices,
|
||||||
req,
|
req,
|
||||||
dst_info,
|
dst_info,
|
||||||
queue,
|
queue,
|
||||||
@@ -801,22 +1076,40 @@ class NixlKVManager(CommonKVManager):
|
|||||||
if self.is_mla_backend or (
|
if self.is_mla_backend or (
|
||||||
decode_tp_size == self.attn_tp_size
|
decode_tp_size == self.attn_tp_size
|
||||||
):
|
):
|
||||||
|
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(
|
kv_xfer_handle = self.send_kvcache(
|
||||||
req.agent_name,
|
req.agent_name,
|
||||||
kv_chunk.prefill_kv_indices,
|
src_prefill_kv_indices,
|
||||||
dst_info.dst_kv_ptrs,
|
dst_info.dst_kv_ptrs,
|
||||||
chunked_dst_kv_indice,
|
chunked_dst_kv_indice,
|
||||||
dst_info.gpu_id,
|
dst_info.gpu_id,
|
||||||
notif,
|
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:
|
else:
|
||||||
kv_xfer_handle = self.send_kvcache_slice(
|
kv_xfer_handle = self.send_kvcache_slice(
|
||||||
req.agent_name,
|
req.agent_name,
|
||||||
kv_chunk.prefill_kv_indices,
|
src_prefill_kv_indices,
|
||||||
chunked_dst_kv_indice,
|
chunked_dst_kv_indice,
|
||||||
notif,
|
notif,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if kv_xfer_handle is not None:
|
||||||
handles.append(kv_xfer_handle)
|
handles.append(kv_xfer_handle)
|
||||||
|
|
||||||
if kv_chunk.is_last_chunk:
|
if kv_chunk.is_last_chunk:
|
||||||
@@ -863,10 +1156,16 @@ class NixlKVManager(CommonKVManager):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
while handles:
|
while handles:
|
||||||
states = [self.agent.check_xfer_state(h) for h in handles]
|
all_done = True
|
||||||
if any(s == "ERR" for s in states):
|
for handle in handles:
|
||||||
raise RuntimeError(f"NIXL transfer encountered ERR room={room}")
|
state = self.agent.check_xfer_state(handle)
|
||||||
if all(s == "DONE" for s in states):
|
if state == "ERR":
|
||||||
|
raise RuntimeError(
|
||||||
|
f"NIXL transfer encountered ERR room={room}"
|
||||||
|
)
|
||||||
|
if state != "DONE":
|
||||||
|
all_done = False
|
||||||
|
if all_done:
|
||||||
break
|
break
|
||||||
time.sleep(0)
|
time.sleep(0)
|
||||||
|
|
||||||
@@ -899,13 +1198,34 @@ class NixlKVManager(CommonKVManager):
|
|||||||
self.update_status(room, KVPoll.Failed)
|
self.update_status(room, KVPoll.Failed)
|
||||||
|
|
||||||
def register_buffer_to_engine(self):
|
def register_buffer_to_engine(self):
|
||||||
kv_addrs = []
|
self.kv_descs = []
|
||||||
for kv_data_ptr, kv_data_len in zip(
|
kv_addrs_by_mem_kind = {"VRAM": [], "DRAM": []}
|
||||||
self.kv_args.kv_data_ptrs, self.kv_args.kv_data_lens
|
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, ""))
|
kv_addrs_by_mem_kind[kv_mem_kind].append(
|
||||||
self.kv_descs = self.agent.register_memory(kv_addrs, "VRAM")
|
(
|
||||||
logger.debug(f"Register kv tensors, len(kv_addr)= {len(kv_addrs)}")
|
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:
|
if not self.kv_descs:
|
||||||
raise Exception("NIXL memory registration failed for kv tensors")
|
raise Exception("NIXL memory registration failed for kv tensors")
|
||||||
aux_addrs = []
|
aux_addrs = []
|
||||||
@@ -957,9 +1277,12 @@ class NixlKVManager(CommonKVManager):
|
|||||||
dst_data_indices: npt.NDArray[np.int32],
|
dst_data_indices: npt.NDArray[np.int32],
|
||||||
dst_gpu_id: int,
|
dst_gpu_id: int,
|
||||||
notif: str,
|
notif: str,
|
||||||
|
src_mem_kind: str = "VRAM",
|
||||||
|
dst_mem_kind: str = "VRAM",
|
||||||
):
|
):
|
||||||
"""Generic KV cache transfer supporting both MHA and MLA architectures.
|
"""Generic KV cache transfer supporting both MHA and MLA architectures.
|
||||||
Used by both send_kvcache and maybe_send_extra."""
|
Used by both send_kvcache and maybe_send_extra."""
|
||||||
|
|
||||||
# Prepped path (KV only; state transfers use the non-prepped path below).
|
# Prepped path (KV only; state transfers use the non-prepped path below).
|
||||||
if (
|
if (
|
||||||
src_data_ptrs is self.kv_args.kv_data_ptrs
|
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)
|
src_reqs = make_req_array(
|
||||||
dst_reqs = make_req_array(dst_addrs, dst_lens, dst_gpu_id)
|
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(
|
logger.debug(
|
||||||
f"len(src_addrs): before group: {len(prefill_data_indices)}, after group: {len(src_addrs)}"
|
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")
|
src_descs = self.agent.get_xfer_descs(src_reqs, src_mem_kind)
|
||||||
dst_descs = self.agent.get_xfer_descs(dst_reqs, "VRAM")
|
dst_descs = self.agent.get_xfer_descs(dst_reqs, dst_mem_kind)
|
||||||
# Transfer data
|
# Transfer data
|
||||||
xfer_handle = self.agent.initialize_xfer(
|
xfer_handle = self.agent.initialize_xfer(
|
||||||
"WRITE",
|
"WRITE",
|
||||||
@@ -1110,7 +1437,9 @@ class NixlKVManager(CommonKVManager):
|
|||||||
dst_kv_indices: npt.NDArray[np.int32],
|
dst_kv_indices: npt.NDArray[np.int32],
|
||||||
dst_gpu_id: int,
|
dst_gpu_id: int,
|
||||||
notif: str,
|
notif: str,
|
||||||
|
dst_mem_kind: str = "VRAM",
|
||||||
):
|
):
|
||||||
|
assert self.src_mem_kind is not None
|
||||||
return self._send_kvcache_generic(
|
return self._send_kvcache_generic(
|
||||||
peer_name=peer_name,
|
peer_name=peer_name,
|
||||||
src_data_ptrs=self.kv_args.kv_data_ptrs,
|
src_data_ptrs=self.kv_args.kv_data_ptrs,
|
||||||
@@ -1120,8 +1449,50 @@ class NixlKVManager(CommonKVManager):
|
|||||||
dst_data_indices=dst_kv_indices,
|
dst_data_indices=dst_kv_indices,
|
||||||
dst_gpu_id=dst_gpu_id,
|
dst_gpu_id=dst_gpu_id,
|
||||||
notif=notif,
|
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(
|
def send_kvcache_slice(
|
||||||
self,
|
self,
|
||||||
peer_name: str,
|
peer_name: str,
|
||||||
@@ -1300,6 +1671,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
self,
|
self,
|
||||||
staging_strategy,
|
staging_strategy,
|
||||||
kv_chunk: TransferKVChunk,
|
kv_chunk: TransferKVChunk,
|
||||||
|
src_prefill_kv_indices: npt.NDArray[np.int32],
|
||||||
req: TransferInfo,
|
req: TransferInfo,
|
||||||
dst_info: KVArgsRegisterInfo,
|
dst_info: KVArgsRegisterInfo,
|
||||||
queue: FastQueue,
|
queue: FastQueue,
|
||||||
@@ -1348,7 +1720,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
)
|
)
|
||||||
handle = self.send_kvcache_staged(
|
handle = self.send_kvcache_staged(
|
||||||
req.agent_name,
|
req.agent_name,
|
||||||
kv_chunk.prefill_kv_indices,
|
src_prefill_kv_indices,
|
||||||
dst_info.staging.base_ptr + c_offset,
|
dst_info.staging.base_ptr + c_offset,
|
||||||
dst_info.staging.total_size - c_offset,
|
dst_info.staging.total_size - c_offset,
|
||||||
dst_info.gpu_id,
|
dst_info.gpu_id,
|
||||||
@@ -1693,6 +2065,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
for msg in messages:
|
for msg in messages:
|
||||||
# Notification tag layouts (underscore-separated):
|
# Notification tag layouts (underscore-separated):
|
||||||
# kv: {room}_kv_{chunk_id}_{is_last}_{pp_rank} -> 5 fields
|
# 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}
|
# stg: {room}_stg_{chunk_id}_{is_last}_{pp_rank}_{chunk_idx}
|
||||||
# _{page_start}_{num_pages}_{agent_name} -> 9 fields
|
# _{page_start}_{num_pages}_{agent_name} -> 9 fields
|
||||||
# aux: {room}_aux -> 2 fields
|
# aux: {room}_aux -> 2 fields
|
||||||
@@ -1707,6 +2080,16 @@ class NixlKVManager(CommonKVManager):
|
|||||||
chunk_id = int(components[2])
|
chunk_id = int(components[2])
|
||||||
is_last_chunk = bool(int(components[3]))
|
is_last_chunk = bool(int(components[3]))
|
||||||
pp_rank = int(components[4]) if len(components) > 4 else 0
|
pp_rank = int(components[4]) if len(components) > 4 else 0
|
||||||
|
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)
|
self._track_kv_arrival(room, chunk_id, is_last_chunk, pp_rank)
|
||||||
elif tag == "stg":
|
elif tag == "stg":
|
||||||
self._handle_stg_notification(components, room)
|
self._handle_stg_notification(components, room)
|
||||||
@@ -1778,6 +2161,45 @@ class NixlKVManager(CommonKVManager):
|
|||||||
):
|
):
|
||||||
self._maybe_submit_last_scatter(room)
|
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(
|
def _handle_staging_chunk_arrived(
|
||||||
self,
|
self,
|
||||||
room: int,
|
room: int,
|
||||||
@@ -2095,6 +2517,13 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
packed_kv_data_ptrs = b"".join(
|
packed_kv_data_ptrs = b"".join(
|
||||||
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.kv_data_ptrs
|
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(
|
packed_aux_data_ptrs = b"".join(
|
||||||
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.aux_data_ptrs
|
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,
|
packed_staging_base_ptr,
|
||||||
staging_total_size_str,
|
staging_total_size_str,
|
||||||
str(dst_num_slots).encode("ascii"),
|
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:
|
def finalize_bootstrap(self, req: Req) -> bool:
|
||||||
"""Initialize the sender after bootstrap completes.
|
"""Initialize the sender after bootstrap completes.
|
||||||
Returns False if no metadata buffer is available (non-terminal)."""
|
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):
|
if not self.ensure_metadata_buffer(req):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -737,7 +737,6 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
undone_reqs: List[Req] = []
|
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
|
# 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):
|
for req, poll in zip(self.disagg_prefill_inflight_queue, polls):
|
||||||
|
|
||||||
if rids_to_check is not None:
|
if rids_to_check is not None:
|
||||||
if req.rid not in rids_to_check:
|
if req.rid not in rids_to_check:
|
||||||
undone_reqs.append(req)
|
undone_reqs.append(req)
|
||||||
|
|||||||
@@ -57,12 +57,14 @@ class HiSparseCoordinator:
|
|||||||
device: str,
|
device: str,
|
||||||
tp_group,
|
tp_group,
|
||||||
host_to_device_ratio: int = 2,
|
host_to_device_ratio: int = 2,
|
||||||
|
swap_in_block_size: int = 960,
|
||||||
):
|
):
|
||||||
self.req_to_token_pool = req_to_token_pool
|
self.req_to_token_pool = req_to_token_pool
|
||||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
self.device_buffer_size = device_buffer_size
|
self.device_buffer_size = device_buffer_size
|
||||||
self.device = device
|
self.device = device
|
||||||
|
self.swap_in_block_size = swap_in_block_size
|
||||||
self.compress_ratio = self.token_to_kv_pool_allocator.compress_ratio
|
self.compress_ratio = self.token_to_kv_pool_allocator.compress_ratio
|
||||||
|
|
||||||
self.is_dsv4_hisparse = isinstance(
|
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 = self.top_k_device_locs_buffer[:num_reqs]
|
||||||
top_k_indices.fill_(-1)
|
top_k_indices.fill_(-1)
|
||||||
|
|
||||||
# todo, adjustable for performance
|
|
||||||
block_size = 1024
|
|
||||||
swap_in_fn = (
|
swap_in_fn = (
|
||||||
load_cache_to_device_buffer_dsv4_mla
|
load_cache_to_device_buffer_dsv4_mla
|
||||||
if self.is_dsv4_hisparse
|
if self.is_dsv4_hisparse
|
||||||
@@ -837,7 +837,7 @@ class HiSparseCoordinator:
|
|||||||
num_top_k=self.top_k,
|
num_top_k=self.top_k,
|
||||||
hot_buffer_size=self.device_buffer_size,
|
hot_buffer_size=self.device_buffer_size,
|
||||||
page_size=1,
|
page_size=1,
|
||||||
block_size=block_size,
|
block_size=self.swap_in_block_size,
|
||||||
num_real_reqs=self.num_real_reqs,
|
num_real_reqs=self.num_real_reqs,
|
||||||
)
|
)
|
||||||
return top_k_indices
|
return top_k_indices
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ class SparseConfig:
|
|||||||
top_k: int = 2048
|
top_k: int = 2048
|
||||||
device_buffer_size: int = 4096
|
device_buffer_size: int = 4096
|
||||||
host_to_device_ratio: int = 2
|
host_to_device_ratio: int = 2
|
||||||
|
swap_in_block_size: int = 960
|
||||||
algorithm: Optional[str] = None
|
algorithm: Optional[str] = None
|
||||||
backend: Optional[str] = None
|
backend: Optional[str] = None
|
||||||
page_size: Optional[int] = None
|
page_size: Optional[int] = None
|
||||||
|
|||||||
@@ -62,7 +62,7 @@ def _parse_sparse_config(server_args) -> SparseConfig:
|
|||||||
"""Parse hierarchical sparse config from JSON string.
|
"""Parse hierarchical sparse config from JSON string.
|
||||||
|
|
||||||
Required fields with defaults: top_k (2048), device_buffer_size (2*top_k),
|
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,
|
Optional fields (default None): algorithm, backend, min_sparse_prompt_len,
|
||||||
page_size. All remaining fields go to sparse_extra_config.
|
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)
|
top_k = extra_config.pop("top_k", 2048)
|
||||||
device_buffer_size = extra_config.pop("device_buffer_size", 2 * top_k)
|
device_buffer_size = extra_config.pop("device_buffer_size", 2 * top_k)
|
||||||
host_to_device_ratio = extra_config.pop("host_to_device_ratio", 2)
|
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:
|
if device_buffer_size < top_k:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"device_buffer_size ({device_buffer_size}) must be no smaller than top_k ({top_k})"
|
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)
|
algorithm = extra_config.pop("algorithm", None)
|
||||||
backend = extra_config.pop("backend", None)
|
backend = extra_config.pop("backend", None)
|
||||||
@@ -93,6 +102,7 @@ def _parse_sparse_config(server_args) -> SparseConfig:
|
|||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
device_buffer_size=device_buffer_size,
|
device_buffer_size=device_buffer_size,
|
||||||
host_to_device_ratio=host_to_device_ratio,
|
host_to_device_ratio=host_to_device_ratio,
|
||||||
|
swap_in_block_size=swap_in_block_size,
|
||||||
algorithm=algorithm,
|
algorithm=algorithm,
|
||||||
backend=backend,
|
backend=backend,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
|
|||||||
@@ -856,6 +856,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
else self.tp_group.cpu_group
|
else self.tp_group.cpu_group
|
||||||
),
|
),
|
||||||
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
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()
|
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"
|
DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8"
|
||||||
|
|
||||||
DEEPEP_CONFIG = '{"normal_dispatch":{"num_sms":96},"normal_combine":{"num_sms":96}}'
|
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 = {
|
DSV4_FLASH_ENV = {
|
||||||
"SGLANG_DSV4_FP4_EXPERTS": "0",
|
"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__":
|
if __name__ == "__main__":
|
||||||
unittest.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"),
|
pack_int_lists(state_dims, "I"),
|
||||||
struct.pack("Q", staging_ptr),
|
struct.pack("Q", staging_ptr),
|
||||||
b"1048576",
|
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)
|
info = KVArgsRegisterInfo.from_zmq(msg)
|
||||||
@@ -228,6 +231,9 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
|
|||||||
self.assertEqual(info.decode_tp_size, 4)
|
self.assertEqual(info.decode_tp_size, 4)
|
||||||
self.assertEqual(info.decode_tp_rank, 1)
|
self.assertEqual(info.decode_tp_rank, 1)
|
||||||
self.assertEqual(info.dst_kv_item_len, 1024)
|
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_item_lens, state_item_lens)
|
||||||
self.assertEqual(info.dst_state_dim_per_tensor, state_dims)
|
self.assertEqual(info.dst_state_dim_per_tensor, state_dims)
|
||||||
self.assertIsNotNone(info.staging)
|
self.assertIsNotNone(info.staging)
|
||||||
@@ -255,6 +261,7 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase):
|
|||||||
self.assertEqual(info.dst_state_data_ptrs, [])
|
self.assertEqual(info.dst_state_data_ptrs, [])
|
||||||
self.assertEqual(info.dst_state_item_lens, [])
|
self.assertEqual(info.dst_state_item_lens, [])
|
||||||
self.assertEqual(info.dst_state_dim_per_tensor, [])
|
self.assertEqual(info.dst_state_dim_per_tensor, [])
|
||||||
|
self.assertEqual(info.dst_kv_item_lens, [256])
|
||||||
self.assertIsNone(info.staging)
|
self.assertIsNone(info.staging)
|
||||||
|
|
||||||
|
|
||||||
@@ -523,6 +530,33 @@ class TestNixlStaging(CustomTestCase):
|
|||||||
mgr.server_args = SimpleNamespace(chunked_prefill_size=4)
|
mgr.server_args = SimpleNamespace(chunked_prefill_size=4)
|
||||||
return mgr
|
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):
|
def test_register_staging_memory_uses_vram_and_fails_on_empty_descs(self):
|
||||||
agent = StagingFakeAgent(register_result=["staging"])
|
agent = StagingFakeAgent(register_result=["staging"])
|
||||||
mgr = self._make_manager(agent)
|
mgr = self._make_manager(agent)
|
||||||
@@ -601,7 +635,12 @@ class TestNixlStaging(CustomTestCase):
|
|||||||
},
|
},
|
||||||
):
|
):
|
||||||
handle, deferred = mgr._do_staging_transfer(
|
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)
|
self.assertIsNone(handle)
|
||||||
@@ -640,6 +679,7 @@ class TestNixlStaging(CustomTestCase):
|
|||||||
mgr._do_staging_transfer(
|
mgr._do_staging_transfer(
|
||||||
strategy,
|
strategy,
|
||||||
kv_chunk,
|
kv_chunk,
|
||||||
|
kv_chunk.prefill_kv_indices,
|
||||||
SimpleNamespace(room=3, agent_name="decode_agent"),
|
SimpleNamespace(room=3, agent_name="decode_agent"),
|
||||||
SimpleNamespace(),
|
SimpleNamespace(),
|
||||||
FakeQueue(),
|
FakeQueue(),
|
||||||
@@ -666,12 +706,14 @@ class TestNixlStaging(CustomTestCase):
|
|||||||
agent_name="decode_agent",
|
agent_name="decode_agent",
|
||||||
agent_metadata=b"",
|
agent_metadata=b"",
|
||||||
dst_kv_ptrs=[],
|
dst_kv_ptrs=[],
|
||||||
|
dst_kv_mem_kinds=[],
|
||||||
dst_aux_ptrs=[],
|
dst_aux_ptrs=[],
|
||||||
dst_state_data_ptrs=[],
|
dst_state_data_ptrs=[],
|
||||||
gpu_id=5,
|
gpu_id=5,
|
||||||
decode_tp_size=1,
|
decode_tp_size=1,
|
||||||
decode_tp_rank=0,
|
decode_tp_rank=0,
|
||||||
dst_kv_item_len=128,
|
dst_kv_item_len=128,
|
||||||
|
dst_kv_item_lens=[],
|
||||||
staging=SimpleNamespace(base_ptr=0x8000, total_size=4096),
|
staging=SimpleNamespace(base_ptr=0x8000, total_size=4096),
|
||||||
)
|
)
|
||||||
calls = []
|
calls = []
|
||||||
@@ -682,6 +724,7 @@ class TestNixlStaging(CustomTestCase):
|
|||||||
handle, deferred = mgr._do_staging_transfer(
|
handle, deferred = mgr._do_staging_transfer(
|
||||||
strategy,
|
strategy,
|
||||||
kv_chunk,
|
kv_chunk,
|
||||||
|
kv_chunk.prefill_kv_indices,
|
||||||
SimpleNamespace(room=3, agent_name="decode_agent"),
|
SimpleNamespace(room=3, agent_name="decode_agent"),
|
||||||
dst_info,
|
dst_info,
|
||||||
FakeQueue(),
|
FakeQueue(),
|
||||||
|
|||||||
Reference in New Issue
Block a user