diff --git a/tensorrt_llm/_torch/disaggregation/orchestration/admission.py b/tensorrt_llm/_torch/disaggregation/orchestration/admission.py index f940a7e09378..07182aa707d2 100644 --- a/tensorrt_llm/_torch/disaggregation/orchestration/admission.py +++ b/tensorrt_llm/_torch/disaggregation/orchestration/admission.py @@ -23,6 +23,30 @@ def is_blocked_by_active_transfers(self) -> bool: ) +@dataclasses.dataclass +class DisaggTransferBudget: + """One scheduling pass's transfer budget, charged only for KV-fitting requests.""" + + max_transfer_blocks: Optional[int] + active_transfer_blocks: int = 0 + admitted_transfer_blocks: int = 0 + admitted_request_count: int = 0 + + def allows(self, request_blocks: int) -> bool: + if self.max_transfer_blocks is None: + return True + used_blocks = self.active_transfer_blocks + self.admitted_transfer_blocks + return used_blocks + request_blocks <= self.max_transfer_blocks or ( + self.admitted_request_count == 0 + and self.active_transfer_blocks == 0 + and request_blocks > self.max_transfer_blocks + ) + + def commit(self, request_blocks: int) -> None: + self.admitted_transfer_blocks += request_blocks + self.admitted_request_count += 1 + + class DisaggTransferAdmissionController: """FCFS admission gate for disaggregated generation KV transfers.""" @@ -62,53 +86,40 @@ def _get_request_transfer_token_count(self, request: LlmRequest) -> int: return token_count return 0 - def _estimate_request_blocks(self, request: LlmRequest) -> int: + 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) + self.estimate_request_blocks(request) for request in active_requests if request.is_disagg_generation_transmission_in_progress ) + def create_budget(self, active_requests: Iterable[LlmRequest]) -> DisaggTransferBudget: + """Snapshot active transfers once for an entire scheduling pass.""" + return DisaggTransferBudget( + self.max_transfer_blocks, self._estimate_active_transfer_blocks(active_requests) + ) + 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 + budget = self.create_budget(active_requests) + result = DisaggTransferAdmissionResult( + admitted_requests=[], active_transfer_blocks=budget.active_transfer_blocks + ) 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: + request_blocks = self.estimate_request_blocks(request) + if not budget.allows(request_blocks): result.limited_by_budget = True break - result.admitted_requests.append(request) - used_blocks += request_blocks - result.admitted_transfer_blocks += request_blocks + budget.commit(request_blocks) + result.admitted_transfer_blocks = budget.admitted_transfer_blocks result.deferred_request_count = len(candidates) - len(result.admitted_requests) return result diff --git a/tensorrt_llm/_torch/disaggregation/orchestration/coordinator.py b/tensorrt_llm/_torch/disaggregation/orchestration/coordinator.py index acd4d55a5930..4bbfff477cca 100644 --- a/tensorrt_llm/_torch/disaggregation/orchestration/coordinator.py +++ b/tensorrt_llm/_torch/disaggregation/orchestration/coordinator.py @@ -42,14 +42,11 @@ def uses_async_gen_transfer() -> bool: ) -def transfer_window_bypass_eligible(transceiver, dist, is_kv_manager_v2: bool) -> bool: - """Whether this runtime may skip the executor-level transfer window. - - ``max_tokens_in_buffer`` describes the C++ transceiver's physical buffer. - The asynchronous Python transceiver does not consume it, and with KV cache - manager V2 and PP1 its generation requests stay bounded by the scheduler's - inline KV admission, so no second budget is applied. Other configurations - keep the window. +def early_transfer_window_eligible(transceiver, dist, is_kv_manager_v2: bool) -> bool: + """Whether the V2 scheduler can apply the transfer window before KV allocation. + + PP followers reconcile their local allocations with the canonical schedule, + so PP and synchronous transfers retain post-scheduling admission. """ return ( transceiver is not None @@ -312,13 +309,7 @@ def revert_deferred_gen_init( def _transfer_window_is_active(self) -> bool: """Whether the executor-level transfer window bounds admission.""" - return ( - self._admission_controller is not None - and self._admission_controller.enabled() - and not transfer_window_bypass_eligible( - self._transceiver, self._dist, self._is_kv_manager_v2 - ) - ) + return self._admission_controller is not None and self._admission_controller.enabled() @nvtx_range("receive_gen_init") def receive_gen_init(self, admitted: List[LlmRequest]) -> None: diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 2f210c81df05..f0ec43b8f67d 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -60,7 +60,7 @@ DisaggTransferAdmissionController from ..disaggregation.orchestration.coordinator import ( DisaggTransferCoordinator, NoopDisaggCoordinator, attach_ctx_usage, - transfer_window_bypass_eligible) + early_transfer_window_eligible) from ..disaggregation.orchestration.pp_termination import \ DisaggPPTerminationHandler from ..disaggregation.orchestration.transfer_manager import AsyncTransferManager @@ -1021,18 +1021,12 @@ 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 transfer_window_bypass_eligible(self.kv_cache_transceiver, - self.dist, - self._is_kv_manager_v2)): - 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") + if (self._disagg_transfer_admission_controller.enabled() + and early_transfer_window_eligible(self.kv_cache_transceiver, + self.dist, + self._is_kv_manager_v2)): + self.scheduler.set_disagg_transfer_admission_controller( + self._disagg_transfer_admission_controller) 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 @@ -4025,6 +4019,8 @@ def _prepare_and_schedule_batch(self): wait_for_disagg_gen_transfer_progress = False admitted_disagg_gen_init_requests, wait_for_disagg_gen_transfer_progress = ( self.disagg.admit(scheduler_fitting_disagg_gen_init_requests)) + wait_for_disagg_gen_transfer_progress |= ( + scheduled_batch.disagg_transfer_budget_blocked) # Prepare KV cache manager resources only for requests admitted # into the transfer window this iteration. self.disagg.receive_gen_init(admitted_disagg_gen_init_requests) @@ -6678,6 +6674,8 @@ def _schedule(self): scheduled_requests.scheduled_mm_encoder_items = ( scheduler_output.scheduled_mm_encoder_items) scheduled_requests.recompute_paused_requests = scheduler_output.recompute_paused_requests + scheduled_requests.disagg_transfer_budget_blocked = ( + scheduler_output.disagg_transfer_budget_blocked) self._maybe_record_hang_diagnostic_phase("scheduled", scheduled_requests) diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py index ff180cca75fb..abe0be4dd55f 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py @@ -9,7 +9,7 @@ from collections import namedtuple from dataclasses import dataclass from enum import Enum -from typing import Any, Callable, Optional, TypeAlias, TypeVar +from typing import TYPE_CHECKING, Any, Callable, Optional, TypeAlias, TypeVar from strenum import StrEnum @@ -25,6 +25,11 @@ is_multimodal_encoder_ready, ) +if TYPE_CHECKING: + from tensorrt_llm._torch.disaggregation.orchestration.admission import ( + DisaggTransferAdmissionController, + ) + RequestList = list[LlmRequest] PrefixReuseSummary: TypeAlias = tb_internal.batch_manager.PrefixReuseSummary PrefixSummaryCache: TypeAlias = dict[int, PrefixReuseSummary] @@ -79,6 +84,7 @@ class SchedulerOutput( "num_fitting_requests", "scheduled_mm_encoder_items", "recompute_paused_requests", + "disagg_transfer_budget_blocked", ], ) ): @@ -87,6 +93,8 @@ class SchedulerOutput( ``scheduled_mm_encoder_items`` defaults to ``None``. The V2-only ``recompute_paused_requests`` defaults to a fresh empty list so existing V1 schedulers can keep constructing the original six-field output. + ``disagg_transfer_budget_blocked`` distinguishes early transfer deferral + from KV exhaustion when the V2 scheduler returns no INIT candidates. """ __slots__ = () @@ -101,6 +109,7 @@ def __new__( num_fitting_requests: int, scheduled_mm_encoder_items: dict[int, list[int]] | None = None, recompute_paused_requests: RequestList | None = None, + disagg_transfer_budget_blocked: bool = False, ): return super(SchedulerOutput, cls).__new__( cls, @@ -112,6 +121,7 @@ def __new__( num_fitting_requests, scheduled_mm_encoder_items, [] if recompute_paused_requests is None else recompute_paused_requests, + disagg_transfer_budget_blocked, ) @@ -227,6 +237,8 @@ def __init__(self): self.recompute_paused_requests: RequestList = [] self.added_inflight_req_ids: list[int] = [] self.scheduled_mm_encoder_items: dict[int, list[int]] | None = None + # An empty INIT list can reflect transfer backpressure rather than KV exhaustion. + self.disagg_transfer_budget_blocked: bool = False @property def is_generation_only(self) -> bool: @@ -620,6 +632,12 @@ def __init__( scheduler, "micro_batch_scheduler" ) + def set_disagg_transfer_admission_controller( + self, controller: "DisaggTransferAdmissionController" + ) -> None: + """Forward early admission configuration to the underlying V2 scheduler.""" + self.scheduler.set_disagg_transfer_admission_controller(controller) + @property def scheduling_state_range(self) -> tuple[LlmRequestState, LlmRequestState]: return self.scheduler.scheduling_state_range diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py index c11210f6b5ed..2844cc0437c1 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py @@ -18,6 +18,9 @@ from collections import Counter from typing import Callable, Optional +from tensorrt_llm._torch.disaggregation.orchestration.admission import ( + DisaggTransferAdmissionController, +) from tensorrt_llm.llmapi.llm_args import CapacitySchedulerPolicy, ContextChunkingPolicy from tensorrt_llm.logger import logger @@ -259,6 +262,15 @@ def __init__( # Registered by PyExecutor; see set_async_transfer_manager. self._async_transfer_manager = None + self._disagg_transfer_admission_controller: Optional[DisaggTransferAdmissionController] = ( + None + ) + + def set_disagg_transfer_admission_controller( + self, controller: DisaggTransferAdmissionController + ) -> None: + """Enable early transfer budgeting for an executor-selected runtime.""" + self._disagg_transfer_admission_controller = controller def set_async_transfer_manager(self, mgr) -> None: """Register the AsyncTransferManager the deadlock detector consults. @@ -290,6 +302,7 @@ def schedule_request( recompute_paused, disagg_candidates, has_chunking, + transfer_budget_blocked, ) = self._schedule_loop(active_requests, inflight_request_ids) # Sort by LoRA task ID @@ -303,6 +316,7 @@ def schedule_request( paused_requests=evicted, recompute_paused_requests=recompute_paused, fitting_disagg_gen_init_requests=disagg_candidates, + disagg_transfer_budget_blocked=transfer_budget_blocked, num_fitting_requests=(len(scheduled_encoder) + len(scheduled_ctx) + len(scheduled_gen)), ) @@ -339,6 +353,31 @@ def _schedule_loop(self, active_requests, inflight_request_ids): ) ) + transfer_controller = self._disagg_transfer_admission_controller + transfer_budget = None + transfer_costs: list[int] = [] + suffix_min_costs: list[Optional[int]] = [] + transfer_gate_closed = False + transfer_budget_blocked = False + if transfer_controller is not None and transfer_controller.enabled(): + transfer_budget = transfer_controller.create_budget(requests_list) + transfer_costs = [0] * len(requests_list) + suffix_min_costs = [None] * len(requests_list) + minimum = None + # A budget-rejected head can fail KV admission, allowing a smaller + # tail request through. Skip allocation only if no remaining INIT + # could use the budget. One suffix pass keeps this check O(1). + for index in range(len(requests_list) - 1, -1, -1): + candidate = requests_list[index] + if ( + candidate.state_value == self._disagg_gen_init_state_value + and candidate.request_id not in inflight_request_ids + ): + cost = transfer_controller.estimate_request_blocks(candidate) + transfer_costs[index] = cost + minimum = cost if minimum is None else min(minimum, cost) + suffix_min_costs[index] = minimum + req_it_end = len(requests_list) recompute_pause_state = _RecomputePauseState(req_it_end) req_it = 0 @@ -390,6 +429,23 @@ def _schedule_loop(self, active_requests, inflight_request_ids): # no free slots remain, so the request is skipped and retried next # iteration. PEFT budget is still checked and committed. if req_state_value == self._disagg_gen_init_state_value: + if transfer_gate_closed: + req_it += 1 + continue + fits_transfer_budget = True + if transfer_budget is not None: + fits_transfer_budget = transfer_budget.allows(transfer_costs[req_it]) + minimum = suffix_min_costs[req_it] + assert minimum is not None + if not fits_transfer_budget and not transfer_budget.allows(minimum): + transfer_gate_closed = True + transfer_budget_blocked = ( + transfer_budget.active_transfer_blocks > 0 + and transfer_budget.admitted_request_count == 0 + ) + req_it += 1 + continue + peft_pages = budget.peft_pages_needed(req) if peft_pages is None: break @@ -401,6 +457,19 @@ def _schedule_loop(self, active_requests, inflight_request_ids): req_it += 1 continue disagg_candidates.append(req) + if transfer_budget is not None: + if not fits_transfer_budget: + # This KV-fitting head is the FCFS barrier. Keep just + # this allocation for coordinator.admit() to revert; + # later INITs cannot pass it, but decode must continue. + transfer_gate_closed = True + transfer_budget_blocked = ( + transfer_budget.active_transfer_blocks > 0 + and transfer_budget.admitted_request_count == 0 + ) + req_it += 1 + continue + transfer_budget.commit(transfer_costs[req_it]) # Disagg requests only commit PEFT (not num_requests/num_tokens) # because they don't participate in the forward pass. Counting # them toward num_requests would steal batch slots from gen/ctx @@ -592,6 +661,7 @@ def preempt_for_pages(req: LlmRequest) -> bool: recompute_paused, disagg_candidates, has_chunking, + transfer_budget_blocked, ) # ---- Prefix-aware skip ---- diff --git a/tests/unittest/_torch/disaggregation/test_benchmark_disagg.py b/tests/unittest/_torch/disaggregation/test_benchmark_disagg.py index c77b12cacf08..6e45a71d1ea1 100644 --- a/tests/unittest/_torch/disaggregation/test_benchmark_disagg.py +++ b/tests/unittest/_torch/disaggregation/test_benchmark_disagg.py @@ -1370,6 +1370,21 @@ def test_fill_with_no_init_requests_does_not_kill(self): assert result is not None ex._handle_errors.assert_not_called() + def test_early_transfer_budget_block_with_empty_fitting_list_does_not_kill(self) -> None: + """Skipping KV allocation under an active window is not terminal KV exhaustion.""" + ex = self._make_executor(fill_phase_active=True, fitting_init_requests=[]) + ex.active_requests.append(_make_active_request(in_transfer=True)) + scheduled = ex._schedule.return_value[0] + scheduled.disagg_transfer_budget_blocked = True + ex.disagg.admit = Mock(return_value=([], False)) + + result, _ = ex._prepare_and_schedule_batch() + + assert result is scheduled + ex.disagg.receive_gen_init.assert_called_once_with([]) + ex.disagg.reap_context_sends.assert_called_once_with(0) + ex._handle_errors.assert_not_called() + def test_transfer_admission_backpressure_does_not_kill(self, monkeypatch): """NVBug 6438658: admission backpressure is not KV exhaustion. diff --git a/tests/unittest/_torch/disaggregation/test_disagg_coordinator_admission.py b/tests/unittest/_torch/disaggregation/test_disagg_coordinator_admission.py index c562a297a2a8..94e7d7d28439 100644 --- a/tests/unittest/_torch/disaggregation/test_disagg_coordinator_admission.py +++ b/tests/unittest/_torch/disaggregation/test_disagg_coordinator_admission.py @@ -17,7 +17,7 @@ DisaggTransferAdmissionController, ) from tensorrt_llm._torch.disaggregation.orchestration.coordinator import ( - transfer_window_bypass_eligible, + early_transfer_window_eligible, ) from tensorrt_llm.bindings import LlmRequestState @@ -105,18 +105,16 @@ def test_the_window_admits_the_head_and_reverts_only_the_deferred_tail() -> None assert h.effects.reverted == [[second]] -def test_async_python_v2_pp1_bypasses_the_window() -> None: - """The asynchronous Python transceiver does not consume the C++ transfer - buffer; with KV cache manager V2 and PP1 the scheduler's own KV admission - bounds it, so the executor-level window is skipped.""" +def test_async_python_v2_pp1_keeps_the_window() -> None: + """Coordinator admission also enforces the window for the early-budget runtime.""" h = _harness(is_kv_manager_v2=True, consumes_transfer_buffer=False) _receiving(h, 1) candidates = [_candidate(2), _candidate(3)] admitted, blocked = h.coordinator.admit(candidates) - assert (admitted, blocked) == (candidates, False) - assert h.effects.reverted == [] + assert (admitted, blocked) == ([], True) + assert h.effects.reverted == [candidates] @pytest.mark.parametrize( @@ -220,25 +218,33 @@ def test_revert_is_a_no_op_without_deferred_v2_candidates( assert h.effects.reverted == [] -# -- window bypass predicate --------------------------------------------------- +# -- early window predicate --------------------------------------------------- -def test_window_bypass_needs_the_async_python_runtime_with_v2_and_pp1() -> None: +def test_early_window_needs_the_async_python_runtime_with_v2_and_pp1() -> None: pp1, pp2 = SimpleNamespace(pp_size=1), SimpleNamespace(pp_size=2) async_python = SimpleNamespace(consumes_transfer_buffer=False) cpp = SimpleNamespace(consumes_transfer_buffer=True) - assert transfer_window_bypass_eligible(async_python, pp1, True) - assert not transfer_window_bypass_eligible(cpp, pp1, True) - assert not transfer_window_bypass_eligible(async_python, pp1, False) - assert not transfer_window_bypass_eligible(async_python, pp2, True) - assert not transfer_window_bypass_eligible(None, pp1, True) + assert early_transfer_window_eligible(async_python, pp1, True) + assert not early_transfer_window_eligible(cpp, pp1, True) + assert not early_transfer_window_eligible(async_python, pp1, False) + assert not early_transfer_window_eligible(async_python, pp2, True) + assert not early_transfer_window_eligible(None, pp1, True) -def test_window_bypass_does_not_assume_pp1_when_dist_lacks_pp_size() -> None: - """A dist without ``pp_size`` must not silently count as PP1: bypassing - the window on a misconfigured executor would drop a real budget.""" +def test_early_window_does_not_assume_pp1_when_dist_lacks_pp_size() -> None: + """Missing PP metadata must not enable local pre-allocation admission.""" async_python = SimpleNamespace(consumes_transfer_buffer=False) with pytest.raises(AttributeError): - transfer_window_bypass_eligible(async_python, SimpleNamespace(), True) + early_transfer_window_eligible(async_python, SimpleNamespace(), True) + + +@pytest.mark.parametrize( + "mode", ["TRTLLM_DISAGG_BENCHMARK_GEN_ONLY", "TRTLLM_DISABLE_KV_CACHE_TRANSFER_OVERLAP"] +) +def test_early_window_is_disabled_without_async_transfers(monkeypatch, mode: str) -> None: + monkeypatch.setenv(mode, "1") + transceiver = SimpleNamespace(consumes_transfer_buffer=False) + assert not early_transfer_window_eligible(transceiver, SimpleNamespace(pp_size=1), True) diff --git a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py index cc663b218a4d..f8c86a52b45c 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py +++ b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py @@ -18,10 +18,14 @@ No GPU required. """ +from typing import Callable from unittest.mock import Mock, call, patch import pytest +from tensorrt_llm._torch.disaggregation.orchestration.admission import ( + DisaggTransferAdmissionController, +) from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import BlockReusePolicy from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState from tensorrt_llm.llmapi.llm_args import CapacitySchedulerPolicy, ContextChunkingPolicy @@ -4056,3 +4060,112 @@ def test_parked_connector_load_keeps_kv_pressure_retryable() -> None: output = scheduler.schedule_request([loading, generation], set()) assert output.generation_requests == [] assert output.recompute_paused_requests == [] + + +class TestDisaggEarlyTransferBudget: + @staticmethod + def _request( + request_id: int, tokens: int, *, receiving: bool = False, global_tokens: int | None = None + ) -> Mock: + req = make_disagg_request(request_id, prompt_len=tokens) + req.py_prompt_len = tokens + req.total_input_len_cp = tokens if global_tokens is None else global_tokens + req.is_disagg_generation_transmission_in_progress = receiving + if receiving: + req.state_value = DISAGG_GEN_TRANS_IN_PROGRESS + return req + + @staticmethod + def _scheduler( + *, + max_tokens: int = 64, + prepare: Callable[[Mock], bool] | None = None, + max_batch_size: int = 100, + ) -> tuple[object, Mock, DisaggTransferAdmissionController]: + manager = make_kv_cache_manager(prepare_disagg_gen_init_fn=prepare) + scheduler = make_scheduler(manager, max_batch_size=max_batch_size, max_num_tokens=100) + controller = DisaggTransferAdmissionController(max_tokens, tokens_per_block=32) + scheduler.set_disagg_transfer_admission_controller(controller) + return scheduler, manager, controller + + def test_full_window_skips_all_init_allocation_and_keeps_decode(self) -> None: + scheduler, manager, _ = self._scheduler(max_tokens=9216, max_batch_size=1) + active = self._request(0, 8192, receiving=True) + pending = [self._request(index, 8192) for index in range(1, 513)] + decode = make_gen_request(1000) + decode.is_disagg_generation_transmission_in_progress = False + + output = scheduler.schedule_request([active, *pending, decode], set()) + + manager.prepare_disagg_gen_init.assert_not_called() + assert output.fitting_disagg_gen_init_requests == [] + assert ids(output.generation_requests) == [1000] + assert output.disagg_transfer_budget_blocked + + def test_idle_window_prepares_only_one_uniform_prompt(self) -> None: + scheduler, manager, controller = self._scheduler(max_tokens=9216) + requests = [self._request(index, 8192) for index in range(512)] + estimator = Mock(wraps=controller.estimate_request_blocks) + controller.estimate_request_blocks = estimator + + output = scheduler.schedule_request(requests, set()) + + manager.prepare_disagg_gen_init.assert_called_once_with(requests[0]) + assert ids(output.fitting_disagg_gen_init_requests) == [0] + assert not output.disagg_transfer_budget_blocked + assert estimator.call_count == len(requests) + + @pytest.mark.parametrize("head_fits_kv", [False, True]) + @pytest.mark.parametrize("tail_tokens", [0, 32]) + def test_budget_rejected_head_preserves_kv_fit_fcfs( + self, head_fits_kv: bool, tail_tokens: int + ) -> None: + scheduler, manager, controller = self._scheduler( + prepare=lambda req: head_fits_kv or req.py_request_id != 1 + ) + active = self._request(0, 32, receiving=True) + head = self._request(1, 64) + tail = self._request(2, tail_tokens) + requests = [active, head, tail] + + output = scheduler.schedule_request(requests, set()) + selected = controller.select(requests, output.fitting_disagg_gen_init_requests) + + if head_fits_kv: + manager.prepare_disagg_gen_init.assert_called_once_with(head) + assert output.fitting_disagg_gen_init_requests == [head] + assert selected.admitted_requests == [] + assert selected.deferred_request_count == 1 + assert output.disagg_transfer_budget_blocked + else: + assert manager.prepare_disagg_gen_init.call_args_list == [call(head), call(tail)] + assert selected.admitted_requests == [tail] + assert not output.disagg_transfer_budget_blocked + + def test_kv_failure_does_not_spend_idle_oversized_exception(self) -> None: + scheduler, manager, controller = self._scheduler( + max_tokens=32, prepare=lambda req: req.py_request_id != 0 + ) + rejected = self._request(0, 32) + oversized = self._request(1, 96) + tail = self._request(2, 0) + requests = [rejected, oversized, tail] + + output = scheduler.schedule_request(requests, set()) + selected = controller.select(requests, output.fitting_disagg_gen_init_requests) + + assert manager.prepare_disagg_gen_init.call_args_list == [call(rejected), call(oversized)] + assert selected.admitted_requests == [oversized] + assert not output.disagg_transfer_budget_blocked + + def test_budget_snapshot_resets_after_transfer_completes(self) -> None: + scheduler, manager, _ = self._scheduler(max_tokens=32) + active = self._request(0, 32, receiving=True) + pending = self._request(1, 32) + blocked = scheduler.schedule_request([active, pending], set()) + admitted = scheduler.schedule_request([pending], set()) + + assert blocked.disagg_transfer_budget_blocked + assert not admitted.disagg_transfer_budget_blocked + manager.prepare_disagg_gen_init.assert_called_once_with(pending) + assert admitted.fitting_disagg_gen_init_requests == [pending]