Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 18 additions & 15 deletions tests/v1/kv_connector/unit/test_handshake_pp_aggregation.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ class _Metadata(KVConnectorHandshakeMetadata):

class _FakeExecutor:
handshake_metadata_src: (
list[dict[tuple[int, int], KVConnectorHandshakeMetadata] | None] | None
list[dict[tuple[int, int, int], KVConnectorHandshakeMetadata] | None] | None
)
last_instance: "_FakeExecutor | None" = None

Expand All @@ -36,7 +36,7 @@ def __init__(

def get_kv_connector_handshake_metadata(
self,
) -> list[dict[tuple[int, int], KVConnectorHandshakeMetadata] | None] | None:
) -> list[dict[tuple[int, int, int], KVConnectorHandshakeMetadata] | None] | None:
self.handshake_calls += 1
return self.handshake_metadata

Expand All @@ -49,7 +49,7 @@ def _run_engine_core_handshake(
connector: KVConnectorBase_V1,
*,
handshake_metadata: (
list[dict[tuple[int, int], KVConnectorHandshakeMetadata] | None] | None
list[dict[tuple[int, int, int], KVConnectorHandshakeMetadata] | None] | None
),
) -> _FakeExecutor:
class _FakeScheduler:
Expand Down Expand Up @@ -161,21 +161,22 @@ class _PPAwareConnector(_LegacyConnector):
def __init__(self) -> None:
super().__init__()
self.pp_aware_metadata: (
dict[tuple[int, int], KVConnectorHandshakeMetadata] | None
dict[tuple[int, int, int], KVConnectorHandshakeMetadata] | None
) = None

def set_xfer_handshake_metadata_pp_aware(
self, metadata: dict[tuple[int, int], KVConnectorHandshakeMetadata]
self,
metadata: dict[tuple[int, int, int], KVConnectorHandshakeMetadata],
) -> None:
self.pp_aware_metadata = metadata


def test_engine_unwraps_handshake_metadata_for_legacy_connector(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Engine core always asks workers for `(pp_rank, tp_rank)`-keyed metadata,
"""Engine core asks workers for `(pp_rank, tp_rank, pcp_rank)` metadata,
then unwraps to `{tp_rank: metadata}` for a connector that has not opted
into PP-aware handshake (single-PP producer, all `pp_rank == 0`)."""
into PP-aware handshake (all PP and PCP ranks are zero)."""
metadata_0 = _Metadata()
metadata_1 = _Metadata()
connector = _LegacyConnector()
Expand All @@ -184,9 +185,9 @@ def test_engine_unwraps_handshake_metadata_for_legacy_connector(
monkeypatch,
connector,
handshake_metadata=[
{(0, 0): metadata_0},
{(0, 0, 0): metadata_0},
None,
{(0, 1): metadata_1},
{(0, 1, 0): metadata_1},
],
)

Expand All @@ -205,28 +206,30 @@ def test_engine_rejects_pp_producer_for_legacy_connector(
_run_engine_core_handshake(
monkeypatch,
connector,
handshake_metadata=[{(0, 0): _Metadata()}, {(1, 0): _Metadata()}],
handshake_metadata=[
{(0, 0, 0): _Metadata()},
{(1, 0, 0): _Metadata()},
],
)


def test_engine_passes_handshake_metadata_through_for_pp_aware_connector(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A PP-aware connector receives the full `(pp_rank, tp_rank)`-keyed dict
unchanged."""
"""A PP-aware connector receives the full worker-keyed dict unchanged."""
metadata_0 = _Metadata()
metadata_1 = _Metadata()
connector = _PPAwareConnector()

executor = _run_engine_core_handshake(
monkeypatch,
connector,
handshake_metadata=[{(0, 0): metadata_0}, {(1, 0): metadata_1}],
handshake_metadata=[{(0, 0, 0): metadata_0}, {(1, 0, 0): metadata_1}],
)

assert executor.handshake_calls == 1
assert connector.legacy_metadata is None
assert connector.pp_aware_metadata == {
(0, 0): metadata_0,
(1, 0): metadata_1,
(0, 0, 0): metadata_0,
(1, 0, 0): metadata_1,
}
110 changes: 92 additions & 18 deletions tests/v1/kv_connector/unit/test_nixl_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -408,10 +408,9 @@ def test_kv_transfer_handshake(dist_init):
decoder = msgspec.msgpack.Decoder(NixlAgentMetadata)
expected_agent_metadata = decoder.decode(metadata.agent_metadata_bytes)

# The scheduler connector expects metadata keyed by
# (pp_rank, tp_rank).
# The scheduler connector expects metadata keyed by worker rank.
scheduler_connector = scheduler.get_kv_connector()
scheduler_connector.set_xfer_handshake_metadata_pp_aware({(0, 0): metadata})
scheduler_connector.set_xfer_handshake_metadata_pp_aware({(0, 0, 0): metadata})

# Simulate a request that finishes prefill, which returns
# corresponding NixlConnectorMetadata for decode instance.
Expand Down Expand Up @@ -495,6 +494,8 @@ def __init__(
total_num_kv_heads=self.model_config.get_total_num_kv_heads(),
attn_backends=self.attn_backends,
tensor_shape=test_shape,
dcp_rank=self.dcp_rank,
dcp_size=self.dcp_size,
)

self.compat_hash = compute_nixl_compatibility_hash(
Expand All @@ -508,8 +509,10 @@ def _nixl_handshake(
remote_tp_size: int,
expected_engine_id: str,
remote_pp_size: int = 1,
remote_dcp_size: int = 1,
remote_pcp_size: int = 1,
notif_agents_only: bool = False,
) -> tuple[dict[tuple[int, int], str], float]:
) -> tuple[dict[tuple[int, int, int], str], float]:
# Mimic slow _nixl_handshake, as well as bypass zmq communication.
time.sleep(self._hand_shake_latency)
# These should've been done in register_kv_caches(), called by
Expand Down Expand Up @@ -539,7 +542,7 @@ def _nixl_handshake(
# When remote tp_size > local tp_size, handshake with multiple
# remote ranks.
num_handshakes = 1 if tp_ratio > 0 else -tp_ratio
remote_agents: dict[tuple[int, int], str] = {}
remote_agents: dict[tuple[int, int, int], str] = {}
for remote_tp_rank in range(num_handshakes):
remote_agent_name = self.add_remote_agent(
NixlAgentMetadata(
Expand All @@ -559,13 +562,77 @@ def _nixl_handshake(
),
remote_tp_rank=remote_tp_rank,
remote_tp_size=remote_tp_size,
remote_dcp_size=remote_dcp_size,
remote_pcp_size=remote_pcp_size,
)
remote_agents[(0, remote_tp_rank)] = remote_agent_name
remote_agents[(0, remote_tp_rank, 0)] = remote_agent_name
# Handshake bypasses zmq, so report a zero clock offset to the peer.
return remote_agents, 0.0


class TestNixlHandshake:
def test_consumer_rejects_pcp(self):
vllm_config = create_vllm_config(kv_role="kv_consumer")
vllm_config.parallel_config.prefill_context_parallel_size = 2

with pytest.raises(
NotImplementedError,
match="consumer.*prefill_context_parallel_size",
):
NixlConnector(
vllm_config,
KVConnectorRole.SCHEDULER,
make_kv_cache_config(block_size=16),
)

@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
FakeNixlWrapper,
)
def test_pcp_producer_publishes_one_replica_per_dcp_rank(
self, default_vllm_config, dist_init
):
vllm_config = create_vllm_config(kv_role="kv_producer")
connector = NixlConnector(
vllm_config,
KVConnectorRole.WORKER,
make_kv_cache_config(block_size=16),
)
payload = MagicMock(spec=NixlHandshakePayload)
worker = connector.connector_worker
assert worker is not None
worker.xfer_handshake_metadata = payload

worker.pcp_rank = 0
assert connector.get_handshake_metadata() is payload

worker.pcp_rank = 1
assert connector.get_handshake_metadata() is None

vllm_config.parallel_config.decode_context_parallel_size = 2
assert connector.get_handshake_metadata() is payload

worker.pcp_rank = 2
assert connector.get_handshake_metadata() is None

def test_pcp_producer_waits_only_for_published_replicas(
self, default_vllm_config, dist_init
):
vllm_config = create_vllm_config(kv_role="kv_producer")
vllm_config.parallel_config.tensor_parallel_size = 2
vllm_config.parallel_config.prefill_context_parallel_size = 2
vllm_config.parallel_config.pipeline_parallel_size = 2
connector = NixlConnector(
vllm_config,
KVConnectorRole.SCHEDULER,
make_kv_cache_config(block_size=16),
)

assert connector.get_finished_count() == 4

vllm_config.parallel_config.decode_context_parallel_size = 2
assert connector.get_finished_count() == 8

@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
FakeNixlWrapper,
Expand Down Expand Up @@ -759,20 +826,24 @@ def test_prefill_tp_size_greater_than_decode_tp_size(

def check_handshake(remote_tp_size: int):
tp_ratio = remote_tp_size // local_tp_size
assert set(remote_agents.keys()) == {(0, r) for r in range(tp_ratio)}
expected_worker_keys = {(rank, 0) for rank in range(tp_ratio)}
expected_agent_keys = {(0, *key) for key in expected_worker_keys}
assert set(remote_agents) == expected_agent_keys

remote_engine_id = worker.REMOTE_ENGINE_ID
remote_info = worker.transfer_topo.get_engine_info(remote_engine_id)
assert remote_info.remote_tp_size == remote_tp_size
assert -tp_ratio == worker.transfer_topo.tp_ratio(remote_tp_size)
# ensure src_xfer_handles_by_tp_ratio is populated with tpratio chunks
# Each remote TP rank has an explicitly addressed split handle.
split_key = (-tp_ratio, worker.block_size)
assert split_key in worker.src_xfer_handles_by_tp_ratio
assert len(worker.src_xfer_handles_by_tp_ratio[split_key]) == tp_ratio
assert remote_engine_id in worker.dst_xfer_side_handles
assert set(worker.dst_xfer_side_handles[remote_engine_id].keys()) == set(
assert set(worker.src_xfer_handles_by_tp_ratio[split_key]) == set(
range(tp_ratio)
)
assert remote_engine_id in worker.dst_xfer_side_handles
assert set(worker.dst_xfer_side_handles[remote_engine_id]) == (
expected_worker_keys
)

remote_agents, _ = worker._nixl_handshake(
host="localhost",
Expand Down Expand Up @@ -2140,9 +2211,9 @@ def test_shutdown_cleans_up_resources(default_vllm_config, dist_init):
# Mock register_kv_cache which registers local handle
worker.src_xfer_handles_by_block_size = {worker.block_size: 455}
# P TP = 2 * D TP case, we should register 2 local handles
worker.src_xfer_handles_by_tp_ratio = {(-2, 16): [456, 457]}
worker.src_xfer_handles_by_tp_ratio = {(-2, 16): {0: 456, 1: 457}}
worker.dst_xfer_side_handles = {"engine1": {0: 789}}
worker._remote_agents = {"engine1": {(0, 0): "agent1"}}
worker._remote_agents = {"engine1": {(0, 0, 0): "agent1"}}
# _cleanup_remote_engine (called by shutdown) also clears these:
worker.kv_caches_base_addr["engine1"] = {0: [0xABC]}
worker.dst_num_blocks["engine1"] = 50
Expand Down Expand Up @@ -2195,7 +2266,10 @@ def _setup_worker_with_remote_engine(
)

engine_id = "remote-engine-1"
worker._remote_agents[engine_id] = {(0, 0): "agent_0", (0, 1): "agent_1"}
worker._remote_agents[engine_id] = {
(0, 0, 0): "agent_0",
(0, 1, 0): "agent_1",
}
worker.dst_xfer_side_handles[engine_id] = {0: 100, 1: 200}
worker.kv_caches_base_addr[engine_id] = {0: [0xABC]}
worker.dst_num_blocks[engine_id] = 50
Expand Down Expand Up @@ -3097,10 +3171,10 @@ def test_mla_broadcast_notif_uses_remote_request_id(
local_block_len=worker.block_size * 4096,
)
worker._remote_agents[remote_engine_id] = {
(0, rank): f"agent_p{rank}" for rank in range(prefill_tp_size)
(0, rank, 0): f"agent_p{rank}" for rank in range(prefill_tp_size)
}
worker.dst_xfer_side_handles = {
remote_engine_id: {rank: 100 + rank for rank in range(prefill_tp_size)}
remote_engine_id: {(rank, 0): 100 + rank for rank in range(prefill_tp_size)}
}
# Sanity: D TP=1, P TP=4 => tp_ratio = -4 (P > D).
assert worker.transfer_topo.tp_ratio(prefill_tp_size) == -prefill_tp_size
Expand Down Expand Up @@ -3141,14 +3215,14 @@ def test_mla_broadcast_notif_uses_remote_request_id(

# MLA: read once from rank 0 and broadcast to the other ranks.
worker._read_blocks.assert_called_once()
assert worker._read_blocks.call_args.kwargs["remote_rank"] == 0
assert worker._read_blocks.call_args.kwargs["read_spec"].remote_rank == 0
assert (
worker._read_blocks.call_args.kwargs["remote_request_id"] == prefill_req_id
)

# Broadcast goes to ranks {1, 2, 3} only, never to the read target.
expected_recipients = {
worker._remote_agents[remote_engine_id][(0, r)]
worker._remote_agents[remote_engine_id][(0, r, 0)]
for r in range(1, prefill_tp_size)
}
assert {agent for agent, _ in send_notif_calls} == expected_recipients
Expand Down
5 changes: 3 additions & 2 deletions tests/v1/kv_connector/unit/test_nixl_connector_hma.py
Original file line number Diff line number Diff line change
Expand Up @@ -212,6 +212,7 @@ def test_read_blocks_for_req_expands_remote_ids(
worker._physical_blocks_per_logical_kv_block = local_physical_per_logical
worker._engine_last_active = {}
worker._bidirectional_kv_xfer_enabled = False
worker._done_recving_without_xfer = set()

has_mamba = any(t is MambaSpec for t in resolved_types)
has_swa = any(t is SlidingWindowSpec for t in resolved_types)
Expand Down Expand Up @@ -1634,8 +1635,8 @@ def test_push_write_hybrid_mla_replicates_attention():
rank_offset_factor=0,
)
}
worker.dst_xfer_side_handles = {engine_id: {0: 100, 1: 101}}
worker.src_xfer_handles_by_tp_ratio = {(-2, 4): [200, 201]}
worker.dst_xfer_side_handles = {engine_id: {(0, 0): 100, (1, 0): 101}}
worker.src_xfer_handles_by_tp_ratio = {(-2, 4): {0: 200, 1: 201}}
worker.src_xfer_handles_by_block_size = {4: 300}
worker._sending_transfers = defaultdict(list)
worker._sending_transfers_lock = threading.Lock()
Expand Down
4 changes: 2 additions & 2 deletions tests/v1/kv_connector/unit/test_nixl_push_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -1058,9 +1058,9 @@ def _mla_worker_writing_to(d_ranks):
)
}
w._logical_to_kernel_block_ids = lambda block_ids, ratio: block_ids
w.dst_xfer_side_handles = {engine_id: {r: 1000 + r for r in d_ranks}}
w.dst_xfer_side_handles = {engine_id: {(r, 0): 1000 + r for r in d_ranks}}
w.src_xfer_handles_by_block_size = {16: 2000}
w._remote_agents = {engine_id: {(0, r): f"agent-{r}" for r in d_ranks}}
w._remote_agents = {engine_id: {(0, r, 0): f"agent-{r}" for r in d_ranks}}
return w, engine_id

def test_mla_hetero_tp_writes_every_d_rank(self):
Expand Down
Loading
Loading