diff --git a/tests/ut/kv_offload/test_mooncake_connector.py b/tests/ut/kv_offload/test_mooncake_connector.py index 82dc816fca51..2422aa0f528e 100644 --- a/tests/ut/kv_offload/test_mooncake_connector.py +++ b/tests/ut/kv_offload/test_mooncake_connector.py @@ -2059,7 +2059,8 @@ def get_block_ids(self): class MockSchedulerOutput: - pass + def __init__(self): + self.num_scheduled_tokens: dict[str, int] = {} class MockForwardContext: @@ -2198,12 +2199,76 @@ def test_update_state_after_alloc_with_remote_prefill(self): self.assertEqual(self.scheduler._reqs_need_recv["req1"][1], ([4, 5, 6],)) self.assertEqual(self.scheduler._reqs_need_recv["req1"][2], ([1, 2, 4, 5, 6],)) + def test_pd_decode_recompute_marker_waits_for_first_forward_step(self): + request = MockRequest( + "req1", + kv_transfer_params={ + "do_remote_prefill": True, + "remote_block_ids": [1, 2, 3], + "remote_engine_id": "remote", + "remote_request_id": "remote_req1", + "remote_host": "localhost", + "remote_port": 5000, + }, + ) + self.scheduler.update_state_after_alloc(request, MockKVCacheBlocks(), 3) + + receive_step = MockSchedulerOutput() + receive_step.num_scheduled_tokens = {"req1": 1} + receive_meta = self.scheduler.build_connector_meta(receive_step) + self.assertEqual(receive_meta.pd_decode_recompute_req_ids, set()) + self.assertEqual(self.scheduler._pd_decode_recompute_req_ids, {"req1"}) + + self.scheduler.update_state_after_alloc(request, MockKVCacheBlocks(), 0) + forward_step = MockSchedulerOutput() + forward_step.num_scheduled_tokens = {"req1": 1, "decode": 1} + forward_meta = self.scheduler.build_connector_meta(forward_step) + self.assertEqual(forward_meta.pd_decode_recompute_req_ids, {"req1"}) + self.assertEqual(self.scheduler._pd_decode_recompute_req_ids, set()) + + later_step = MockSchedulerOutput() + later_step.num_scheduled_tokens = {"req1": 1} + later_meta = self.scheduler.build_connector_meta(later_step) + self.assertEqual(later_meta.pd_decode_recompute_req_ids, set()) + def test_request_finished_no_remote_decode(self): request = MockRequest("req1") delay_free, params = self.scheduler.request_finished(request, [1, 2, 3]) self.assertFalse(delay_free) self.assertIsNone(params) + def test_request_finished_clears_pd_decode_recompute_marker(self): + for phase in ("pending", "ready"): + with self.subTest(phase=phase): + request = MockRequest("req1") + self.scheduler._pd_decode_recompute_req_ids = {"req1", "other"} + self.scheduler._pd_decode_recompute_ready_req_ids = {"req1", "other"} if phase == "ready" else {"other"} + + delay_free, params = self.scheduler.request_finished(request, [1, 2, 3]) + + self.assertFalse(delay_free) + self.assertIsNone(params) + self.assertEqual(self.scheduler._pd_decode_recompute_req_ids, {"other"}) + self.assertEqual( + self.scheduler._pd_decode_recompute_ready_req_ids, + {"other"}, + ) + + metadata = self.scheduler.build_connector_meta(MockSchedulerOutput()) + self.assertEqual(metadata.pd_decode_recompute_req_ids, {"other"}) + + def test_finished_marker_does_not_contaminate_reused_request_id(self): + finished = MockRequest("req1") + self.scheduler._pd_decode_recompute_req_ids.add("req1") + self.scheduler._pd_decode_recompute_ready_req_ids.add("req1") + self.scheduler.request_finished(finished, [1, 2, 3]) + + reused = MockRequest("req1") + self.scheduler.update_state_after_alloc(reused, MockKVCacheBlocks(), 0) + metadata = self.scheduler.build_connector_meta(MockSchedulerOutput()) + + self.assertEqual(metadata.pd_decode_recompute_req_ids, set()) + def test_request_finished_rejected_remote_prefill_enqueues_empty_recv(self): request = MockRequest( "req1", diff --git a/tests/ut/worker/test_model_runner_v2.py b/tests/ut/worker/test_model_runner_v2.py index 211a82139606..7b84c09e0c33 100644 --- a/tests/ut/worker/test_model_runner_v2.py +++ b/tests/ut/worker/test_model_runner_v2.py @@ -10,7 +10,7 @@ from vllm.config import CUDAGraphMode from vllm.v1.kv_cache_interface import KVCacheConfig from vllm.v1.worker.gpu import model_runner as vllm_model_runner -from vllm.v1.worker.gpu.model_runner import GPUModelRunner +from vllm.v1.worker.gpu.model_runner import BatchReqState, GPUModelRunner from vllm_ascend.ascend_forward_context import MoECommType from vllm_ascend.utils import vllm_version_is @@ -36,6 +36,182 @@ def _make_runner(need_timing: bool = True): return runner +def test_pd_decode_tail_recompute_keeps_uniform_decode_batch(): + runner = _make_runner() + runner.decode_query_len = 1 + runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=True)) + batch_state = BatchReqState( + req_ids=["migrating", "decoding"], + num_scheduled_tokens=np.array([1, 1], dtype=np.int32), + num_tokens=2, + idx_mapping_np=np.array([0, 1], dtype=np.intp), + prefill_len_np=np.array([128, 64], dtype=np.int32), + num_computed_prefill_tokens_np=np.array([127, 64], dtype=np.int32), + is_prefilling_np=np.array([True, False]), + has_prefill=True, + ) + + transfer_step = SimpleNamespace( + kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids={"migrating"}), + finished_req_ids=set(), + preempted_req_ids=None, + ) + runner._update_pd_decode_recompute_requests(transfer_step) + + with patch.object( + GPUModelRunner, + "gather_batch_req_state", + return_value=(batch_state, None), + ): + scheduler_output = SimpleNamespace(kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids=set())) + gathered, uniform = runner.gather_batch_req_state(scheduler_output, False) + + assert gathered is not batch_state + np.testing.assert_array_equal(gathered.is_prefilling_np, [False, False]) + assert not gathered.has_prefill + assert uniform == 1 + assert not runner._pd_decode_recompute_req_ids + + +@pytest.mark.parametrize("computed,scheduled", [(64, 8), (127, 1)]) +def test_pd_decode_real_prefill_still_blocks_decode_graph(computed, scheduled): + runner = _make_runner() + runner.decode_query_len = 1 + runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=True)) + batch_state = BatchReqState( + req_ids=["fallback-prefill", "decoding"], + num_scheduled_tokens=np.array([scheduled, 1], dtype=np.int32), + num_tokens=scheduled + 1, + idx_mapping_np=np.array([0, 1], dtype=np.intp), + prefill_len_np=np.array([128, 64], dtype=np.int32), + num_computed_prefill_tokens_np=np.array([computed, 64], dtype=np.int32), + is_prefilling_np=np.array([True, False]), + has_prefill=True, + ) + + with patch.object( + GPUModelRunner, + "gather_batch_req_state", + return_value=(batch_state, None), + ): + scheduler_output = SimpleNamespace(kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids=set())) + gathered, uniform = runner.gather_batch_req_state(scheduler_output, False) + + np.testing.assert_array_equal(gathered.is_prefilling_np, [True, False]) + assert gathered.has_prefill + assert uniform is None + + +def test_pd_decode_only_reclassifies_transferred_requests_in_mixed_prefill_batch(): + runner = _make_runner() + runner.decode_query_len = 1 + runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=True)) + batch_state = BatchReqState( + req_ids=["migrating", "local-final-prefill", "decoding"], + num_scheduled_tokens=np.array([1, 1, 1], dtype=np.int32), + num_tokens=3, + idx_mapping_np=np.array([0, 1, 2], dtype=np.intp), + prefill_len_np=np.array([128, 128, 64], dtype=np.int32), + num_computed_prefill_tokens_np=np.array([127, 127, 64], dtype=np.int32), + is_prefilling_np=np.array([True, True, False]), + has_prefill=True, + ) + + with patch.object( + GPUModelRunner, + "gather_batch_req_state", + return_value=(batch_state, None), + ): + scheduler_output = SimpleNamespace( + kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids={"migrating"}) + ) + gathered, uniform = runner.gather_batch_req_state(scheduler_output, False) + + np.testing.assert_array_equal(gathered.is_prefilling_np, [False, True, False]) + assert gathered.has_prefill + assert uniform is None + + +def test_pd_decode_transfer_tracking_clears_finished_and_preempted_requests(): + runner = _make_runner() + runner._pd_decode_recompute_req_ids = {"finished", "preempted", "active"} + scheduler_output = SimpleNamespace( + kv_connector_metadata=SimpleNamespace( + metadata=( + SimpleNamespace(pd_decode_recompute_req_ids={"new"}), + SimpleNamespace(pd_decode_recompute_req_ids=set()), + ) + ), + finished_req_ids={"finished"}, + preempted_req_ids={"preempted"}, + ) + + runner._update_pd_decode_recompute_requests(scheduler_output) + + assert runner._pd_decode_recompute_req_ids == {"active", "new"} + + +def test_pd_decode_tail_recompute_supports_multi_token_decode_query(): + runner = _make_runner() + runner.decode_query_len = 2 + runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=True)) + batch_state = BatchReqState( + req_ids=["migrating", "decoding"], + num_scheduled_tokens=np.array([2, 2], dtype=np.int32), + num_tokens=4, + idx_mapping_np=np.array([0, 1], dtype=np.intp), + prefill_len_np=np.array([128, 64], dtype=np.int32), + num_computed_prefill_tokens_np=np.array([127, 64], dtype=np.int32), + is_prefilling_np=np.array([True, False]), + has_prefill=True, + ) + + with patch.object( + GPUModelRunner, + "gather_batch_req_state", + return_value=(batch_state, None), + ): + scheduler_output = SimpleNamespace( + kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids={"migrating"}) + ) + gathered, uniform = runner.gather_batch_req_state(scheduler_output, False) + + np.testing.assert_array_equal(gathered.is_prefilling_np, [False, False]) + assert not gathered.has_prefill + assert uniform == 2 + + +def test_pd_decode_non_consumer_does_not_reclassify_transferred_request(): + runner = _make_runner() + runner.decode_query_len = 1 + runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=False)) + batch_state = BatchReqState( + req_ids=["transferred"], + num_scheduled_tokens=np.array([1], dtype=np.int32), + num_tokens=1, + idx_mapping_np=np.array([0], dtype=np.intp), + prefill_len_np=np.array([128], dtype=np.int32), + num_computed_prefill_tokens_np=np.array([127], dtype=np.int32), + is_prefilling_np=np.array([True]), + has_prefill=True, + ) + + with patch.object( + GPUModelRunner, + "gather_batch_req_state", + return_value=(batch_state, None), + ): + scheduler_output = SimpleNamespace( + kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids={"transferred"}) + ) + gathered, uniform = runner.gather_batch_req_state(scheduler_output, False) + + assert gathered is batch_state + np.testing.assert_array_equal(gathered.is_prefilling_np, [True]) + assert gathered.has_prefill + assert uniform is None + + def test_execute_model_records_profiling_time(): runner = _make_runner() scheduler_output = SimpleNamespace(disable_profiling_timing=False) diff --git a/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py b/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py index a26b60c54403..e6cf905ca1f1 100644 --- a/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py +++ b/vllm_ascend/distributed/kv_transfer/kv_p2p/mooncake_connector.py @@ -1576,6 +1576,9 @@ def __init__(self): self.requests: dict[str, ReqMeta] = {} self.requests_to_send: dict[str, float] = {} self.reqs_in_batch: set[str] = set() + # Requests whose first local forward after a remote prefill is the + # decode-side tail-token recomputation. + self.pd_decode_recompute_req_ids: set[str] = set() def add_new_req( self, @@ -1768,6 +1771,8 @@ def __init__(self, vllm_config: VllmConfig, engine_id: str, kv_cache_config: KVC self._reqs_need_recv: dict[str, tuple[Request, BlockIds, BlockIds, int]] = {} self._reqs_need_send: dict[str, float] = {} self._reqs_in_batch: set[str] = set() + self._pd_decode_recompute_req_ids: set[str] = set() + self._pd_decode_recompute_ready_req_ids: set[str] = set() # master-slave meta information for cross-nodes self.multi_nodes_meta_mapping: dict[str, dict[str, Any]] = {} @@ -1948,6 +1953,12 @@ def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", if params is not None and (params.get("do_remote_prefill", False) or params.get("do_remote_decode", False)): self._reqs_in_batch.add(request.request_id) + if ( + request.request_id in self._pd_decode_recompute_req_ids + and params is not None + and not params.get("do_remote_prefill", False) + ): + self._pd_decode_recompute_ready_req_ids.add(request.request_id) if params is not None and params.get("do_remote_prefill"): if params.get("remote_block_ids"): if all(p in params for p in ("remote_engine_id", "remote_host", "remote_port", "remote_request_id")): @@ -1960,6 +1971,7 @@ def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", local_full_block_ids, num_external_tokens, ) + self._pd_decode_recompute_req_ids.add(request.request_id) else: logger.warning("Got invalid KVTransferParams. params=%s. ", params) else: @@ -1994,6 +2006,18 @@ def build_connector_meta( meta.reqs_in_batch = self._reqs_in_batch self._reqs_in_batch = set() + # KV receive is scheduled in a no-forward step. update_state_after_alloc + # marks the same request ready when the scheduler later allocates its + # first local forward, which is the decode-side tail recomputation. + meta.pd_decode_recompute_req_ids = self._pd_decode_recompute_ready_req_ids + if meta.pd_decode_recompute_req_ids: + logger.debug( + "Emitting PD decode recompute requests: %s", + sorted(meta.pd_decode_recompute_req_ids), + ) + self._pd_decode_recompute_req_ids.difference_update(meta.pd_decode_recompute_req_ids) + self._pd_decode_recompute_ready_req_ids = set() + return meta def request_finished( @@ -2006,6 +2030,13 @@ def request_finished( should be freed now or will be sent asynchronously and freed later. """ + # A request can finish or be aborted while its remote prefill is still + # pending, or after it becomes ready but before metadata is built. + # Clear both phases before any early return so stale request IDs cannot + # leak or be applied to a later request that reuses the same ID. + self._pd_decode_recompute_req_ids.discard(request.request_id) + self._pd_decode_recompute_ready_req_ids.discard(request.request_id) + params = request.kv_transfer_params logger.debug( "MooncakeConnector request_finished, request_status=%s, kv_transfer_params=%s", request.status, params diff --git a/vllm_ascend/worker/v2/model_runner.py b/vllm_ascend/worker/v2/model_runner.py index 59f3f4f9f5a2..fb903f978989 100644 --- a/vllm_ascend/worker/v2/model_runner.py +++ b/vllm_ascend/worker/v2/model_runner.py @@ -85,6 +85,17 @@ from vllm.v1.worker.gpu.cp_utils import prepare_dcp_local_seq_lens +def _get_kv_transfer_req_ids(metadata: object | None) -> set[str]: + """Return requests explicitly marked for PD tail recomputation.""" + if metadata is None: + return set() + + req_ids = set(getattr(metadata, "pd_decode_recompute_req_ids", ())) + for child_metadata in getattr(metadata, "metadata", ()): + req_ids.update(_get_kv_transfer_req_ids(child_metadata)) + return req_ids + + class NPUModelRunner(GPUModelRunner): """Model runner for Ascend NPUs.""" @@ -300,6 +311,8 @@ def execute_model( context_len: int = 0, valid_dummy_state_slots: bool = False, ): + if not dummy_run: + self._update_pd_decode_recompute_requests(scheduler_output) self._cpp_execution_time_ms = None profiling_config = self.ascend_config.scheduler_config.profiling_chunk_config execution_start_time = _start_profiling_chunk_timing( @@ -334,6 +347,16 @@ def execute_model( ) return output + def _update_pd_decode_recompute_requests(self, scheduler_output: SchedulerOutput) -> None: + pending_req_ids: set[str] = getattr(self, "_pd_decode_recompute_req_ids", set()) + pending_req_ids.difference_update(getattr(scheduler_output, "finished_req_ids", set())) + pending_req_ids.difference_update(getattr(scheduler_output, "preempted_req_ids", None) or set()) + new_req_ids = _get_kv_transfer_req_ids(getattr(scheduler_output, "kv_connector_metadata", None)) + if new_req_ids: + logger.debug("Tracking PD decode KV transfer requests: %s", sorted(new_req_ids)) + pending_req_ids.update(new_req_ids) + self._pd_decode_recompute_req_ids = pending_req_ids + @torch.inference_mode() def profile_run(self) -> None: """Override GPUModelRunner.profile_run for Ascend NPUs. @@ -356,6 +379,47 @@ def profile_run(self) -> None: def gather_batch_req_state(self, scheduler_output: SchedulerOutput, dummy_run: bool): batch_state, uniform_token_count = super().gather_batch_req_state(scheduler_output, dummy_run) + kv_transfer_config = getattr(self.vllm_config, "kv_transfer_config", None) + if batch_state is not None and kv_transfer_config is not None and kv_transfer_config.is_kv_consumer: + # A PD decode consumer recomputes the last prompt token after the + # transferred hybrid state is installed. Upstream MRV2 derives + # ``is_prefilling`` solely from the prompt boundary, so that + # one-token recompute row incorrectly turns a mixed decode batch + # into a prefill batch and disables FULL_DECODE_ONLY replay. + self._update_pd_decode_recompute_requests(scheduler_output) + kv_transfer_req_ids = self._pd_decode_recompute_req_ids + is_kv_transfer_req = np.fromiter( + (req_id in kv_transfer_req_ids for req_id in batch_state.req_ids), + dtype=np.bool_, + count=len(batch_state.req_ids), + ) + pd_decode_recompute = ( + batch_state.is_prefilling_np + & is_kv_transfer_req + & (batch_state.num_computed_prefill_tokens_np > 0) + & (batch_state.num_scheduled_tokens == self.decode_query_len) + & ( + batch_state.num_computed_prefill_tokens_np + batch_state.num_scheduled_tokens + >= batch_state.prefill_len_np + ) + ) + if np.any(pd_decode_recompute): + matched_req_ids = [ + req_id for req_id, matched in zip(batch_state.req_ids, pd_decode_recompute) if matched + ] + logger.debug( + "Treating PD decode tail recompute as decode for requests: %s", + matched_req_ids, + ) + self._pd_decode_recompute_req_ids.difference_update(matched_req_ids) + batch_state.is_prefilling_np[pd_decode_recompute] = False + batch_state = batch_state._replace(has_prefill=bool(batch_state.is_prefilling_np.any())) + uniform_token_count = vllm_model_runner.get_uniform_decode_token_count( + len(batch_state.req_ids), + batch_state.num_tokens, + int(batch_state.num_scheduled_tokens.max()), + batch_state.has_prefill, + ) num_tokens = None if vllm_version_is("0.28.0") and self.pcp_manager is not None and batch_state is not None: num_tokens = self.pcp_manager.get_num_tokens_for_dispatch(