diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 6cf98b759f02..85642d772cb0 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -250,108 +250,6 @@ def _strip_py_multimodal_data_post_prefill(request: LlmRequest) -> None: strip_mm_data_for_generation(mm_data) -@dataclasses.dataclass -class DisaggTransferAdmissionResult: - admitted_requests: List[LlmRequest] - active_transfer_blocks: int = 0 - admitted_transfer_blocks: int = 0 - deferred_request_count: int = 0 - limited_by_budget: bool = False - - def is_blocked_by_active_transfers(self) -> bool: - return (self.limited_by_budget and not self.admitted_requests - and self.active_transfer_blocks > 0) - - -class DisaggTransferAdmissionController: - """FCFS admission gate for disaggregated generation KV transfers.""" - - def __init__(self, max_tokens_in_buffer: Optional[int], - tokens_per_block: Optional[int]) -> None: - self.max_transfer_blocks = self._to_block_budget( - max_tokens_in_buffer, tokens_per_block) - self.tokens_per_block = tokens_per_block or 0 - - def enabled(self) -> bool: - return self.max_transfer_blocks is not None - - @staticmethod - def _to_block_budget(max_tokens_in_buffer: Optional[int], - tokens_per_block: Optional[int]) -> Optional[int]: - if (max_tokens_in_buffer is None or max_tokens_in_buffer == 0 - or tokens_per_block is None or tokens_per_block <= 0): - return None - return (max_tokens_in_buffer + tokens_per_block - 1) // tokens_per_block - - @staticmethod - def _to_nonnegative_int(value) -> Optional[int]: - try: - return max(int(value), 0) - except (TypeError, ValueError): - return None - - def _get_request_transfer_token_count(self, request: LlmRequest) -> int: - for attr_name in ("total_input_len_cp", "py_prompt_len", "prompt_len"): - token_count = self._to_nonnegative_int( - getattr(request, attr_name, None)) - if token_count is not None: - return token_count - return 0 - - def _estimate_request_blocks(self, request: LlmRequest) -> int: - if self.tokens_per_block <= 0: - return 0 - prompt_len = self._get_request_transfer_token_count(request) - return (prompt_len + self.tokens_per_block - 1) // self.tokens_per_block - - def _estimate_requests_blocks(self, requests: Iterable[LlmRequest]) -> int: - return sum( - self._estimate_request_blocks(request) for request in requests) - - def _estimate_active_transfer_blocks( - self, active_requests: Iterable[LlmRequest]) -> int: - return sum( - self._estimate_request_blocks(request) - for request in active_requests - if request.is_disagg_generation_transmission_in_progress) - - def select(self, active_requests: Iterable[LlmRequest], - candidates: List[LlmRequest]) -> DisaggTransferAdmissionResult: - if not self.enabled(): - return DisaggTransferAdmissionResult( - admitted_requests=list(candidates), - active_transfer_blocks=self._estimate_active_transfer_blocks( - active_requests), - admitted_transfer_blocks=self._estimate_requests_blocks( - candidates), - ) - - result = DisaggTransferAdmissionResult(admitted_requests=[]) - result.active_transfer_blocks = self._estimate_active_transfer_blocks( - active_requests) - - used_blocks = result.active_transfer_blocks - max_transfer_blocks = self.max_transfer_blocks - assert max_transfer_blocks is not None - for request in candidates: - request_blocks = self._estimate_request_blocks(request) - fits_budget = used_blocks + request_blocks <= max_transfer_blocks - admit_oversized_head = (not result.admitted_requests - and result.active_transfer_blocks == 0 - and request_blocks > max_transfer_blocks) - if not fits_budget and not admit_oversized_head: - result.limited_by_budget = True - break - - result.admitted_requests.append(request) - used_blocks += request_blocks - result.admitted_transfer_blocks += request_blocks - - result.deferred_request_count = len(candidates) - len( - result.admitted_requests) - return result - - @dataclasses.dataclass class ScheduledBatchStats: # None means the counter was not captured and _update_iter_stats should @@ -935,14 +833,6 @@ def on_detected(): if kv_cache_transceiver is not None: self.hang_detector.register_status_provider( kv_cache_transceiver.get_status_dump) - cache_transceiver_config = getattr(self.llm_args, - "cache_transceiver_config", None) - max_tokens_in_buffer = getattr(cache_transceiver_config, - "max_tokens_in_buffer", None) - tokens_per_block = getattr(self.kv_cache_manager, "tokens_per_block", - None) - self._disagg_transfer_admission_controller = DisaggTransferAdmissionController( - max_tokens_in_buffer, tokens_per_block) 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 @@ -2440,19 +2330,14 @@ def _pp_schedule_and_propagate(self, microbatch_id: int): # For DP cases, the first PP rank schedules the requests. scheduled_batch = None serializable_schedule = None - wait_for_disagg_gen_transfer_progress = False is_dp_broadcast = self.dist.tp_size > 1 and self.enable_attention_dp if self.dist.rank == 0 or (self.dist.is_first_pp_rank and is_dp_broadcast): scheduled_batch, fitting_disagg_gen_init_requests, num_fitting_reqs = self._schedule( ) - if self.kv_cache_transceiver: - fitting_disagg_gen_init_requests, wait_for_disagg_gen_transfer_progress = ( - self._apply_disagg_transfer_admission( - fitting_disagg_gen_init_requests)) serializable_schedule = SerializableSchedulerOutput.from_scheduler_result( scheduled_batch, fitting_disagg_gen_init_requests, - num_fitting_reqs, wait_for_disagg_gen_transfer_progress) + num_fitting_reqs) # Broadcast within first tp+cp group before send/recv chain to other tp+cp groups if self.dist.is_first_pp_rank: @@ -2484,10 +2369,8 @@ def _pp_schedule_and_propagate(self, microbatch_id: int): if scheduled_batch is None: scheduled_batch, fitting_disagg_gen_init_requests, num_fitting_reqs = serializable_schedule.to_scheduler_result( self.active_requests) - wait_for_disagg_gen_transfer_progress = ( - serializable_schedule.wait_for_disagg_gen_transfer_progress) return (scheduled_batch, fitting_disagg_gen_init_requests, - num_fitting_reqs, wait_for_disagg_gen_transfer_progress) + num_fitting_reqs) def _pp_retry_until_can_schedule(self, scheduled_batch): """ @@ -2568,7 +2451,7 @@ def _executor_loop_pp(self): # Stage 0: first PP rank schedules requests and propagates the result to all other PP ranks. (scheduled_batch, fitting_disagg_gen_init_requests, - num_fitting_reqs, wait_for_disagg_gen_transfer_progress + num_fitting_reqs ) = self._pp_schedule_and_propagate(microbatch_id) if self.dist.rank != 0: # Retry until current rank can run first PP's schedule result. @@ -2578,15 +2461,8 @@ def _executor_loop_pp(self): "prepare_expect_snapshot_points"): self.kv_cache_manager.prepare_expect_snapshot_points( self.active_requests) - local_scheduler_output = self.scheduler.schedule_request( - self.active_requests, self.inflight_req_ids) - if self.kv_cache_transceiver: - local_disagg_candidates = getattr( - local_scheduler_output, - "fitting_disagg_gen_init_requests", []) - self._revert_deferred_disagg_gen_init_alloc( - local_disagg_candidates, - fitting_disagg_gen_init_requests) + self.scheduler.schedule_request(self.active_requests, + self.inflight_req_ids) # For requests that are fitting disagg gen init, also prepare resources for KV cache manager if self.kv_cache_transceiver: @@ -2600,7 +2476,7 @@ def _executor_loop_pp(self): for req in self.active_requests) self._check_disagg_transfer_progress_when_idle( num_fitting_reqs, fitting_disagg_gen_init_requests, - wait_for_disagg_gen_transfer_progress, all_gen_first) + all_gen_first) self.num_scheduled_requests = scheduled_batch.batch_size @@ -3370,11 +3246,6 @@ def _finalize_adp_dummy_allocation(self, can_queue: bool) -> None: self.kv_cache_manager.free_resources(dummy_request) self.active_requests.remove(dummy_request) - def _revert_ctx_alloc(self, dropped_context_requests): - """Revert V2 context KV growth for requests deferred after scheduling.""" - for req in dropped_context_requests: - self.kv_cache_manager.revert_allocate_context(req) - @nvtx_range("_prefetch_for_context_requests") def _prefetch_for_context_requests(self) -> None: """Pre-stage disk blocks to host for upcoming context requests with block reuse.""" @@ -3405,21 +3276,6 @@ def _commit_kv_cache_stats(self, self.kv_cache_manager.commit_scheduled_kv_cache_stats( scheduled_batch) - def _get_disagg_transfer_admission_controller( - self) -> DisaggTransferAdmissionController: - controller = getattr(self, "_disagg_transfer_admission_controller", - None) - if controller is not None: - return controller - - cache_transceiver_config = getattr(getattr(self, "llm_args", None), - "cache_transceiver_config", None) - kv_cache_manager = getattr(self, "kv_cache_manager", None) - return DisaggTransferAdmissionController( - getattr(cache_transceiver_config, "max_tokens_in_buffer", None), - getattr(kv_cache_manager, "tokens_per_block", None), - ) - @staticmethod def _is_disagg_gen_only_no_context_benchmark() -> bool: """Return whether ``gen_only_no_context`` skips KV transfer.""" @@ -3430,62 +3286,6 @@ def _uses_async_disagg_gen_transfer(self) -> bool: return (not self._is_disagg_gen_only_no_context_benchmark() and os.getenv("TRTLLM_DISABLE_KV_CACHE_TRANSFER_OVERLAP") != "1") - def _uses_kv_manager_v2(self) -> bool: - explicit_flag = getattr(self, "_is_kv_manager_v2", None) - if explicit_flag is not None: - return bool(explicit_flag) - return isinstance(getattr(self, "kv_cache_manager", None), - KVCacheManagerV2) - - def _apply_disagg_transfer_admission( - self, fitting_disagg_gen_init_requests: List[LlmRequest] - ) -> Tuple[List[LlmRequest], bool]: - # gen_only_no_context has no CTX worker and does not transfer data. - # Real synchronous gen_only transfers still honor the budget to bound - # the number of blocking transfers started in one executor iteration. - if self._is_disagg_gen_only_no_context_benchmark(): - 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): - return fitting_disagg_gen_init_requests, False - - admission_result = controller.select(self.active_requests, - fitting_disagg_gen_init_requests) - if admission_result.deferred_request_count > 0: - logger.debug("Disagg transfer admission deferred " - f"{admission_result.deferred_request_count} requests; " - f"active transfer blocks=" - f"{admission_result.active_transfer_blocks}, " - f"admitted transfer blocks=" - f"{admission_result.admitted_transfer_blocks}, " - f"budget={controller.max_transfer_blocks}") - - self._revert_deferred_disagg_gen_init_alloc( - 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( - self, candidates: List[LlmRequest], - admitted_requests: List[LlmRequest]) -> None: - if not (self._uses_kv_manager_v2() and candidates): - return - - admitted_request_ids = { - request.py_request_id - for request in admitted_requests - } - deferred_requests = [ - request for request in candidates - if request.py_request_id not in admitted_request_ids - ] - if deferred_requests: - self._revert_ctx_alloc(deferred_requests) - @staticmethod def _dist_size(dist, name: str) -> int: try: @@ -3514,11 +3314,6 @@ def _allgather_model_parallel_status( return self.dist.tp_allgather(local_status) return [local_status] - def _sync_disagg_gen_status_entry(self, local_need_check: bool) -> int: - if self._dist_size(self.dist, "world_size") > 1: - return self.dist.allreduce(int(local_need_check), op=ReduceOp.MAX) - return int(local_need_check) - def _sync_disagg_ctx_status_entry(self, local_need_check: bool) -> int: if self._dist_size(self.dist, "cp_size") > 1: return int(any(self.dist.tp_cp_allgather(int(local_need_check)))) @@ -3530,30 +3325,16 @@ def _sync_disagg_ctx_status_entry(self, local_need_check: bool) -> int: def _check_disagg_transfer_progress_when_idle( self, num_fitting_reqs: int, fitting_disagg_gen_init_requests: List[LlmRequest], - wait_for_disagg_gen_transfer_progress: bool, all_gen_first: bool) -> None: local_need_check = (num_fitting_reqs == 0 and not fitting_disagg_gen_init_requests) # A synchronous GEN receive is rank-local and blocking. One rank can - # still be receiving while another is idle, so entering either the - # generation or context progress collective here is unsafe. + # still be receiving while another is idle, so entering the context + # progress collective here is unsafe. if not self._uses_async_disagg_gen_transfer(): return - local_need_gen_check = (local_need_check - and wait_for_disagg_gen_transfer_progress) - - any_need_gen_check = self._sync_disagg_gen_status_entry( - local_need_gen_check) - if any_need_gen_check > 0: - if local_need_gen_check: - logger.debug( - "Waiting for generation KV cache transfer progress to " - "free disagg admission budget") - self._check_disagg_gen_cache_transfer_status(1) - return - any_need_check = self._sync_disagg_ctx_status_entry(local_need_check) if any_need_check > 0: if local_need_check and not all_gen_first: @@ -3573,8 +3354,7 @@ def _check_disagg_transfer_progress_when_idle( self._check_disagg_ctx_cache_transfer_status(0) def _sync_gen_only_benchmark_has_insufficient_kv( - self, scheduler_fitting_disagg_gen_init_requests: List[LlmRequest], - wait_for_disagg_gen_transfer_progress: bool) -> bool: + self, fitting_disagg_gen_init_requests: List[LlmRequest]) -> bool: """Return whether benchmark fill has terminal KV exhaustion. Model-parallel ranks can make different local scheduling decisions. @@ -3584,18 +3364,13 @@ def _sync_gen_only_benchmark_has_insufficient_kv( to every decode iteration after the gate opens. Args: - scheduler_fitting_disagg_gen_init_requests: Generation INIT - requests that fit KV capacity before transfer admission. A - nonempty list means KV capacity exists even if transfer - admission temporarily defers every request. - wait_for_disagg_gen_transfer_progress: Whether active generation - transfers are consuming the admission budget and transfer - progress can unblock a deferred request. + fitting_disagg_gen_init_requests: Generation INIT requests that + fit KV capacity. A nonempty list means KV capacity exists. Returns: True when every TP+CP rank has fetched its full benchmark queue and - at least one rank has an INIT request that cannot fit KV capacity - and has no transfer progress that can unblock it; otherwise False. + at least one rank has an INIT request that cannot fit KV capacity; + otherwise False. """ if (self.benchmark_req_queues_size <= 0 or self.is_warmup or not self._benchmark_fill_phase_active): @@ -3605,9 +3380,8 @@ def _sync_gen_only_benchmark_has_insufficient_kv( for req in self.active_requests) local_all_fetched = (self.num_fetch_requests >= self.benchmark_req_queues_size) - local_terminal_no_fit = (local_has_stuck and - not scheduler_fitting_disagg_gen_init_requests - and not wait_for_disagg_gen_transfer_progress) + local_terminal_no_fit = (local_has_stuck + and not fitting_disagg_gen_init_requests) local_status = (local_all_fetched, local_terminal_no_fit) all_rank_status = self._allgather_model_parallel_status(local_status) @@ -3699,7 +3473,7 @@ def _prepare_and_schedule_batch(self): continue request.draft_tokens = [0] * self.max_total_draft_tokens - scheduled_batch, scheduler_fitting_disagg_gen_init_requests, num_fitting_reqs = self._schedule( + scheduled_batch, fitting_disagg_gen_init_requests, num_fitting_reqs = self._schedule( ) if self.drafter is not None and not self.use_spec_decode: @@ -3707,33 +3481,24 @@ def _prepare_and_schedule_batch(self): request.py_disable_speculative_decoding = True if self.kv_cache_transceiver: - wait_for_disagg_gen_transfer_progress = False - admitted_disagg_gen_init_requests, wait_for_disagg_gen_transfer_progress = ( - self._apply_disagg_transfer_admission( - scheduler_fitting_disagg_gen_init_requests)) - # Prepare KV cache manager resources only for requests admitted - # into the transfer window this iteration. - self._prepare_disagg_gen_init(admitted_disagg_gen_init_requests) + # For requests that are fitting disagg gen init, also prepare resources for KV cache manager + self._prepare_disagg_gen_init(fitting_disagg_gen_init_requests) all_gen_first = self.active_requests and all( req.py_disaggregated_params and req.py_disaggregated_params. schedule_style == DisaggScheduleStyle.GENERATION_FIRST for req in self.active_requests) self._check_disagg_transfer_progress_when_idle( - num_fitting_reqs, admitted_disagg_gen_init_requests, - wait_for_disagg_gen_transfer_progress, all_gen_first) + num_fitting_reqs, fitting_disagg_gen_init_requests, + all_gen_first) # In gen-only benchmark mode, all requests must fit in KV cache # simultaneously. If some requests are stuck in INIT state and the # scheduler could not allocate KV for any of them, the benchmark # will hang forever because in-progress generation requests won't # release their KV cache. - # Check the scheduler result from before transfer admission. An - # empty admitted list can mean that active transfers are - # temporarily consuming the transfer budget. has_insufficient_kv = self._sync_gen_only_benchmark_has_insufficient_kv( - scheduler_fitting_disagg_gen_init_requests, - wait_for_disagg_gen_transfer_progress) + fitting_disagg_gen_init_requests) if has_insufficient_kv: error_msg = ( f"Insufficient KV cache for gen-only benchmark mode: " diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py index caa8e3cb3de1..deb5ae4baca2 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py @@ -256,7 +256,6 @@ class SerializableSchedulerOutput: int ] # request ids of fitting disaggregated generation initialization requests num_fitting_requests: int # number of fitting requests - wait_for_disagg_gen_transfer_progress: bool = False @classmethod def from_scheduler_result( @@ -264,7 +263,6 @@ def from_scheduler_result( scheduled_requests: ScheduledRequests, fitting_disagg_gen_init_requests: RequestList, num_fitting_requests: int, - wait_for_disagg_gen_transfer_progress: bool = False, ) -> "SerializableSchedulerOutput": return cls( encoder_requests=[req.request_id for req in scheduled_requests.encoder_requests], @@ -280,7 +278,6 @@ def from_scheduler_result( req.request_id for req in fitting_disagg_gen_init_requests ], num_fitting_requests=num_fitting_requests, - wait_for_disagg_gen_transfer_progress=wait_for_disagg_gen_transfer_progress, ) def to_scheduler_result( diff --git a/tests/unittest/_torch/executor/test_benchmark_disagg.py b/tests/unittest/_torch/executor/test_benchmark_disagg.py index b8cb399b1203..d9dc3d6c6249 100644 --- a/tests/unittest/_torch/executor/test_benchmark_disagg.py +++ b/tests/unittest/_torch/executor/test_benchmark_disagg.py @@ -1154,25 +1154,6 @@ def test_healthy_fill_phase_does_not_kill(self): ) ex._handle_errors.assert_not_called() - def test_partial_transfer_admission_uses_only_admitted_requests(self): - """The admitted subset is prepared and passed to the idle check.""" - admitted_req = _make_active_request(in_init=True) - deferred_req = _make_active_request(in_init=True) - candidates = [admitted_req, deferred_req] - ex = self._make_executor(fill_phase_active=True, fitting_init_requests=candidates) - ex._apply_disagg_transfer_admission = Mock(return_value=([admitted_req], False)) - ex._check_disagg_transfer_progress_when_idle = Mock() - - result, _ = ex._prepare_and_schedule_batch() - - assert result is not None - ex._apply_disagg_transfer_admission.assert_called_once_with(candidates) - ex._prepare_disagg_gen_init.assert_called_once_with([admitted_req]) - ex._check_disagg_transfer_progress_when_idle.assert_called_once_with( - 0, [admitted_req], False, False - ) - ex._handle_errors.assert_not_called() - def test_fill_with_no_init_requests_does_not_kill(self): """The final fill iteration is ready for the gate, not terminal.""" ex = self._make_executor(fill_phase_active=True, num_init_requests=0) @@ -1182,31 +1163,6 @@ def test_fill_with_no_init_requests_does_not_kill(self): assert result is not None ex._handle_errors.assert_not_called() - def test_transfer_admission_backpressure_does_not_kill(self, monkeypatch): - """NVBug 6438658: admission backpressure is not KV exhaustion. - - Args: - monkeypatch: Pytest fixture used to select asynchronous transfer - behavior. - """ - monkeypatch.delenv("TRTLLM_DISAGG_BENCHMARK_GEN_ONLY", raising=False) - monkeypatch.delenv("TRTLLM_DISABLE_KV_CACHE_TRANSFER_OVERLAP", raising=False) - fitting_req = _make_active_request(in_init=True) - ex = self._make_executor(fill_phase_active=True, fitting_init_requests=[fitting_req]) - ex._apply_disagg_transfer_admission = Mock(return_value=([], True)) - - result, _ = ex._prepare_and_schedule_batch() - - assert result is not None, ( - "Fail-fast should NOT fire when the scheduler fit an INIT request " - "that transfer admission temporarily deferred" - ) - ex._apply_disagg_transfer_admission.assert_called_once_with([fitting_req]) - ex._prepare_disagg_gen_init.assert_called_once_with([]) - ex._check_disagg_gen_cache_transfer_status.assert_called_once_with(1) - ex._check_disagg_ctx_cache_transfer_status.assert_not_called() - ex._handle_errors.assert_not_called() - @pytest.mark.parametrize( "enable_attention_dp, tp_size, cp_size, gather_name", [ @@ -1236,7 +1192,6 @@ def test_model_parallel_peer_terminal_no_fit_kills_all_ranks( all_rank_status[-1] = (True, True) gather = getattr(ex.dist, gather_name) gather.return_value = all_rank_status - ex._apply_disagg_transfer_admission = Mock(return_value=([], True)) ex._check_disagg_transfer_progress_when_idle = Mock() result, _ = ex._prepare_and_schedule_batch() @@ -1248,8 +1203,8 @@ def test_model_parallel_peer_terminal_no_fit_kills_all_ranks( ex._handle_errors.assert_called_once() assert "one or more requests" in ex._handle_errors.call_args.args[0] - def test_attention_dp_backpressure_without_terminal_peer_does_not_kill(self): - """Admission backpressure stays non-terminal on every rank.""" + def test_attention_dp_without_terminal_peer_does_not_kill(self): + """A fitting INIT request stays non-terminal on every rank.""" fitting_req = _make_active_request(in_init=True) ex = self._make_executor(fill_phase_active=True, fitting_init_requests=[fitting_req]) ex.enable_attention_dp = True @@ -1259,7 +1214,6 @@ def test_attention_dp_backpressure_without_terminal_peer_does_not_kill(self): (True, False), (True, False), ] - ex._apply_disagg_transfer_admission = Mock(return_value=([], True)) ex._check_disagg_transfer_progress_when_idle = Mock() result, _ = ex._prepare_and_schedule_batch() diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index bd9f4ea2f634..57972b17b534 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -20,19 +20,14 @@ import pytest -from tensorrt_llm._torch.distributed.communicator import ReduceOp from tensorrt_llm._torch.pyexecutor.executor_request_queue import ( SHUTDOWN_REQUEST_ID, RequestQueueItem, ) from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestState, SamplingConfig -from tensorrt_llm._torch.pyexecutor.py_executor import DisaggTransferAdmissionController, PyExecutor +from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor from tensorrt_llm._torch.pyexecutor.resource_manager import NoFreeSlotsError, ResourceManagerType -from tensorrt_llm._torch.pyexecutor.scheduler import ( - FCFSWaitingQueue, - ScheduledRequests, - SerializableSchedulerOutput, -) +from tensorrt_llm._torch.pyexecutor.scheduler import FCFSWaitingQueue, ScheduledRequests pytestmark = pytest.mark.cpu_only @@ -445,168 +440,6 @@ def _clear_disagg_transfer_mode_env(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("TRTLLM_DISABLE_KV_CACHE_TRANSFER_OVERLAP", raising=False) -@pytest.mark.usefixtures("_clear_disagg_transfer_mode_env") -class TestDisaggTransferAdmissionController: - def test_disabled_preserves_candidates(self): - controller = DisaggTransferAdmissionController( - max_tokens_in_buffer=None, tokens_per_block=32 - ) - candidate = _make_disagg_transfer_request(1, 64) - - result = controller.select(active_requests=[], candidates=[candidate]) - - assert result.admitted_requests == [candidate] - assert result.deferred_request_count == 0 - assert not result.is_blocked_by_active_transfers() - - def test_fcfs_budget_counts_active_transfers(self): - controller = DisaggTransferAdmissionController(max_tokens_in_buffer=64, tokens_per_block=32) - active = _make_disagg_transfer_request(1, 32, in_progress=True) - admitted = _make_disagg_transfer_request(2, 32) - deferred = _make_disagg_transfer_request(3, 32) - - result = controller.select(active_requests=[active], candidates=[admitted, deferred]) - - assert result.admitted_requests == [admitted] - assert result.active_transfer_blocks == 1 - assert result.admitted_transfer_blocks == 1 - assert result.deferred_request_count == 1 - assert result.limited_by_budget - assert not result.is_blocked_by_active_transfers() - - def test_reports_active_transfer_budget_block(self): - controller = DisaggTransferAdmissionController(max_tokens_in_buffer=32, tokens_per_block=32) - active = _make_disagg_transfer_request(1, 32, in_progress=True) - candidate = _make_disagg_transfer_request(2, 32) - - result = controller.select(active_requests=[active], candidates=[candidate]) - - assert result.admitted_requests == [] - assert result.active_transfer_blocks == 1 - assert result.deferred_request_count == 1 - assert result.is_blocked_by_active_transfers() - - def test_admits_oversized_head_when_idle(self): - controller = DisaggTransferAdmissionController(max_tokens_in_buffer=32, tokens_per_block=32) - oversized = _make_disagg_transfer_request(1, 96) - deferred = _make_disagg_transfer_request(2, 32) - - result = controller.select(active_requests=[], candidates=[oversized, deferred]) - - assert result.admitted_requests == [oversized] - assert result.admitted_transfer_blocks == 3 - assert result.deferred_request_count == 1 - assert result.limited_by_budget - assert not result.is_blocked_by_active_transfers() - - def test_uses_global_cp_prompt_length_for_transfer_cost(self): - controller = DisaggTransferAdmissionController( - max_tokens_in_buffer=128, tokens_per_block=32 - ) - request = _make_disagg_transfer_request(1, 32, total_input_len_cp=96) - - result = controller.select(active_requests=[], candidates=[request]) - - assert result.admitted_requests == [request] - assert result.admitted_transfer_blocks == 3 - - def test_apply_reverts_deferred_v2_allocations(self): - executor = object.__new__(PyExecutor) - executor.kv_cache_transceiver = Mock() - 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 - ) - 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_called_once_with([candidate]) - - def test_apply_missing_controller_preserves_candidates(self): - executor = object.__new__(PyExecutor) - executor.kv_cache_transceiver = Mock() - executor.active_requests = [] - candidate = _make_disagg_transfer_request(1, 32) - - admitted, wait_for_progress = PyExecutor._apply_disagg_transfer_admission( - executor, [candidate] - ) - - assert admitted == [candidate] - assert not wait_for_progress - - def test_apply_missing_v2_flag_defaults_to_non_v2(self): - executor = object.__new__(PyExecutor) - executor.kv_cache_transceiver = Mock() - 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_sync_mode_retains_transfer_budget(self, monkeypatch): - monkeypatch.setenv("TRTLLM_DISABLE_KV_CACHE_TRANSFER_OVERLAP", "1") - executor = object.__new__(PyExecutor) - executor.kv_cache_transceiver = Mock() - executor._is_kv_manager_v2 = True - executor._revert_ctx_alloc = Mock() - executor.active_requests = [] - 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[0]] - assert not wait_for_progress - executor._revert_ctx_alloc.assert_called_once_with([candidates[1]]) - - def test_gen_only_no_context_bypasses_transfer_budget(self, monkeypatch): - monkeypatch.setenv("TRTLLM_DISAGG_BENCHMARK_GEN_ONLY", "1") - executor = object.__new__(PyExecutor) - executor.kv_cache_transceiver = Mock() - 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() - - @pytest.mark.usefixtures("_clear_disagg_transfer_mode_env") class TestDisaggTransferIdleProgress: def test_gen_transfer_status_polls_active_transfers(self): @@ -636,43 +469,7 @@ def test_gen_transfer_status_skips_sync_mode(self, monkeypatch): executor._check_disagg_gen_cache_transfer_status.assert_not_called() - def test_polls_generation_transfer_when_admission_blocked(self): - executor = object.__new__(PyExecutor) - executor.dist = Mock(tp_size=1) - executor._check_disagg_gen_cache_transfer_status = Mock() - executor._check_disagg_ctx_cache_transfer_status = Mock() - - PyExecutor._check_disagg_transfer_progress_when_idle( - executor, - num_fitting_reqs=0, - fitting_disagg_gen_init_requests=[], - wait_for_disagg_gen_transfer_progress=True, - all_gen_first=False, - ) - - executor._check_disagg_gen_cache_transfer_status.assert_called_once_with(1) - executor._check_disagg_ctx_cache_transfer_status.assert_not_called() - - def test_peer_rank_enters_bounded_progress_poll(self): - executor = object.__new__(PyExecutor) - executor.dist = Mock(tp_size=1, cp_size=4, world_size=4) - executor.dist.allreduce.return_value = 1 - executor._check_disagg_gen_cache_transfer_status = Mock() - executor._check_disagg_ctx_cache_transfer_status = Mock() - - PyExecutor._check_disagg_transfer_progress_when_idle( - executor, - num_fitting_reqs=1, - fitting_disagg_gen_init_requests=[], - wait_for_disagg_gen_transfer_progress=True, - all_gen_first=False, - ) - - executor._check_disagg_gen_cache_transfer_status.assert_called_once_with(1) - executor._check_disagg_ctx_cache_transfer_status.assert_not_called() - executor.dist.allreduce.assert_called_once_with(0, op=ReduceOp.MAX) - - def test_falls_back_to_context_transfer_when_not_generation_blocked(self): + def test_falls_back_to_context_transfer_when_idle(self): executor = object.__new__(PyExecutor) executor.dist = Mock(tp_size=1) executor._check_disagg_gen_cache_transfer_status = Mock() @@ -682,7 +479,6 @@ def test_falls_back_to_context_transfer_when_not_generation_blocked(self): executor, num_fitting_reqs=0, fitting_disagg_gen_init_requests=[], - wait_for_disagg_gen_transfer_progress=False, all_gen_first=False, ) @@ -701,7 +497,6 @@ def test_sync_benchmark_skips_idle_transfer_collectives(self, monkeypatch): executor, num_fitting_reqs=0, fitting_disagg_gen_init_requests=[], - wait_for_disagg_gen_transfer_progress=True, all_gen_first=False, ) @@ -722,7 +517,6 @@ def test_sync_non_benchmark_skips_idle_transfer_collectives(self, monkeypatch): executor, num_fitting_reqs=0, fitting_disagg_gen_init_requests=[], - wait_for_disagg_gen_transfer_progress=True, all_gen_first=False, ) @@ -816,7 +610,6 @@ def complete_or_error(req): def test_peer_cp_rank_enters_context_progress_poll(self): executor = object.__new__(PyExecutor) executor.dist = Mock(tp_size=1, cp_size=4, world_size=4) - executor.dist.allreduce.return_value = 0 executor.dist.tp_cp_allgather.return_value = [0, 1, 0, 0] executor._check_disagg_gen_cache_transfer_status = Mock() executor._check_disagg_ctx_cache_transfer_status = Mock() @@ -825,7 +618,6 @@ def test_peer_cp_rank_enters_context_progress_poll(self): executor, num_fitting_reqs=1, fitting_disagg_gen_init_requests=[], - wait_for_disagg_gen_transfer_progress=False, all_gen_first=False, ) @@ -834,67 +626,6 @@ def test_peer_cp_rank_enters_context_progress_poll(self): executor.dist.tp_cp_allgather.assert_called_once_with(0) -@pytest.mark.usefixtures("_clear_disagg_transfer_mode_env") -class TestDisaggTransferAdmissionPP: - def test_pp_schedule_applies_gate_before_serializing(self): - executor = object.__new__(PyExecutor) - 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 - ) - 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 - - def test_pp_schedule_restores_propagated_gate_decision(self): - executor = object.__new__(PyExecutor) - executor.dist = Mock( - rank=1, - is_first_pp_rank=False, - is_last_pp_rank=True, - prev_pp_rank=0, - tp_size=1, - cp_size=1, - ) - executor.enable_attention_dp = False - executor.active_requests = [ - _make_disagg_transfer_request(1, 32, in_progress=True), - _make_disagg_transfer_request(2, 32), - ] - serializable_schedule = SerializableSchedulerOutput( - encoder_requests=[], - context_requests_chunking=[], - context_requests_last_chunk=[], - generation_requests=[], - paused_requests=[], - fitting_disagg_gen_init_requests=[2], - num_fitting_requests=0, - wait_for_disagg_gen_transfer_progress=True, - ) - executor.dist.recv_object = Mock(return_value=serializable_schedule) - - _, fitting, _, wait_for_progress = PyExecutor._pp_schedule_and_propagate( - executor, microbatch_id=0 - ) - - assert [req.py_request_id for req in fitting] == [2] - assert wait_for_progress - - def test_nonzero_pp_rank_prepares_snapshot_points_before_local_schedule( monkeypatch, ): @@ -916,7 +647,7 @@ class StopLocalSchedule(RuntimeError): executor.kv_cache_transceiver = None executor._pad_attention_dp_dummy_request = Mock() scheduled_batch = Mock() - executor._pp_schedule_and_propagate = Mock(return_value=(scheduled_batch, [], 0, False)) + executor._pp_schedule_and_propagate = Mock(return_value=(scheduled_batch, [], 0)) executor._pp_retry_until_can_schedule = Mock() request = Mock() executor.active_requests = [request] diff --git a/tests/unittest/_torch/executor/test_scheduler_serializable_output.py b/tests/unittest/_torch/executor/test_scheduler_serializable_output.py index 904123518dce..550fce0e626f 100644 --- a/tests/unittest/_torch/executor/test_scheduler_serializable_output.py +++ b/tests/unittest/_torch/executor/test_scheduler_serializable_output.py @@ -40,7 +40,6 @@ def test_serializable_scheduler_output_round_trip(): scheduled_requests, fitting_disagg_gen_init_requests, num_fitting_requests, - wait_for_disagg_gen_transfer_progress=True, ) # Serialize and deserialize the serializable scheduler output @@ -55,7 +54,6 @@ def test_serializable_scheduler_output_round_trip(): # Verify the restored scheduler result is correct assert restored_num_fitting == num_fitting_requests - assert restored_output.wait_for_disagg_gen_transfer_progress assert _request_ids(restored_schedule.encoder_requests) == _request_ids( scheduled_requests.encoder_requests )