[P/D] [NixlConnector] kv load recovery integration (#26171)
Signed-off-by: Will Eaton <weaton@redhat.com>
This commit is contained in:
@@ -190,7 +190,6 @@ def _make_fake_nixl_pkg():
|
||||
# Copy of FakeNixlWrapper implementation for Ray workers
|
||||
import uuid
|
||||
from collections import defaultdict
|
||||
from typing import Optional
|
||||
|
||||
{fake_nixl_source}
|
||||
|
||||
@@ -1143,3 +1142,145 @@ def test_aborted_request_removed_from_worker_in_batch(dist_init):
|
||||
# After abort, the worker should not keep tracking it as "in-batch"
|
||||
assert req.request_id not in connector.connector_worker._reqs_to_process
|
||||
#### Model Runner end ####
|
||||
|
||||
|
||||
class FailingNixlWrapper(FakeNixlWrapper):
|
||||
"""Mock NixlWrapper that fails on specific operations."""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.fail_handshake = False
|
||||
self.fail_transfer_setup = False
|
||||
self.fail_send_notif = False
|
||||
|
||||
def add_remote_agent(self, agent_metadata: bytes) -> str:
|
||||
if self.fail_handshake:
|
||||
from zmq.error import Again
|
||||
|
||||
raise Again("Simulated timeout failure")
|
||||
return super().add_remote_agent(agent_metadata)
|
||||
|
||||
def make_prepped_xfer(
|
||||
self,
|
||||
xfer_type: str,
|
||||
local_xfer_side_handle: int,
|
||||
local_block_descs_ids: list[int],
|
||||
remote_xfer_side_handle: int,
|
||||
remote_block_descs_ids: list[int],
|
||||
notif_msg: bytes | None = None,
|
||||
) -> int:
|
||||
if self.fail_transfer_setup:
|
||||
# classic RuntimeError to simulate failure
|
||||
raise RuntimeError("BAD STATUS")
|
||||
return super().make_prepped_xfer(
|
||||
xfer_type,
|
||||
local_xfer_side_handle,
|
||||
local_block_descs_ids,
|
||||
remote_xfer_side_handle,
|
||||
remote_block_descs_ids,
|
||||
notif_msg,
|
||||
)
|
||||
|
||||
def send_notif(self, agent_name: str, notif_msg: bytes) -> None:
|
||||
if self.fail_send_notif:
|
||||
raise RuntimeError("Simulated send_notif failure")
|
||||
return super().send_notif(agent_name, notif_msg)
|
||||
|
||||
|
||||
@patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.nixl_connector.NixlWrapper",
|
||||
FailingNixlWrapper,
|
||||
)
|
||||
def test_handshake_failure_returns_finished(dist_init):
|
||||
"""Test that handshake failures mark blocks invalid and return via get_finished."""
|
||||
vllm_config = create_vllm_config()
|
||||
|
||||
connector = NixlConnector(vllm_config, KVConnectorRole.WORKER)
|
||||
connector.connector_worker = FakeNixlConnectorWorker(
|
||||
vllm_config, connector.engine_id, hand_shake_latency=0.1
|
||||
)
|
||||
connector.connector_worker.nixl_wrapper.fail_handshake = True
|
||||
|
||||
request_id = "test_handshake_fail"
|
||||
metadata = NixlConnectorMetadata()
|
||||
metadata.add_new_req(
|
||||
request_id=request_id,
|
||||
local_block_ids=[1, 2, 3],
|
||||
kv_transfer_params={
|
||||
"remote_block_ids": [4, 5, 6],
|
||||
"remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
|
||||
"remote_host": "localhost",
|
||||
"remote_port": 1234,
|
||||
"remote_tp_size": 1,
|
||||
},
|
||||
)
|
||||
connector.bind_connector_metadata(metadata)
|
||||
|
||||
dummy_ctx = ForwardContext(
|
||||
no_compile_layers={},
|
||||
attn_metadata={},
|
||||
virtual_engine=0,
|
||||
)
|
||||
connector.start_load_kv(dummy_ctx)
|
||||
|
||||
# Wait for handshake to fail
|
||||
time.sleep(0.3)
|
||||
|
||||
# Check that blocks were marked invalid
|
||||
invalid_blocks = connector.get_block_ids_with_load_errors()
|
||||
assert invalid_blocks == {1, 2, 3}
|
||||
|
||||
# Check that request appears in get_finished
|
||||
_, done_recving = connector.get_finished(finished_req_ids=set())
|
||||
assert request_id in done_recving
|
||||
|
||||
|
||||
@patch(
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.nixl_connector.NixlWrapper",
|
||||
FailingNixlWrapper,
|
||||
)
|
||||
def test_transfer_setup_failure_returns_finished(dist_init):
|
||||
"""Test that transfer setup failures mark blocks invalid
|
||||
and return via get_finished."""
|
||||
vllm_config = create_vllm_config()
|
||||
|
||||
connector = NixlConnector(vllm_config, KVConnectorRole.WORKER)
|
||||
connector.connector_worker = FakeNixlConnectorWorker(
|
||||
vllm_config, connector.engine_id, hand_shake_latency=0
|
||||
)
|
||||
connector.connector_worker.nixl_wrapper.fail_transfer_setup = True
|
||||
|
||||
request_id = "test_transfer_fail"
|
||||
metadata = NixlConnectorMetadata()
|
||||
metadata.add_new_req(
|
||||
request_id=request_id,
|
||||
local_block_ids=[7, 8, 9],
|
||||
kv_transfer_params={
|
||||
"remote_block_ids": [10, 11, 12],
|
||||
"remote_engine_id": FakeNixlConnectorWorker.REMOTE_ENGINE_ID,
|
||||
"remote_host": "localhost",
|
||||
"remote_port": 1234,
|
||||
"remote_tp_size": 1,
|
||||
},
|
||||
)
|
||||
connector.bind_connector_metadata(metadata)
|
||||
|
||||
dummy_ctx = ForwardContext(
|
||||
no_compile_layers={},
|
||||
attn_metadata={},
|
||||
virtual_engine=0,
|
||||
)
|
||||
connector.start_load_kv(dummy_ctx)
|
||||
|
||||
# Wait for handshake to complete and process ready_requests
|
||||
connector.bind_connector_metadata(NixlConnectorMetadata())
|
||||
time.sleep(0.1)
|
||||
connector.start_load_kv(dummy_ctx)
|
||||
|
||||
# check that blocks were marked invalid
|
||||
invalid_blocks = connector.get_block_ids_with_load_errors()
|
||||
assert invalid_blocks == {7, 8, 9}
|
||||
|
||||
# ensure request appears in get_finished
|
||||
_, done_recving = connector.get_finished(finished_req_ids=set())
|
||||
assert request_id in done_recving
|
||||
|
||||
Reference in New Issue
Block a user