[PD] Handle NIXL abort notifications (#30352)
This commit is contained in:
@@ -631,6 +631,12 @@ class NixlKVManager(CommonKVManager):
|
||||
def check_status(self, bootstrap_room: int):
|
||||
return self.request_status.get(bootstrap_room, KVPoll.WaitingForInput)
|
||||
|
||||
def update_status(self, bootstrap_room: int, status: KVPoll):
|
||||
# Keep Failed sticky until the sender clears the room.
|
||||
if self.request_status.get(bootstrap_room) == KVPoll.Failed:
|
||||
return
|
||||
super().update_status(bootstrap_room, status)
|
||||
|
||||
def _prep_equal_tp_dlist(
|
||||
self,
|
||||
peer_name: str,
|
||||
@@ -2316,6 +2322,40 @@ class NixlKVManager(CommonKVManager):
|
||||
return False
|
||||
return self.transfer_statuses[room].is_done()
|
||||
|
||||
def _handle_abort_notification(self, msg: List[bytes]) -> bool:
|
||||
if not msg or msg[0] != b"ABORT":
|
||||
return False
|
||||
|
||||
try:
|
||||
room_to_be_aborted = int(msg[1].decode("ascii"))
|
||||
except Exception as e:
|
||||
logger.debug(f"Ignoring malformed abort notification: {e}")
|
||||
return True
|
||||
|
||||
if (
|
||||
room_to_be_aborted in self.request_status
|
||||
and self.check_status(room_to_be_aborted) != KVPoll.Success
|
||||
):
|
||||
self.record_failure(
|
||||
room_to_be_aborted,
|
||||
"Aborted by decode-side abort notification.",
|
||||
)
|
||||
self.update_status(room_to_be_aborted, KVPoll.Failed)
|
||||
logger.debug(
|
||||
f"Received abort notification for room {room_to_be_aborted}, "
|
||||
f"marked as Failed"
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
f"Received abort notification for room {room_to_be_aborted}, "
|
||||
f"ignoring (already completed or unknown)"
|
||||
)
|
||||
|
||||
# TODO: Define real ACK/deferred-release semantics if decode-side buffer
|
||||
# release needs to wait for prefill-side NIXL transfer quiescence.
|
||||
|
||||
return True
|
||||
|
||||
def _start_bootstrap_thread(self):
|
||||
def bootstrap_thread():
|
||||
"""This thread recvs transfer info from the decode engine"""
|
||||
@@ -2345,6 +2385,9 @@ class NixlKVManager(CommonKVManager):
|
||||
handle_staging_rsp(waiting_req_bytes, self.transfer_infos)
|
||||
continue
|
||||
|
||||
if self._handle_abort_notification(waiting_req_bytes):
|
||||
continue
|
||||
|
||||
assert (
|
||||
waiting_req_bytes[0] == GUARD
|
||||
), f"First message should be {GUARD}. Foreign traffic?"
|
||||
|
||||
@@ -330,6 +330,185 @@ class TestNixlKVSenderChunkPolicy(CustomTestCase):
|
||||
self.assertTrue(sender.should_send_kv_chunk(3, last_chunk=False))
|
||||
|
||||
|
||||
class TestNixlAbortHandling(CustomTestCase):
|
||||
def _make_manager(self, request_status=None):
|
||||
mgr = object.__new__(NixlKVManager)
|
||||
mgr.request_status = dict(request_status or {})
|
||||
mgr._connect = MagicMock()
|
||||
mgr.failure_lock = threading.Lock()
|
||||
mgr.failure_records = {}
|
||||
return mgr
|
||||
|
||||
def test_given_known_incomplete_room_when_abort_arrives_then_room_fails_without_ack(
|
||||
self,
|
||||
):
|
||||
mgr = self._make_manager({11: KVPoll.WaitingForInput})
|
||||
|
||||
handled = mgr._handle_abort_notification(
|
||||
[b"ABORT", b"11", b"127.0.0.1", b"5555"]
|
||||
)
|
||||
|
||||
self.assertTrue(handled)
|
||||
self.assertEqual(mgr.request_status[11], KVPoll.Failed)
|
||||
self.assertEqual(
|
||||
mgr.failure_records[11],
|
||||
"Aborted by decode-side abort notification.",
|
||||
)
|
||||
mgr._connect.assert_not_called()
|
||||
|
||||
def test_given_successful_room_when_abort_arrives_then_status_is_preserved(self):
|
||||
mgr = self._make_manager({12: KVPoll.Success})
|
||||
|
||||
handled = mgr._handle_abort_notification(
|
||||
[b"ABORT", b"12", b"127.0.0.1", b"5556"]
|
||||
)
|
||||
|
||||
self.assertTrue(handled)
|
||||
self.assertEqual(mgr.request_status[12], KVPoll.Success)
|
||||
self.assertEqual(mgr.failure_records, {})
|
||||
mgr._connect.assert_not_called()
|
||||
|
||||
def test_given_unknown_room_when_abort_arrives_then_status_remains_absent(self):
|
||||
mgr = self._make_manager()
|
||||
|
||||
handled = mgr._handle_abort_notification(
|
||||
[b"ABORT", b"14", b"127.0.0.1", b"5557"]
|
||||
)
|
||||
|
||||
self.assertTrue(handled)
|
||||
self.assertNotIn(14, mgr.request_status)
|
||||
self.assertEqual(mgr.failure_records, {})
|
||||
mgr._connect.assert_not_called()
|
||||
|
||||
def test_given_malformed_abort_when_handled_then_no_exception_or_ack(self):
|
||||
mgr = self._make_manager({13: KVPoll.WaitingForInput})
|
||||
|
||||
handled = mgr._handle_abort_notification(
|
||||
[b"ABORT", b"invalid-room", b"127.0.0.1", b"5558"]
|
||||
)
|
||||
|
||||
self.assertTrue(handled)
|
||||
self.assertEqual(mgr.request_status[13], KVPoll.WaitingForInput)
|
||||
self.assertEqual(mgr.failure_records, {})
|
||||
mgr._connect.assert_not_called()
|
||||
|
||||
|
||||
class TestNixlUpdateStatus(CustomTestCase):
|
||||
def _make_manager(self, request_status):
|
||||
mgr = object.__new__(NixlKVManager)
|
||||
mgr.request_status = dict(request_status)
|
||||
return mgr
|
||||
|
||||
def test_given_failed_room_when_status_is_promoted_then_failed_is_preserved(self):
|
||||
for status in (KVPoll.Transferring, KVPoll.Success):
|
||||
with self.subTest(status=status):
|
||||
mgr = self._make_manager({17: KVPoll.Failed})
|
||||
|
||||
mgr.update_status(17, status)
|
||||
|
||||
self.assertEqual(mgr.request_status[17], KVPoll.Failed)
|
||||
|
||||
def test_given_missing_room_when_failed_update_arrives_then_room_is_not_resurrected(
|
||||
self,
|
||||
):
|
||||
mgr = self._make_manager({})
|
||||
|
||||
mgr.update_status(18, KVPoll.Failed)
|
||||
|
||||
self.assertNotIn(18, mgr.request_status)
|
||||
|
||||
|
||||
class TestNixlTransferWorker(CustomTestCase):
|
||||
def _make_manager(self, room):
|
||||
mgr = object.__new__(NixlKVManager)
|
||||
mgr.request_status = {room: KVPoll.WaitingForInput}
|
||||
mgr.transfer_infos = {
|
||||
room: {
|
||||
"agent": TransferInfo(
|
||||
room=room,
|
||||
endpoint="127.0.0.1",
|
||||
dst_port=5555,
|
||||
agent_name="agent",
|
||||
dst_kv_indices=np.array([2], dtype=np.int32),
|
||||
dst_aux_index=0,
|
||||
required_dst_info_num=1,
|
||||
dst_state_indices=[],
|
||||
)
|
||||
}
|
||||
}
|
||||
mgr.decode_kv_args_table = {
|
||||
"agent": SimpleNamespace(
|
||||
decode_tp_size=1,
|
||||
dst_kv_ptrs=[0],
|
||||
dst_aux_ptrs=[0],
|
||||
gpu_id=0,
|
||||
staging=None,
|
||||
kv_xfer_segments=None,
|
||||
dst_homogeneous_mem_kind="VRAM",
|
||||
)
|
||||
}
|
||||
mgr.req_to_decode_prefix_len = {room: 4}
|
||||
mgr.enable_staging = False
|
||||
mgr._staging_ctx = None
|
||||
mgr.is_mla_backend = False
|
||||
mgr.attn_tp_size = 1
|
||||
mgr.kv_args = SimpleNamespace(engine_rank=0)
|
||||
mgr.exceptions = {}
|
||||
mgr.failure_lock = threading.Lock()
|
||||
mgr.failure_records = {}
|
||||
|
||||
def check_xfer_state(_handle):
|
||||
mgr.update_status(room, KVPoll.Failed)
|
||||
return "DONE"
|
||||
|
||||
mgr.agent = SimpleNamespace(check_xfer_state=check_xfer_state)
|
||||
return mgr
|
||||
|
||||
def _make_chunk(self, room, prefill_kv_indices, is_last_chunk):
|
||||
return TransferKVChunk(
|
||||
room=room,
|
||||
prefill_kv_indices=np.array(prefill_kv_indices, dtype=np.int32),
|
||||
index_slice=slice(0, len(prefill_kv_indices)),
|
||||
is_last_chunk=is_last_chunk,
|
||||
chunk_id=0,
|
||||
prefill_aux_index=0 if is_last_chunk else None,
|
||||
state_indices=None,
|
||||
)
|
||||
|
||||
def _run_worker_once(self, mgr, chunk):
|
||||
queue = SimpleNamespace(get=MagicMock(side_effect=[chunk, SystemExit()]))
|
||||
with self.assertRaises(SystemExit):
|
||||
mgr.transfer_worker(queue)
|
||||
|
||||
def test_given_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved(
|
||||
self,
|
||||
):
|
||||
room = 21
|
||||
mgr = self._make_manager(room)
|
||||
mgr.send_aux = MagicMock(return_value="aux_handle")
|
||||
chunk = self._make_chunk(room, [], is_last_chunk=True)
|
||||
|
||||
self._run_worker_once(mgr, chunk)
|
||||
|
||||
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
|
||||
self.assertNotIn(room, mgr.transfer_infos)
|
||||
self.assertNotIn(room, mgr.req_to_decode_prefix_len)
|
||||
|
||||
def test_given_non_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved(
|
||||
self,
|
||||
):
|
||||
room = 22
|
||||
mgr = self._make_manager(room)
|
||||
mgr.send_kvcache = MagicMock(return_value="kv_handle")
|
||||
chunk = self._make_chunk(room, [1], is_last_chunk=False)
|
||||
|
||||
self._run_worker_once(mgr, chunk)
|
||||
|
||||
self.assertEqual(mgr.request_status[room], KVPoll.Failed)
|
||||
self.assertIn(room, mgr.transfer_infos)
|
||||
self.assertIn(room, mgr.req_to_decode_prefix_len)
|
||||
|
||||
|
||||
class TestNixlNotifications(CustomTestCase):
|
||||
def _make_manager(self, messages, required=None):
|
||||
mgr = object.__new__(NixlKVManager)
|
||||
|
||||
Reference in New Issue
Block a user