Skip to content
Open
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
70 changes: 69 additions & 1 deletion tests/v1/kv_connector/unit/test_nixl_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -722,8 +722,76 @@ def test_pcp_producer_exposes_dcp_shards_or_canonical_replica(
worker.get_transfer_results = MagicMock(
return_value=KVConnectorTransferResults(finished_sending={"sent"})
)
# The runner reads completions through get_transfer_results, so it must

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This comment describes the old bug. I think it fits better in the PR description.

# agree with get_finished: replicated ranks report their synthetic
# completions too, or the world_size aggregation never finishes.
results = connector.get_transfer_results(set())
assert results.finished_sending == ({"sent"} if expected_tracked else set())
assert results.finished_sending == {"sent"}

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

With get_transfer_results mocked this only checks a passthrough, so it passes regardless of the fix. The new test below covers the real path, so I think this block can go.


@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
FakeNixlWrapper,
)
def test_replicated_pcp_producer_send_aggregation_completes(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

please add a similar test for test_nixl_push_connector.py as well.

self, default_vllm_config, dist_init
):
"""PCP=2 replicated producer: D reads from the canonical rank 0 and
notifies only it; rank 1 reports a synthetic completion. The
world_size=2 aggregation must report the request as sent so the
scheduler frees its blocks on P."""
from vllm.v1.attention.backends.flash_attn import FlashAttentionBackend

vllm_config = create_vllm_config(kv_role="kv_producer")
vllm_config.parallel_config.prefill_context_parallel_size = 2
vllm_config.parallel_config.decode_context_parallel_size = 1
req_id = "req"
connectors = []
for pcp_rank in (0, 1):
with (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This repeats the connector setup from the parametrized test above. Could it be pulled into a small helper both tests share?

patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl."
"base_worker.get_current_attn_backends",
return_value=[FlashAttentionBackend],
),
patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl."
"base_worker.get_pcp_group"
) as mock_get_pcp_group,
):
mock_get_pcp_group.return_value.rank_in_group = pcp_rank
connector = NixlConnector(
vllm_config,
KVConnectorRole.WORKER,
make_kv_cache_config(block_size=16),
)
worker = connector.connector_worker
worker.transfer_topo = MagicMock()
metadata = NixlConnectorMetadata()
metadata.reqs_in_batch.add(req_id)
metadata.reqs_to_send[req_id] = time.perf_counter() + 30
worker.start_load_kv(metadata)
connectors.append(connector)
connectors[0].connector_worker._get_new_notifs = MagicMock(
return_value={req_id}
)

aggregator = KVOutputAggregator.from_connector(connectors[0], world_size=2)
outputs = [
ModelRunnerOutput(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

create_model_runner_output from tests/v1/kv_connector/unit/utils.py could replace the hand-built ModelRunnerOutput here.

req_ids=[],
req_id_to_index={},
sampled_token_ids=[],
logprobs=None,
prompt_logprobs_dict={},
pooler_output=[],
kv_connector_output=KVConnectorOutput(
finished_sending=c.get_transfer_results(set()).finished_sending
),
)
for c in connectors
]
aggregated = aggregator.aggregate(outputs)
assert aggregated.kv_connector_output.finished_sending == {req_id}

@patch(
"vllm.distributed.kv_transfer.kv_connector.v1.nixl.base_worker.NixlWrapper",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -244,14 +244,7 @@ def get_transfer_results(
self, finished_req_ids: set[str]
) -> KVConnectorTransferResults:
assert self.connector_worker is not None
results = self.connector_worker.get_transfer_results()
if (
self.kv_transfer_config.kv_role == "kv_producer"
and self.connector_worker.pcp_rank > 0
and not self.connector_worker.pcp_dcp_sharded
):
results.finished_sending.clear()
return results
return self.connector_worker.get_transfer_results()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

it's good to add a note in both producer and consumer logic about the contract for replicated-KV PCP scenario.


def get_block_ids_with_load_errors(self) -> set[int]:
"""Get block IDs that failed to load via NIXL."""
Expand Down
Loading