diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 1585ac07fdff..0f92d90dc6ec 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -74,6 +74,10 @@ def _find_consensus_request_ids(request_ids_all_ranks, sync_size): class KvCacheTransceiverV2(KvCacheTransceiver): + @property + def consumes_transfer_buffer(self) -> bool: + return False + def __init__( self, mapping: Mapping, diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py index 65f613337fff..c0bfa3ee3c87 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py @@ -232,6 +232,11 @@ def create_kv_cache_transceiver( class KvCacheTransceiver(ABC): + @property + def consumes_transfer_buffer(self) -> bool: + """Return whether this runtime consumes the C++ CacheTransBuffer budget.""" + return True + @abstractmethod def respond_and_send_async(self, req: LlmRequest): raise NotImplementedError diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index a27a2a417275..07e9da1a2f44 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -893,6 +893,16 @@ def on_detected(): None) self._disagg_transfer_admission_controller = DisaggTransferAdmissionController( max_tokens_in_buffer, tokens_per_block) + if (self.global_rank == 0 + and self._disagg_transfer_admission_controller.enabled() + and self._is_disagg_transfer_window_bypass_eligible()): + logger.warning_once( + f"[PyExecutor] Bypassing the executor transfer window " + f"configured by max_tokens_in_buffer={max_tokens_in_buffer} " + "for asynchronous Python generation with KV cache manager " + "V2 and pp_size=1; " + "scheduler KV cache capacity admission remains active.", + key="disagg_transfer_window_bypass") self.is_benchmark_disagg = (self.benchmark_req_queues_size > 0 and self.kv_cache_transceiver is not None) # True while the benchmark disagg fill phase is in progress (waiting @@ -2618,7 +2628,7 @@ def _executor_loop_pp(self): local_disagg_candidates = getattr( local_scheduler_output, "fitting_disagg_gen_init_requests", []) - self._revert_deferred_disagg_gen_init_alloc( + self._reconcile_disagg_gen_init_allocations( local_disagg_candidates, fitting_disagg_gen_init_requests) @@ -3492,6 +3502,27 @@ def _get_disagg_transfer_admission_controller( getattr(kv_cache_manager, "tokens_per_block", None), ) + def _is_disagg_transfer_window_bypass_eligible(self) -> bool: + """Return whether this runtime may bypass an enabled transfer window.""" + transceiver = getattr(self, "kv_cache_transceiver", None) + return (transceiver is not None + and transceiver.consumes_transfer_buffer is False + and self._uses_async_disagg_gen_transfer() + and self.dist.pp_size == 1 and self._uses_kv_manager_v2()) + + def _disagg_transfer_window_is_active(self) -> bool: + """Return whether the executor-level transfer window is active. + + ``max_tokens_in_buffer`` describes the C++ transceiver's physical + buffer. The asynchronous Python transceiver does not consume that + buffer. With KV cache manager V2 and PP1, its generation requests + remain constrained by inline scheduler KV admission without a second + executor-level budget. Other configurations retain the window. + """ + return (getattr(self, "kv_cache_transceiver", None) is not None + and self._get_disagg_transfer_admission_controller().enabled() + and not self._is_disagg_transfer_window_bypass_eligible()) + @staticmethod def _is_disagg_gen_only_no_context_benchmark() -> bool: """Return whether ``gen_only_no_context`` skips KV transfer.""" @@ -3519,8 +3550,8 @@ def _apply_disagg_transfer_admission( return fitting_disagg_gen_init_requests, False controller = self._get_disagg_transfer_admission_controller() - if not (getattr(self, "kv_cache_transceiver", None) - and controller.enabled() and fitting_disagg_gen_init_requests): + if not (self._disagg_transfer_window_is_active() + and fitting_disagg_gen_init_requests): return fitting_disagg_gen_init_requests, False admission_result = controller.select(self.active_requests, @@ -3534,29 +3565,36 @@ def _apply_disagg_transfer_admission( f"{admission_result.admitted_transfer_blocks}, " f"budget={controller.max_transfer_blocks}") - self._revert_deferred_disagg_gen_init_alloc( + self._reconcile_disagg_gen_init_allocations( fitting_disagg_gen_init_requests, admission_result.admitted_requests) return (admission_result.admitted_requests, admission_result.is_blocked_by_active_transfers()) - def _revert_deferred_disagg_gen_init_alloc( + def _reconcile_disagg_gen_init_allocations( self, candidates: List[LlmRequest], - admitted_requests: List[LlmRequest]) -> None: + selected_requests: List[LlmRequest]) -> None: + """Revert Scheduler V2 allocations absent from a selected request set. + + Scheduler V2 allocates KV while evaluating generation-init requests. + This reconciliation is required both after transfer-window admission + and when a PP follower's local candidates differ from the canonical + schedule propagated by rank 0. + """ if not (self._uses_kv_manager_v2() and candidates): return - admitted_request_ids = { + selected_request_ids = { request.py_request_id - for request in admitted_requests + for request in selected_requests } - deferred_requests = [ + unselected_requests = [ request for request in candidates - if request.py_request_id not in admitted_request_ids + if request.py_request_id not in selected_request_ids ] - if deferred_requests: - self._revert_ctx_alloc(deferred_requests) + if unselected_requests: + self._revert_ctx_alloc(unselected_requests) @staticmethod def _dist_size(dist, name: str) -> int: diff --git a/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py b/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py index 0294bcd0ecf4..9e0bf5c944e6 100644 --- a/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py +++ b/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py @@ -644,7 +644,13 @@ def test_flag_unset_preserves_legacy_backend_env_precedence(monkeypatch): ("UCX", "CPP", "UCX", False), ], ) -def test_cpp_capability_is_config_scoped(monkeypatch, backend, runtime, nixl_backend, expected): +def test_cpp_capability_is_config_scoped( + monkeypatch: pytest.MonkeyPatch, + backend: str, + runtime: str | None, + nixl_backend: str | None, + expected: bool, +) -> None: if nixl_backend is not None: monkeypatch.setenv(transceiver_module._NIXL_KVCACHE_BACKEND_ENV, nixl_backend) config = CacheTransceiverConfig(backend=backend, transceiver_runtime=runtime) @@ -668,15 +674,17 @@ def test_cpp_capability_is_config_scoped(monkeypatch, backend, runtime, nixl_bac transceiver = BindKvCacheTransceiver(Mock(), dist, kv_cache_manager, Mock(), config) + assert transceiver.consumes_transfer_buffer assert transceiver.supports_inflight_request_cancellation() is expected constructor.assert_called_once() -def test_python_transceiver_capability_defaults_to_unsupported(): +def test_python_transceiver_capability_defaults_to_unsupported() -> None: from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 transceiver = object.__new__(KvCacheTransceiverV2) + assert not transceiver.consumes_transfer_buffer assert not transceiver.supports_inflight_request_cancellation() assert not transceiver.has_poisoned_transfer_buffer() diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index 607f2b0871bc..82318666f6a5 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -765,6 +765,13 @@ def _make_disagg_transfer_request( return req +def _set_disagg_transceiver_capability( + executor: PyExecutor, *, consumes_transfer_buffer: bool +) -> None: + executor.kv_cache_transceiver = Mock() + executor.kv_cache_transceiver.consumes_transfer_buffer = consumes_transfer_buffer + + @pytest.fixture def _clear_disagg_transfer_mode_env(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("TRTLLM_DISAGG_BENCHMARK_GEN_ONLY", raising=False) @@ -838,7 +845,7 @@ def test_uses_global_cp_prompt_length_for_transfer_cost(self): def test_apply_reverts_deferred_v2_allocations(self): executor = object.__new__(PyExecutor) - executor.kv_cache_transceiver = Mock() + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=True) executor._is_kv_manager_v2 = True executor._revert_ctx_alloc = Mock() executor.active_requests = [_make_disagg_transfer_request(1, 32, in_progress=True)] @@ -855,6 +862,89 @@ def test_apply_reverts_deferred_v2_allocations(self): assert wait_for_progress executor._revert_ctx_alloc.assert_called_once_with([candidate]) + def test_async_python_v2_pp1_bypasses_transfer_budget(self) -> None: + executor = object.__new__(PyExecutor) + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=False) + executor.dist = Mock(pp_size=1) + executor._is_kv_manager_v2 = True + executor._revert_ctx_alloc = Mock() + executor.active_requests = [_make_disagg_transfer_request(1, 32, in_progress=True)] + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + candidates = [ + _make_disagg_transfer_request(2, 32), + _make_disagg_transfer_request(3, 32), + ] + + admitted, wait_for_progress = PyExecutor._apply_disagg_transfer_admission( + executor, candidates + ) + + assert admitted == candidates + assert not wait_for_progress + executor._revert_ctx_alloc.assert_not_called() + + def test_async_python_v1_pp1_retains_transfer_budget(self) -> None: + executor = object.__new__(PyExecutor) + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=False) + executor.dist = Mock(pp_size=1) + executor._is_kv_manager_v2 = False + executor._revert_ctx_alloc = Mock() + executor.active_requests = [_make_disagg_transfer_request(1, 32, in_progress=True)] + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + candidate = _make_disagg_transfer_request(2, 32) + + admitted, wait_for_progress = PyExecutor._apply_disagg_transfer_admission( + executor, [candidate] + ) + + assert admitted == [] + assert wait_for_progress + executor._revert_ctx_alloc.assert_not_called() + + def test_disabled_transfer_window_is_inactive(self) -> None: + executor = object.__new__(PyExecutor) + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=False) + executor.dist = Mock(pp_size=1) + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=0, tokens_per_block=32 + ) + + assert not PyExecutor._disagg_transfer_window_is_active(executor) + + def test_transfer_window_without_transceiver_is_inactive(self) -> None: + executor = object.__new__(PyExecutor) + executor.kv_cache_transceiver = None + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + + assert not PyExecutor._disagg_transfer_window_is_active(executor) + + def test_active_window_check_requires_initialized_dist(self) -> None: + executor = object.__new__(PyExecutor) + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=False) + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + + with pytest.raises(AttributeError): + PyExecutor._disagg_transfer_window_is_active(executor) + + def test_active_window_check_requires_pp_size(self) -> None: + executor = object.__new__(PyExecutor) + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=False) + executor.dist = types.SimpleNamespace() + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + + with pytest.raises(AttributeError): + PyExecutor._disagg_transfer_window_is_active(executor) + def test_apply_missing_controller_preserves_candidates(self): executor = object.__new__(PyExecutor) executor.kv_cache_transceiver = Mock() @@ -870,7 +960,7 @@ def test_apply_missing_controller_preserves_candidates(self): def test_apply_missing_v2_flag_defaults_to_non_v2(self): executor = object.__new__(PyExecutor) - executor.kv_cache_transceiver = Mock() + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=True) executor._revert_ctx_alloc = Mock() executor.active_requests = [_make_disagg_transfer_request(1, 32, in_progress=True)] executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( @@ -886,10 +976,13 @@ def test_apply_missing_v2_flag_defaults_to_non_v2(self): assert wait_for_progress executor._revert_ctx_alloc.assert_not_called() - def test_sync_mode_retains_transfer_budget(self, monkeypatch): + def test_sync_python_runtime_retains_transfer_budget( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: monkeypatch.setenv("TRTLLM_DISABLE_KV_CACHE_TRANSFER_OVERLAP", "1") executor = object.__new__(PyExecutor) - executor.kv_cache_transceiver = Mock() + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=False) + executor.dist = Mock(pp_size=1) executor._is_kv_manager_v2 = True executor._revert_ctx_alloc = Mock() executor.active_requests = [] @@ -1209,13 +1302,13 @@ def test_pp_ring_drained_only_when_no_microbatch_is_outstanding( @pytest.mark.usefixtures("_clear_disagg_transfer_mode_env") class TestDisaggTransferAdmissionPP: - def test_pp_schedule_applies_gate_before_serializing(self): + def test_pp_schedule_applies_gate_before_serializing(self) -> None: executor = object.__new__(PyExecutor) + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=True) executor.dist = Mock( rank=0, is_first_pp_rank=True, is_last_pp_rank=True, tp_size=1, cp_size=1 ) executor.enable_attention_dp = False - executor.kv_cache_transceiver = Mock() executor.active_requests = [_make_disagg_transfer_request(1, 32, in_progress=True)] executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( max_tokens_in_buffer=32, tokens_per_block=32 @@ -1233,6 +1326,44 @@ def test_pp_schedule_applies_gate_before_serializing(self): assert num_fitting == 0 assert wait_for_progress + def test_pp_schedule_async_python_retains_transfer_window(self) -> None: + executor = object.__new__(PyExecutor) + _set_disagg_transceiver_capability(executor, consumes_transfer_buffer=False) + executor.dist = Mock( + rank=0, + is_first_pp_rank=True, + is_last_pp_rank=False, + tp_size=1, + cp_size=1, + pp_size=2, + next_pp_rank=1, + ) + executor._is_kv_manager_v2 = True + executor.enable_attention_dp = False + executor.send_schedule_handles = [None] + executor.wait_on_pp_send_handles = Mock() + executor.active_requests = [_make_disagg_transfer_request(1, 32, in_progress=True)] + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + scheduled_batch = ScheduledRequests() + candidate = _make_disagg_transfer_request(2, 32) + executor._schedule = Mock(return_value=(scheduled_batch, [candidate], 0)) + + scheduled, fitting, num_fitting, wait_for_progress = PyExecutor._pp_schedule_and_propagate( + executor, microbatch_id=0 + ) + + assert scheduled is scheduled_batch + assert fitting == [] + assert num_fitting == 0 + assert wait_for_progress + executor.wait_on_pp_send_handles.assert_called_once() + wait_args = executor.wait_on_pp_send_handles.call_args.args + assert wait_args[0] is executor.send_schedule_handles + assert wait_args[1] == 0 + executor.dist.isend_object.assert_called_once() + def test_pp_schedule_restores_propagated_gate_decision(self): executor = object.__new__(PyExecutor) executor.dist = Mock( @@ -1321,6 +1452,69 @@ def stop_after_schedule(requests, inflight_req_ids): ] +def test_nonzero_pp_rank_reconciles_local_only_disagg_allocations( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class StopAfterReconciliation(RuntimeError): + pass + + executor = object.__new__(PyExecutor) + executor.dist = Mock(pp_rank=1, rank=1) + executor.device_id = 0 + profiler = MagicMock() + profiler.__enter__.return_value = Mock() + executor._profiler = Mock(return_value=profiler) + executor.hang_detector = MagicMock() + executor.enable_iter_perf_stats = False + executor._is_kv_manager_v2 = True + executor._pp_rebalance_drain_iters = None + executor._can_pause_for_rebalance = Mock(return_value=False) + executor._handle_disagg_cache_errors_synced = Mock() + executor._fetch_and_activate_new_requests = Mock(return_value=[]) + executor.is_shutdown = False + executor._handle_control_request = Mock() + executor.kv_cache_transceiver = Mock() + executor._check_disagg_ctx_schedulable_status = Mock() + executor._check_disagg_gen_transfer_status = Mock() + executor._pad_attention_dp_dummy_request = Mock() + executor._pp_retry_until_can_schedule = Mock() + executor._mm_encoder_item_scheduling_enabled = False + executor._terminate_recompute_paused_requests = Mock() + executor._pause_recompute_paused_requests = Mock() + executor._revert_ctx_alloc = Mock() + executor.inflight_req_ids = set() + executor.kv_cache_manager = Mock() + executor.scheduler = Mock() + + canonical = _make_disagg_transfer_request(1, 32) + local_canonical = _make_disagg_transfer_request(1, 32) + local_only = _make_disagg_transfer_request(2, 32) + executor.active_requests = [canonical, local_only] + scheduled_batch = ScheduledRequests() + executor._pp_schedule_and_propagate = Mock( + return_value=(scheduled_batch, [canonical], 0, False) + ) + executor.scheduler.schedule_request.return_value = types.SimpleNamespace( + fitting_disagg_gen_init_requests=[local_canonical, local_only] + ) + executor._prepare_disagg_gen_init = Mock(side_effect=StopAfterReconciliation) + + monkeypatch.setattr( + "tensorrt_llm._torch.pyexecutor.py_executor.torch.cuda.set_device", + Mock(), + ) + monkeypatch.setattr( + "tensorrt_llm._torch.pyexecutor.py_executor.cudart.cudaSetDevice", + Mock(), + ) + monkeypatch.setattr("tensorrt_llm._torch.pyexecutor.py_executor.CUASSERT", Mock()) + + with pytest.raises(StopAfterReconciliation): + PyExecutor._executor_loop_pp(executor) + + executor._revert_ctx_alloc.assert_called_once_with([local_only]) + + def test_schedule_prepares_snapshot_points_before_scheduling(): class StopSchedule(RuntimeError): pass