diff --git a/tests/v1/worker/test_gpu_model_runner_v2.py b/tests/v1/worker/test_gpu_model_runner_v2.py index 308c9fe9eaab..2e0c7271ce62 100644 --- a/tests/v1/worker/test_gpu_model_runner_v2.py +++ b/tests/v1/worker/test_gpu_model_runner_v2.py @@ -3,6 +3,7 @@ import contextlib from types import SimpleNamespace +from unittest.mock import Mock import pytest import torch @@ -18,7 +19,43 @@ UniformTypeKVCacheSpecs, ) from vllm.v1.worker.gpu.block_table import BlockTables -from vllm.v1.worker.gpu.model_runner import GPUModelRunner +from vllm.v1.worker.gpu.model_runner import ExecuteModelState, GPUModelRunner + + +def test_non_last_pp_rank_uses_global_batch_for_sample_feedback(): + runner = GPUModelRunner.__new__(GPUModelRunner) + runner.is_last_pp_rank = False + local_batch = object() + global_batch = SimpleNamespace(idx_mapping=object()) + runner.pcp_manager = SimpleNamespace( + global_batch=global_batch, + restore_for_sampling=Mock(), + ) + runner.pp_handler = SimpleNamespace(receive=Mock(return_value=False)) + runner.postprocess_num_computed_tokens = Mock() + runner.model_state = SimpleNamespace(postprocess_state=Mock()) + runner.kv_connector = SimpleNamespace(post_forward=Mock(return_value=None)) + runner.eplb = SimpleNamespace(step=Mock()) + runner.execute_model_state = ExecuteModelState( + input_batch=local_batch, + attn_metadata=None, + slot_mappings_by_layer=None, + hidden_states=None, + aux_hidden_states=None, + dp_sync=None, + finished_req_ids=set(), + ec_connector_output=None, + cudagraph_stats=None, + ) + + runner.sample_tokens(None) + + runner.pp_handler.receive.assert_called_once_with(global_batch) + runner.postprocess_num_computed_tokens.assert_called_once_with(global_batch) + runner.model_state.postprocess_state.assert_called_once_with( + global_batch.idx_mapping, 0 + ) + runner.pcp_manager.restore_for_sampling.assert_not_called() def test_qsa_circular_group_uses_custom_slot_mapping(monkeypatch): diff --git a/tests/v1/worker/test_gpu_model_runner_v2_eplb.py b/tests/v1/worker/test_gpu_model_runner_v2_eplb.py index 7c759894fe1c..bbde3295ec62 100644 --- a/tests/v1/worker/test_gpu_model_runner_v2_eplb.py +++ b/tests/v1/worker/test_gpu_model_runner_v2_eplb.py @@ -79,6 +79,7 @@ def _make_runner(**overrides: Any) -> Any: runner.use_aux_hidden_state_outputs = False runner.speculative_config = None runner.speculator = None + runner.pcp_manager = None runner.num_speculative_steps = 0 runner.encoder_cache = None runner.is_pooling_model = False diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 800ae825d4c0..3bdeda23dd96 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -2064,6 +2064,8 @@ def sample_tokens( # Non-last PP rank: hidden_states is None because this rank produced # IntermediateTensors instead of final hidden states. Receive the # sampled tokens broadcast from the last rank and update local state. + if self.pcp_manager is not None: + input_batch = self.pcp_manager.global_batch assert self.pp_handler is not None all_decode_next = self.pp_handler.receive(input_batch) # Optimistically update num_computed_tokens for entire batch here. diff --git a/vllm/v1/worker/gpu/pcp_manager.py b/vllm/v1/worker/gpu/pcp_manager.py index 11ee7522fc9f..f740e4caf355 100644 --- a/vllm/v1/worker/gpu/pcp_manager.py +++ b/vllm/v1/worker/gpu/pcp_manager.py @@ -131,8 +131,6 @@ def validate_config( if not model_config.use_mla: raise NotImplementedError("MRV2 PCP currently supports MLA models only.") - if parallel_config.pipeline_parallel_size > 1: - raise NotImplementedError("MRV2 PCP does not support PP yet.") if model_config.is_encoder_decoder: raise NotImplementedError( "MRV2 PCP does not support encoder-decoder models yet." @@ -406,6 +404,11 @@ def input_buffers(self) -> InputBuffers: assert self._input_buffers is not None return self._input_buffers + @property + def global_batch(self) -> InputBatch: + assert self._global_batch is not None + return self._global_batch + def partition_batch( self, input_batch: InputBatch, batch_desc: "BatchExecutionDescriptor" ) -> InputBatch: