From 54c4cab324754bc2d9a41cdd30b5c4bdbea3bd2c Mon Sep 17 00:00:00 2001 From: Yi Zhang <187001205+yizhang-nv@users.noreply.github.com> Date: Sun, 4 Oct 2026 20:31:43 -0700 Subject: [PATCH 1/2] [None][fix] Bound V2 disaggregated KV transfer admission Signed-off-by: Yi Zhang <187001205+yizhang-nv@users.noreply.github.com> --- .../disaggregation/orchestration/admission.py | 69 +++--- .../orchestration/coordinator.py | 21 +- tensorrt_llm/_torch/pyexecutor/py_executor.py | 24 +- .../_torch/pyexecutor/scheduler/scheduler.py | 20 +- .../pyexecutor/scheduler/scheduler_v2.py | 70 ++++++ .../disaggregation/test_benchmark_disagg.py | 15 ++ .../test_disagg_coordinator_admission.py | 42 ++-- .../kv_cache/test_kv_cache_v2_scheduler.py | 149 +++++++++++++ .../executor/test_multimodal_scheduler.py | 24 ++ .../disaggregated/test_kv_transfer.py | 206 +++++++++++++++++- 10 files changed, 561 insertions(+), 79 deletions(-) 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..deeded004413 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,148 @@ 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_zero_cost_head_can_enter_a_full_window(self) -> None: + scheduler, manager, controller = self._scheduler(max_tokens=32) + active = self._request(0, 32, receiving=True) + zero = self._request(1, 0) + tail = self._request(2, 32) + requests = [active, zero, tail] + + output = scheduler.schedule_request(requests, set()) + + manager.prepare_disagg_gen_init.assert_called_once_with(zero) + selected = controller.select(requests, output.fitting_disagg_gen_init_requests) + assert selected.admitted_requests == [zero] + assert not output.disagg_transfer_budget_blocked + + def test_inflight_init_is_excluded_and_cp_uses_global_cost(self) -> None: + scheduler, manager, _ = self._scheduler(max_tokens=64) + active = self._request(0, 16, receiving=True, global_tokens=32) + inflight = self._request(1, 0) + pending = self._request(2, 16, global_tokens=64) + + output = scheduler.schedule_request([active, inflight, pending], {1}) + + manager.prepare_disagg_gen_init.assert_not_called() + assert output.fitting_disagg_gen_init_requests == [] + assert 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] + + def test_disabled_window_preserves_allocation_candidates(self) -> None: + scheduler, manager, _ = self._scheduler(max_tokens=0) + requests = [self._request(index, 128) for index in range(3)] + + output = scheduler.schedule_request(requests, set()) + + assert manager.prepare_disagg_gen_init.call_args_list == [call(req) for req in requests] + assert output.fitting_disagg_gen_init_requests == requests + assert not output.disagg_transfer_budget_blocked diff --git a/tests/unittest/_torch/executor/test_multimodal_scheduler.py b/tests/unittest/_torch/executor/test_multimodal_scheduler.py index f67ad06095f4..c3469229b4a6 100644 --- a/tests/unittest/_torch/executor/test_multimodal_scheduler.py +++ b/tests/unittest/_torch/executor/test_multimodal_scheduler.py @@ -758,3 +758,27 @@ def test_mm_encoder_state_publishes_the_buffer_without_copying(): assert state.progress is MultimodalEncoderProgress.READY assert state.pending_item_indices() == [] assert state.resident_output_bytes(4) == (2 + 3) * 4 + + +@pytest.mark.parametrize("wrapper_type", [MultimodalScheduler, MultimodalEagerEncoderScheduler]) +def test_mm_wrapper_forwards_early_transfer_budget_and_preserves_blocked_result( + wrapper_type, +) -> None: + from tensorrt_llm._torch.disaggregation.orchestration.admission import ( + DisaggTransferAdmissionController, + ) + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler import SchedulerOutput + + output = SchedulerOutput([], [], [], [], [], 0, disagg_transfer_budget_blocked=True) + base = SimpleNamespace( + schedule_request=Mock(return_value=output), + set_disagg_transfer_admission_controller=Mock(), + ) + wrapper = wrapper_type(base, max_batch_size=1, max_num_tokens=32) + controller = DisaggTransferAdmissionController(32, 32) + + wrapper.set_disagg_transfer_admission_controller(controller) + result = wrapper.schedule_request([], set()) + + base.set_disagg_transfer_admission_controller.assert_called_once_with(controller) + assert result.disagg_transfer_budget_blocked diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index 92b7a836032f..6c7e4c432cca 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -55,6 +55,7 @@ from tensorrt_llm import DisaggregatedParams, Mapping, SamplingParams from tensorrt_llm._torch.disaggregation.base import CacheKind, Chunk, TokenRange from tensorrt_llm._torch.disaggregation.base.transfer import SessionStatus, WaitResult +from tensorrt_llm._torch.disaggregation.native.bounce.config import Config as BounceConfig from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 @@ -448,6 +449,7 @@ def create_transfer_worker_setup( is_mla: bool = False, use_v2: bool = False, max_attention_window_vec: Optional[List[int]] = None, + bounce_config: BounceConfig | None = None, ): """Helper function to set up transfer workers for testing. @@ -628,6 +630,7 @@ def create_transfer_worker_setup( TransferWorker( TransferWorkerConfig( kv_cache_manager=ctx_kv_cache_manager, + bounce=bounce_config, device_id=device_id, instance_name=ctx_instance_name, max_concurrent_sessions=max_batch_size * 2, @@ -744,6 +747,7 @@ def create_transfer_worker_setup( TransferWorker( TransferWorkerConfig( kv_cache_manager=gen_kv_cache_manager, + bounce=bounce_config, device_id=device_id, instance_name=gen_instance_name, max_concurrent_sessions=max_batch_size * 2, @@ -1089,10 +1093,31 @@ def add_and_verify_request( for ctx_transfer_worker in valid_ctx_transfer_workers ] - time.sleep(0.1) - + readiness_start = time.monotonic() + readiness_deadline = readiness_start + 10.0 + while not all(session.receiver_ready for session in sender_sessions): + if time.monotonic() >= readiness_deadline: + break + time.sleep(0.005) + readiness_elapsed = time.monotonic() - readiness_start + sender_readiness = [ + { + "status": session.status.name, + "receiver_ready": session.receiver_ready, + "metadata_count": len( + session._sender._get_req_info(session.disagg_request_id) or {} + ), + } + for session in sender_sessions + ] + readiness_diagnostics = ( + f"Receiver-first readiness: elapsed={readiness_elapsed:.6f}s, " + f"senders={sender_readiness}, " + f"receivers={[session.status.name for session in receiver_sessions]}" + ) + print(readiness_diagnostics, flush=True) for sender_session in sender_sessions: - assert sender_session.status != SessionStatus.INIT + assert sender_session.status == SessionStatus.READY, readiness_diagnostics send_kv_slices = [ Chunk( @@ -2236,5 +2261,180 @@ def test_transfer_worker_pipelined_unaligned_chunk_boundaries( worker.shutdown() +@pytest.mark.timeout(240) +@pytest.mark.parametrize( + "ctx_tp,gen_tp", [(1, 1), (2, 1), (1, 4)], ids=["single", "fanin", "fanout"] +) +@pytest.mark.parametrize("mode", ["coalesced", "receiver_full", "one_sender_full"]) +@pytest.mark.parametrize("window", [None, [1024, 24]], ids=["full", "vswa"]) +def test_transfer_worker_v2_native_bounce_data( + monkeypatch: pytest.MonkeyPatch, + ctx_tp: int, + gen_tp: int, + mode: str, + window: list[int] | None, +) -> None: + """Check real KV bytes through native bounce and real arena-pressure fallback.""" + from tensorrt_llm._torch.disaggregation.native.bounce import impl as bounce_impl + from tensorrt_llm._torch.disaggregation.native.bounce.config import FixedSizing + + reservations = [] + coalesced = [] + fragments = [] + scattered = [] + scatter_owners = {} + reserve = bounce_impl.VmmBounceTransport.reserve + build_request = bounce_impl.VmmBounceTransport.build_request + make_fragment_request = transfer_mod.Sender._make_agent_request + scatter = bounce_impl.scatter_contiguous + + def observe_reserve( + transport: bounce_impl.VmmBounceTransport, + recv_req: transfer_mod.RecvReqInfo, + num_writers: int = 1, + **kwargs, + ) -> bool: + key = (recv_req.unique_rid, recv_req.slice_id) + assert recv_req.bounce_dst_base is None + assert key not in transport._reserved_map + accepted = reserve(transport, recv_req, num_writers, **kwargs) + if accepted: + context = transport._reserved_map[key] + for writer_index in range(num_writers): + scatter_owners[context.writer_base(writer_index)] = (transport, key) + else: + assert recv_req.bounce_dst_base is None + assert key not in transport._reserved_map + reservations.append((accepted, num_writers)) + return accepted + + def observe_build( + transport: bounce_impl.VmmBounceTransport, write_meta: transfer_mod.WriteMeta + ) -> tuple | None: + built = build_request(transport, write_meta) + if built is not None: + coalesced.append( + ( + int(write_meta.bounce_dst_base), + tuple(map(int, write_meta.dst_ptrs)), + tuple(map(int, write_meta.sizes)), + ) + ) + return built + + def observe_fragment(write_meta: transfer_mod.WriteMeta, device_id: int) -> object: + request = make_fragment_request(write_meta, device_id) + if write_meta.meta_type != transfer_mod.WriteMetaType.AUX: + fragments.append(int(write_meta.sizes.size)) + return request + + def observe_scatter( + src_base: int, + dst_ptrs: np.ndarray, + sizes: np.ndarray, + offsets: np.ndarray, + *, + stream, + ) -> None: + transport, key = scatter_owners[int(src_base)] + owner = next(worker for worker in gen_workers if worker._bounce is transport) + session = owner._receiver._get_session(key[0]) + assert session is not None + task = next(item for item in session._kv_tasks if item.slice_id == key[1]) + assert task.status != transfer_mod.TaskStatus.TRANSFERRED + scatter(src_base, dst_ptrs, sizes, offsets, stream=stream) + # The scatter worker synchronizes only after this real copy launch returns. + assert task.status != transfer_mod.TaskStatus.TRANSFERRED + scattered.append((int(src_base), tuple(map(int, dst_ptrs)), tuple(map(int, sizes)))) + + monkeypatch.setattr(bounce_impl.VmmBounceTransport, "reserve", observe_reserve) + monkeypatch.setattr(bounce_impl.VmmBounceTransport, "build_request", observe_build) + monkeypatch.setattr(transfer_mod.Sender, "_make_agent_request", staticmethod(observe_fragment)) + monkeypatch.setattr(bounce_impl, "scatter_contiguous", observe_scatter) + + config = BounceConfig(sizing=FixedSizing(capacity_mb=32), min_blocks=1, min_bytes=1) + setup = create_transfer_worker_setup( + ctx_tp=ctx_tp, + ctx_pp=1, + ctx_enable_dp=False, + gen_tp=gen_tp, + gen_pp=1, + gen_enable_dp=False, + use_v2=True, + max_attention_window_vec=window, + bounce_config=config, + ) + ctx_workers = setup["ctx_transfer_workers"] + gen_workers = setup["gen_transfer_workers"] + workers = ctx_workers + gen_workers + held = [] + try: + # Allocation failure must fail this test instead of silently testing NoBounce. + for worker in workers: + assert isinstance(worker._bounce, bounce_impl.VmmBounceTransport) + assert worker._bounce.enabled + assert worker._config.agent_buffer_size_mb == 0 + + if mode == "receiver_full": + pressure_allocators = [worker._bounce._recv_alloc for worker in gen_workers] + elif mode == "one_sender_full": + pressure_allocators = [ctx_workers[0]._bounce._send_alloc] + else: + pressure_allocators = [] + for allocator in pressure_allocators: + allocation = allocator.reserve(allocator.capacity, timeout=0) + assert allocation is not None + held.append((allocator, allocation[0])) + + expected_writers = max(ctx_tp // gen_tp, 1) + total_writes = max(ctx_tp, gen_tp) + for phase, request_len in enumerate((64, 32)): + if phase: + while held: + allocator, slot = held.pop() + allocator.release(slot) + reservations.clear() + coalesced.clear() + fragments.clear() + scattered.clear() + scatter_owners.clear() + + # The helper waits for real sessions, verifies every KV group with .equal(), + # and closes request caches. Second transfer also exercises arena reuse. + add_and_verify_request( + setup, 2 * phase, 2 * phase + 1, request_len, send_first=bool(phase) + ) + receiver_full = mode == "receiver_full" and phase == 0 + expected_fallback = ( + total_writes + if receiver_full + else max(gen_tp // ctx_tp, 1) + if mode == "one_sender_full" and phase == 0 + else 0 + ) + assert reservations == [(not receiver_full, expected_writers)] * gen_tp + assert len(fragments) == expected_fallback + assert all(count > 1 for count in fragments) + assert len(coalesced) == total_writes - expected_fallback + assert all(len(sizes) > 1 for _, _, sizes in coalesced) + # Each successful writer must scatter from its own base, including when + # its fan-in sibling fell back and wrote directly to final destinations. + assert sorted(scattered) == sorted(coalesced) + for worker in workers: + assert not worker._bounce._reserved_map + + for worker in workers: + for allocator in (worker._bounce._send_alloc, worker._bounce._recv_alloc): + allocation = allocator.reserve(allocator.capacity, timeout=0) + assert allocation is not None, "completed transfer leaked an arena slot" + allocator.release(allocation[0]) + finally: + while held: + allocator, slot = held.pop() + allocator.release(slot) + for worker in workers: + worker.shutdown() + + if __name__ == "__main__": test_transfer_worker_v1(1, 1, False, 1, 1, False, False) From bbf95c9b0a8a54f7ffc3d4b3e42d4850e33a6368 Mon Sep 17 00:00:00 2001 From: Yi Zhang <187001205+yizhang-nv@users.noreply.github.com> Date: Sun, 4 Oct 2026 20:41:09 -0700 Subject: [PATCH 2/2] [None][test] Limit regression coverage to transfer admission Signed-off-by: Yi Zhang <187001205+yizhang-nv@users.noreply.github.com> --- .../kv_cache/test_kv_cache_v2_scheduler.py | 36 --- .../executor/test_multimodal_scheduler.py | 24 -- .../disaggregated/test_kv_transfer.py | 206 +----------------- 3 files changed, 3 insertions(+), 263 deletions(-) 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 deeded004413..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 @@ -4158,32 +4158,6 @@ def test_kv_failure_does_not_spend_idle_oversized_exception(self) -> None: assert selected.admitted_requests == [oversized] assert not output.disagg_transfer_budget_blocked - def test_zero_cost_head_can_enter_a_full_window(self) -> None: - scheduler, manager, controller = self._scheduler(max_tokens=32) - active = self._request(0, 32, receiving=True) - zero = self._request(1, 0) - tail = self._request(2, 32) - requests = [active, zero, tail] - - output = scheduler.schedule_request(requests, set()) - - manager.prepare_disagg_gen_init.assert_called_once_with(zero) - selected = controller.select(requests, output.fitting_disagg_gen_init_requests) - assert selected.admitted_requests == [zero] - assert not output.disagg_transfer_budget_blocked - - def test_inflight_init_is_excluded_and_cp_uses_global_cost(self) -> None: - scheduler, manager, _ = self._scheduler(max_tokens=64) - active = self._request(0, 16, receiving=True, global_tokens=32) - inflight = self._request(1, 0) - pending = self._request(2, 16, global_tokens=64) - - output = scheduler.schedule_request([active, inflight, pending], {1}) - - manager.prepare_disagg_gen_init.assert_not_called() - assert output.fitting_disagg_gen_init_requests == [] - assert 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) @@ -4195,13 +4169,3 @@ def test_budget_snapshot_resets_after_transfer_completes(self) -> None: 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] - - def test_disabled_window_preserves_allocation_candidates(self) -> None: - scheduler, manager, _ = self._scheduler(max_tokens=0) - requests = [self._request(index, 128) for index in range(3)] - - output = scheduler.schedule_request(requests, set()) - - assert manager.prepare_disagg_gen_init.call_args_list == [call(req) for req in requests] - assert output.fitting_disagg_gen_init_requests == requests - assert not output.disagg_transfer_budget_blocked diff --git a/tests/unittest/_torch/executor/test_multimodal_scheduler.py b/tests/unittest/_torch/executor/test_multimodal_scheduler.py index c3469229b4a6..f67ad06095f4 100644 --- a/tests/unittest/_torch/executor/test_multimodal_scheduler.py +++ b/tests/unittest/_torch/executor/test_multimodal_scheduler.py @@ -758,27 +758,3 @@ def test_mm_encoder_state_publishes_the_buffer_without_copying(): assert state.progress is MultimodalEncoderProgress.READY assert state.pending_item_indices() == [] assert state.resident_output_bytes(4) == (2 + 3) * 4 - - -@pytest.mark.parametrize("wrapper_type", [MultimodalScheduler, MultimodalEagerEncoderScheduler]) -def test_mm_wrapper_forwards_early_transfer_budget_and_preserves_blocked_result( - wrapper_type, -) -> None: - from tensorrt_llm._torch.disaggregation.orchestration.admission import ( - DisaggTransferAdmissionController, - ) - from tensorrt_llm._torch.pyexecutor.scheduler.scheduler import SchedulerOutput - - output = SchedulerOutput([], [], [], [], [], 0, disagg_transfer_budget_blocked=True) - base = SimpleNamespace( - schedule_request=Mock(return_value=output), - set_disagg_transfer_admission_controller=Mock(), - ) - wrapper = wrapper_type(base, max_batch_size=1, max_num_tokens=32) - controller = DisaggTransferAdmissionController(32, 32) - - wrapper.set_disagg_transfer_admission_controller(controller) - result = wrapper.schedule_request([], set()) - - base.set_disagg_transfer_admission_controller.assert_called_once_with(controller) - assert result.disagg_transfer_budget_blocked diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index 6c7e4c432cca..92b7a836032f 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -55,7 +55,6 @@ from tensorrt_llm import DisaggregatedParams, Mapping, SamplingParams from tensorrt_llm._torch.disaggregation.base import CacheKind, Chunk, TokenRange from tensorrt_llm._torch.disaggregation.base.transfer import SessionStatus, WaitResult -from tensorrt_llm._torch.disaggregation.native.bounce.config import Config as BounceConfig from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 @@ -449,7 +448,6 @@ def create_transfer_worker_setup( is_mla: bool = False, use_v2: bool = False, max_attention_window_vec: Optional[List[int]] = None, - bounce_config: BounceConfig | None = None, ): """Helper function to set up transfer workers for testing. @@ -630,7 +628,6 @@ def create_transfer_worker_setup( TransferWorker( TransferWorkerConfig( kv_cache_manager=ctx_kv_cache_manager, - bounce=bounce_config, device_id=device_id, instance_name=ctx_instance_name, max_concurrent_sessions=max_batch_size * 2, @@ -747,7 +744,6 @@ def create_transfer_worker_setup( TransferWorker( TransferWorkerConfig( kv_cache_manager=gen_kv_cache_manager, - bounce=bounce_config, device_id=device_id, instance_name=gen_instance_name, max_concurrent_sessions=max_batch_size * 2, @@ -1093,31 +1089,10 @@ def add_and_verify_request( for ctx_transfer_worker in valid_ctx_transfer_workers ] - readiness_start = time.monotonic() - readiness_deadline = readiness_start + 10.0 - while not all(session.receiver_ready for session in sender_sessions): - if time.monotonic() >= readiness_deadline: - break - time.sleep(0.005) - readiness_elapsed = time.monotonic() - readiness_start - sender_readiness = [ - { - "status": session.status.name, - "receiver_ready": session.receiver_ready, - "metadata_count": len( - session._sender._get_req_info(session.disagg_request_id) or {} - ), - } - for session in sender_sessions - ] - readiness_diagnostics = ( - f"Receiver-first readiness: elapsed={readiness_elapsed:.6f}s, " - f"senders={sender_readiness}, " - f"receivers={[session.status.name for session in receiver_sessions]}" - ) - print(readiness_diagnostics, flush=True) + time.sleep(0.1) + for sender_session in sender_sessions: - assert sender_session.status == SessionStatus.READY, readiness_diagnostics + assert sender_session.status != SessionStatus.INIT send_kv_slices = [ Chunk( @@ -2261,180 +2236,5 @@ def test_transfer_worker_pipelined_unaligned_chunk_boundaries( worker.shutdown() -@pytest.mark.timeout(240) -@pytest.mark.parametrize( - "ctx_tp,gen_tp", [(1, 1), (2, 1), (1, 4)], ids=["single", "fanin", "fanout"] -) -@pytest.mark.parametrize("mode", ["coalesced", "receiver_full", "one_sender_full"]) -@pytest.mark.parametrize("window", [None, [1024, 24]], ids=["full", "vswa"]) -def test_transfer_worker_v2_native_bounce_data( - monkeypatch: pytest.MonkeyPatch, - ctx_tp: int, - gen_tp: int, - mode: str, - window: list[int] | None, -) -> None: - """Check real KV bytes through native bounce and real arena-pressure fallback.""" - from tensorrt_llm._torch.disaggregation.native.bounce import impl as bounce_impl - from tensorrt_llm._torch.disaggregation.native.bounce.config import FixedSizing - - reservations = [] - coalesced = [] - fragments = [] - scattered = [] - scatter_owners = {} - reserve = bounce_impl.VmmBounceTransport.reserve - build_request = bounce_impl.VmmBounceTransport.build_request - make_fragment_request = transfer_mod.Sender._make_agent_request - scatter = bounce_impl.scatter_contiguous - - def observe_reserve( - transport: bounce_impl.VmmBounceTransport, - recv_req: transfer_mod.RecvReqInfo, - num_writers: int = 1, - **kwargs, - ) -> bool: - key = (recv_req.unique_rid, recv_req.slice_id) - assert recv_req.bounce_dst_base is None - assert key not in transport._reserved_map - accepted = reserve(transport, recv_req, num_writers, **kwargs) - if accepted: - context = transport._reserved_map[key] - for writer_index in range(num_writers): - scatter_owners[context.writer_base(writer_index)] = (transport, key) - else: - assert recv_req.bounce_dst_base is None - assert key not in transport._reserved_map - reservations.append((accepted, num_writers)) - return accepted - - def observe_build( - transport: bounce_impl.VmmBounceTransport, write_meta: transfer_mod.WriteMeta - ) -> tuple | None: - built = build_request(transport, write_meta) - if built is not None: - coalesced.append( - ( - int(write_meta.bounce_dst_base), - tuple(map(int, write_meta.dst_ptrs)), - tuple(map(int, write_meta.sizes)), - ) - ) - return built - - def observe_fragment(write_meta: transfer_mod.WriteMeta, device_id: int) -> object: - request = make_fragment_request(write_meta, device_id) - if write_meta.meta_type != transfer_mod.WriteMetaType.AUX: - fragments.append(int(write_meta.sizes.size)) - return request - - def observe_scatter( - src_base: int, - dst_ptrs: np.ndarray, - sizes: np.ndarray, - offsets: np.ndarray, - *, - stream, - ) -> None: - transport, key = scatter_owners[int(src_base)] - owner = next(worker for worker in gen_workers if worker._bounce is transport) - session = owner._receiver._get_session(key[0]) - assert session is not None - task = next(item for item in session._kv_tasks if item.slice_id == key[1]) - assert task.status != transfer_mod.TaskStatus.TRANSFERRED - scatter(src_base, dst_ptrs, sizes, offsets, stream=stream) - # The scatter worker synchronizes only after this real copy launch returns. - assert task.status != transfer_mod.TaskStatus.TRANSFERRED - scattered.append((int(src_base), tuple(map(int, dst_ptrs)), tuple(map(int, sizes)))) - - monkeypatch.setattr(bounce_impl.VmmBounceTransport, "reserve", observe_reserve) - monkeypatch.setattr(bounce_impl.VmmBounceTransport, "build_request", observe_build) - monkeypatch.setattr(transfer_mod.Sender, "_make_agent_request", staticmethod(observe_fragment)) - monkeypatch.setattr(bounce_impl, "scatter_contiguous", observe_scatter) - - config = BounceConfig(sizing=FixedSizing(capacity_mb=32), min_blocks=1, min_bytes=1) - setup = create_transfer_worker_setup( - ctx_tp=ctx_tp, - ctx_pp=1, - ctx_enable_dp=False, - gen_tp=gen_tp, - gen_pp=1, - gen_enable_dp=False, - use_v2=True, - max_attention_window_vec=window, - bounce_config=config, - ) - ctx_workers = setup["ctx_transfer_workers"] - gen_workers = setup["gen_transfer_workers"] - workers = ctx_workers + gen_workers - held = [] - try: - # Allocation failure must fail this test instead of silently testing NoBounce. - for worker in workers: - assert isinstance(worker._bounce, bounce_impl.VmmBounceTransport) - assert worker._bounce.enabled - assert worker._config.agent_buffer_size_mb == 0 - - if mode == "receiver_full": - pressure_allocators = [worker._bounce._recv_alloc for worker in gen_workers] - elif mode == "one_sender_full": - pressure_allocators = [ctx_workers[0]._bounce._send_alloc] - else: - pressure_allocators = [] - for allocator in pressure_allocators: - allocation = allocator.reserve(allocator.capacity, timeout=0) - assert allocation is not None - held.append((allocator, allocation[0])) - - expected_writers = max(ctx_tp // gen_tp, 1) - total_writes = max(ctx_tp, gen_tp) - for phase, request_len in enumerate((64, 32)): - if phase: - while held: - allocator, slot = held.pop() - allocator.release(slot) - reservations.clear() - coalesced.clear() - fragments.clear() - scattered.clear() - scatter_owners.clear() - - # The helper waits for real sessions, verifies every KV group with .equal(), - # and closes request caches. Second transfer also exercises arena reuse. - add_and_verify_request( - setup, 2 * phase, 2 * phase + 1, request_len, send_first=bool(phase) - ) - receiver_full = mode == "receiver_full" and phase == 0 - expected_fallback = ( - total_writes - if receiver_full - else max(gen_tp // ctx_tp, 1) - if mode == "one_sender_full" and phase == 0 - else 0 - ) - assert reservations == [(not receiver_full, expected_writers)] * gen_tp - assert len(fragments) == expected_fallback - assert all(count > 1 for count in fragments) - assert len(coalesced) == total_writes - expected_fallback - assert all(len(sizes) > 1 for _, _, sizes in coalesced) - # Each successful writer must scatter from its own base, including when - # its fan-in sibling fell back and wrote directly to final destinations. - assert sorted(scattered) == sorted(coalesced) - for worker in workers: - assert not worker._bounce._reserved_map - - for worker in workers: - for allocator in (worker._bounce._send_alloc, worker._bounce._recv_alloc): - allocation = allocator.reserve(allocator.capacity, timeout=0) - assert allocation is not None, "completed transfer leaked an arena slot" - allocator.release(allocation[0]) - finally: - while held: - allocator, slot = held.pop() - allocator.release(slot) - for worker in workers: - worker.shutdown() - - if __name__ == "__main__": test_transfer_worker_v1(1, 1, False, 1, 1, False, False)