add hybrid model PD to NIXL connector (#16229)
Signed-off-by: Neal Vaidya <nealv@nvidia.com> Co-authored-by: ishandhanani <82981111+ishandhanani@users.noreply.github.com>
This commit is contained in:
co-authored by
ishandhanani
parent
c850a8a41a
commit
f1ff697494
@@ -41,12 +41,19 @@ class TransferInfo:
|
|||||||
dst_kv_indices: npt.NDArray[np.int32]
|
dst_kv_indices: npt.NDArray[np.int32]
|
||||||
dst_aux_index: int
|
dst_aux_index: int
|
||||||
required_dst_info_num: int
|
required_dst_info_num: int
|
||||||
|
dst_state_indices: List[int]
|
||||||
|
|
||||||
def is_dummy(self):
|
def is_dummy(self):
|
||||||
return self.dst_kv_indices.size == 0
|
return self.dst_kv_indices.size == 0
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_zmq(cls, msg: List[bytes]):
|
def from_zmq(cls, msg: List[bytes]):
|
||||||
|
# Parse state_indices from msg[7] if present
|
||||||
|
if len(msg) > 7 and msg[7] != b"":
|
||||||
|
dst_state_indices = list(np.frombuffer(msg[7], dtype=np.int32))
|
||||||
|
else:
|
||||||
|
dst_state_indices = []
|
||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
room=int(msg[0].decode("ascii")),
|
room=int(msg[0].decode("ascii")),
|
||||||
endpoint=msg[1].decode("ascii"),
|
endpoint=msg[1].decode("ascii"),
|
||||||
@@ -55,6 +62,7 @@ class TransferInfo:
|
|||||||
dst_kv_indices=np.frombuffer(msg[4], dtype=np.int32),
|
dst_kv_indices=np.frombuffer(msg[4], dtype=np.int32),
|
||||||
dst_aux_index=int(msg[5].decode("ascii")),
|
dst_aux_index=int(msg[5].decode("ascii")),
|
||||||
required_dst_info_num=int(msg[6].decode("ascii")),
|
required_dst_info_num=int(msg[6].decode("ascii")),
|
||||||
|
dst_state_indices=dst_state_indices,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -69,6 +77,7 @@ class KVArgsRegisterInfo:
|
|||||||
agent_metadata: bytes
|
agent_metadata: bytes
|
||||||
dst_kv_ptrs: list[int]
|
dst_kv_ptrs: list[int]
|
||||||
dst_aux_ptrs: list[int]
|
dst_aux_ptrs: list[int]
|
||||||
|
dst_state_data_ptrs: list[int]
|
||||||
gpu_id: int
|
gpu_id: int
|
||||||
decode_tp_size: int
|
decode_tp_size: int
|
||||||
decode_tp_rank: int
|
decode_tp_rank: int
|
||||||
@@ -76,6 +85,12 @@ class KVArgsRegisterInfo:
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_zmq(cls, msg: List[bytes]):
|
def from_zmq(cls, msg: List[bytes]):
|
||||||
|
# Parse state_data_ptrs from msg[7] if present
|
||||||
|
if len(msg) > 7 and msg[7] != b"":
|
||||||
|
dst_state_data_ptrs = list(struct.unpack(f"{len(msg[7]) // 8}Q", msg[7]))
|
||||||
|
else:
|
||||||
|
dst_state_data_ptrs = []
|
||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
room=str(msg[0].decode("ascii")),
|
room=str(msg[0].decode("ascii")),
|
||||||
endpoint=msg[1].decode("ascii"),
|
endpoint=msg[1].decode("ascii"),
|
||||||
@@ -84,10 +99,11 @@ class KVArgsRegisterInfo:
|
|||||||
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=list(struct.unpack(f"{len(msg[5]) // 8}Q", msg[5])),
|
||||||
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])),
|
||||||
gpu_id=int(msg[7].decode("ascii")),
|
dst_state_data_ptrs=dst_state_data_ptrs,
|
||||||
decode_tp_size=int(msg[8].decode("ascii")),
|
gpu_id=int(msg[8].decode("ascii")),
|
||||||
decode_tp_rank=int(msg[9].decode("ascii")),
|
decode_tp_size=int(msg[9].decode("ascii")),
|
||||||
dst_kv_item_len=int(msg[10].decode("ascii")),
|
decode_tp_rank=int(msg[10].decode("ascii")),
|
||||||
|
dst_kv_item_len=int(msg[11].decode("ascii")),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -105,6 +121,10 @@ class TransferStatus:
|
|||||||
num_pp_ranks_expected: Optional[int] = None
|
num_pp_ranks_expected: Optional[int] = None
|
||||||
# Whether aux data has been received.
|
# Whether aux data has been received.
|
||||||
received_aux: bool = False
|
received_aux: bool = False
|
||||||
|
# PP ranks that have sent state data (state is layer-specific, each PP rank sends its portion).
|
||||||
|
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
|
||||||
# Mark as failed
|
# Mark as failed
|
||||||
is_failure: bool = False
|
is_failure: bool = False
|
||||||
|
|
||||||
@@ -113,6 +133,12 @@ class TransferStatus:
|
|||||||
return True
|
return True
|
||||||
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:
|
||||||
return False
|
return False
|
||||||
|
# If state data is expected, check all PP ranks have sent it
|
||||||
|
if (
|
||||||
|
self.expects_state
|
||||||
|
and len(self.received_state_per_pp) < self.num_pp_ranks_expected
|
||||||
|
):
|
||||||
|
return False
|
||||||
# All PP ranks must have reported their expected count
|
# All PP ranks must have reported their expected count
|
||||||
if len(self.expected_kvs_per_pp) < self.num_pp_ranks_expected:
|
if len(self.expected_kvs_per_pp) < self.num_pp_ranks_expected:
|
||||||
return False
|
return False
|
||||||
@@ -318,6 +344,22 @@ class NixlKVManager(CommonKVManager):
|
|||||||
if not self.aux_descs:
|
if not self.aux_descs:
|
||||||
raise Exception("NIXL memory registration failed for aux tensors")
|
raise Exception("NIXL memory registration failed for aux tensors")
|
||||||
|
|
||||||
|
# Register state/extra pool data buffers if present
|
||||||
|
if self.kv_args.state_data_ptrs and self.kv_args.state_data_lens:
|
||||||
|
state_addrs = []
|
||||||
|
for state_data_ptr, state_data_len in zip(
|
||||||
|
self.kv_args.state_data_ptrs, self.kv_args.state_data_lens
|
||||||
|
):
|
||||||
|
state_addrs.append(
|
||||||
|
(state_data_ptr, state_data_len, self.kv_args.gpu_id, "")
|
||||||
|
)
|
||||||
|
self.state_descs = self.agent.register_memory(state_addrs, "VRAM")
|
||||||
|
logger.debug(
|
||||||
|
f"Register state tensors, len(state_addrs)= {len(state_addrs)}"
|
||||||
|
)
|
||||||
|
if not self.state_descs:
|
||||||
|
raise Exception("NIXL memory registration failed for state tensors")
|
||||||
|
|
||||||
def _add_remote_peer(self, decode_kv_args: KVArgsRegisterInfo):
|
def _add_remote_peer(self, decode_kv_args: KVArgsRegisterInfo):
|
||||||
agent_name = decode_kv_args.agent_name
|
agent_name = decode_kv_args.agent_name
|
||||||
if agent_name in self.decode_kv_args_table:
|
if agent_name in self.decode_kv_args_table:
|
||||||
@@ -595,6 +637,53 @@ class NixlKVManager(CommonKVManager):
|
|||||||
raise Exception("KVSender failed to post transfer")
|
raise Exception("KVSender failed to post transfer")
|
||||||
return xfer_handle
|
return xfer_handle
|
||||||
|
|
||||||
|
def _send_mamba_state(
|
||||||
|
self,
|
||||||
|
peer_name: str,
|
||||||
|
prefill_state_indices: List[int],
|
||||||
|
dst_state_data_ptrs: list[int],
|
||||||
|
dst_state_indices: List[int],
|
||||||
|
dst_gpu_id: int,
|
||||||
|
notif: str,
|
||||||
|
):
|
||||||
|
"""Transfer Mamba states via RDMA."""
|
||||||
|
assert len(prefill_state_indices) == 1, "Mamba should have single state index"
|
||||||
|
assert len(dst_state_indices) == len(
|
||||||
|
prefill_state_indices
|
||||||
|
), "State indices count mismatch between Prefill and Decode"
|
||||||
|
|
||||||
|
src_addrs = []
|
||||||
|
dst_addrs = []
|
||||||
|
|
||||||
|
prefill_state_data_ptrs = self.kv_args.state_data_ptrs
|
||||||
|
prefill_state_item_lens = self.kv_args.state_item_lens
|
||||||
|
|
||||||
|
for i, dst_state_ptr in enumerate(dst_state_data_ptrs):
|
||||||
|
length = prefill_state_item_lens[i]
|
||||||
|
src_addr = prefill_state_data_ptrs[i] + length * int(
|
||||||
|
prefill_state_indices[0]
|
||||||
|
)
|
||||||
|
dst_addr = dst_state_ptr + length * int(dst_state_indices[0])
|
||||||
|
src_addrs.append((src_addr, length, self.kv_args.gpu_id))
|
||||||
|
dst_addrs.append((dst_addr, length, dst_gpu_id))
|
||||||
|
|
||||||
|
src_descs = self.agent.get_xfer_descs(src_addrs, "VRAM")
|
||||||
|
dst_descs = self.agent.get_xfer_descs(dst_addrs, "VRAM")
|
||||||
|
|
||||||
|
xfer_handle = self.agent.initialize_xfer(
|
||||||
|
"WRITE",
|
||||||
|
src_descs,
|
||||||
|
dst_descs,
|
||||||
|
peer_name,
|
||||||
|
notif.encode("ascii"),
|
||||||
|
)
|
||||||
|
if not xfer_handle:
|
||||||
|
raise Exception("Failed to create Mamba state transfer")
|
||||||
|
state = self.agent.transfer(xfer_handle)
|
||||||
|
if state == "ERR":
|
||||||
|
raise Exception("Failed to post Mamba state transfer")
|
||||||
|
return xfer_handle
|
||||||
|
|
||||||
def add_transfer_request(
|
def add_transfer_request(
|
||||||
self,
|
self,
|
||||||
bootstrap_room: int,
|
bootstrap_room: int,
|
||||||
@@ -603,6 +692,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
is_last: bool,
|
is_last: bool,
|
||||||
chunk_id: int,
|
chunk_id: int,
|
||||||
aux_index: Optional[int] = None,
|
aux_index: Optional[int] = None,
|
||||||
|
state_indices: Optional[List[int]] = None,
|
||||||
):
|
):
|
||||||
assert self.disaggregation_mode == DisaggregationMode.PREFILL
|
assert self.disaggregation_mode == DisaggregationMode.PREFILL
|
||||||
assert not is_last or (is_last and aux_index is not None)
|
assert not is_last or (is_last and aux_index is not None)
|
||||||
@@ -651,13 +741,36 @@ class NixlKVManager(CommonKVManager):
|
|||||||
handles.append(kv_xfer_handle)
|
handles.append(kv_xfer_handle)
|
||||||
# Only the last chunk we need to send the aux data.
|
# Only the last chunk we need to send the aux data.
|
||||||
if is_last:
|
if is_last:
|
||||||
|
if state_indices is not None:
|
||||||
|
state_type = getattr(self.kv_args, "state_type", "none")
|
||||||
|
if (
|
||||||
|
self.attn_tp_size
|
||||||
|
!= self.decode_kv_args_table[req.agent_name].decode_tp_size
|
||||||
|
):
|
||||||
|
raise RuntimeError(
|
||||||
|
"PD Disaggregation does NOT support PD different TP sizes for hybrid mamba models yet."
|
||||||
|
)
|
||||||
|
|
||||||
|
if state_type == "mamba":
|
||||||
|
state_xfer_handle = self._send_mamba_state(
|
||||||
|
req.agent_name,
|
||||||
|
state_indices,
|
||||||
|
self.decode_kv_args_table[
|
||||||
|
req.agent_name
|
||||||
|
].dst_state_data_ptrs,
|
||||||
|
req.dst_state_indices,
|
||||||
|
self.decode_kv_args_table[req.agent_name].gpu_id,
|
||||||
|
f"{req.room}_state_{self.kv_args.pp_rank}",
|
||||||
|
)
|
||||||
|
handles.append(state_xfer_handle)
|
||||||
|
|
||||||
assert aux_index is not None
|
assert aux_index is not None
|
||||||
aux_xfer_handle = self.send_aux(
|
aux_xfer_handle = self.send_aux(
|
||||||
req.agent_name,
|
req.agent_name,
|
||||||
aux_index,
|
aux_index,
|
||||||
self.decode_kv_args_table[req.agent_name].dst_aux_ptrs,
|
self.decode_kv_args_table[req.agent_name].dst_aux_ptrs,
|
||||||
req.dst_aux_index,
|
req.dst_aux_index,
|
||||||
str(req.room) + "_aux",
|
f"{req.room}_aux",
|
||||||
)
|
)
|
||||||
handles.append(aux_xfer_handle)
|
handles.append(aux_xfer_handle)
|
||||||
if is_last:
|
if is_last:
|
||||||
@@ -694,6 +807,9 @@ class NixlKVManager(CommonKVManager):
|
|||||||
)
|
)
|
||||||
elif components[1] == "aux":
|
elif components[1] == "aux":
|
||||||
self.transfer_statuses[room].received_aux = True
|
self.transfer_statuses[room].received_aux = True
|
||||||
|
elif components[1] == "state":
|
||||||
|
pp_rank = int(components[2]) if len(components) > 2 else 0
|
||||||
|
self.transfer_statuses[room].received_state_per_pp.add(pp_rank)
|
||||||
|
|
||||||
def check_transfer_done(self, room: int):
|
def check_transfer_done(self, room: int):
|
||||||
if room not in self.transfer_statuses:
|
if room not in self.transfer_statuses:
|
||||||
@@ -768,6 +884,7 @@ class NixlKVSender(CommonKVSender):
|
|||||||
is_last,
|
is_last,
|
||||||
self.chunk_id,
|
self.chunk_id,
|
||||||
self.aux_index,
|
self.aux_index,
|
||||||
|
state_indices,
|
||||||
)
|
)
|
||||||
self.xfer_handles.extend(new_xfer_handles)
|
self.xfer_handles.extend(new_xfer_handles)
|
||||||
self.chunk_id += 1
|
self.chunk_id += 1
|
||||||
@@ -841,9 +958,18 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
kv_indices.tobytes() if not is_dummy else b"",
|
kv_indices.tobytes() if not is_dummy else b"",
|
||||||
str(aux_index).encode("ascii"),
|
str(aux_index).encode("ascii"),
|
||||||
str(self.required_dst_info_num).encode("ascii"),
|
str(self.required_dst_info_num).encode("ascii"),
|
||||||
|
(
|
||||||
|
np.array(state_indices, dtype=np.int32).tobytes()
|
||||||
|
if not is_dummy and state_indices is not None
|
||||||
|
else b""
|
||||||
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Mark that we expect state data if state_indices was provided
|
||||||
|
if state_indices is not None:
|
||||||
|
self.kv_mgr.transfer_statuses[self.bootstrap_room].expects_state = True
|
||||||
|
|
||||||
self.started_transfer = True
|
self.started_transfer = True
|
||||||
self.init_time = time.time()
|
self.init_time = time.time()
|
||||||
|
|
||||||
@@ -895,6 +1021,9 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
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
|
||||||
)
|
)
|
||||||
|
packed_state_data_ptrs = b"".join(
|
||||||
|
struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.state_data_ptrs
|
||||||
|
)
|
||||||
|
|
||||||
with lock:
|
with lock:
|
||||||
sock.send_multipart(
|
sock.send_multipart(
|
||||||
@@ -907,6 +1036,7 @@ class NixlKVReceiver(CommonKVReceiver):
|
|||||||
self.kv_mgr.agent.get_agent_metadata(),
|
self.kv_mgr.agent.get_agent_metadata(),
|
||||||
packed_kv_data_ptrs,
|
packed_kv_data_ptrs,
|
||||||
packed_aux_data_ptrs,
|
packed_aux_data_ptrs,
|
||||||
|
packed_state_data_ptrs,
|
||||||
str(self.kv_mgr.kv_args.gpu_id).encode("ascii"),
|
str(self.kv_mgr.kv_args.gpu_id).encode("ascii"),
|
||||||
str(self.kv_mgr.kv_args.decode_tp_size).encode("ascii"),
|
str(self.kv_mgr.kv_args.decode_tp_size).encode("ascii"),
|
||||||
str(self.kv_mgr.kv_args.engine_rank).encode("ascii"),
|
str(self.kv_mgr.kv_args.engine_rank).encode("ascii"),
|
||||||
|
|||||||
Reference in New Issue
Block a user