Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 40 additions & 29 deletions tensorrt_llm/_torch/disaggregation/orchestration/admission.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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
21 changes: 6 additions & 15 deletions tensorrt_llm/_torch/disaggregation/orchestration/coordinator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
24 changes: 11 additions & 13 deletions tensorrt_llm/_torch/pyexecutor/py_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
20 changes: 19 additions & 1 deletion tensorrt_llm/_torch/pyexecutor/scheduler/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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]
Expand Down Expand Up @@ -79,6 +84,7 @@ class SchedulerOutput(
"num_fitting_requests",
"scheduled_mm_encoder_items",
"recompute_paused_requests",
"disagg_transfer_budget_blocked",
],
)
):
Expand All @@ -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__ = ()
Expand All @@ -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,
Expand All @@ -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,
)


Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
70 changes: 70 additions & 0 deletions tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand All @@ -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)),
)

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -592,6 +661,7 @@ def preempt_for_pages(req: LlmRequest) -> bool:
recompute_paused,
disagg_candidates,
has_chunking,
transfer_budget_blocked,
)

# ---- Prefix-aware skip ----
Expand Down
Loading
Loading