Repository navigation
[Bugfix][NIXL] Stop clearing finished_sending for replicated-PCP producer ranks #59325
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 | ||
| # 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"} | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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( | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. please add a similar test for |
||
| 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 ( | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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( | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
| 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", | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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() | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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.""" | ||
|
|
||
There was a problem hiding this comment.
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.