Skip to content
Closed
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
52 changes: 33 additions & 19 deletions tensorrt_llm/_torch/disaggregation/native/transfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -537,6 +537,7 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta):
assert write_meta.src_ptrs.size == write_meta.dst_ptrs.size == write_meta.sizes.size, (
f"WriteMeta ptr/size mismatch for unique_rid={write_meta.unique_rid}"
)
assert write_meta.slice_id is not None

with self._sessions_lock:
session = self._get_session(write_meta.unique_rid)
Expand All @@ -546,8 +547,12 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta):
)
logger.error(msg)
write_meta.task.fail(RuntimeError(msg))
# The session can be deregistered (cancel_request/ctx timeout) while
# this slice is still queued. Without a result frame the peer's RX
# task stays unresolved for the whole kv_transfer_timeout_ms with its
# KV pages pinned, so abort it here as the sibling exits below do.
self._abort_receiver_slice(write_meta)
return
assert write_meta.slice_id is not None
task = session.kv_tasks[write_meta.slice_id]
timer = task._perf_timer
if timer:
Expand All @@ -573,15 +578,7 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta):
task.fail(
RuntimeError(f"session {write_meta.unique_rid} {status.value}, transfer aborted")
)
self._get_or_connect_dealer(write_meta.peer_endpoint).send(
_make_kv_result_msg(
self._instance_rank,
write_meta.unique_rid,
write_meta.slice_id,
True, # is_last_slice — ensures receiver resolves its task future
AgentResult.FAILED,
)
)
self._abort_receiver_slice(write_meta)
return

from .bounce import build_send_request, encode_result_tail
Expand All @@ -603,15 +600,7 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta):
f"{write_meta.unique_rid} slice={write_meta.slice_id}: {e}"
)
task.fail(RuntimeError(f"build_send_request failed: {e}"))
self._get_or_connect_dealer(write_meta.peer_endpoint).send(
_make_kv_result_msg(
self._instance_rank,
write_meta.unique_rid,
write_meta.slice_id,
True, # is_last_slice — ensures receiver resolves its task future
AgentResult.FAILED,
)
)
self._abort_receiver_slice(write_meta)
return
if timer:
timer.record_transfer_start(write_meta.peer_rank)
Expand Down Expand Up @@ -1131,6 +1120,31 @@ def _respond_with_kv(self, _send_id: bytes, message: list[bytes]):
task._perf_timer.record_push_start(trans_meta.peer_rank)
self._enqueue(trans_meta)

def _abort_receiver_slice(self, write_meta: WriteMeta):
"""Tell the peer this slice failed so it resolves its RX task future now.

Called from _deliver_kv_to_agent on a _process_task_queue worker thread,
hence the thread-local DEALER cache: self._dealers is unsynchronized and
listener-thread-only. A send failure is swallowed like in
_send_failed_result_to_receiver — the local task is already failed and a
dead peer must not turn into a second, unhandled failure.
"""
try:
self._get_or_connect_thread_dealer(write_meta.peer_endpoint).send(
_make_kv_result_msg(
self._instance_rank,
write_meta.unique_rid,
write_meta.slice_id,
True, # is_last_slice — ensures receiver resolves its task future
AgentResult.FAILED,
)
)
except Exception as e:
logger.warning(
f"_abort_receiver_slice: failed to abort receiver slice for "
f"rid={write_meta.unique_rid} slice={write_meta.slice_id}: {e}"
)

def _send_failed_result_to_receiver(self, info: RecvReqInfo):
try:
peer_ri = self._registrar.get_peer_rank_info(info.instance_name, info.instance_rank)
Expand Down
5 changes: 5 additions & 0 deletions tensorrt_llm/_torch/disaggregation/transceiver.py
Original file line number Diff line number Diff line change
Expand Up @@ -965,6 +965,11 @@ def _assert_disagg_history_declared(self, req: LlmRequest) -> None:
f"_try_schedule_disagg_gen_init."
)

def context_transfer_is_waiting_for_peer(self, req: LlmRequest) -> bool:
# The send session stays in SessionStatus.INIT until every peer rank's
# request info has arrived; only then can any KV be written.
return not self._transfer_worker.has_all_peer_req_infos_for_send(get_unique_rid(req))

def cancel_request(self, req: LlmRequest) -> bool:
"""Cancel the transfer for the given request.

Expand Down
11 changes: 11 additions & 0 deletions tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py
Original file line number Diff line number Diff line change
Expand Up @@ -263,6 +263,17 @@ def cancel_request(self, req: LlmRequest):
def supports_inflight_request_cancellation(self) -> bool:
return False

def context_transfer_is_waiting_for_peer(self, req: LlmRequest) -> bool:
"""Whether a context send is still waiting for its peer to ask for the data.

respond_and_send_async() only creates the send session; the KV write
cannot start until the generation peer requests it. Runtimes that can
observe that boundary report it here so the transfer timeout measures
the transfer instead of the peer's queueing delay. Default False keeps
the timeout unchanged for runtimes that cannot distinguish the two.
"""
return False

def has_poisoned_transfer_buffer(self) -> bool:
return False

Expand Down
3 changes: 3 additions & 0 deletions tensorrt_llm/_torch/pyexecutor/llm_request.py
Original file line number Diff line number Diff line change
Expand Up @@ -991,6 +991,9 @@ def __init__(
self.is_cuda_graph_dummy = False
self.py_kv_transfer_start_time = None
self.py_kv_transfer_timed_out = False
# Set alongside py_kv_transfer_start_time for a context send and never
# rebased, so waiting for a peer that never asks stays bounded.
self.py_kv_transfer_peer_wait_start = None

# Encoder-decoder runtime state. ``py_encoder_output`` holds the
# packed encoder hidden states produced by the encoder iteration as
Expand Down
63 changes: 58 additions & 5 deletions tensorrt_llm/_torch/pyexecutor/py_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -3626,16 +3626,27 @@ def _check_disagg_transfer_progress_when_idle(
and not self._is_disagg_gen_only_no_context_benchmark()):
return

local_need_gen_check = (uses_async_gen_transfer and local_needs_progress
and wait_for_disagg_gen_transfer_progress)
# An in-flight receive must vote too, not just admission-budget
# pressure. The context wait is a no-op for a receive-only worker:
# check_context_transfer_status() returns on `not
# _ever_had_send_session` *above* its poll, so it never sleeps. Such a
# worker would spin the scheduler loop and starve the GIL from the
# transfer threads that alone finish the receive it is waiting on,
# whereas the generation wait does reach its poll interval.
local_need_gen_check = (
uses_async_gen_transfer and local_needs_progress
and (wait_for_disagg_gen_transfer_progress
or any(req.is_disagg_generation_transmission_in_progress
for req in self.active_requests)))

any_need_gen_check = self._sync_disagg_gen_status_entry(
local_need_gen_check)
if any_need_gen_check > 0:
if local_need_gen_check:
logger.debug(
"Waiting for generation KV cache transfer progress to "
"free disagg admission budget")
"complete an in-flight receive or free disagg admission "
"budget")
self._check_disagg_gen_cache_transfer_status(1)
return

Expand Down Expand Up @@ -6713,6 +6724,38 @@ def _check_gen_cache_transfer_errors_consensus(self) -> None:
requests=error_requests,
charge_budget=False)

# How much of kv_transfer_timeout_ms a context send may spend waiting for
# its generation peer to ask for the data before the deadline stops being
# rebased. The peer arrives late for a structural reason (generation
# max_batch_size below the offered concurrency), so the wait must be
# tolerated -- but a peer that never asks has to stay reclaimable, since
# this deadline is the only path that ends such a transfer. 3x covers the
# measured peer-wait spread on the 8K/512 stress run (max 182.7s) and lands
# at the disaggregated router's own req_timeout_secs=180 default, past
# which no peer will ask.
_CTX_PEER_WAIT_TIMEOUT_MULTIPLIER = 3

def _context_transfer_peer_wait_is_within_ceiling(
self, req: LlmRequest, current_time: float,
ceiling_ms: float) -> bool:
"""Whether a context send may keep rebasing its transfer deadline.

The clock is stamped when the send session is created, but the write
only starts once the peer requests the data, so charging the peer wait
to the transfer times out transfers that never got to run. The local
stamps are checked first: the transceiver query walks per-request peer
bookkeeping under a lock, and a request past the ceiling cannot be
rebased regardless of what it reports.
"""
if req.py_kv_transfer_peer_wait_start is None:
return False
peer_wait_ms = (current_time -
req.py_kv_transfer_peer_wait_start) * 1000
if peer_wait_ms > ceiling_ms:
return False
return self.kv_cache_transceiver.context_transfer_is_waiting_for_peer(
req)

Comment thread
coderabbitai[bot] marked this conversation as resolved.
@nvtx_range("_check_kv_transfer_timeout")
def _check_kv_transfer_timeout(self):
if not self.kv_cache_transceiver:
Expand All @@ -6721,8 +6764,10 @@ def _check_kv_transfer_timeout(self):
if timeout_ms is None:
return

current_time = time.monotonic()
peer_wait_ceiling_ms = timeout_ms * self._CTX_PEER_WAIT_TIMEOUT_MULTIPLIER

def flag_if_kv_transfer_timed_out(req: LlmRequest, type: str) -> None:
current_time = time.monotonic()
if req.py_kv_transfer_start_time is None:
return
elapsed_time = (current_time - req.py_kv_transfer_start_time) * 1000
Expand All @@ -6737,6 +6782,10 @@ def flag_if_kv_transfer_timed_out(req: LlmRequest, type: str) -> None:
req.py_kv_transfer_timed_out = True

for req in self.async_transfer_manager.requests_in_transfer().values():
if self._context_transfer_peer_wait_is_within_ceiling(
req, current_time, peer_wait_ceiling_ms):
req.py_kv_transfer_start_time = current_time
continue
flag_if_kv_transfer_timed_out(req, "context")

for req in self.active_requests:
Expand Down Expand Up @@ -7202,6 +7251,7 @@ def _prepare_disagg_gen_transmission_complete(self, scheduled_batch):
req.decoding_iter = 1
req.py_decoding_iter = 1
req.py_kv_transfer_start_time = None
req.py_kv_transfer_peer_wait_start = None
req.py_kv_transfer_timed_out = False
first_gen_tokens = req.context_phase_params.first_gen_tokens
ctx_draft_tokens = req.context_phase_params.draft_tokens
Expand Down Expand Up @@ -7446,7 +7496,9 @@ def kv_connector_request_finished(req: LlmRequest):
self.kv_cache_transceiver.respond_and_send_async(req)

if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None:
req.py_kv_transfer_start_time = time.monotonic()
transfer_start = time.monotonic()
req.py_kv_transfer_start_time = transfer_start
req.py_kv_transfer_peer_wait_start = transfer_start

if self.kv_connector_manager:
if not self.disable_overlap_scheduler:
Expand Down Expand Up @@ -7546,6 +7598,7 @@ def _check_disagg_ctx_cache_transfer_status(self, atLeastNum: int = 0):
# cancellation is disabled: a queued transfer that can be
# cancelled is immediately released from the async manager.
request.py_kv_transfer_start_time = None
request.py_kv_transfer_peer_wait_start = None
request.state = LlmRequestState.DISAGG_CONTEXT_COMPLETE
self._end_transfer_and_maybe_terminate(request)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,13 @@ context_servers:
print_iter_log: true
cache_transceiver_config:
backend: DEFAULT
max_tokens_in_buffer: 16384
# DisaggTransferAdmissionController spends max_tokens_in_buffer /
# tokens_per_block as an aggregate budget for concurrent generation KV
# transfers, so an arena-sized value (~one sequence) admits ~2 of the 512
# requests this test offers at 8K ISL and the rest die on the router's 180s
# timeout. Sized to the batch declared above (max_batch_size * 8K ISL);
# this model runs the Python transceiver, which allocates no such arena.
max_tokens_in_buffer: 1048576
generation_servers:
num_instances: 1
tensor_parallel_size: 1
Expand Down Expand Up @@ -47,4 +53,6 @@ generation_servers:
print_iter_log: true
cache_transceiver_config:
backend: DEFAULT
max_tokens_in_buffer: 16384
# Kept in step with the context server above: the generation side is the one
# whose admission gate defers the transfers.
max_tokens_in_buffer: 1048576
1 change: 0 additions & 1 deletion tests/integration/test_lists/waives.txt
Original file line number Diff line number Diff line change
Expand Up @@ -108,7 +108,6 @@ disaggregated/test_disaggregated.py::test_disaggregated_genbs1[TinyLlama-1.1B-Ch
disaggregated/test_disaggregated.py::test_disaggregated_qwen3_32b_fp8[Qwen3/Qwen3-32B-FP8] SKIP (https://nvbugs/6566734)
disaggregated/test_disaggregated.py::test_disaggregated_stress_test[input8k-output1k-conc512-deepseek_r1_v2_fp4_stress] SKIP (https://nvbugs/6621358)
disaggregated/test_disaggregated.py::test_disaggregated_stress_test[input8k-output1k-conc512-gpt_oss_120b_eagle_triton_stress] SKIP (https://nvbugs/6621362)
disaggregated/test_disaggregated.py::test_disaggregated_stress_test[input8k-output1k-conc512-qwen3_5_4b_fp8_stress] SKIP (https://nvbugs/6621362)
disaggregated/test_workers.py::test_workers_conversation_router[TinyLlama-1.1B-Chat-v1.0] SKIP (https://nvbugs/6162322)
disaggregated/test_workers.py::test_workers_kv_cache_aware_router_deepseek_v3_lite_bf16[DeepSeek-V3-Lite-bf16] SKIP (https://nvbugs/6162322)
disaggregated/test_workers.py::test_workers_kv_cache_aware_router_eviction[TinyLlama-1.1B-Chat-v1.0] SKIP (https://nvbugs/6162322)
Expand Down
Loading
Loading