From 20570013f046f8e5963227f3709f9ef1567f0947 Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Fri, 17 Jul 2026 16:36:39 -0700 Subject: [PATCH 1/7] [NVBUG 6312828][test] instrument disaggregated admission telemetry Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- .../batch_manager/baseTransBuffer.h | 5 + .../batch_manager/dataTransceiver.cpp | 143 +- scripts/disagg_admission_telemetry.py | 1494 +++++++++++++++++ .../_torch/disaggregation/native/transfer.py | 35 +- .../_torch/disaggregation/transceiver.py | 55 + tensorrt_llm/_torch/pyexecutor/py_executor.py | 223 ++- ...tx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml | 2 +- .../_torch/executor/test_py_executor.py | 189 +++ .../test_transceiver_bounded_polling.py | 99 +- .../tools/test_disagg_admission_telemetry.py | 298 ++++ 10 files changed, 2538 insertions(+), 5 deletions(-) create mode 100644 scripts/disagg_admission_telemetry.py create mode 100644 tests/unittest/tools/test_disagg_admission_telemetry.py diff --git a/cpp/tensorrt_llm/batch_manager/baseTransBuffer.h b/cpp/tensorrt_llm/batch_manager/baseTransBuffer.h index 88585818d4e5..df5b8db29c65 100644 --- a/cpp/tensorrt_llm/batch_manager/baseTransBuffer.h +++ b/cpp/tensorrt_llm/batch_manager/baseTransBuffer.h @@ -112,6 +112,11 @@ class BufferIndexHolder return mHeld; } + [[nodiscard]] BaseTransBufferManager const* manager() const noexcept + { + return mMgr; + } + [[nodiscard]] bool isBoundTo(BaseTransBufferManager const& manager) const noexcept { return mMgr == &manager; diff --git a/cpp/tensorrt_llm/batch_manager/dataTransceiver.cpp b/cpp/tensorrt_llm/batch_manager/dataTransceiver.cpp index 0f8ded65613f..e64831a0e4dc 100644 --- a/cpp/tensorrt_llm/batch_manager/dataTransceiver.cpp +++ b/cpp/tensorrt_llm/batch_manager/dataTransceiver.cpp @@ -31,11 +31,14 @@ #include "tensorrt_llm/runtime/utils/mpiUtils.h" #include #include +#include #include #include #include +#include #include #include +#include #include #include @@ -44,6 +47,33 @@ namespace tensorrt_llm::batch_manager using BlockRange = tensorrt_llm::batch_manager::kv_cache_manager::BlockRange; +namespace +{ + +bool isDisaggTransferDiagnosticsEnabled() +{ + static bool const enabled = [] + { + auto const* value = std::getenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS"); + return value != nullptr && std::string_view{value} == "1"; + }(); + return enabled; +} + +double getSteadyClockTimeSeconds() +{ + using Seconds = std::chrono::duration; + return Seconds(std::chrono::steady_clock::now().time_since_epoch()).count(); +} + +std::mutex& getDisaggDiagnosticsLogMutex() +{ + static std::mutex mutex; + return mutex; +} + +} // namespace + std::vector const& TransferSession::getConnections() const { return mConnections; @@ -160,16 +190,68 @@ bool TransferSession::releaseReservedRecvBuffer(BaseTransBufferManager const& ma { return false; } + auto const diagnosticsEnabled = isDisaggTransferDiagnosticsEnabled(); + auto const wasHeld = diagnosticsEnabled && holderIt->held(); + auto const bufferIndex = diagnosticsEnabled ? holderIt->index() : std::nullopt; + // Sample before release so a waiter awakened by release() cannot appear + // to acquire the slot before this event. This is a conservative upper + // bound on the true release-to-refill interval by the release call cost. + auto const releaseTime = wasHeld ? getSteadyClockTimeSeconds() : 0.0; holderIt->release(); + if (diagnosticsEnabled && mRequest != nullptr && wasHeld) + { + try + { + auto const contextRequestId = mRequest->getContextPhaseParams().has_value() + ? mRequest->getContextPhaseParams().value().getReqId() + : 0; + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO( + "[DISAGG_DIAG][receiver-slot] t=%.9f rank=%d action=released request=%zu context_request=%zu " + "manager=%p buffer=%d release_reason=formatter", + releaseTime, mpi::MpiComm::world().getRank(), mRequest->mRequestId, contextRequestId, + static_cast(&manager), bufferIndex.value_or(-1)); + } + catch (...) + { + // This method is noexcept; diagnostics must not alter release semantics. + } + } mReservedRecvBuffers.erase(holderIt); return true; } void TransferSession::releaseReservedRecvBuffers() noexcept { + auto const diagnosticsEnabled = isDisaggTransferDiagnosticsEnabled(); for (auto& holder : mReservedRecvBuffers) { + auto const wasHeld = diagnosticsEnabled && holder.held(); + auto const bufferIndex = diagnosticsEnabled ? holder.index() : std::nullopt; + auto const* manager = diagnosticsEnabled ? holder.manager() : nullptr; + // See releaseReservedRecvBuffer(): preserve causal ordering with a + // waiter that may acquire the slot as soon as release() notifies it. + auto const releaseTime = wasHeld ? getSteadyClockTimeSeconds() : 0.0; holder.release(); + if (diagnosticsEnabled && mRequest != nullptr && wasHeld) + { + try + { + auto const contextRequestId = mRequest->getContextPhaseParams().has_value() + ? mRequest->getContextPhaseParams().value().getReqId() + : 0; + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO( + "[DISAGG_DIAG][receiver-slot] t=%.9f rank=%d action=released request=%zu " + "context_request=%zu manager=%p buffer=%d release_reason=session", + releaseTime, mpi::MpiComm::world().getRank(), mRequest->mRequestId, contextRequestId, + static_cast(manager), bufferIndex.value_or(-1)); + } + catch (...) + { + // This method is noexcept; diagnostics must not alter release semantics. + } + } } mReservedRecvBuffers.clear(); } @@ -1268,10 +1350,37 @@ class CacheReceiver::Impl auto const& managers = agentConnectionManager->getCacheTransBufferManagers(); recvHolders.reserve(managers.size()); cacheBufferIds.reserve(managers.size()); - for (auto& cacheTransBufferManager : managers) + auto const diagnosticsEnabled = isDisaggTransferDiagnosticsEnabled(); + for (size_t managerIdx = 0; managerIdx < managers.size(); ++managerIdx) { + auto* cacheTransBufferManager = managers[managerIdx]; + auto const waitStart + = diagnosticsEnabled ? std::chrono::steady_clock::now() : std::chrono::steady_clock::time_point{}; auto rawIdx = cacheTransBufferManager->assignBufferIndexForRecv(bufferCancel); recvHolders.emplace_back(*cacheTransBufferManager, rawIdx, /*isRecv=*/true); + if (diagnosticsEnabled) + { + try + { + using Milliseconds = std::chrono::duration; + using Seconds = std::chrono::duration; + auto const acquiredTime = std::chrono::steady_clock::now(); + auto const waitMs = Milliseconds(acquiredTime - waitStart).count(); + auto const* action = rawIdx.has_value() ? "acquired" : "not-acquired"; + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO( + "[DISAGG_DIAG][receiver-slot] t=%.9f wait_start_t=%.9f rank=%d action=%s " + "request=%zu context_request=%zu manager_index=%zu manager=%p buffer=%d wait_ms=%.3f", + Seconds(acquiredTime.time_since_epoch()).count(), + Seconds(waitStart.time_since_epoch()).count(), mpi::MpiComm::world().getRank(), action, + llmRequest.mRequestId, requestId, managerIdx, + static_cast(cacheTransBufferManager), rawIdx.value_or(-1), waitMs); + } + catch (...) + { + // Diagnostics must not alter transfer semantics. + } + } if (rawIdx.has_value()) { cacheBufferIds.push_back(static_cast(rawIdx.value())); @@ -1686,6 +1795,22 @@ class CacheReceiver::Impl session->poisonReservedRecvBuffers(); } llmRequest.setKvCacheTransferEnd(LlmRequest::getSteadyClockNow()); + if (isDisaggTransferDiagnosticsEnabled()) + { + try + { + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO( + "[DISAGG_DIAG][receiver-transfer] t=%.9f rank=%d action=failed request=%zu " + "context_request=%zu phase=%s", + getSteadyClockTimeSeconds(), mpi::MpiComm::world().getRank(), requestId, contextRequestId, + phase); + } + catch (...) + { + // Preserve the original transfer exception if diagnostics fail. + } + } TLLM_LOG_ERROR("KV cache receive request %zu, context request %zu failed in phase=%s: %s", requestId, contextRequestId, phase, err.what()); throw; @@ -1697,6 +1822,22 @@ class CacheReceiver::Impl session->poisonReservedRecvBuffers(); } llmRequest.setKvCacheTransferEnd(LlmRequest::getSteadyClockNow()); + if (isDisaggTransferDiagnosticsEnabled()) + { + try + { + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO( + "[DISAGG_DIAG][receiver-transfer] t=%.9f rank=%d action=failed request=%zu " + "context_request=%zu phase=%s", + getSteadyClockTimeSeconds(), mpi::MpiComm::world().getRank(), requestId, contextRequestId, + phase); + } + catch (...) + { + // Preserve the original transfer exception if diagnostics fail. + } + } TLLM_LOG_ERROR( "KV cache receive request %zu, context request %zu failed in phase=%s with an unknown " "exception", diff --git a/scripts/disagg_admission_telemetry.py b/scripts/disagg_admission_telemetry.py new file mode 100644 index 000000000000..8c2a391629df --- /dev/null +++ b/scripts/disagg_admission_telemetry.py @@ -0,0 +1,1494 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Analyze opt-in disaggregated KV-transfer admission diagnostics. + +The analyzer intentionally depends only on the Python standard library so it +can run directly against CI log artifacts without importing TensorRT-LLM. +Malformed, incomplete, and unmatched diagnostic events are ignored. +""" + +from __future__ import annotations + +import argparse +import json +import math +import re +from collections import Counter, defaultdict, deque +from dataclasses import dataclass +from pathlib import Path +from statistics import fmean +from typing import Iterable, Sequence + +_CATEGORY_PATTERN = re.compile(r"\[DISAGG_DIAG\]\[([^]]+)]") +_FIELD_PATTERN = re.compile(r"([A-Za-z_][A-Za-z0-9_]*)=([^\s]+)") +_RANK_PATTERN = re.compile(r"\[RANK\s+(\d+)]") + + +@dataclass(frozen=True) +class DiagnosticEvent: + """One parsed ``[DISAGG_DIAG]`` event.""" + + category: str + time_s: float + rank: str + fields: dict[str, str] + + +@dataclass(frozen=True) +class Admission: + """An admission decision usable by the offline model.""" + + time_s: float + sequence: str | None + admitted: int + deferred: int + budget_blocks: float | None + active_blocks: float | None + candidate_requests: tuple[tuple[str, float], ...] + admitted_requests: tuple[str, ...] + deferred_requests: tuple[str, ...] + + +@dataclass(frozen=True) +class Decision: + """A lightweight admission-controller invocation.""" + + time_s: float + sequence: str | None + admitted: int + deferred: int + budget_blocks: float | None + + +@dataclass(frozen=True) +class PointEvent: + """A timestamp associated with a request.""" + + time_s: float + request: str + service_start_s: float | None = None + logged_time_s: float | None = None + call_ms: float | None = None + + +@dataclass(frozen=True) +class ServiceInterval: + """A completed request service interval.""" + + request: str + start_s: float + end_s: float + blocks: float | None + start_kind: str + end_kind: str + + +@dataclass(frozen=True) +class SlotInterval: + """A matched receiver-slot acquisition and release.""" + + request: str + manager: str + buffer: str + manager_index: str | None + start_s: float + end_s: float + wait_ms: float | None + + +@dataclass(frozen=True) +class ReleasePoint: + """A point where transfer capacity may be reusable.""" + + time_s: float + request: str + source: str + + +def parse_diagnostic_line(line: str) -> DiagnosticEvent | None: + """Parse one diagnostic line, returning ``None`` when it is unusable. + + Args: + line: An arbitrary application log line. + + Returns: + The parsed event when the category, timestamp, and rank are valid. + """ + category_match = _CATEGORY_PATTERN.search(line) + if category_match is None: + return None + + fields = dict(_FIELD_PATTERN.findall(line)) + event_time = _as_float(fields.get("t")) + if event_time is None: + return None + + rank = fields.get("rank") + if rank is None: + rank_match = _RANK_PATTERN.search(line) + rank = rank_match.group(1) if rank_match is not None else "unknown" + return DiagnosticEvent(category_match.group(1), event_time, rank, fields) + + +def read_diagnostic_events(paths: Iterable[str | Path]) -> list[DiagnosticEvent]: + """Read parseable diagnostics from one or more log paths. + + Args: + paths: Text log paths. Unreadable paths are skipped. + + Returns: + Parsed events in file/line order. + """ + events: list[DiagnosticEvent] = [] + for path_like in paths: + try: + with Path(path_like).open(errors="replace") as log_file: + for line in log_file: + event = parse_diagnostic_line(line) + if event is not None: + events.append(event) + except OSError: + continue + return events + + +def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: + """Calculate admission-window measurements from parsed events. + + The reported throughput is completed blocks divided by the union of + request-service intervals on each rank. This avoids double-counting time + when transfers overlap. The shadow multiplier is observational only and + never changes runtime admission or physical memory allocation. + + Args: + events: Parsed diagnostic events. + + Returns: + A JSON-serializable analysis dictionary. + """ + sorted_events = sorted(events, key=lambda event: (event.rank, event.time_s)) + category_counts = Counter(event.category for event in sorted_events) + events_by_rank: dict[str, list[DiagnosticEvent]] = defaultdict(list) + for event in sorted_events: + events_by_rank[event.rank].append(event) + + global_blocks = _collect_global_request_blocks(sorted_events) + ranks: dict[str, object] = {} + aggregate_service_intervals: list[ServiceInterval] = [] + aggregate_selected_gaps: list[dict[str, object]] = [] + aggregate_gaps_by_source: dict[str, list[dict[str, object]]] = defaultdict(list) + aggregate_slot_refill_gaps: list[float] = [] + aggregate_progress_credits: list[float] = [] + aggregate_fixed_multipliers: list[float] = [] + aggregate_poll_durations_ms: list[float] = [] + aggregate_progress_poll_durations_ms: list[float] = [] + aggregate_no_progress_poll_durations_ms: list[float] = [] + aggregate_busy_s = 0.0 + aggregate_completed_blocks = 0.0 + + for rank in sorted(events_by_rank, key=_rank_sort_key): + rank_analysis, rank_intervals, selected_gaps = _analyze_rank( + events_by_rank[rank], global_blocks + ) + ranks[rank] = rank_analysis + aggregate_service_intervals.extend(rank_intervals) + aggregate_selected_gaps.extend(selected_gaps) + release_analysis = rank_analysis["release_to_admission"] + if isinstance(release_analysis, dict): + by_source = release_analysis["by_source"] + if isinstance(by_source, dict): + for source, source_analysis in by_source.items(): + if isinstance(source_analysis, dict): + aggregate_gaps_by_source[source].extend(source_analysis["samples"]) + + service = rank_analysis["service"] + if isinstance(service, dict): + aggregate_busy_s += float(service["busy_s"]) + aggregate_completed_blocks += float(service["completed_blocks"]) + receiver_slots = rank_analysis["receiver_slots"] + if isinstance(receiver_slots, dict): + aggregate_slot_refill_gaps.extend(receiver_slots["backlog_refill_gap_samples_s"]) + progress = rank_analysis["linear_progress_credit"] + if isinstance(progress, dict): + aggregate_progress_credits.extend(progress["credit_samples_blocks"]) + counterfactual = rank_analysis["fixed_multiplier_counterfactual"] + if isinstance(counterfactual, dict): + aggregate_fixed_multipliers.extend( + counterfactual["next_deferred_required_multiplier_samples"] + ) + status_poll = rank_analysis["status_poll"] + if isinstance(status_poll, dict): + aggregate_poll_durations_ms.extend(status_poll["duration_samples_ms"]) + aggregate_progress_poll_durations_ms.extend(status_poll["progress_duration_samples_ms"]) + aggregate_no_progress_poll_durations_ms.extend( + status_poll["no_progress_duration_samples_ms"] + ) + + aggregate_throughput = _safe_ratio(aggregate_completed_blocks, aggregate_busy_s) + selected_decision_gaps = [ + float(sample["decision_gap_s"]) + for sample in aggregate_selected_gaps + if sample.get("decision_gap_s") is not None + ] + selected_successful_admission_gaps = [ + float(sample["successful_admission_gap_s"]) + for sample in aggregate_selected_gaps + if sample.get("successful_admission_gap_s") is not None + ] + selected_refill_gaps = [ + float(sample["refill_gap_s"]) + for sample in aggregate_selected_gaps + if sample.get("refill_gap_s") is not None + ] + shadow_samples = [ + float(sample["shadow_multiplier"]) + for sample in aggregate_selected_gaps + if sample.get("shadow_multiplier") is not None + ] + aggregate_release_bounds = { + source: { + "decision_gap_s": _summary([float(sample["decision_gap_s"]) for sample in samples]), + "successful_admission_gap_s": _summary( + [ + float(sample["successful_admission_gap_s"]) + for sample in samples + if sample.get("successful_admission_gap_s") is not None + ] + ), + "refill_gap_s": _summary( + [ + float(sample["refill_gap_s"]) + for sample in samples + if sample.get("refill_gap_s") is not None + ] + ), + "shadow_multiplier": _summary( + [ + float(sample["shadow_multiplier"]) + for sample in samples + if sample.get("shadow_multiplier") is not None + ] + ), + } + for source, samples in sorted(aggregate_gaps_by_source.items()) + } + + return { + "schema_version": 1, + "parsed_event_count": len(sorted_events), + "event_counts": dict(sorted(category_counts.items())), + "ranks": ranks, + "aggregate": { + "completed_service_intervals": len(aggregate_service_intervals), + "completed_blocks": aggregate_completed_blocks, + "busy_rank_seconds": aggregate_busy_s, + "throughput_blocks_per_s": aggregate_throughput, + "service_latency_s": _summary( + [interval.end_s - interval.start_s for interval in aggregate_service_intervals] + ), + "selected_physical_release_to_next_decision_gap_s": _summary(selected_decision_gaps), + "selected_physical_release_to_successful_admission_gap_s": _summary( + selected_successful_admission_gaps + ), + "selected_physical_release_to_refill_gap_s": _summary(selected_refill_gaps), + "release_bounds_by_source": aggregate_release_bounds, + "receiver_slot_refill_gap_s": _summary(aggregate_slot_refill_gaps), + "selected_physical_shadow_multiplier": _summary(shadow_samples), + "next_deferred_required_fixed_multiplier": _summary(aggregate_fixed_multipliers), + "linear_progress_credit_blocks": _summary(aggregate_progress_credits), + "status_poll": { + "duration_ms": _summary(aggregate_poll_durations_ms), + "progress_duration_ms": _summary(aggregate_progress_poll_durations_ms), + "no_progress_duration_ms": _summary(aggregate_no_progress_poll_durations_ms), + }, + }, + "model": { + "shadow_multiplier": "1 + throughput_blocks_per_s * refill_gap_s / budget_blocks", + "fixed_multiplier_counterfactual": ( + "max(1, (active_blocks + FCFS_prefix_blocks) / budget_blocks)" + ), + "linear_progress_credit": ( + "sum(request_blocks * elapsed_service_s / realized_service_s)" + ), + "caveat": ( + "Retrospective service and progress use completed intervals; they are validation " + "estimates, not online remaining-work measurements. Python local-ready is a " + "rank-local bound and reap is scheduler-visible; runtime control requires " + "conservative cross-rank aggregation or global-ready semantics." + ), + }, + } + + +def analyze_log_paths(paths: Iterable[str | Path]) -> dict[str, object]: + """Read and analyze diagnostic log paths.""" + return analyze_events(read_diagnostic_events(paths)) + + +def _analyze_rank( + events: list[DiagnosticEvent], global_blocks: dict[str, float] +) -> tuple[dict[str, object], list[ServiceInterval], list[dict[str, object]]]: + admissions, request_blocks = _collect_admissions(events) + decisions = _collect_decisions(events, admissions) + unsuccessful_requests = _collect_unsuccessful_requests(events) + for request, blocks in global_blocks.items(): + request_blocks.setdefault(request, blocks) + + submits = _collect_points(events, "submit") + local_ready = _collect_points( + events, + "python-transfer", + action="local-ready", + excluded_requests=unsuccessful_requests, + completed_only=True, + ) + reaps = _collect_points( + events, + "reap", + excluded_requests=unsuccessful_requests, + completed_only=True, + ) + for event in events: + if event.category == "submit": + blocks = _as_float(event.fields.get("blocks")) + request = event.fields.get("request") + if request is not None and blocks is not None and blocks >= 0.0: + request_blocks[request] = blocks + elif event.category == "reap": + blocks = _as_float(event.fields.get("blocks")) + request = event.fields.get("request") + if request is not None and blocks is not None and blocks >= 0.0: + request_blocks.setdefault(request, blocks) + + raw_slot_intervals, unmatched_acquires, unmatched_releases = _match_slot_intervals(events) + slot_intervals = [ + interval for interval in raw_slot_intervals if interval.request not in unsuccessful_requests + ] + physical_service_intervals = _build_request_slot_intervals(slot_intervals, request_blocks) + physical_queue_samples = _submit_to_interval_start_gaps(submits, physical_service_intervals) + service_intervals = _build_service_intervals( + submits, local_ready, reaps, physical_service_intervals, request_blocks + ) + busy_s = _union_duration(service_intervals) + completed_blocks = sum(interval.blocks or 0.0 for interval in service_intervals) + throughput = _safe_ratio(completed_blocks, busy_s) + + release_points = { + "local-ready": [ + ReleasePoint(point.time_s, point.request, "local-ready") for point in local_ready + ], + "reap": [ReleasePoint(point.time_s, point.request, "reap") for point in reaps], + "receiver-slot": [ + ReleasePoint(interval.end_s, interval.request, "receiver-slot") + for interval in physical_service_intervals + ], + } + gaps_by_source = { + source: _match_release_gaps(points, decisions, admissions, submits, throughput) + for source, points in release_points.items() + } + selected_source = _select_release_source(release_points) + selected_gaps = gaps_by_source[selected_source] if selected_source is not None else [] + + slot_refill_gaps = _slot_refill_gaps( + raw_slot_intervals, + decisions, + admissions, + unsuccessful_requests, + ) + progress_samples = _linear_progress_credit(admissions, service_intervals) + fixed_multiplier_samples = _fixed_multiplier_counterfactual(admissions) + ready_to_reap_samples = _point_pair_gaps(local_ready, reaps) + submit_to_service_start_samples = _submit_to_service_start_gaps(submits, local_ready) + status_poll_samples = _status_poll_samples(events) + progress_poll_durations = [ + float(sample["duration_ms"]) for sample in status_poll_samples if sample["made_progress"] + ] + no_progress_poll_durations = [ + float(sample["duration_ms"]) + for sample in status_poll_samples + if not sample["made_progress"] + ] + slot_latencies = [interval.end_s - interval.start_s for interval in slot_intervals] + wait_samples = [interval.wait_ms for interval in slot_intervals if interval.wait_ms is not None] + + analysis: dict[str, object] = { + "admission": { + "invocations": len(decisions), + "detailed_snapshots": len(admissions), + "deferred_invocations": sum(decision.deferred > 0 for decision in decisions), + "successful_invocations": sum(decision.admitted > 0 for decision in decisions), + "admitted_requests": sum(decision.admitted for decision in decisions), + "max_deferred": max((decision.deferred for decision in decisions), default=0), + "budgets_blocks": sorted( + { + decision.budget_blocks + for decision in decisions + if decision.budget_blocks is not None + } + ), + }, + "service": { + "intervals": [_service_interval_json(interval) for interval in service_intervals], + "excluded_unsuccessful_requests": sorted(unsuccessful_requests), + "latency_s": _summary( + [interval.end_s - interval.start_s for interval in service_intervals] + ), + "busy_s": busy_s, + "completed_blocks": completed_blocks, + "throughput_blocks_per_s": throughput, + }, + "python_transfer": { + "submit_to_service_start_samples_s": [ + float(sample["gap_s"]) for sample in submit_to_service_start_samples + ], + "submit_to_service_start_s": _summary( + [float(sample["gap_s"]) for sample in submit_to_service_start_samples] + ), + "submit_to_service_start_pairs": submit_to_service_start_samples, + "ready_to_reap_samples_s": [float(sample["gap_s"]) for sample in ready_to_reap_samples], + "ready_to_reap_s": _summary( + [float(sample["gap_s"]) for sample in ready_to_reap_samples] + ), + "pairs": ready_to_reap_samples, + }, + "status_poll": { + "samples": status_poll_samples, + "duration_samples_ms": [float(sample["duration_ms"]) for sample in status_poll_samples], + "progress_duration_samples_ms": progress_poll_durations, + "no_progress_duration_samples_ms": no_progress_poll_durations, + "duration_ms": _summary( + [float(sample["duration_ms"]) for sample in status_poll_samples] + ), + "progress_duration_ms": _summary(progress_poll_durations), + "no_progress_duration_ms": _summary(no_progress_poll_durations), + }, + "receiver_slots": { + "submit_to_service_start_samples_s": [ + float(sample["gap_s"]) for sample in physical_queue_samples + ], + "submit_to_service_start_s": _summary( + [float(sample["gap_s"]) for sample in physical_queue_samples] + ), + "submit_to_service_start_pairs": physical_queue_samples, + "intervals": [_slot_interval_json(interval) for interval in slot_intervals], + "service_latency_s": _summary(slot_latencies), + "wait_ms": _summary(wait_samples), + "unmatched_acquisitions": unmatched_acquires, + "unmatched_releases": unmatched_releases, + "excluded_unsuccessful_intervals": len(raw_slot_intervals) - len(slot_intervals), + "backlog_refill_gap_samples_s": slot_refill_gaps, + "backlog_refill_gap_s": _summary(slot_refill_gaps), + }, + "release_to_admission": { + "selected_release_source": selected_source, + "selected_samples": selected_gaps, + "selected_decision_gap_s": _summary( + [ + float(sample["decision_gap_s"]) + for sample in selected_gaps + if sample.get("decision_gap_s") is not None + ] + ), + "selected_refill_gap_s": _summary( + [ + float(sample["refill_gap_s"]) + for sample in selected_gaps + if sample.get("refill_gap_s") is not None + ] + ), + "selected_successful_admission_gap_s": _summary( + [ + float(sample["successful_admission_gap_s"]) + for sample in selected_gaps + if sample.get("successful_admission_gap_s") is not None + ] + ), + "by_source": { + source: { + "samples": samples, + "decision_gap_s": _summary( + [ + float(sample["decision_gap_s"]) + for sample in samples + if sample.get("decision_gap_s") is not None + ] + ), + "refill_gap_s": _summary( + [ + float(sample["refill_gap_s"]) + for sample in samples + if sample.get("refill_gap_s") is not None + ] + ), + "successful_admission_gap_s": _summary( + [ + float(sample["successful_admission_gap_s"]) + for sample in samples + if sample.get("successful_admission_gap_s") is not None + ] + ), + } + for source, samples in gaps_by_source.items() + }, + }, + "shadow_multiplier": { + "fitted_source": selected_source, + "by_source": { + source: { + "samples": [ + sample["shadow_multiplier"] + for sample in samples + if sample.get("shadow_multiplier") is not None + ], + "summary": _summary( + [ + float(sample["shadow_multiplier"]) + for sample in samples + if sample.get("shadow_multiplier") is not None + ] + ), + } + for source, samples in gaps_by_source.items() + }, + "policy_note": ( + "Python local-ready is a rank-local idle-opportunity bound; reap is a " + "conservative scheduler-visible bound. An adaptive policy must aggregate " + "conservatively across ranks or use a global-ready signal." + ), + }, + "fixed_multiplier_counterfactual": { + "samples": fixed_multiplier_samples, + "next_deferred_required_multiplier_samples": [ + float(sample["next_deferred_required_multiplier"]) + for sample in fixed_multiplier_samples + ], + "next_deferred_required_multiplier": _summary( + [ + float(sample["next_deferred_required_multiplier"]) + for sample in fixed_multiplier_samples + ] + ), + "all_prefix_required_multiplier": _summary( + [ + float(prefix["required_multiplier"]) + for sample in fixed_multiplier_samples + for prefix in sample["prefixes"] + ] + ), + }, + "linear_progress_credit": { + "samples": progress_samples, + "credit_samples_blocks": [ + float(sample["estimated_progress_credit_blocks"]) for sample in progress_samples + ], + "credit_blocks": _summary( + [float(sample["estimated_progress_credit_blocks"]) for sample in progress_samples] + ), + "credit_fraction": _summary( + [float(sample["estimated_progress_fraction"]) for sample in progress_samples] + ), + }, + } + return analysis, service_intervals, selected_gaps + + +def _collect_admissions(events: list[DiagnosticEvent]) -> tuple[list[Admission], dict[str, float]]: + admissions: list[Admission] = [] + request_blocks: dict[str, float] = {} + for event in events: + if event.category != "admission": + continue + candidate_requests = _parse_request_blocks(event.fields.get("candidate_requests")) + admitted_request_blocks = _parse_request_blocks(event.fields.get("admitted_requests")) + deferred_request_blocks = _parse_request_blocks(event.fields.get("deferred_requests")) + for request, blocks in ( + candidate_requests + admitted_request_blocks + deferred_request_blocks + ): + request_blocks[request] = blocks + + admitted = _as_int(event.fields.get("admitted")) + deferred = _as_int(event.fields.get("deferred")) + if admitted is None: + admitted = len(admitted_request_blocks) + if deferred is None: + deferred = len(deferred_request_blocks) + if admitted < 0 or deferred < 0: + continue + budget = _as_float(event.fields.get("budget")) + if budget is not None and budget <= 0.0: + budget = None + active_blocks = _as_float(event.fields.get("active_blocks")) + admissions.append( + Admission( + time_s=event.time_s, + sequence=event.fields.get("sequence"), + admitted=admitted, + deferred=deferred, + budget_blocks=budget, + active_blocks=active_blocks, + candidate_requests=tuple(candidate_requests), + admitted_requests=tuple(request for request, _ in admitted_request_blocks), + deferred_requests=tuple(request for request, _ in deferred_request_blocks), + ) + ) + admissions.sort(key=lambda admission: admission.time_s) + return admissions, request_blocks + + +def _collect_decisions( + events: list[DiagnosticEvent], admissions: list[Admission] +) -> list[Decision]: + decisions: list[Decision] = [] + for event in events: + if event.category != "decision": + continue + admitted = _as_int(event.fields.get("admitted")) + deferred = _as_int(event.fields.get("deferred")) + if admitted is None or deferred is None or admitted < 0 or deferred < 0: + continue + budget = _as_float(event.fields.get("budget")) + if budget is not None and budget <= 0.0: + budget = None + decisions.append( + Decision( + time_s=event.time_s, + sequence=event.fields.get("sequence"), + admitted=admitted, + deferred=deferred, + budget_blocks=budget, + ) + ) + if not decisions: + decisions = [ + Decision( + time_s=admission.time_s, + sequence=admission.sequence, + admitted=admission.admitted, + deferred=admission.deferred, + budget_blocks=admission.budget_blocks, + ) + for admission in admissions + ] + return sorted(decisions, key=lambda decision: decision.time_s) + + +def _collect_global_request_blocks(events: list[DiagnosticEvent]) -> dict[str, float]: + blocks_by_request: dict[str, float] = {} + conflicts: set[str] = set() + for event in events: + pairs: list[tuple[str, float]] = [] + if event.category == "admission": + for field in ("candidate_requests", "admitted_requests", "deferred_requests"): + pairs.extend(_parse_request_blocks(event.fields.get(field))) + elif event.category in {"submit", "reap"}: + request = event.fields.get("request") + blocks = _as_float(event.fields.get("blocks")) + if request is not None and blocks is not None and blocks >= 0.0: + pairs.append((request, blocks)) + for request, blocks in pairs: + previous = blocks_by_request.get(request) + if previous is not None and previous != blocks: + conflicts.add(request) + else: + blocks_by_request[request] = blocks + for request in conflicts: + blocks_by_request.pop(request, None) + return blocks_by_request + + +def _collect_points( + events: list[DiagnosticEvent], + category: str, + action: str | None = None, + excluded_requests: set[str] | None = None, + completed_only: bool = False, +) -> list[PointEvent]: + excluded_requests = excluded_requests or set() + points: list[PointEvent] = [] + for event in events: + if event.category != category: + continue + if action is not None and event.fields.get("action") != action: + continue + request = event.fields.get("request") + if ( + request is not None + and request not in excluded_requests + and (not completed_only or _event_outcome(event) is not False) + ): + service_start = _as_float(event.fields.get("service_start_t")) + if service_start is not None and (service_start < 0.0 or service_start > event.time_s): + service_start = None + point_time = event.time_s + if category == "submit": + submit_start = _as_float(event.fields.get("submit_start_t")) + if ( + submit_start is not None + and submit_start >= 0.0 + and submit_start <= event.time_s + ): + point_time = submit_start + points.append( + PointEvent( + point_time, + request, + service_start, + event.time_s, + _as_float(event.fields.get("submit_call_ms")), + ) + ) + return sorted(points, key=lambda point: point.time_s) + + +def _collect_unsuccessful_requests(events: list[DiagnosticEvent]) -> set[str]: + return { + request + for event in events + if (request := event.fields.get("request")) is not None and _event_outcome(event) is False + } + + +def _event_outcome(event: DiagnosticEvent) -> bool | None: + outcome = event.fields.get("outcome", "").lower() + if outcome in {"completed", "complete", "success", "successful", "succeeded", "ok"}: + return True + if outcome in { + "failed", + "failure", + "error", + "cancelled", + "canceled", + "aborted", + "timeout", + "timed-out", + }: + return False + + action = event.fields.get("action", "").lower() + if event.category == "receiver-transfer" and action in { + "failed", + "cancelled", + "canceled", + "aborted", + "timeout", + }: + return False + + state = event.fields.get("state", "").upper() + if any(marker in state for marker in ("ERROR", "FAIL", "CANCEL", "TIMEOUT")): + return False + if "COMPLETE" in state: + return True + return None + + +def _status_poll_samples(events: list[DiagnosticEvent]) -> list[dict[str, object]]: + samples: list[dict[str, object]] = [] + for event in events: + if event.category != "status-poll": + continue + duration_ms = _as_float(event.fields.get("poll_call_ms")) + completed = _as_int(event.fields.get("completed")) + failed = _as_int(event.fields.get("failed")) + cancelled = _as_int(event.fields.get("cancelled")) + if ( + duration_ms is None + or duration_ms < 0.0 + or completed is None + or failed is None + or cancelled is None + or min(completed, failed, cancelled) < 0 + ): + continue + samples.append( + { + "t": event.time_s, + "poll_start_t": _as_float(event.fields.get("poll_start_t")), + "duration_ms": duration_ms, + "at_least_num": _as_int(event.fields.get("at_least_num")), + "tracked": _as_int(event.fields.get("tracked")), + "completed": completed, + "failed": failed, + "cancelled": cancelled, + "made_progress": completed + failed + cancelled > 0, + } + ) + return samples + + +def _match_slot_intervals( + events: list[DiagnosticEvent], +) -> tuple[list[SlotInterval], int, int]: + acquisitions: dict[tuple[str, str], deque[DiagnosticEvent]] = defaultdict(deque) + intervals: list[SlotInterval] = [] + unmatched_releases = 0 + for event in sorted(events, key=lambda item: item.time_s): + if event.category != "receiver-slot": + continue + action = event.fields.get("action") + manager = event.fields.get("manager") + buffer = event.fields.get("buffer") + if manager is None or buffer in (None, "-1"): + continue + key = (manager, buffer) + if action in {"acquire", "acquired"}: + acquisitions[key].append(event) + elif action in {"release", "released"}: + if not acquisitions[key]: + unmatched_releases += 1 + continue + acquired = acquisitions[key].popleft() + if acquired.time_s > event.time_s: + unmatched_releases += 1 + continue + request = acquired.fields.get("request") or event.fields.get("request") + if request is None: + continue + intervals.append( + SlotInterval( + request=request, + manager=manager, + buffer=buffer, + manager_index=acquired.fields.get("manager_index"), + start_s=acquired.time_s, + end_s=event.time_s, + wait_ms=_as_float(acquired.fields.get("wait_ms")), + ) + ) + unmatched_acquires = sum(len(queue) for queue in acquisitions.values()) + intervals.sort(key=lambda interval: (interval.start_s, interval.end_s)) + return intervals, unmatched_acquires, unmatched_releases + + +def _build_service_intervals( + submits: list[PointEvent], + local_ready: list[PointEvent], + reaps: list[PointEvent], + physical_intervals: list[ServiceInterval], + request_blocks: dict[str, float], +) -> list[ServiceInterval]: + ready_by_request: dict[str, list[PointEvent]] = defaultdict(list) + reap_by_request: dict[str, list[PointEvent]] = defaultdict(list) + for point in local_ready: + ready_by_request[point.request].append(point) + for point in reaps: + reap_by_request[point.request].append(point) + + # Receiver-slot timestamps directly measure the C++ physical service + # interval. Prefer them over Python submit/reap observations, which include + # different parts of the lifecycle and can exist for the same request. + intervals = list(physical_intervals) + requests_with_physical_interval = {interval.request for interval in physical_intervals} + for submit in submits: + if submit.request in requests_with_physical_interval: + continue + endpoint = _first_point_after(ready_by_request.get(submit.request, []), submit.time_s) + end_kind = "local-ready" + if endpoint is None: + endpoint = _first_point_after(reap_by_request.get(submit.request, []), submit.time_s) + end_kind = "reap" + if endpoint is None: + continue + service_start = ( + endpoint.service_start_s if endpoint.service_start_s is not None else submit.time_s + ) + intervals.append( + ServiceInterval( + request=submit.request, + start_s=service_start, + end_s=endpoint.time_s, + blocks=request_blocks.get(submit.request), + start_kind=( + "python-service-start" if endpoint.service_start_s is not None else "submit" + ), + end_kind=end_kind, + ) + ) + intervals.sort(key=lambda interval: (interval.start_s, interval.end_s, interval.request)) + return intervals + + +def _build_request_slot_intervals( + slots: list[SlotInterval], request_blocks: dict[str, float] +) -> list[ServiceInterval]: + slots_by_request: dict[str, list[SlotInterval]] = defaultdict(list) + for slot in slots: + slots_by_request[slot.request].append(slot) + intervals = [ + ServiceInterval( + request=request, + start_s=min(slot.start_s for slot in request_slots), + end_s=max(slot.end_s for slot in request_slots), + blocks=request_blocks.get(request), + start_kind="receiver-slot-acquired", + end_kind="receiver-slot-released", + ) + for request, request_slots in slots_by_request.items() + ] + return sorted( + intervals, key=lambda interval: (interval.start_s, interval.end_s, interval.request) + ) + + +def _first_point_after(points: list[PointEvent], start_s: float) -> PointEvent | None: + return next((point for point in points if point.time_s >= start_s), None) + + +def _point_pair_gaps(starts: list[PointEvent], ends: list[PointEvent]) -> list[dict[str, object]]: + ends_by_request: dict[str, deque[PointEvent]] = defaultdict(deque) + for point in ends: + ends_by_request[point.request].append(point) + samples: list[dict[str, object]] = [] + for start in starts: + candidates = ends_by_request[start.request] + while candidates and candidates[0].time_s < start.time_s: + candidates.popleft() + if not candidates: + continue + end = candidates.popleft() + samples.append( + { + "request": start.request, + "ready_t": start.time_s, + "reap_t": end.time_s, + "gap_s": end.time_s - start.time_s, + } + ) + return samples + + +def _submit_to_service_start_gaps( + submits: list[PointEvent], local_ready: list[PointEvent] +) -> list[dict[str, object]]: + submits_by_request: dict[str, list[PointEvent]] = defaultdict(list) + for submit in submits: + submits_by_request[submit.request].append(submit) + samples: list[dict[str, object]] = [] + for ready in local_ready: + if ready.service_start_s is None: + continue + submit = next( + ( + candidate + for candidate in reversed(submits_by_request[ready.request]) + if candidate.time_s <= ready.service_start_s + ), + None, + ) + if submit is None: + submit = next( + ( + candidate + for candidate in reversed(submits_by_request[ready.request]) + if candidate.time_s <= ready.time_s + ), + None, + ) + if submit is None: + continue + samples.append( + { + "request": ready.request, + "submit_t": submit.time_s, + "submit_return_t": submit.logged_time_s, + "submit_call_ms": submit.call_ms, + "service_start_t": ready.service_start_s, + "gap_s": ready.service_start_s - submit.time_s, + } + ) + return samples + + +def _submit_to_interval_start_gaps( + submits: list[PointEvent], intervals: list[ServiceInterval] +) -> list[dict[str, object]]: + submits_by_request: dict[str, list[PointEvent]] = defaultdict(list) + for submit in submits: + submits_by_request[submit.request].append(submit) + samples: list[dict[str, object]] = [] + for interval in intervals: + submit = next( + ( + candidate + for candidate in reversed(submits_by_request[interval.request]) + if candidate.time_s <= interval.start_s + ), + None, + ) + if submit is None: + submit = next( + ( + candidate + for candidate in reversed(submits_by_request[interval.request]) + if candidate.time_s <= interval.end_s + ), + None, + ) + if submit is None: + continue + samples.append( + { + "request": interval.request, + "submit_t": submit.time_s, + "submit_return_t": submit.logged_time_s, + "submit_call_ms": submit.call_ms, + "service_start_t": interval.start_s, + "gap_s": interval.start_s - submit.time_s, + } + ) + return samples + + +def _match_release_gaps( + releases: list[ReleasePoint], + decisions: list[Decision], + admissions: list[Admission], + submits: list[PointEvent], + throughput_blocks_per_s: float | None, +) -> list[dict[str, object]]: + samples: list[dict[str, object]] = [] + for release in sorted(releases, key=lambda point: point.time_s): + prior = _latest_backlog_signal(decisions, admissions, release.time_s) + if prior is None or prior[0] <= 0: + continue + next_decision = next( + (decision for decision in decisions if decision.time_s > release.time_s), + None, + ) + if next_decision is None: + continue + backlog_requests = _backlog_request_ids_at(admissions, release.time_s) + backlog_identity_unknown = not backlog_requests + successful_admission = None + matched_backlog_requests: set[str] = set() + for decision in decisions: + if decision.time_s <= release.time_s or decision.admitted <= 0: + continue + if backlog_identity_unknown: + successful_admission = decision + break + detailed_admission = _matching_admission(decision, admissions) + if detailed_admission is None: + backlog_identity_unknown = True + successful_admission = decision + break + matched = set(detailed_admission.admitted_requests).intersection(backlog_requests) + if matched: + successful_admission = decision + matched_backlog_requests = matched + break + refill = ( + _find_refill_submit( + submits, + successful_admission, + admissions, + matched_backlog_requests or None, + ) + if successful_admission is not None + else None + ) + budget = prior[1] or ( + successful_admission.budget_blocks + if successful_admission is not None + else next_decision.budget_blocks + ) + refill_gap = refill.time_s - release.time_s if refill is not None else None + shadow_multiplier = None + eligible_for_multiplier_fit = ( + not backlog_identity_unknown and bool(matched_backlog_requests) and refill is not None + ) + if ( + eligible_for_multiplier_fit + and throughput_blocks_per_s is not None + and refill_gap is not None + and budget is not None + and budget > 0.0 + ): + shadow_multiplier = 1.0 + throughput_blocks_per_s * refill_gap / budget + samples.append( + { + "release_source": release.source, + "release_request": release.request, + "release_t": release.time_s, + "backlog_request_ids": sorted(backlog_requests), + "backlog_identity_unknown": backlog_identity_unknown, + "matched_backlog_request_ids": sorted(matched_backlog_requests), + "eligible_for_multiplier_fit": eligible_for_multiplier_fit, + "decision_t": next_decision.time_s, + "decision_sequence": next_decision.sequence, + "decision_gap_s": next_decision.time_s - release.time_s, + "successful_admission_t": ( + successful_admission.time_s if successful_admission is not None else None + ), + "successful_admission_sequence": ( + successful_admission.sequence if successful_admission is not None else None + ), + "successful_admission_gap_s": ( + successful_admission.time_s - release.time_s + if successful_admission is not None + else None + ), + "refill_t": refill.time_s if refill is not None else None, + "refill_submit_return_t": (refill.logged_time_s if refill is not None else None), + "refill_submit_call_ms": (refill.call_ms if refill is not None else None), + "refill_request": refill.request if refill is not None else None, + "refill_gap_s": refill_gap, + "budget_blocks": budget, + "throughput_blocks_per_s": throughput_blocks_per_s, + "shadow_multiplier": shadow_multiplier, + } + ) + return samples + + +def _find_refill_submit( + submits: list[PointEvent], + decision: Decision, + admissions: list[Admission], + required_requests: set[str] | None = None, +) -> PointEvent | None: + candidates = [submit for submit in submits if submit.time_s >= decision.time_s] + if required_requests: + return next( + (submit for submit in candidates if submit.request in required_requests), + None, + ) + admission = _matching_admission(decision, admissions) + if admission is not None and admission.admitted_requests: + admitted = set(admission.admitted_requests) + return next((submit for submit in candidates if submit.request in admitted), None) + return candidates[0] if candidates else None + + +def _backlog_request_ids_at(admissions: list[Admission], time_s: float) -> set[str]: + admission = next( + (candidate for candidate in reversed(admissions) if candidate.time_s <= time_s), + None, + ) + if admission is None or admission.deferred <= 0: + return set() + return set(admission.deferred_requests) + + +def _matching_admission(decision: Decision, admissions: list[Admission]) -> Admission | None: + if decision.sequence is not None: + match = next( + (admission for admission in admissions if admission.sequence == decision.sequence), + None, + ) + if match is not None: + return match + return next( + ( + admission + for admission in admissions + if math.isclose( + admission.time_s, + decision.time_s, + rel_tol=0.0, + abs_tol=1e-9, + ) + ), + None, + ) + + +def _latest_backlog_signal( + decisions: list[Decision], admissions: list[Admission], time_s: float +) -> tuple[int, float | None] | None: + signals = [ + (decision.time_s, 1, decision.deferred, decision.budget_blocks) + for decision in decisions + if decision.time_s <= time_s + ] + signals.extend( + (admission.time_s, 0, admission.deferred, admission.budget_blocks) + for admission in admissions + if admission.time_s <= time_s + ) + if not signals: + return None + _, _, deferred, budget = max(signals, key=lambda signal: (signal[0], signal[1])) + return deferred, budget + + +def _select_release_source(release_points: dict[str, list[ReleasePoint]]) -> str | None: + # Only the C++ path has a directly observed physical release signal. + # Python local-ready and consensus reap are complementary bounds, so the + # report intentionally does not collapse them into one selected source. + return "receiver-slot" if release_points["receiver-slot"] else None + + +def _slot_refill_gaps( + intervals: list[SlotInterval], + decisions: list[Decision], + admissions: list[Admission], + excluded_requests: set[str], +) -> list[float]: + by_slot: dict[tuple[str, str], list[SlotInterval]] = defaultdict(list) + for interval in intervals: + by_slot[(interval.manager, interval.buffer)].append(interval) + gaps: list[float] = [] + for slot_intervals in by_slot.values(): + slot_intervals.sort(key=lambda interval: interval.start_s) + for current, following in zip(slot_intervals, slot_intervals[1:]): + if current.request in excluded_requests or following.request in excluded_requests: + continue + prior = _latest_backlog_signal(decisions, admissions, current.end_s) + if prior is not None and prior[0] > 0 and following.start_s >= current.end_s: + gaps.append(following.start_s - current.end_s) + return gaps + + +def _fixed_multiplier_counterfactual( + admissions: list[Admission], +) -> list[dict[str, object]]: + samples: list[dict[str, object]] = [] + for admission in admissions: + if ( + admission.deferred <= 0 + or admission.active_blocks is None + or admission.budget_blocks is None + or not admission.candidate_requests + ): + continue + + deferred_requests = set(admission.deferred_requests) + first_deferred_index = next( + ( + index + for index, (request, _) in enumerate(admission.candidate_requests) + if request in deferred_requests + ), + None, + ) + if first_deferred_index is None and admission.admitted < len(admission.candidate_requests): + first_deferred_index = admission.admitted + if first_deferred_index is None: + continue + + prefix_blocks = 0.0 + prefixes: list[dict[str, object]] = [] + for index, (request, blocks) in enumerate(admission.candidate_requests): + prefix_blocks += blocks + required_multiplier = max( + 1.0, + (admission.active_blocks + prefix_blocks) / admission.budget_blocks, + ) + prefixes.append( + { + "prefix_length": index + 1, + "last_request": request, + "prefix_blocks": prefix_blocks, + "required_multiplier": required_multiplier, + "minimum_integer_multiplier": math.ceil(required_multiplier), + "observed_status": ( + "deferred" + if request in deferred_requests or index >= admission.admitted + else "admitted" + ), + } + ) + + next_deferred = prefixes[first_deferred_index] + samples.append( + { + "decision_t": admission.time_s, + "active_blocks": admission.active_blocks, + "budget_blocks": admission.budget_blocks, + "next_deferred_request": next_deferred["last_request"], + "next_deferred_prefix_blocks": next_deferred["prefix_blocks"], + "next_deferred_required_multiplier": next_deferred["required_multiplier"], + "next_deferred_minimum_integer_multiplier": next_deferred[ + "minimum_integer_multiplier" + ], + "prefixes": prefixes, + } + ) + return samples + + +def _linear_progress_credit( + admissions: list[Admission], intervals: list[ServiceInterval] +) -> list[dict[str, object]]: + samples: list[dict[str, object]] = [] + for admission in admissions: + if admission.deferred <= 0: + continue + in_progress = [ + interval + for interval in intervals + if interval.blocks is not None + and interval.start_s <= admission.time_s < interval.end_s + and interval.end_s > interval.start_s + ] + if not in_progress: + continue + original_blocks = sum(interval.blocks or 0.0 for interval in in_progress) + credit = sum( + (interval.blocks or 0.0) + * (admission.time_s - interval.start_s) + / (interval.end_s - interval.start_s) + for interval in in_progress + ) + fraction = _safe_ratio(credit, original_blocks) or 0.0 + samples.append( + { + "decision_t": admission.time_s, + "in_progress_requests": len(in_progress), + "logged_active_blocks": admission.active_blocks, + "original_in_progress_blocks": original_blocks, + "estimated_progress_credit_blocks": credit, + "estimated_remaining_blocks": original_blocks - credit, + "estimated_progress_fraction": fraction, + } + ) + return samples + + +def _union_duration(intervals: list[ServiceInterval]) -> float: + ranges = sorted((interval.start_s, interval.end_s) for interval in intervals) + if not ranges: + return 0.0 + merged: list[list[float]] = [] + for start_s, end_s in ranges: + if end_s < start_s: + continue + if not merged or start_s > merged[-1][1]: + merged.append([start_s, end_s]) + else: + merged[-1][1] = max(merged[-1][1], end_s) + return sum(end_s - start_s for start_s, end_s in merged) + + +def _summary(values: Iterable[float | None]) -> dict[str, float | int | None]: + samples = sorted(value for value in values if value is not None and math.isfinite(value)) + if not samples: + return { + "count": 0, + "min": None, + "mean": None, + "p50": None, + "p95": None, + "p99": None, + "max": None, + } + return { + "count": len(samples), + "min": samples[0], + "mean": fmean(samples), + "p50": _percentile(samples, 0.50), + "p95": _percentile(samples, 0.95), + "p99": _percentile(samples, 0.99), + "max": samples[-1], + } + + +def _percentile(sorted_values: list[float], quantile: float) -> float: + position = (len(sorted_values) - 1) * quantile + lower = math.floor(position) + upper = math.ceil(position) + if lower == upper: + return sorted_values[lower] + weight = position - lower + return sorted_values[lower] * (1.0 - weight) + sorted_values[upper] * weight + + +def _parse_request_blocks(value: str | None) -> list[tuple[str, float]]: + if value in (None, "", "-"): + return [] + pairs: list[tuple[str, float]] = [] + for item in value.split(","): + request, separator, blocks_text = item.partition(":") + blocks = _as_float(blocks_text) if separator else None + if request and blocks is not None and blocks >= 0.0: + pairs.append((request, blocks)) + return pairs + + +def _as_float(value: str | None) -> float | None: + if value is None: + return None + try: + parsed = float(value) + except ValueError: + return None + return parsed if math.isfinite(parsed) else None + + +def _as_int(value: str | None) -> int | None: + if value is None: + return None + try: + return int(value) + except ValueError: + return None + + +def _safe_ratio(numerator: float, denominator: float) -> float | None: + return numerator / denominator if denominator > 0.0 else None + + +def _service_interval_json(interval: ServiceInterval) -> dict[str, object]: + return { + "request": interval.request, + "start_t": interval.start_s, + "end_t": interval.end_s, + "latency_s": interval.end_s - interval.start_s, + "blocks": interval.blocks, + "start_kind": interval.start_kind, + "end_kind": interval.end_kind, + } + + +def _slot_interval_json(interval: SlotInterval) -> dict[str, object]: + return { + "request": interval.request, + "manager": interval.manager, + "manager_index": interval.manager_index, + "buffer": interval.buffer, + "acquired_t": interval.start_s, + "released_t": interval.end_s, + "service_s": interval.end_s - interval.start_s, + "wait_ms": interval.wait_ms, + } + + +def _rank_sort_key(rank: str) -> tuple[int, int | str]: + try: + return 0, int(rank) + except ValueError: + return 1, rank + + +def _build_argument_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description="Analyze [DISAGG_DIAG] admission and KV-transfer events." + ) + parser.add_argument("logs", nargs="+", help="Worker or preserved diagnostic log paths") + parser.add_argument("--indent", type=int, default=2, help="JSON indentation (default: 2)") + return parser + + +def main(argv: Sequence[str] | None = None) -> int: + """Run the command-line analyzer.""" + args = _build_argument_parser().parse_args(argv) + print(json.dumps(analyze_log_paths(args.logs), indent=args.indent, sort_keys=True)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 353444eeca6f..98e273e702d8 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -78,6 +78,7 @@ # Number of worker threads for KV transfer queues (default: 1) KV_TRANSFER_NUM_THREADS = int(os.environ.get("TRTLLM_KV_TRANSFER_NUM_THREADS", "1")) +_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED = os.getenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS") == "1" @dataclass @@ -1424,6 +1425,17 @@ def __init__( self._exception: Optional[Exception] = None self._aux_slot = aux_slot self._perf_timer = PerfTimer() if perf_log_manager.enabled else None + if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + self.transfer_start_time_s: Optional[float] = None + self.completion_time_s: Optional[float] = None + + def mark_transferring(self) -> None: + if ( + _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + and getattr(self, "transfer_start_time_s", None) is None + ): + self.transfer_start_time_s = tensorrt_llm.bindings.steady_clock_now().total_seconds() + self.status = TaskStatus.TRANSFERRING def fail(self, exc: Exception) -> None: self._exception = exc @@ -1431,6 +1443,11 @@ def fail(self, exc: Exception) -> None: self._event.set() def complete(self) -> None: + if ( + _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + and getattr(self, "completion_time_s", None) is None + ): + self.completion_time_s = tensorrt_llm.bindings.steady_clock_now().total_seconds() self.status = TaskStatus.TRANSFERRED self._event.set() @@ -1862,7 +1879,23 @@ def status(self) -> SessionStatus: def mark_transferring(self, slice_id: int): with self.lock: - self._kv_tasks[slice_id].status = TaskStatus.TRANSFERRING + self._kv_tasks[slice_id].mark_transferring() + + @property + def kv_transfer_start_time_s(self) -> Optional[float]: + start_times = [ + getattr(task, "transfer_start_time_s", None) + for task in self._kv_tasks + if getattr(task, "transfer_start_time_s", None) is not None + ] + return min(start_times) if start_times else None + + @property + def kv_ready_time_s(self) -> Optional[float]: + completion_times = [getattr(task, "completion_time_s", None) for task in self._kv_tasks] + if not completion_times or any(time_s is None for time_s in completion_times): + return None + return max(time_s for time_s in completion_times if time_s is not None) def receive(self, slice: KVSlice) -> None: if self.transfer_start_time is None: diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 2c257a39f731..4789d81258ff 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -1,3 +1,6 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + import os import time import uuid @@ -42,6 +45,12 @@ from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig from tensorrt_llm.mapping import Mapping +_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED = os.getenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS") == "1" + + +def _is_disagg_transfer_diagnostics_enabled() -> bool: + return _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + def _find_consensus_request_ids(request_ids_all_ranks, sync_size): frequency_map = defaultdict(int) @@ -102,6 +111,7 @@ def __init__( self._recv_sessions: Dict[int, RxSessionBase] = {} self._send_reqs = {} self._recv_reqs = {} + self._diagnostic_ready_rids = set() if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED else None self._wait_reqs = {} self._page_table = self._transfer_worker.page_table # _slice_num_bytes() is this rank's KV shard, so scale by tp_size to get the request total (kv_cache_size), @@ -671,6 +681,44 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): self._poll_gen_sessions_for_poll_interval(wait_num) local_completed, local_failed = self._collect_done(self._recv_sessions, self._recv_reqs) + diagnostic_ready_rids = getattr(self, "_diagnostic_ready_rids", None) + if _is_disagg_transfer_diagnostics_enabled(): + assert diagnostic_ready_rids is not None + for rid in local_completed: + if rid in diagnostic_ready_rids: + continue + session = self._recv_sessions[rid] + service_start_time = getattr(session, "kv_transfer_start_time_s", None) + ready_time = getattr(session, "kv_ready_time_s", None) + if ready_time is None: + diagnostic_ready_rids.add(rid) + logger.warning( + "[DISAGG_DIAG][python-transfer] " + f"rank={self._dist.rank} action=missing-native-timestamp " + f"request={self._recv_reqs[rid].py_request_id}" + ) + continue + req = self._recv_reqs[rid] + if service_start_time is not None: + req.py_kv_transfer_service_start_time_s = service_start_time + req.py_kv_transfer_ready_time_s = ready_time + diagnostic_ready_rids.add(rid) + service_ms = ( + (ready_time - service_start_time) * 1000 + if service_start_time is not None + else -1.0 + ) + diagnostic_service_start_time = ( + service_start_time if service_start_time is not None else -1.0 + ) + logger.info( + "[DISAGG_DIAG][python-transfer] " + f"t={ready_time:.9f} rank={self._dist.rank} " + f"action=local-ready request={req.py_request_id} " + f"bytes={getattr(req, 'py_kv_cache_xfer_bytes', 0)} " + f"service_start_t={diagnostic_service_start_time:.9f} " + f"service_ms={service_ms:.3f}" + ) to_process = self._build_to_process( self._recv_sessions, self._gen_consensus(local_completed + local_failed), @@ -710,6 +758,8 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): self._recv_sessions[rid].close() del self._recv_reqs[rid] del self._recv_sessions[rid] + if diagnostic_ready_rids is not None: + diagnostic_ready_rids.discard(rid) # Log gen-side transfer summary after consensus. if completed and os.getenv("TRTLLM_KVCACHE_TIME_OUTPUT_PATH"): @@ -737,12 +787,17 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): session.close() del self._recv_reqs[rid] del self._recv_sessions[rid] + if diagnostic_ready_rids is not None: + diagnostic_ready_rids.discard(rid) if failed: logger.warning( f"Disagg gen transfer FAILED rank={self._dist.rank} " f"rids={failed} gen_need_sync={self._gen_need_sync}" ) self._close_failed_sessions(self._recv_sessions, self._recv_reqs, failed) + if diagnostic_ready_rids is not None: + for rid in failed: + diagnostic_ready_rids.discard(rid) return completed, failed, cancelled_reqs diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index feb02319a3e3..0e79e03e3d06 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -114,6 +114,22 @@ def _stats_buffer_is_unbounded(max_stats_len: int) -> bool: # Default: "0" (only rank 0 prints, matching existing behavior). PROFILE_LOG_RANKS_ENV_VAR_NAME = "TLLM_PROFILE_LOG_RANKS" +_DISAGG_TRANSFER_DIAGNOSTICS_ENV = "TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS" +_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED = ( + os.getenv(_DISAGG_TRANSFER_DIAGNOSTICS_ENV) == "1") +_DISAGG_STATUS_POLL_LOG_THRESHOLD_MS = 1.0 + + +def _is_disagg_transfer_diagnostics_enabled() -> bool: + return _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED + + +def _format_disagg_diag_request_blocks( + request_blocks: Iterable[Tuple[int, int]]) -> str: + encoded = ",".join(f"{request_id}:{blocks}" + for request_id, blocks in request_blocks) + return encoded or "-" + class PPCommTag(IntEnum): """ @@ -3408,6 +3424,30 @@ def _uses_kv_manager_v2(self) -> bool: return isinstance(getattr(self, "kv_cache_manager", None), KVCacheManagerV2) + def _log_disagg_transfer_diagnostic(self, category: str, **fields) -> None: + if not _is_disagg_transfer_diagnostics_enabled(): + return + timestamp = fields.pop("t", None) + if timestamp is None: + timestamp = get_steady_clock_now_in_seconds() + rank = getattr(getattr(self, "dist", None), "rank", -1) + encoded_fields = " ".join(f"{key}={value}" + for key, value in fields.items()) + logger.info(f"[DISAGG_DIAG][{category}] t={timestamp:.9f} " + f"rank={rank} {encoded_fields}") + + @staticmethod + def _disagg_diag_request_id(request: LlmRequest) -> int: + return request.py_request_id + + def _disagg_diag_request_blocks( + self, requests: Iterable[LlmRequest], + controller: DisaggTransferAdmissionController + ) -> List[Tuple[int, int]]: + return [(self._disagg_diag_request_id(request), + controller._estimate_request_blocks(request)) + for request in requests] + def _apply_disagg_transfer_admission( self, fitting_disagg_gen_init_requests: List[LlmRequest] ) -> Tuple[List[LlmRequest], bool]: @@ -3424,6 +3464,76 @@ def _apply_disagg_transfer_admission( admission_result = controller.select(self.active_requests, fitting_disagg_gen_init_requests) + if _is_disagg_transfer_diagnostics_enabled(): + decision_time = get_steady_clock_now_in_seconds() + decision_sequence = getattr( + self, "_disagg_diag_admission_decision_sequence", 0) + 1 + self._disagg_diag_admission_decision_sequence = decision_sequence + active_requests = [ + request for request in self.active_requests + if request.is_disagg_generation_transmission_in_progress + ] + candidate_request_blocks = self._disagg_diag_request_blocks( + fitting_disagg_gen_init_requests, controller) + admitted_request_ids = { + self._disagg_diag_request_id(request) + for request in admission_result.admitted_requests + } + admitted_request_blocks = [ + request_block for request_block in candidate_request_blocks + if request_block[0] in admitted_request_ids + ] + deferred_request_blocks = [ + request_block for request_block in candidate_request_blocks + if request_block[0] not in admitted_request_ids + ] + active_request_blocks = self._disagg_diag_request_blocks( + active_requests, controller) + candidate_transfer_blocks = sum( + blocks for _, blocks in candidate_request_blocks) + deferred_transfer_blocks = sum( + blocks for _, blocks in deferred_request_blocks) + self._log_disagg_transfer_diagnostic( + "decision", + t=decision_time, + sequence=decision_sequence, + runtime=type(self.kv_cache_transceiver).__name__, + active_blocks=admission_result.active_transfer_blocks, + candidates=len(candidate_request_blocks), + candidate_blocks=candidate_transfer_blocks, + admitted=len(admitted_request_blocks), + admitted_blocks=admission_result.admitted_transfer_blocks, + deferred=admission_result.deferred_request_count, + deferred_blocks=deferred_transfer_blocks, + budget=controller.max_transfer_blocks) + snapshot = (tuple(active_request_blocks), + tuple(candidate_request_blocks), + tuple(admitted_request_blocks)) + last_snapshot = getattr(self, "_last_disagg_diag_admission", None) + if admission_result.admitted_requests or snapshot != last_snapshot: + self._log_disagg_transfer_diagnostic( + "admission", + t=decision_time, + sequence=decision_sequence, + runtime=type(self.kv_cache_transceiver).__name__, + active=len(active_request_blocks), + active_blocks=admission_result.active_transfer_blocks, + active_requests=_format_disagg_diag_request_blocks( + active_request_blocks), + candidates=len(candidate_request_blocks), + candidate_blocks=candidate_transfer_blocks, + candidate_requests=_format_disagg_diag_request_blocks( + candidate_request_blocks), + admitted=len(admitted_request_blocks), + admitted_blocks=(admission_result.admitted_transfer_blocks), + admitted_requests=_format_disagg_diag_request_blocks( + admitted_request_blocks), + deferred=admission_result.deferred_request_count, + deferred_blocks=deferred_transfer_blocks, + deferred_requests=_format_disagg_diag_request_blocks( + deferred_request_blocks), + budget=controller.max_transfer_blocks) + self._last_disagg_diag_admission = snapshot if admission_result.deferred_request_count > 0: logger.debug("Disagg transfer admission deferred " f"{admission_result.deferred_request_count} requests; " @@ -6039,8 +6149,27 @@ def _recv_disagg_gen_cache(self, new_gen_reqs): self._check_cache_transfer_errors("generation requests") return + diagnostics_enabled = _is_disagg_transfer_diagnostics_enabled() + controller = (self._get_disagg_transfer_admission_controller() + if diagnostics_enabled else None) for req in new_gen_reqs: + if diagnostics_enabled: + submit_start = get_steady_clock_now_in_seconds() self.kv_cache_transceiver.request_and_receive_async(req) + if diagnostics_enabled: + submit_end = get_steady_clock_now_in_seconds() + assert controller is not None + self._log_disagg_transfer_diagnostic( + "submit", + t=submit_end, + runtime=type(self.kv_cache_transceiver).__name__, + request=self._disagg_diag_request_id(req), + blocks=controller._estimate_request_blocks(req), + bytes=getattr(req, "py_kv_cache_xfer_bytes", 0), + submit_start_t=f"{submit_start:.9f}", + submit_call_ms=( + f"{(submit_end - submit_start) * 1000:.6f}"), + state=getattr(req.state, "name", str(req.state))) if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None: for req in new_gen_reqs: @@ -6190,14 +6319,106 @@ def _check_disagg_ctx_cache_transfer_status(self, atLeastNum: int = 0): @nvtx_range("_check_disagg_gen_cache_transfer_status") def _check_disagg_gen_cache_transfer_status(self, atLeastNum: int = 0): + diagnostics_enabled = _is_disagg_transfer_diagnostics_enabled() + tracked_requests = [] + poll_start = 0.0 + if diagnostics_enabled: + tracked_requests = [ + request for request in self.active_requests + if request.is_disagg_generation_transmission_in_progress + ] + poll_start = get_steady_clock_now_in_seconds() result = self.kv_cache_transceiver.check_gen_transfer_status(atLeastNum) + completed_count = 0 + failed_count = 0 + cancelled_count = 0 if isinstance(result, tuple): - _, _, cancelled_reqs = result + completed_reqs, failed_reqs, cancelled_reqs = result + completed_count = len(completed_reqs) + failed_count = len(failed_reqs) + cancelled_count = len(cancelled_reqs) user_canceled_set = set(self.canceled_req_ids) for req in cancelled_reqs: req_id = req.py_request_id if not req.is_child else req.parent_request_id if req_id not in user_canceled_set: req.state = LlmRequestState.DISAGG_TRANS_ERROR + else: + # The C++ binding reports progress by mutating request state and + # returns no outcome tuple. Derive the same poll counters from the + # before/after request snapshot so no-progress samples are valid + # for both transceiver runtimes. + for request in tracked_requests: + if request.is_disagg_generation_transmission_in_progress: + continue + if (request.state == + LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE): + completed_count += 1 + elif request.state == LlmRequestState.DISAGG_TRANS_ERROR: + failed_count += 1 + else: + cancelled_count += 1 + poll_end = 0.0 + if diagnostics_enabled: + poll_end = get_steady_clock_now_in_seconds() + poll_call_ms = (poll_end - poll_start) * 1000 + should_log_poll = ( + atLeastNum is None or atLeastNum > 0 + or poll_call_ms >= _DISAGG_STATUS_POLL_LOG_THRESHOLD_MS + or completed_count + failed_count + cancelled_count > 0) + if should_log_poll: + self._log_disagg_transfer_diagnostic( + "status-poll", + t=poll_end, + runtime=type(self.kv_cache_transceiver).__name__, + poll_start_t=f"{poll_start:.9f}", + poll_call_ms=f"{poll_call_ms:.6f}", + at_least_num=("all" if atLeastNum is None else atLeastNum), + tracked=len(tracked_requests), + completed=completed_count, + failed=failed_count, + cancelled=cancelled_count) + if tracked_requests: + reap_time = poll_end + controller = self._get_disagg_transfer_admission_controller() + offset = LlmRequest.global_steady_clock_offset + offset_seconds = (offset.total_seconds() + if offset is not None else 0.0) + for request in tracked_requests: + if request.is_disagg_generation_transmission_in_progress: + continue + request_state = getattr(request.state, "name", + str(request.state)) + if (request.state == + LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE): + outcome = "completed" + elif request.state == LlmRequestState.DISAGG_TRANS_ERROR: + outcome = "failed" + else: + outcome = "cancelled" + ready_time = getattr(request, "py_kv_transfer_ready_time_s", + 0.0) + if not ready_time: + transfer_end = getattr(request, "kv_cache_transfer_end", + None) + if transfer_end is not None: + ready_time = max( + 0.0, + transfer_end.total_seconds() - offset_seconds) + ready_to_reap_ms = (-1.0 if not ready_time else + (reap_time - ready_time) * 1000) + self._log_disagg_transfer_diagnostic( + "reap", + t=reap_time, + runtime=type(self.kv_cache_transceiver).__name__, + request=self._disagg_diag_request_id(request), + blocks=controller._estimate_request_blocks(request), + bytes=getattr(request, "py_kv_cache_xfer_bytes", + getattr(request, "kv_cache_size", 0)), + ready_t=f"{ready_time:.9f}", + ready_to_reap_ms=f"{ready_to_reap_ms:.6f}", + poll_call_ms=f"{(poll_end - poll_start) * 1000:.6f}", + outcome=outcome, + state=request_state) if not self._is_disagg_inflight_cancel_active(): self._check_cache_transfer_errors("generation requests") diff --git a/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_128k8k_con256_ctx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml b/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_128k8k_con256_ctx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml index 92e9e36230a3..6db032ae6c3a 100644 --- a/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_128k8k_con256_ctx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml +++ b/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_128k8k_con256_ctx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml @@ -36,7 +36,7 @@ environment: build_wheel: false work_dir: worker_env_var: TLLM_LOG_LEVEL=INFO TRTLLM_SERVER_DISABLE_GC=1 TRTLLM_WORKER_DISABLE_GC=1 TLLM_SPEC_DECODE_FORCE_NUM_ACCEPTED_TOKENS=1 - TRTLLM_ENABLE_PDL=1 ENROOT_ALLOW_DEV=yes + TRTLLM_ENABLE_PDL=1 TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS=1 ENROOT_ALLOW_DEV=yes server_env_var: TRTLLM_SERVER_DISABLE_GC=1 profiling: nsys_on: false diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index 101edffbe155..9ebd1236db30 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -21,6 +21,7 @@ import pytest from tensorrt_llm._torch.distributed.communicator import ReduceOp +from tensorrt_llm._torch.pyexecutor import py_executor as py_executor_module from tensorrt_llm._torch.pyexecutor.executor_request_queue import ( SHUTDOWN_REQUEST_ID, RequestQueueItem, @@ -441,6 +442,8 @@ def _make_disagg_transfer_request( def _clear_disagg_transfer_mode_env(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("TRTLLM_DISAGG_BENCHMARK_GEN_ONLY", raising=False) monkeypatch.delenv("TRTLLM_DISABLE_KV_CACHE_TRANSFER_OVERLAP", raising=False) + monkeypatch.delenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", raising=False) + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", False) @pytest.mark.usefixtures("_clear_disagg_transfer_mode_env") @@ -527,6 +530,46 @@ def test_apply_reverts_deferred_v2_allocations(self): assert wait_for_progress executor._revert_ctx_alloc.assert_called_once_with([candidate]) + def test_apply_emits_changed_admission_snapshot(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=3) + executor.kv_cache_transceiver = Mock() + executor._is_kv_manager_v2 = False + executor.active_requests = [_make_disagg_transfer_request(1, 32, in_progress=True)] + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + candidate = _make_disagg_transfer_request(2, 32) + + PyExecutor._apply_disagg_transfer_admission(executor, [candidate]) + PyExecutor._apply_disagg_transfer_admission(executor, [candidate]) + + assert log_info.call_count == 3 + decision_messages = [ + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][decision]" in call.args[0] + ] + assert len(decision_messages) == 2 + assert "sequence=1" in decision_messages[0] + assert "sequence=2" in decision_messages[1] + message = next( + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][admission]" in call.args[0] + ) + assert "[DISAGG_DIAG][admission]" in message + assert "sequence=1" in message + assert "rank=3" in message + assert "active_requests=1:1" in message + assert "candidate_requests=2:1" in message + assert "deferred_requests=2:1" in message + assert "budget=1" in message + def test_apply_missing_controller_preserves_candidates(self): executor = object.__new__(PyExecutor) executor.kv_cache_transceiver = Mock() @@ -616,6 +659,152 @@ def test_gen_transfer_status_polls_active_transfers(self): executor._check_disagg_gen_cache_transfer_status.assert_called_once_with(0) + def test_async_receive_emits_submit_work(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + timestamps = iter((10.0, 10.002)) + monkeypatch.setattr( + py_executor_module, "get_steady_clock_now_in_seconds", lambda: next(timestamps) + ) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor.kv_cache_transceiver = Mock(kv_transfer_timeout_ms=None) + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=64, tokens_per_block=32 + ) + executor._check_disagg_gen_cache_transfer_status = Mock() + request = _make_disagg_transfer_request(7, 64) + request.state = LlmRequestState.DISAGG_GENERATION_INIT + request.py_kv_cache_xfer_bytes = 4096 + + def mark_in_progress(req): + req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + + executor.kv_cache_transceiver.request_and_receive_async.side_effect = mark_in_progress + + PyExecutor._recv_disagg_gen_cache(executor, [request]) + + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][submit]" in message + assert "request=7" in message + assert "blocks=2" in message + assert "bytes=4096" in message + assert "submit_call_ms=2.000000" in message + + def test_transfer_status_emits_ready_to_reap_delay(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + timestamps = iter((10.0, 10.1)) + monkeypatch.setattr( + py_executor_module, "get_steady_clock_now_in_seconds", lambda: next(timestamps) + ) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + executor._is_disagg_inflight_cancel_active = Mock(return_value=False) + executor._check_cache_transfer_errors = Mock() + executor.canceled_req_ids = [] + request = _make_disagg_transfer_request(8, 32, in_progress=True) + request.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + request.py_kv_transfer_ready_time_s = 9.5 + request.py_kv_cache_xfer_bytes = 2048 + executor.active_requests = [request] + executor.kv_cache_transceiver = Mock() + + def complete(_at_least_num): + request.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + request.is_disagg_generation_transmission_in_progress = False + return [], [], [] + + executor.kv_cache_transceiver.check_gen_transfer_status.side_effect = complete + + PyExecutor._check_disagg_gen_cache_transfer_status(executor, 0) + + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][reap]" in message + assert "request=8" in message + assert "ready_t=9.500000000" in message + assert "ready_to_reap_ms=600.000000" in message + assert "poll_call_ms=100.000000" in message + assert "outcome=completed" in message + + def test_blocking_status_poll_emits_no_progress_duration(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + timestamps = iter((10.0, 15.0)) + monkeypatch.setattr( + py_executor_module, "get_steady_clock_now_in_seconds", lambda: next(timestamps) + ) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor._is_disagg_inflight_cancel_active = Mock(return_value=False) + executor._check_cache_transfer_errors = Mock() + executor.canceled_req_ids = [] + request = _make_disagg_transfer_request(9, 32, in_progress=True) + executor.active_requests = [request] + executor.kv_cache_transceiver = Mock() + executor.kv_cache_transceiver.check_gen_transfer_status.return_value = ([], [], []) + + PyExecutor._check_disagg_gen_cache_transfer_status(executor, 1) + + log_info.assert_called_once() + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][status-poll]" in message + assert "poll_call_ms=5000.000000" in message + assert "at_least_num=1" in message + assert "tracked=1" in message + assert "completed=0" in message + + def test_cpp_status_poll_derives_progress_from_request_state(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + timestamps = iter((20.0, 20.01)) + monkeypatch.setattr( + py_executor_module, "get_steady_clock_now_in_seconds", lambda: next(timestamps) + ) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + executor._is_disagg_inflight_cancel_active = Mock(return_value=False) + executor._check_cache_transfer_errors = Mock() + request = _make_disagg_transfer_request(10, 32, in_progress=True) + request.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + request.py_kv_transfer_ready_time_s = 20.005 + executor.active_requests = [request] + executor.kv_cache_transceiver = Mock() + + def complete_without_result(_at_least_num): + request.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + request.is_disagg_generation_transmission_in_progress = False + return None + + executor.kv_cache_transceiver.check_gen_transfer_status.side_effect = ( + complete_without_result + ) + + PyExecutor._check_disagg_gen_cache_transfer_status(executor, 1) + + poll_message = next( + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][status-poll]" in call.args[0] + ) + assert "completed=1" in poll_message + assert "failed=0" in poll_message + assert "cancelled=0" in poll_message + def test_gen_transfer_status_enters_without_local_active_transfers(self): executor = object.__new__(PyExecutor) executor.active_requests = [] diff --git a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py index 68b38bf896ae..c8a75f6d7f21 100644 --- a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py +++ b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py @@ -17,13 +17,21 @@ from __future__ import annotations from dataclasses import dataclass +from datetime import timedelta from typing import Optional from unittest.mock import Mock import pytest +from tensorrt_llm._torch.disaggregation import transceiver as transceiver_module from tensorrt_llm._torch.disaggregation.base.transfer import SessionStatus, WaitResult -from tensorrt_llm._torch.disaggregation.native.transfer import TaskStatus, TxSession +from tensorrt_llm._torch.disaggregation.native import transfer as native_transfer_module +from tensorrt_llm._torch.disaggregation.native.transfer import ( + KVRecvTask, + RxSession, + TaskStatus, + TxSession, +) from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 from tensorrt_llm.bindings import LlmRequestState @@ -50,12 +58,16 @@ def __init__( status: SessionStatus = SessionStatus.READY, is_completed: bool = False, has_failed: bool = False, + kv_transfer_start_time_s: Optional[float] = None, + kv_ready_time_s: Optional[float] = None, ) -> None: self._rid = rid self._wait_result = wait_result self._status = status self._is_completed = is_completed self._has_failed = has_failed + self.kv_transfer_start_time_s = kv_transfer_start_time_s + self.kv_ready_time_s = kv_ready_time_s self.blocking_calls: list[bool] = [] self.closed = False @@ -247,6 +259,91 @@ def test_gen_transfer_status_enters_consensus_when_sync_required() -> None: transceiver._gen_consensus.assert_called_once_with([]) +def test_gen_transfer_status_stamps_first_local_ready_time( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(transceiver_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(transceiver_module.logger, "info", log_info) + session = _FakeSession( + rid=21, + wait_result=WaitResult.COMPLETED, + is_completed=True, + kv_transfer_start_time_s=10.25, + kv_ready_time_s=12.5, + ) + request = Mock( + py_request_id=21, + py_kv_cache_xfer_bytes=8192, + state=LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS, + ) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = True + transceiver._gen_need_sync = False + transceiver._recv_sessions = {21: session} + transceiver._recv_reqs = {21: request} + transceiver._diagnostic_ready_rids = set() + transceiver._dist = Mock(rank=2) + transceiver._collect_done = Mock(return_value=([21], [])) + transceiver._gen_consensus = Mock(return_value=[21]) + transceiver._build_to_process = Mock(return_value=[21]) + transceiver._gen_consensus_outcome = Mock(return_value=([], [], [21])) + transceiver._need_aux_transfer = Mock(return_value=False) + transceiver._assert_disagg_history_declared = Mock() + transceiver._close_failed_sessions = Mock() + + completed, failed, cancelled = transceiver.check_gen_transfer_status(at_least_request_num=0) + + assert completed == [21] + assert failed == [] + assert cancelled == [] + assert request.py_kv_transfer_service_start_time_s == 10.25 + assert request.py_kv_transfer_ready_time_s == 12.5 + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][python-transfer]" in message + assert "rank=2" in message + assert "request=21" in message + assert "bytes=8192" in message + assert "service_start_t=10.250000000" in message + assert "service_ms=2250.000" in message + + +def test_kv_recv_task_records_native_transfer_boundaries( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(native_transfer_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + timestamps = iter((timedelta(seconds=3.25), timedelta(seconds=4.75))) + monkeypatch.setattr( + native_transfer_module.tensorrt_llm.bindings, + "steady_clock_now", + lambda: next(timestamps), + ) + task = KVRecvTask(17, Mock(), 0, Mock(), None) + + task.mark_transferring() + task.mark_transferring() + task.complete() + + assert task.transfer_start_time_s == 3.25 + assert task.completion_time_s == 4.75 + assert task.status == TaskStatus.TRANSFERRED + + +def test_rx_session_aggregates_native_task_timestamps() -> None: + session = object.__new__(RxSession) + first_task = Mock(transfer_start_time_s=2.5, completion_time_s=7.0) + second_task = Mock(transfer_start_time_s=3.0, completion_time_s=8.25) + session._kv_tasks = [first_task, second_task] + + assert session.kv_transfer_start_time_s == 2.5 + assert session.kv_ready_time_s == 8.25 + + second_task.completion_time_s = None + assert session.kv_ready_time_s is None + + def test_consensus_outcome_uses_single_batched_allgather() -> None: # The cancelled/failed/completed id lists are exchanged with ONE allgather # (packed as a list-of-lists) instead of three; verify a single call and that diff --git a/tests/unittest/tools/test_disagg_admission_telemetry.py b/tests/unittest/tools/test_disagg_admission_telemetry.py new file mode 100644 index 000000000000..7c1f37242455 --- /dev/null +++ b/tests/unittest/tools/test_disagg_admission_telemetry.py @@ -0,0 +1,298 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import importlib.util +import json +import sys +from pathlib import Path + +import pytest + +_SCRIPT_PATH = Path(__file__).parents[3] / "scripts" / "disagg_admission_telemetry.py" +_SPEC = importlib.util.spec_from_file_location("disagg_admission_telemetry", _SCRIPT_PATH) +if _SPEC is None or _SPEC.loader is None: + raise ImportError(f"Cannot load telemetry analyzer from {_SCRIPT_PATH}") +_TELEMETRY = importlib.util.module_from_spec(_SPEC) +sys.modules[_SPEC.name] = _TELEMETRY +_SPEC.loader.exec_module(_TELEMETRY) + +analyze_events = _TELEMETRY.analyze_events +main = _TELEMETRY.main +parse_diagnostic_line = _TELEMETRY.parse_diagnostic_line + + +def _parse_lines(lines: list[str]): + return [event for line in lines if (event := parse_diagnostic_line(line)) is not None] + + +def test_parse_diagnostic_line_accepts_rank_prefix_and_ignores_malformed_lines(): + event = parse_diagnostic_line( + "INFO [RANK 3] [DISAGG_DIAG][admission] t=12.5 active_blocks=8 " + "candidate_requests=101:4,102:4 admitted=1 deferred=1 budget=16" + ) + + assert event is not None + assert event.category == "admission" + assert event.time_s == 12.5 + assert event.rank == "3" + assert event.fields["candidate_requests"] == "101:4,102:4" + assert parse_diagnostic_line("ordinary log line") is None + assert parse_diagnostic_line("[DISAGG_DIAG][submit] t=not-a-number rank=0") is None + + +def test_python_transfer_analysis_derives_refill_multiplier_and_progress_credit(): + events = _parse_lines( + [ + "[DISAGG_DIAG][decision] t=0.0 rank=0 sequence=1 runtime=Python " + "active_blocks=0 candidates=2 candidate_blocks=20 admitted=1 " + "admitted_blocks=10 deferred=1 deferred_blocks=10 budget=10", + "[DISAGG_DIAG][admission] t=0.0 rank=0 active_blocks=0 " + "candidate_requests=1:10,2:10 admitted=1 admitted_requests=1:10 " + "deferred=1 deferred_requests=2:10 budget=10 sequence=1", + "[DISAGG_DIAG][submit] t=0.1 rank=0 request=1 blocks=10 " + "submit_start_t=0.05 submit_call_ms=50", + "[DISAGG_DIAG][decision] t=0.5 rank=0 sequence=2 runtime=Python " + "active_blocks=10 candidates=1 candidate_blocks=10 admitted=0 " + "admitted_blocks=0 deferred=1 deferred_blocks=10 budget=10", + "[DISAGG_DIAG][admission] t=0.5 rank=0 active_blocks=10 " + "candidate_requests=2:10 admitted=0 admitted_requests=- deferred=1 " + "deferred_requests=2:10 budget=10 sequence=2", + "[DISAGG_DIAG][python-transfer] t=1.1 rank=0 action=local-ready request=1 " + "bytes=4096 service_start_t=0.3 outcome=completed", + "[DISAGG_DIAG][reap] t=1.2 rank=0 request=1 blocks=10 ready_t=1.1 outcome=completed", + "[DISAGG_DIAG][decision] t=1.25 rank=0 sequence=3 runtime=Python " + "active_blocks=10 candidates=1 candidate_blocks=10 admitted=0 " + "admitted_blocks=0 deferred=1 deferred_blocks=10 budget=10", + "[DISAGG_DIAG][decision] t=1.3 rank=0 sequence=4 runtime=Python " + "active_blocks=0 candidates=1 candidate_blocks=10 admitted=1 " + "admitted_blocks=10 deferred=0 deferred_blocks=0 budget=10", + "[DISAGG_DIAG][admission] t=1.3 rank=0 active_blocks=0 " + "candidate_requests=2:10 admitted=1 admitted_requests=2:10 deferred=0 " + "deferred_requests=- budget=10 sequence=4", + "[DISAGG_DIAG][submit] t=1.4 rank=0 request=2 blocks=10 " + "submit_start_t=1.35 submit_call_ms=50", + "[DISAGG_DIAG][python-transfer] t=2.4 rank=0 action=local-ready request=2 " + "bytes=4096 service_start_t=1.6 outcome=completed", + "[DISAGG_DIAG][reap] t=2.5 rank=0 request=2 blocks=10 ready_t=2.4 outcome=completed", + "[DISAGG_DIAG][status-poll] t=2.6 rank=0 poll_start_t=2.598 " + "poll_call_ms=2.0 at_least_num=1 tracked=1 completed=0 failed=0 " + "cancelled=0", + "[DISAGG_DIAG][status-poll] t=2.7 rank=0 poll_start_t=2.6995 " + "poll_call_ms=0.5 at_least_num=1 tracked=1 completed=1 failed=0 " + "cancelled=0", + "[DISAGG_DIAG][submit] t=bad rank=0 request=broken blocks=10", + ] + ) + + result = analyze_events(events) + rank = result["ranks"]["0"] + service = rank["service"] + python_transfer = rank["python_transfer"] + status_poll = rank["status_poll"] + release = rank["release_to_admission"] + progress = rank["linear_progress_credit"] + counterfactual = rank["fixed_multiplier_counterfactual"] + + assert result["parsed_event_count"] == 15 + assert service["completed_blocks"] == 20.0 + assert service["busy_s"] == pytest.approx(1.6) + assert service["throughput_blocks_per_s"] == pytest.approx(12.5) + assert service["latency_s"]["p50"] == pytest.approx(0.8) + assert python_transfer["submit_to_service_start_s"]["p50"] == pytest.approx(0.25) + assert python_transfer["ready_to_reap_s"]["p50"] == pytest.approx(0.1) + assert status_poll["no_progress_duration_ms"]["p50"] == pytest.approx(2.0) + assert status_poll["progress_duration_ms"]["p50"] == pytest.approx(0.5) + assert result["aggregate"]["status_poll"]["no_progress_duration_ms"]["p50"] == pytest.approx( + 2.0 + ) + + assert release["selected_release_source"] is None + assert release["by_source"]["reap"]["decision_gap_s"]["p50"] == pytest.approx(0.05) + assert release["by_source"]["reap"]["successful_admission_gap_s"]["p50"] == pytest.approx(0.1) + assert release["by_source"]["reap"]["refill_gap_s"]["p50"] == pytest.approx(0.15) + assert rank["shadow_multiplier"]["by_source"]["reap"]["summary"]["p50"] == pytest.approx(1.1875) + + assert len(progress["samples"]) == 1 + assert progress["samples"][0]["estimated_progress_credit_blocks"] == pytest.approx(2.5) + assert progress["samples"][0]["estimated_remaining_blocks"] == pytest.approx(7.5) + assert progress["samples"][0]["estimated_progress_fraction"] == pytest.approx(0.25) + assert counterfactual["next_deferred_required_multiplier"]["count"] == 2 + assert counterfactual["next_deferred_required_multiplier"]["p50"] == pytest.approx(2.0) + assert counterfactual["samples"][0]["next_deferred_request"] == "2" + assert [ + prefix["required_multiplier"] for prefix in counterfactual["samples"][0]["prefixes"] + ] == pytest.approx([1.0, 2.0]) + + +def test_receiver_slot_analysis_matches_reuse_and_backlog_refill_gap(): + events = _parse_lines( + [ + "[DISAGG_DIAG][admission] t=0.0 rank=2 active_blocks=0 " + "candidate_requests=11:4,12:4 admitted=1 admitted_requests=11:4 " + "deferred=1 deferred_requests=12:4 budget=4", + "[DISAGG_DIAG][submit] t=0.05 rank=2 request=11 blocks=4", + "[DISAGG_DIAG][receiver-slot] t=0.1 rank=2 action=acquire request=11 " + "manager_index=0 manager=0xabc buffer=7 wait_ms=2.5", + "[DISAGG_DIAG][receiver-slot] t=0.12 rank=2 action=acquired request=11 " + "manager_index=1 manager=0xdef buffer=9 wait_ms=3.0", + "[DISAGG_DIAG][receiver-slot] t=0.5 rank=2 action=release request=11 " + "manager=0xabc buffer=7", + "[DISAGG_DIAG][receiver-slot] t=0.6 rank=2 action=released request=11 " + "manager=0xdef buffer=9", + "[DISAGG_DIAG][admission] t=0.7 rank=2 active_blocks=0 " + "candidate_requests=12:4 admitted=1 admitted_requests=12:4 deferred=0 " + "deferred_requests=- budget=4", + "[DISAGG_DIAG][submit] t=0.75 rank=2 request=12 blocks=4", + "[DISAGG_DIAG][receiver-slot] t=0.8 rank=2 action=acquired request=12 " + "manager_index=0 manager=0xabc buffer=7 wait_ms=1.0", + "[DISAGG_DIAG][receiver-slot] t=0.82 rank=2 action=acquired request=12 " + "manager_index=1 manager=0xdef buffer=9 wait_ms=1.5", + "[DISAGG_DIAG][python-transfer] t=0.9 rank=2 action=local-ready request=11 bytes=4096", + "[DISAGG_DIAG][receiver-slot] t=1.2 rank=2 action=released request=12 " + "manager=0xabc buffer=7", + "[DISAGG_DIAG][receiver-slot] t=1.3 rank=2 action=released request=12 " + "manager=0xdef buffer=9", + "[DISAGG_DIAG][python-transfer] t=1.5 rank=2 action=local-ready request=12 bytes=4096", + "[DISAGG_DIAG][receiver-slot] t=1.4 rank=2 action=released request=999 " + "manager=0xmissing buffer=3", + ] + ) + + result = analyze_events(events) + rank = result["ranks"]["2"] + slots = rank["receiver_slots"] + service = rank["service"] + release = rank["release_to_admission"] + + assert release["selected_release_source"] == "receiver-slot" + assert slots["service_latency_s"]["count"] == 4 + assert service["latency_s"]["count"] == 2 + assert service["latency_s"]["p50"] == pytest.approx(0.5) + assert service["completed_blocks"] == 8.0 + assert slots["submit_to_service_start_s"]["p50"] == pytest.approx(0.05) + assert all( + interval["start_kind"] == "receiver-slot-acquired" for interval in service["intervals"] + ) + assert slots["wait_ms"]["p50"] == pytest.approx(2.0) + assert slots["unmatched_releases"] == 1 + assert slots["backlog_refill_gap_s"]["p50"] == pytest.approx(0.26) + assert release["selected_decision_gap_s"]["p50"] == pytest.approx(0.1) + assert release["selected_refill_gap_s"]["p50"] == pytest.approx(0.15) + assert release["selected_samples"][0]["release_t"] == pytest.approx(0.6) + + +def test_reap_release_uses_first_decision_then_matching_deferred_refill(): + events = _parse_lines( + [ + "[DISAGG_DIAG][decision] t=0.0 rank=0 sequence=1 active_blocks=0 " + "candidates=2 candidate_blocks=8 admitted=1 admitted_blocks=4 " + "deferred=1 deferred_blocks=4 budget=4", + "[DISAGG_DIAG][admission] t=0.0 rank=0 sequence=1 active_blocks=0 " + "candidate_requests=1:4,2:4 admitted=1 admitted_requests=1:4 " + "deferred=1 deferred_requests=2:4 budget=4", + "[DISAGG_DIAG][submit] t=0.1 rank=0 request=1 blocks=4", + "[DISAGG_DIAG][python-transfer] t=0.9 rank=0 action=local-ready " + "request=1 service_start_t=0.2 outcome=completed", + "[DISAGG_DIAG][reap] t=1.0 rank=0 request=1 blocks=4 outcome=completed", + "[DISAGG_DIAG][decision] t=1.1 rank=0 sequence=2 active_blocks=4 " + "candidates=1 candidate_blocks=4 admitted=0 admitted_blocks=0 " + "deferred=1 deferred_blocks=4 budget=4", + "[DISAGG_DIAG][decision] t=1.5 rank=0 sequence=3 active_blocks=0 " + "candidates=1 candidate_blocks=4 admitted=1 admitted_blocks=4 " + "deferred=0 deferred_blocks=0 budget=4", + "[DISAGG_DIAG][admission] t=1.5 rank=0 sequence=3 active_blocks=0 " + "candidate_requests=2:4 admitted=1 admitted_requests=2:4 deferred=0 " + "deferred_requests=- budget=4", + "[DISAGG_DIAG][submit] t=1.55 rank=0 request=2 blocks=4 " + "submit_start_t=1.5 submit_call_ms=50", + ] + ) + + result = analyze_events(events) + sample = result["ranks"]["0"]["release_to_admission"]["by_source"]["reap"]["samples"][0] + + assert sample["decision_gap_s"] == pytest.approx(0.1) + assert sample["successful_admission_gap_s"] == pytest.approx(0.5) + assert sample["refill_gap_s"] == pytest.approx(0.5) + assert sample["backlog_identity_unknown"] is False + assert sample["matched_backlog_request_ids"] == ["2"] + assert sample["eligible_for_multiplier_fit"] is True + + +def test_failed_transfer_contributes_no_service_or_release_samples(): + events = _parse_lines( + [ + "[DISAGG_DIAG][submit] t=0.1 rank=0 request=9 blocks=4 outcome=failed", + "[DISAGG_DIAG][receiver-slot] t=0.2 rank=0 action=acquired request=9 " + "manager_index=0 manager=0xabc buffer=1", + "[DISAGG_DIAG][python-transfer] t=0.5 rank=0 action=local-ready " + "request=9 service_start_t=0.3", + "[DISAGG_DIAG][receiver-slot] t=0.6 rank=0 action=released request=9 " + "manager=0xabc buffer=1", + "[DISAGG_DIAG][reap] t=0.7 rank=0 request=9 blocks=4 outcome=failed " + "state=DISAGG_TRANS_ERROR", + "[DISAGG_DIAG][receiver-transfer] t=0.8 rank=0 action=failed request=9", + ] + ) + + rank = analyze_events(events)["ranks"]["0"] + + assert rank["service"]["intervals"] == [] + assert rank["service"]["completed_blocks"] == 0 + assert rank["service"]["throughput_blocks_per_s"] is None + assert rank["receiver_slots"]["excluded_unsuccessful_intervals"] == 1 + assert rank["python_transfer"]["ready_to_reap_s"]["count"] == 0 + assert all( + not source["samples"] for source in rank["release_to_admission"]["by_source"].values() + ) + + +def test_unknown_backlog_identity_is_excluded_from_multiplier_fit(): + events = _parse_lines( + [ + "[DISAGG_DIAG][admission] t=0.0 rank=0 active_blocks=0 " + "candidate_requests=1:4 admitted=0 admitted_requests=- deferred=1 " + "deferred_requests=- budget=4", + "[DISAGG_DIAG][submit] t=0.1 rank=0 request=8 blocks=4", + "[DISAGG_DIAG][python-transfer] t=0.8 rank=0 action=local-ready " + "request=8 service_start_t=0.2 outcome=completed", + "[DISAGG_DIAG][reap] t=1.0 rank=0 request=8 blocks=4 outcome=completed", + "[DISAGG_DIAG][admission] t=1.2 rank=0 active_blocks=0 " + "candidate_requests=99:4 admitted=1 admitted_requests=99:4 deferred=0 " + "deferred_requests=- budget=4", + "[DISAGG_DIAG][submit] t=1.2 rank=0 request=99 blocks=4", + ] + ) + + rank = analyze_events(events)["ranks"]["0"] + sample = rank["release_to_admission"]["by_source"]["reap"]["samples"][0] + + assert sample["backlog_identity_unknown"] is True + assert sample["eligible_for_multiplier_fit"] is False + assert sample["shadow_multiplier"] is None + assert rank["shadow_multiplier"]["by_source"]["reap"]["summary"]["count"] == 0 + + +def test_cli_reads_log_paths_and_prints_json(tmp_path, capsys): + log_path = tmp_path / "worker.log" + log_path.write_text( + "noise\n[DISAGG_DIAG][submit] t=1.0 rank=0 request=5 blocks=2\n", + encoding="utf-8", + ) + + assert main([str(log_path), "--indent", "0"]) == 0 + output = json.loads(capsys.readouterr().out) + assert output["parsed_event_count"] == 1 + assert output["event_counts"] == {"submit": 1} From b0609a84eafe5ffbeac2b10154375ecc0a592383 Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Fri, 17 Jul 2026 19:29:38 -0700 Subject: [PATCH 2/7] [NVBUG 6312828][test] finalize admission telemetry analysis Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- scripts/disagg_admission_telemetry.py | 125 +++++++++++++++++- tensorrt_llm/_torch/pyexecutor/py_executor.py | 16 ++- ...tx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml | 2 +- .../_torch/executor/test_py_executor.py | 24 +++- .../tools/test_disagg_admission_telemetry.py | 43 +++++- 5 files changed, 193 insertions(+), 17 deletions(-) diff --git a/scripts/disagg_admission_telemetry.py b/scripts/disagg_admission_telemetry.py index 8c2a391629df..c648ff08aa33 100644 --- a/scripts/disagg_admission_telemetry.py +++ b/scripts/disagg_admission_telemetry.py @@ -44,6 +44,7 @@ class DiagnosticEvent: time_s: float rank: str fields: dict[str, str] + source: str | None = None @dataclass(frozen=True) @@ -158,7 +159,15 @@ def read_diagnostic_events(paths: Iterable[str | Path]) -> list[DiagnosticEvent] for line in log_file: event = parse_diagnostic_line(line) if event is not None: - events.append(event) + events.append( + DiagnosticEvent( + event.category, + event.time_s, + event.rank, + event.fields, + str(path_like), + ) + ) except OSError: continue return events @@ -178,11 +187,17 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: Returns: A JSON-serializable analysis dictionary. """ - sorted_events = sorted(events, key=lambda event: (event.rank, event.time_s)) + sorted_events = sorted( + events, + key=lambda event: (event.source or "", event.rank, event.time_s), + ) + sources = {event.source for event in sorted_events if event.source is not None} + namespace_by_source = len(sources) > 1 category_counts = Counter(event.category for event in sorted_events) events_by_rank: dict[str, list[DiagnosticEvent]] = defaultdict(list) for event in sorted_events: - events_by_rank[event.rank].append(event) + rank_key = f"{event.source}::rank={event.rank}" if namespace_by_source else event.rank + events_by_rank[rank_key].append(event) global_blocks = _collect_global_request_blocks(sorted_events) ranks: dict[str, object] = {} @@ -195,6 +210,9 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: aggregate_poll_durations_ms: list[float] = [] aggregate_progress_poll_durations_ms: list[float] = [] aggregate_no_progress_poll_durations_ms: list[float] = [] + aggregate_reported_ready_to_reap_ms: list[float] = [] + aggregate_physical_release_to_reap_s: list[float] = [] + aggregate_invalid_ready_to_reap_samples = 0 aggregate_busy_s = 0.0 aggregate_completed_blocks = 0.0 @@ -235,6 +253,17 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: aggregate_no_progress_poll_durations_ms.extend( status_poll["no_progress_duration_samples_ms"] ) + scheduler_visibility = rank_analysis["scheduler_visibility"] + if isinstance(scheduler_visibility, dict): + aggregate_reported_ready_to_reap_ms.extend( + scheduler_visibility["reported_ready_to_reap_samples_ms"] + ) + aggregate_physical_release_to_reap_s.extend( + scheduler_visibility["physical_release_to_reap_samples_s"] + ) + aggregate_invalid_ready_to_reap_samples += int( + scheduler_visibility["invalid_reported_ready_to_reap_samples"] + ) aggregate_throughput = _safe_ratio(aggregate_completed_blocks, aggregate_busy_s) selected_decision_gaps = [ @@ -287,6 +316,7 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: return { "schema_version": 1, + "rank_namespace": "source-path::rank" if namespace_by_source else "rank", "parsed_event_count": len(sorted_events), "event_counts": dict(sorted(category_counts.items())), "ranks": ranks, @@ -313,6 +343,11 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: "progress_duration_ms": _summary(aggregate_progress_poll_durations_ms), "no_progress_duration_ms": _summary(aggregate_no_progress_poll_durations_ms), }, + "scheduler_visibility": { + "reported_ready_to_reap_ms": _summary(aggregate_reported_ready_to_reap_ms), + "physical_release_to_reap_s": _summary(aggregate_physical_release_to_reap_s), + "invalid_reported_ready_to_reap_samples": (aggregate_invalid_ready_to_reap_samples), + }, }, "model": { "shadow_multiplier": "1 + throughput_blocks_per_s * refill_gap_s / budget_blocks", @@ -333,8 +368,22 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: def analyze_log_paths(paths: Iterable[str | Path]) -> dict[str, object]: - """Read and analyze diagnostic log paths.""" - return analyze_events(read_diagnostic_events(paths)) + """Analyze logs without merging overlapping rank IDs across inputs.""" + path_list = list(paths) + events = read_diagnostic_events(path_list) + result = analyze_events(events) + if len(path_list) > 1: + source_aggregates: dict[str, object] = {} + for path_like in path_list: + source = str(path_like) + source_result = analyze_events(event for event in events if event.source == source) + source_aggregates[source] = { + "parsed_event_count": source_result["parsed_event_count"], + "event_counts": source_result["event_counts"], + "aggregate": source_result["aggregate"], + } + result["source_aggregates"] = source_aggregates + return result def _analyze_rank( @@ -411,6 +460,11 @@ def _analyze_rank( progress_samples = _linear_progress_credit(admissions, service_intervals) fixed_multiplier_samples = _fixed_multiplier_counterfactual(admissions) ready_to_reap_samples = _point_pair_gaps(local_ready, reaps) + physical_release_to_reap_samples = _point_pair_gaps( + [PointEvent(interval.end_s, interval.request) for interval in physical_service_intervals], + reaps, + ) + reported_ready_to_reap_samples = _reported_ready_to_reap_samples(events, unsuccessful_requests) submit_to_service_start_samples = _submit_to_service_start_gaps(submits, local_ready) status_poll_samples = _status_poll_samples(events) progress_poll_durations = [ @@ -475,6 +529,25 @@ def _analyze_rank( "progress_duration_ms": _summary(progress_poll_durations), "no_progress_duration_ms": _summary(no_progress_poll_durations), }, + "scheduler_visibility": { + "reported_ready_to_reap_samples_ms": [ + float(sample["duration_ms"]) for sample in reported_ready_to_reap_samples + ], + "reported_ready_to_reap_ms": _summary( + [float(sample["duration_ms"]) for sample in reported_ready_to_reap_samples] + ), + "reported_ready_to_reap_samples": reported_ready_to_reap_samples, + "invalid_reported_ready_to_reap_samples": ( + _invalid_reported_ready_to_reap_sample_count(events, unsuccessful_requests) + ), + "physical_release_to_reap_samples_s": [ + float(sample["gap_s"]) for sample in physical_release_to_reap_samples + ], + "physical_release_to_reap_s": _summary( + [float(sample["gap_s"]) for sample in physical_release_to_reap_samples] + ), + "physical_release_to_reap_pairs": physical_release_to_reap_samples, + }, "receiver_slots": { "submit_to_service_start_samples_s": [ float(sample["gap_s"]) for sample in physical_queue_samples @@ -829,6 +902,48 @@ def _status_poll_samples(events: list[DiagnosticEvent]) -> list[dict[str, object return samples +def _reported_ready_to_reap_samples( + events: list[DiagnosticEvent], excluded_requests: set[str] +) -> list[dict[str, object]]: + """Collect scheduler-visible delay reported by completed reap events.""" + samples: list[dict[str, object]] = [] + for event in events: + if event.category != "reap" or _event_outcome(event) is False: + continue + request = event.fields.get("request") + duration_ms = _as_float(event.fields.get("ready_to_reap_ms")) + if ( + request is None + or request in excluded_requests + or duration_ms is None + or duration_ms < 0.0 + ): + continue + samples.append( + { + "t": event.time_s, + "request": request, + "duration_ms": duration_ms, + } + ) + return samples + + +def _invalid_reported_ready_to_reap_sample_count( + events: list[DiagnosticEvent], excluded_requests: set[str] +) -> int: + """Count negative or unavailable delays excluded from the summary.""" + return sum( + 1 + for event in events + if event.category == "reap" + and _event_outcome(event) is not False + and event.fields.get("request") not in excluded_requests + and (duration_ms := _as_float(event.fields.get("ready_to_reap_ms"))) is not None + and duration_ms < 0.0 + ) + + def _match_slot_intervals( events: list[DiagnosticEvent], ) -> tuple[list[SlotInterval], int, int]: diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 0e79e03e3d06..7f4fe5b553f6 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -6380,9 +6380,6 @@ def _check_disagg_gen_cache_transfer_status(self, atLeastNum: int = 0): if tracked_requests: reap_time = poll_end controller = self._get_disagg_transfer_admission_controller() - offset = LlmRequest.global_steady_clock_offset - offset_seconds = (offset.total_seconds() - if offset is not None else 0.0) for request in tracked_requests: if request.is_disagg_generation_transmission_in_progress: continue @@ -6397,13 +6394,19 @@ def _check_disagg_gen_cache_transfer_status(self, atLeastNum: int = 0): outcome = "cancelled" ready_time = getattr(request, "py_kv_transfer_ready_time_s", 0.0) + ready_time_source = "python-local" if not ready_time: transfer_end = getattr(request, "kv_cache_transfer_end", None) if transfer_end is not None: - ready_time = max( - 0.0, - transfer_end.total_seconds() - offset_seconds) + # This delay is measured entirely inside the local + # executor process. Applying the cross-rank clock + # offset would move the C++ timestamp into rank 0's + # domain and make it incomparable with reap_time. + ready_time = max(0.0, transfer_end.total_seconds()) + ready_time_source = "cpp-local" + else: + ready_time_source = "unavailable" ready_to_reap_ms = (-1.0 if not ready_time else (reap_time - ready_time) * 1000) self._log_disagg_transfer_diagnostic( @@ -6415,6 +6418,7 @@ def _check_disagg_gen_cache_transfer_status(self, atLeastNum: int = 0): bytes=getattr(request, "py_kv_cache_xfer_bytes", getattr(request, "kv_cache_size", 0)), ready_t=f"{ready_time:.9f}", + ready_time_source=ready_time_source, ready_to_reap_ms=f"{ready_to_reap_ms:.6f}", poll_call_ms=f"{(poll_end - poll_start) * 1000:.6f}", outcome=outcome, diff --git a/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_128k8k_con256_ctx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml b/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_128k8k_con256_ctx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml index 6db032ae6c3a..92e9e36230a3 100644 --- a/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_128k8k_con256_ctx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml +++ b/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_128k8k_con256_ctx1_pp4_gen1_dep8_eplb0_mtp1_ccb-NIXL.yaml @@ -36,7 +36,7 @@ environment: build_wheel: false work_dir: worker_env_var: TLLM_LOG_LEVEL=INFO TRTLLM_SERVER_DISABLE_GC=1 TRTLLM_WORKER_DISABLE_GC=1 TLLM_SPEC_DECODE_FORCE_NUM_ACCEPTED_TOKENS=1 - TRTLLM_ENABLE_PDL=1 TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS=1 ENROOT_ALLOW_DEV=yes + TRTLLM_ENABLE_PDL=1 ENROOT_ALLOW_DEV=yes server_env_var: TRTLLM_SERVER_DISABLE_GC=1 profiling: nsys_on: false diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index 9ebd1236db30..7e87745063fd 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -1,6 +1,5 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 - """Tests for PyExecutor request handling functionality. This module tests the request handling logic that was moved from ExecutorRequestQueue @@ -693,7 +692,20 @@ def mark_in_progress(req): assert "bytes=4096" in message assert "submit_call_ms=2.000000" in message - def test_transfer_status_emits_ready_to_reap_delay(self, monkeypatch): + @pytest.mark.parametrize( + ("python_ready_time", "cpp_ready_time", "ready_time_source"), + [ + (9.5, None, "python-local"), + (0.0, 9.5, "cpp-local"), + ], + ) + def test_transfer_status_emits_ready_to_reap_delay( + self, + monkeypatch, + python_ready_time, + cpp_ready_time, + ready_time_source, + ): monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) log_info = Mock() @@ -712,7 +724,12 @@ def test_transfer_status_emits_ready_to_reap_delay(self, monkeypatch): executor.canceled_req_ids = [] request = _make_disagg_transfer_request(8, 32, in_progress=True) request.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS - request.py_kv_transfer_ready_time_s = 9.5 + request.py_kv_transfer_ready_time_s = python_ready_time + request.kv_cache_transfer_end = ( + None + if cpp_ready_time is None + else Mock(total_seconds=Mock(return_value=cpp_ready_time)) + ) request.py_kv_cache_xfer_bytes = 2048 executor.active_requests = [request] executor.kv_cache_transceiver = Mock() @@ -730,6 +747,7 @@ def complete(_at_least_num): assert "[DISAGG_DIAG][reap]" in message assert "request=8" in message assert "ready_t=9.500000000" in message + assert f"ready_time_source={ready_time_source}" in message assert "ready_to_reap_ms=600.000000" in message assert "poll_call_ms=100.000000" in message assert "outcome=completed" in message diff --git a/tests/unittest/tools/test_disagg_admission_telemetry.py b/tests/unittest/tools/test_disagg_admission_telemetry.py index 7c1f37242455..9e0cce7e49cb 100644 --- a/tests/unittest/tools/test_disagg_admission_telemetry.py +++ b/tests/unittest/tools/test_disagg_admission_telemetry.py @@ -71,7 +71,8 @@ def test_python_transfer_analysis_derives_refill_multiplier_and_progress_credit( "deferred_requests=2:10 budget=10 sequence=2", "[DISAGG_DIAG][python-transfer] t=1.1 rank=0 action=local-ready request=1 " "bytes=4096 service_start_t=0.3 outcome=completed", - "[DISAGG_DIAG][reap] t=1.2 rank=0 request=1 blocks=10 ready_t=1.1 outcome=completed", + "[DISAGG_DIAG][reap] t=1.2 rank=0 request=1 blocks=10 ready_t=1.1 " + "ready_to_reap_ms=100 outcome=completed", "[DISAGG_DIAG][decision] t=1.25 rank=0 sequence=3 runtime=Python " "active_blocks=10 candidates=1 candidate_blocks=10 admitted=0 " "admitted_blocks=0 deferred=1 deferred_blocks=10 budget=10", @@ -85,7 +86,8 @@ def test_python_transfer_analysis_derives_refill_multiplier_and_progress_credit( "submit_start_t=1.35 submit_call_ms=50", "[DISAGG_DIAG][python-transfer] t=2.4 rank=0 action=local-ready request=2 " "bytes=4096 service_start_t=1.6 outcome=completed", - "[DISAGG_DIAG][reap] t=2.5 rank=0 request=2 blocks=10 ready_t=2.4 outcome=completed", + "[DISAGG_DIAG][reap] t=2.5 rank=0 request=2 blocks=10 ready_t=2.4 " + "ready_to_reap_ms=-1 outcome=completed", "[DISAGG_DIAG][status-poll] t=2.6 rank=0 poll_start_t=2.598 " "poll_call_ms=2.0 at_least_num=1 tracked=1 completed=0 failed=0 " "cancelled=0", @@ -101,6 +103,7 @@ def test_python_transfer_analysis_derives_refill_multiplier_and_progress_credit( service = rank["service"] python_transfer = rank["python_transfer"] status_poll = rank["status_poll"] + visibility = rank["scheduler_visibility"] release = rank["release_to_admission"] progress = rank["linear_progress_credit"] counterfactual = rank["fixed_multiplier_counterfactual"] @@ -114,6 +117,8 @@ def test_python_transfer_analysis_derives_refill_multiplier_and_progress_credit( assert python_transfer["ready_to_reap_s"]["p50"] == pytest.approx(0.1) assert status_poll["no_progress_duration_ms"]["p50"] == pytest.approx(2.0) assert status_poll["progress_duration_ms"]["p50"] == pytest.approx(0.5) + assert visibility["reported_ready_to_reap_ms"]["p50"] == pytest.approx(100.0) + assert visibility["invalid_reported_ready_to_reap_samples"] == 1 assert result["aggregate"]["status_poll"]["no_progress_duration_ms"]["p50"] == pytest.approx( 2.0 ) @@ -151,6 +156,8 @@ def test_receiver_slot_analysis_matches_reuse_and_backlog_refill_gap(): "manager=0xabc buffer=7", "[DISAGG_DIAG][receiver-slot] t=0.6 rank=2 action=released request=11 " "manager=0xdef buffer=9", + "[DISAGG_DIAG][reap] t=0.65 rank=2 request=11 blocks=4 " + "ready_to_reap_ms=25 outcome=completed", "[DISAGG_DIAG][admission] t=0.7 rank=2 active_blocks=0 " "candidate_requests=12:4 admitted=1 admitted_requests=12:4 deferred=0 " "deferred_requests=- budget=4", @@ -164,6 +171,8 @@ def test_receiver_slot_analysis_matches_reuse_and_backlog_refill_gap(): "manager=0xabc buffer=7", "[DISAGG_DIAG][receiver-slot] t=1.3 rank=2 action=released request=12 " "manager=0xdef buffer=9", + "[DISAGG_DIAG][reap] t=1.35 rank=2 request=12 blocks=4 " + "ready_to_reap_ms=35 outcome=completed", "[DISAGG_DIAG][python-transfer] t=1.5 rank=2 action=local-ready request=12 bytes=4096", "[DISAGG_DIAG][receiver-slot] t=1.4 rank=2 action=released request=999 " "manager=0xmissing buffer=3", @@ -175,6 +184,7 @@ def test_receiver_slot_analysis_matches_reuse_and_backlog_refill_gap(): slots = rank["receiver_slots"] service = rank["service"] release = rank["release_to_admission"] + visibility = rank["scheduler_visibility"] assert release["selected_release_source"] == "receiver-slot" assert slots["service_latency_s"]["count"] == 4 @@ -191,6 +201,11 @@ def test_receiver_slot_analysis_matches_reuse_and_backlog_refill_gap(): assert release["selected_decision_gap_s"]["p50"] == pytest.approx(0.1) assert release["selected_refill_gap_s"]["p50"] == pytest.approx(0.15) assert release["selected_samples"][0]["release_t"] == pytest.approx(0.6) + assert visibility["reported_ready_to_reap_ms"]["p50"] == pytest.approx(30.0) + assert visibility["physical_release_to_reap_s"]["p50"] == pytest.approx(0.05) + assert result["aggregate"]["scheduler_visibility"]["reported_ready_to_reap_ms"][ + "p50" + ] == pytest.approx(30.0) def test_reap_release_uses_first_decision_then_matching_deferred_refill(): @@ -296,3 +311,27 @@ def test_cli_reads_log_paths_and_prints_json(tmp_path, capsys): output = json.loads(capsys.readouterr().out) assert output["parsed_event_count"] == 1 assert output["event_counts"] == {"submit": 1} + + +def test_cli_namespaces_overlapping_ranks_from_distinct_logs(tmp_path, capsys): + ctx_log = tmp_path / "ctx.log" + gen_log = tmp_path / "gen.log" + ctx_log.write_text( + "[DISAGG_DIAG][status-poll] t=1.0 rank=0 poll_call_ms=5 completed=0 failed=0 cancelled=0\n", + encoding="utf-8", + ) + gen_log.write_text( + "[DISAGG_DIAG][submit] t=2.0 rank=0 request=5 blocks=2\n", + encoding="utf-8", + ) + + assert main([str(ctx_log), str(gen_log), "--indent", "0"]) == 0 + output = json.loads(capsys.readouterr().out) + + assert output["rank_namespace"] == "source-path::rank" + assert sorted(output["ranks"]) == [ + f"{ctx_log}::rank=0", + f"{gen_log}::rank=0", + ] + assert output["source_aggregates"][str(ctx_log)]["event_counts"] == {"status-poll": 1} + assert output["source_aggregates"][str(gen_log)]["event_counts"] == {"submit": 1} From c68c06aa8ddcd382a5f57652c0e0ff3471bf4aaa Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Fri, 17 Jul 2026 19:31:44 -0700 Subject: [PATCH 3/7] [NVBUG 6312828][test] isolate multi-log telemetry analysis Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- scripts/disagg_admission_telemetry.py | 17 +++++++--- .../tools/test_disagg_admission_telemetry.py | 32 +++++++++++++++++++ 2 files changed, 45 insertions(+), 4 deletions(-) diff --git a/scripts/disagg_admission_telemetry.py b/scripts/disagg_admission_telemetry.py index c648ff08aa33..ef5d1aa27036 100644 --- a/scripts/disagg_admission_telemetry.py +++ b/scripts/disagg_admission_telemetry.py @@ -200,6 +200,12 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: events_by_rank[rank_key].append(event) global_blocks = _collect_global_request_blocks(sorted_events) + blocks_by_source = { + source: _collect_global_request_blocks( + [event for event in sorted_events if event.source == source] + ) + for source in sources + } ranks: dict[str, object] = {} aggregate_service_intervals: list[ServiceInterval] = [] aggregate_selected_gaps: list[dict[str, object]] = [] @@ -217,9 +223,11 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: aggregate_completed_blocks = 0.0 for rank in sorted(events_by_rank, key=_rank_sort_key): - rank_analysis, rank_intervals, selected_gaps = _analyze_rank( - events_by_rank[rank], global_blocks - ) + rank_events = events_by_rank[rank] + request_blocks = global_blocks + if namespace_by_source and rank_events: + request_blocks = blocks_by_source.get(rank_events[0].source, {}) + rank_analysis, rank_intervals, selected_gaps = _analyze_rank(rank_events, request_blocks) ranks[rank] = rank_analysis aggregate_service_intervals.extend(rank_intervals) aggregate_selected_gaps.extend(selected_gaps) @@ -317,6 +325,7 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: return { "schema_version": 1, "rank_namespace": "source-path::rank" if namespace_by_source else "rank", + "aggregate_scope": "all-input-sources" if namespace_by_source else "single-source", "parsed_event_count": len(sorted_events), "event_counts": dict(sorted(category_counts.items())), "ranks": ranks, @@ -932,7 +941,7 @@ def _reported_ready_to_reap_samples( def _invalid_reported_ready_to_reap_sample_count( events: list[DiagnosticEvent], excluded_requests: set[str] ) -> int: - """Count negative or unavailable delays excluded from the summary.""" + """Count negative reported delays excluded from the summary.""" return sum( 1 for event in events diff --git a/tests/unittest/tools/test_disagg_admission_telemetry.py b/tests/unittest/tools/test_disagg_admission_telemetry.py index 9e0cce7e49cb..d22652508d73 100644 --- a/tests/unittest/tools/test_disagg_admission_telemetry.py +++ b/tests/unittest/tools/test_disagg_admission_telemetry.py @@ -329,9 +329,41 @@ def test_cli_namespaces_overlapping_ranks_from_distinct_logs(tmp_path, capsys): output = json.loads(capsys.readouterr().out) assert output["rank_namespace"] == "source-path::rank" + assert output["aggregate_scope"] == "all-input-sources" assert sorted(output["ranks"]) == [ f"{ctx_log}::rank=0", f"{gen_log}::rank=0", ] assert output["source_aggregates"][str(ctx_log)]["event_counts"] == {"status-poll": 1} assert output["source_aggregates"][str(gen_log)]["event_counts"] == {"submit": 1} + + +def test_multi_log_block_lookup_is_scoped_by_source(tmp_path, capsys): + ctx_log = tmp_path / "ctx.log" + gen_log = tmp_path / "gen.log" + ctx_log.write_text( + "[DISAGG_DIAG][admission] t=0 rank=0 active_blocks=0 " + "candidate_requests=shared:4 admitted=1 admitted_requests=shared:4 " + "deferred=0 deferred_requests=- budget=4\n" + "[DISAGG_DIAG][receiver-slot] t=0.1 rank=1 action=acquired " + "request=shared manager=ctx buffer=0\n" + "[DISAGG_DIAG][receiver-slot] t=0.2 rank=1 action=released " + "request=shared manager=ctx buffer=0\n", + encoding="utf-8", + ) + gen_log.write_text( + "[DISAGG_DIAG][admission] t=0 rank=0 active_blocks=0 " + "candidate_requests=shared:2 admitted=1 admitted_requests=shared:2 " + "deferred=0 deferred_requests=- budget=2\n" + "[DISAGG_DIAG][receiver-slot] t=0.1 rank=1 action=acquired " + "request=shared manager=gen buffer=0\n" + "[DISAGG_DIAG][receiver-slot] t=0.2 rank=1 action=released " + "request=shared manager=gen buffer=0\n", + encoding="utf-8", + ) + + assert main([str(ctx_log), str(gen_log), "--indent", "0"]) == 0 + output = json.loads(capsys.readouterr().out) + + assert output["ranks"][f"{ctx_log}::rank=1"]["service"]["completed_blocks"] == 4 + assert output["ranks"][f"{gen_log}::rank=1"]["service"]["completed_blocks"] == 2 From e7f9121a7afd860b2c274442e1e482ff6b493f20 Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Fri, 24 Jul 2026 09:46:58 -0700 Subject: [PATCH 4/7] [NVBUG 6312828][test] correlate disaggregated transfer lifecycle Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- .../batch_manager/dataTransceiver.cpp | 350 ++++++- scripts/disagg_admission_telemetry.py | 986 +++++++++++++++++- .../_torch/disaggregation/native/transfer.py | 163 +++ .../_torch/disaggregation/transceiver.py | 180 +++- tensorrt_llm/_torch/pyexecutor/llm_request.py | 5 + tensorrt_llm/_torch/pyexecutor/py_executor.py | 311 +++++- ...1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml | 2 +- .../_torch/executor/test_py_executor.py | 252 ++++- .../test_transceiver_bounded_polling.py | 61 +- .../tools/test_disagg_admission_telemetry.py | 108 +- 10 files changed, 2333 insertions(+), 85 deletions(-) diff --git a/cpp/tensorrt_llm/batch_manager/dataTransceiver.cpp b/cpp/tensorrt_llm/batch_manager/dataTransceiver.cpp index e64831a0e4dc..f0dab09f9413 100644 --- a/cpp/tensorrt_llm/batch_manager/dataTransceiver.cpp +++ b/cpp/tensorrt_llm/batch_manager/dataTransceiver.cpp @@ -40,6 +40,7 @@ #include #include #include +#include #include namespace tensorrt_llm::batch_manager @@ -72,6 +73,135 @@ std::mutex& getDisaggDiagnosticsLogMutex() return mutex; } +LlmRequest::RequestIdType getDiagnosticContextRequestId( + LlmRequest const* request, LlmRequest::RequestIdType fallback) noexcept +{ + if (!isDisaggTransferDiagnosticsEnabled() || request == nullptr) + { + return fallback; + } + try + { + auto const& contextPhaseParams = request->getContextPhaseParams(); + if (contextPhaseParams.has_value()) + { + return contextPhaseParams->getReqId(); + } + } + catch (...) + { + // Diagnostics must not alter transfer semantics. + } + return fallback; +} + +void logTransferDiagnostic(char const* category, char const* action, LlmRequest::RequestIdType requestId, + LlmRequest::RequestIdType contextRequestId, char const* phase, double timestamp = 0.0) noexcept +{ + if (!isDisaggTransferDiagnosticsEnabled()) + { + return; + } + try + { + auto const eventTime = timestamp > 0.0 ? timestamp : getSteadyClockTimeSeconds(); + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO("[DISAGG_DIAG][%s-transfer] t=%.9f rank=%d action=%s request=%zu context_request=%zu phase=%s", + category, eventTime, mpi::MpiComm::world().getRank(), action, requestId, contextRequestId, phase); + } + catch (...) + { + // Diagnostics must not alter transfer semantics. + } +} + +void logSenderRequestInfoCompleteDiagnostic(LlmRequest::RequestIdType requestId, + LlmRequest::RequestIdType contextRequestId, size_t counterparts, bool ready) noexcept +{ + if (!isDisaggTransferDiagnosticsEnabled()) + { + return; + } + try + { + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO( + "[DISAGG_DIAG][sender-transfer] t=%.9f rank=%d action=request-info-complete request=%zu " + "context_request=%zu phase=receiver-credit counterparts=%zu ready=%d", + getSteadyClockTimeSeconds(), mpi::MpiComm::world().getRank(), requestId, contextRequestId, counterparts, + static_cast(ready)); + } + catch (...) + { + // Diagnostics must not alter transfer semantics. + } +} + +void logSenderServiceStartDiagnostic(LlmRequest::RequestIdType requestId, LlmRequest::RequestIdType contextRequestId, + char const* source, double timestamp = 0.0) noexcept +{ + if (!isDisaggTransferDiagnosticsEnabled()) + { + return; + } + try + { + auto const eventTime = timestamp > 0.0 ? timestamp : getSteadyClockTimeSeconds(); + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO( + "[DISAGG_DIAG][sender-transfer] t=%.9f rank=%d action=service-start request=%zu context_request=%zu " + "phase=format source=%s", + eventTime, mpi::MpiComm::world().getRank(), requestId, contextRequestId, source); + } + catch (...) + { + // Diagnostics must not alter transfer semantics. + } +} + +void logReceiverRequestInfoSubmittedDiagnostic( + LlmRequest::RequestIdType requestId, LlmRequest::RequestIdType contextRequestId, size_t counterparts) noexcept +{ + if (!isDisaggTransferDiagnosticsEnabled()) + { + return; + } + try + { + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO( + "[DISAGG_DIAG][receiver-transfer] t=%.9f rank=%d action=request-info-submitted request=%zu " + "context_request=%zu phase=request-info counterparts=%zu", + getSteadyClockTimeSeconds(), mpi::MpiComm::world().getRank(), requestId, contextRequestId, counterparts); + } + catch (...) + { + // Diagnostics must not alter transfer semantics. + } +} + +void logReceiverReadyDiagnostic(LlmRequest::RequestIdType requestId, LlmRequest::RequestIdType contextRequestId, + char const* result, double timestamp = 0.0) noexcept +{ + if (!isDisaggTransferDiagnosticsEnabled()) + { + return; + } + try + { + auto const eventTime = timestamp > 0.0 ? timestamp : getSteadyClockTimeSeconds(); + std::lock_guard lock(getDisaggDiagnosticsLogMutex()); + TLLM_LOG_INFO( + "[DISAGG_DIAG][receiver-transfer] t=%.9f rank=%d action=ready request=%zu context_request=%zu " + "phase=ready-signal result=%s", + eventTime, mpi::MpiComm::world().getRank(), requestId, contextRequestId, result); + } + catch (...) + { + // Diagnostics must not alter transfer semantics. + } +} + } // namespace std::vector const& TransferSession::getConnections() const @@ -457,6 +587,7 @@ class CacheSender::Impl { (void) getOrCreateInFlightCancelFlag(llmRequest->mRequestId); } + double queuedTime = 0.0; { std::scoped_lock lock(mSenderMutex); TLLM_CHECK_WITH_INFO( @@ -465,8 +596,11 @@ class CacheSender::Impl = mReadyResponses.emplace(llmRequest->mRequestId, Response{llmRequest, std::move(promise)}); TLLM_CHECK_WITH_INFO( result.second, "Request %zu is already queued for KV cache transfer", llmRequest->mRequestId); + queuedTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; } mSenderCv.notify_all(); + logTransferDiagnostic("sender", "queued", llmRequest->mRequestId, + getDiagnosticContextRequestId(llmRequest.get(), llmRequest->mRequestId), "wait-request-info", queuedTime); return future; } @@ -501,6 +635,42 @@ class CacheSender::Impl return it->second.getConnections().size(); } + void logRequestInfoCompleteDiagnostic( + RequestIdType sessionId, RequestIdType requestId, RequestIdType contextRequestId, bool ready) noexcept + { + if (!isDisaggTransferDiagnosticsEnabled()) + { + return; + } + try + { + logSenderRequestInfoCompleteDiagnostic(requestId, contextRequestId, getCounterpartsCount(sessionId), ready); + } + catch (...) + { + // Diagnostics must not alter transfer semantics. + } + } + + [[nodiscard]] bool isDiagnosticCancellationRequested(RequestIdType requestId) noexcept + { + if (!isDisaggTransferDiagnosticsEnabled() || !common::getEnvDisaggEnableInflightCancel()) + { + return false; + } + try + { + std::lock_guard lock(mInFlightCancelMutex); + auto const it = mInFlightCancelFlags.find(requestId); + return it != mInFlightCancelFlags.end() && it->second->load(std::memory_order_relaxed); + } + catch (...) + { + // Diagnostics must not alter transfer semantics. + return false; + } + } + void release(LlmRequest::RequestIdType requestId) { { @@ -627,7 +797,7 @@ class CacheSender::Impl return info; } - void sendSync(LlmRequest const& llmRequest) + std::pair sendSync(LlmRequest const& llmRequest, bool emitServiceDiagnostic = true) { TransferSession* session = nullptr; { @@ -638,9 +808,26 @@ class CacheSender::Impl } session->setLlmRequest(llmRequest); TLLM_LOG_DEBUG("KV cache transfer request %zu phase=transfer-submit begin.", llmRequest.mRequestId); - mCacheTransferLayer.format(*session); + auto const serviceStartTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + try + { + mCacheTransferLayer.format(*session); + } + catch (...) + { + logSenderServiceStartDiagnostic(llmRequest.mRequestId, + getDiagnosticContextRequestId(&llmRequest, llmRequest.mRequestId), "request", serviceStartTime); + throw; + } + auto const completedTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + if (emitServiceDiagnostic) + { + logSenderServiceStartDiagnostic(llmRequest.mRequestId, + getDiagnosticContextRequestId(&llmRequest, llmRequest.mRequestId), "request", serviceStartTime); + } TLLM_LOG_DEBUG("KV cache transfer request %zu phase=transfer-complete end.", llmRequest.mRequestId); llmRequest.setKvCacheTransferEnd(LlmRequest::getSteadyClockNow()); + return {serviceStartTime, completedTime}; } bool cancelRequest(LlmRequest const& llmRequest) @@ -648,6 +835,8 @@ class CacheSender::Impl bool const inflightCancelEnabled = common::getEnvDisaggEnableInflightCancel(); bool isCancelled = false; bool isCurrentRequest = false; + bool cancelledBeforeReceiverReady = false; + double cancellationTime = 0.0; { std::scoped_lock lock(mSenderMutex); auto it = mReadyResponses.find(llmRequest.mRequestId); @@ -664,12 +853,14 @@ class CacheSender::Impl { // Keep only the request ID as a tombstone so a late peer // receives ready=false without retaining the request. + cancellationTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; failResponse(it->second, std::make_exception_ptr( TLLM_REQUEST_EXCEPTION(llmRequest.mRequestId, common::RequestErrorCode::kNETWORK_ERROR, "Context KV cache request cancelled before a peer was ready for request %zu", llmRequest.mRequestId))); mReadyResponses.erase(it); + cancelledBeforeReceiverReady = true; } } } @@ -692,6 +883,12 @@ class CacheSender::Impl { mSenderCv.notify_all(); } + if (cancelledBeforeReceiverReady) + { + logTransferDiagnostic("sender", "cancelled", llmRequest.mRequestId, + getDiagnosticContextRequestId(&llmRequest, llmRequest.mRequestId), "wait-request-info", + cancellationTime); + } return isCancelled; } @@ -805,41 +1002,57 @@ class CacheSender::Impl void sendAndRemoveResponse(RequestIdType id, Response resp) noexcept { + auto const* llmRequest = resp.getRequest(); + auto const requestId = llmRequest != nullptr ? llmRequest->mRequestId : id; + auto const contextRequestId = getDiagnosticContextRequestId(llmRequest, id); try { TLLM_CUDA_CHECK(cudaSetDevice(mDeviceId)); - if (auto* llmRequest = resp.getRequest(); llmRequest != nullptr) + std::pair transferTimes; + if (llmRequest != nullptr) { - sendSync(*llmRequest); + transferTimes = sendSync(*llmRequest, false); } else { // Reuse tree path — no LlmRequest - sendSyncFromReuseTree(id); + transferTimes = sendSyncFromReuseTree(id); } release(id); resp.mPromise.set_value(); + logSenderServiceStartDiagnostic( + requestId, contextRequestId, llmRequest != nullptr ? "request" : "reuse-tree", transferTimes.first); + logTransferDiagnostic("sender", "completed", requestId, contextRequestId, "format", transferTimes.second); } catch (tensorrt_llm::common::RequestSpecificException const& e) { + auto const terminalTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + auto const* action = isDiagnosticCancellationRequested(id) ? "cancelled" : "failed"; TLLM_LOG_ERROR("Exception in sendAndRemoveResponse: %s ", e.what()); discardTransferState(id); auto new_exception = TLLM_REQUEST_EXCEPTION(id, e.getErrorCode(), "%s", e.what()); failResponse(resp, std::make_exception_ptr(new_exception)); + logTransferDiagnostic("sender", action, requestId, contextRequestId, "format", terminalTime); } catch (std::exception const& e) { + auto const terminalTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + auto const* action = isDiagnosticCancellationRequested(id) ? "cancelled" : "failed"; auto const exception = std::current_exception(); TLLM_LOG_ERROR("Exception in sendAndRemoveResponse: %s request id: %ld", e.what(), id); discardTransferState(id); failResponse(resp, exception); + logTransferDiagnostic("sender", action, requestId, contextRequestId, "format", terminalTime); } catch (...) { + auto const terminalTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + auto const* action = isDiagnosticCancellationRequested(id) ? "cancelled" : "failed"; auto const exception = std::current_exception(); TLLM_LOG_ERROR("Unknown exception in sendAndRemoveResponse for request id: %ld", id); discardTransferState(id); failResponse(resp, exception); + logTransferDiagnostic("sender", action, requestId, contextRequestId, "format", terminalTime); } releasePinnedBlocks(resp); } @@ -854,15 +1067,19 @@ class CacheSender::Impl } catch (std::exception const& err) { + auto const terminalTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; TLLM_LOG_ERROR("Failed to queue asynchronous KV cache send for request %zu: %s", id, err.what()); discardTransferState(id); failResponse(resp, std::current_exception()); + logTransferDiagnostic("sender", "failed", id, id, "async-queue", terminalTime); } catch (...) { + auto const terminalTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; TLLM_LOG_ERROR("Unknown error while queueing asynchronous KV cache send for request %zu", id); discardTransferState(id); failResponse(resp, std::current_exception()); + logTransferDiagnostic("sender", "failed", id, id, "async-queue", terminalTime); } } @@ -871,12 +1088,20 @@ class CacheSender::Impl bool isReady = true; bool allCounterpartsReady = false; std::optional cancelledResponse; + RequestIdType requestId = reqId; + RequestIdType contextRequestId = reqId; { std::scoped_lock lock(mSenderMutex); TLLM_CHECK(mCurrentRequest.has_value() && mCurrentRequest.value() == reqId); auto responseIt = mReadyResponses.find(reqId); bool const isCancelled = mCancelledRequests.find(reqId) != mCancelledRequests.end(); TLLM_CHECK(responseIt != mReadyResponses.end() || isCancelled); + if (responseIt != mReadyResponses.end()) + { + auto const* request = responseIt->second.getRequest(); + requestId = request != nullptr ? request->mRequestId : reqId; + contextRequestId = getDiagnosticContextRequestId(request, reqId); + } auto countIt = mRemainSendCount.find(reqId); TLLM_CHECK(countIt != mRemainSendCount.end()); auto const count = --countIt->second; @@ -900,15 +1125,20 @@ class CacheSender::Impl if (cancelledResponse.has_value()) { + auto const cancellationTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; failResponse(*cancelledResponse, std::make_exception_ptr(TLLM_REQUEST_EXCEPTION(reqId, common::RequestErrorCode::kNETWORK_ERROR, "KV cache transfer for request %zu was cancelled", reqId))); + logTransferDiagnostic( + "sender", "cancelled", requestId, contextRequestId, "wait-request-info", cancellationTime); } if (!allCounterpartsReady) { return; } + logRequestInfoCompleteDiagnostic(reqId, requestId, contextRequestId, isReady); + // Keep mCurrentRequest set while notifying the peer so cancellation cannot change the decision after it has // been made. The network operation must not run under mSenderMutex. sendReadySignal(reqId, isReady); @@ -948,7 +1178,7 @@ class CacheSender::Impl } } - void sendSyncFromReuseTree(RequestIdType requestId) + std::pair sendSyncFromReuseTree(RequestIdType requestId) { TransferSession* session = nullptr; { @@ -958,7 +1188,18 @@ class CacheSender::Impl session = std::addressof(it->second); } // READY was already sent by response(); the receiver consumes exactly one per transfer. - mCacheTransferLayer.format(*session); + auto const serviceStartTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + try + { + mCacheTransferLayer.format(*session); + } + catch (...) + { + logSenderServiceStartDiagnostic(requestId, requestId, "reuse-tree", serviceStartTime); + throw; + } + auto const completedTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + return {serviceStartTime, completedTime}; } // Pin the requested chain in the reuse tree; an empty result means no full match. @@ -1023,15 +1264,18 @@ class CacheSender::Impl auto pinnedIds = pinReuseTreeBlocks(reqId); if (pinnedIds.empty()) { + logRequestInfoCompleteDiagnostic(reqId, reqId, reqId, false); TLLM_LOG_ERROR( "Requested blocks do not exist in the source's reuse tree (request id: %lu). Notifying " "receiver.", reqId); sendReadySignal(reqId, false); discardTransferState(reqId); + logTransferDiagnostic("sender", "failed", reqId, reqId, "reuse-tree"); } else { + logRequestInfoCompleteDiagnostic(reqId, reqId, reqId, true); sendReadySignal(reqId, true); std::promise promise; // Id-only response: the reuse-tree path has no LlmRequest. @@ -1134,7 +1378,12 @@ class CacheSender::Impl = std::make_exception_ptr(std::runtime_error("CacheSender terminated before asynchronous send completed")); for (auto& response : pendingAsyncResponses) { + auto const* request = response.getRequest(); + auto const requestId = response.getRequestId(); + auto const contextRequestId = getDiagnosticContextRequestId(request, requestId); + auto const terminalTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; failResponse(response, exception); + logTransferDiagnostic("sender", "failed", requestId, contextRequestId, "shutdown", terminalTime); } } @@ -1163,7 +1412,11 @@ class CacheSender::Impl } for (auto& entry : pendingResponses) { + auto const* request = entry.second.getRequest(); + auto const contextRequestId = getDiagnosticContextRequestId(request, entry.first); + auto const terminalTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; failResponse(entry.second, exception); + logTransferDiagnostic("sender", "failed", entry.first, contextRequestId, "shutdown", terminalTime); } } @@ -1486,6 +1739,8 @@ class CacheReceiver::Impl } } + logReceiverRequestInfoSubmittedDiagnostic(llmRequest.mRequestId, requestId, allCounterparts.size()); + auto const& resource = getReceiveCacheResource(llmRequest); TransferSession session = perRequestCancel != nullptr ? TransferSession(std::move(allConnections), @@ -1606,6 +1861,11 @@ class CacheReceiver::Impl std::lock_guard lg(mInFlightCancelMutex); mInFlightCancelFlags.erase(*queuedCancelledReqId); } + if (queuedCancelledReqId.has_value()) + { + logTransferDiagnostic("receiver", "cancelled", llmRequest.mRequestId, + getDiagnosticContextRequestId(&llmRequest, llmRequest.mRequestId), "receive-queue"); + } if (!isCancelled && common::getEnvDisaggEnableInflightCancel()) { std::lock_guard lg(mInFlightCancelMutex); @@ -1631,6 +1891,18 @@ class CacheReceiver::Impl kCancelled, }; + static char const* readySignalResultName(ReadySignalResult result) noexcept + { + switch (result) + { + case ReadySignalResult::kReady: return "ready"; + case ReadySignalResult::kNotReady: return "not-ready"; + case ReadySignalResult::kMixed: return "mixed"; + case ReadySignalResult::kCancelled: return "cancelled"; + } + return "unknown"; + } + ReadySignalResult receiveReadySignalDetailed(TransferSession& session, std::atomic const& perRequestCancel) { bool isReady = false; @@ -1725,6 +1997,9 @@ class CacheReceiver::Impl auto const requestId = llmRequest.mRequestId; auto const contextRequestId = llmRequest.getContextPhaseParams().value().getReqId(); char const* phase = "request-info"; + bool cancellationObserved = false; + bool readyDiagnosticPending = false; + double readyTime = 0.0; TLLM_LOG_DEBUG("KV cache receive request %zu, context request %zu started.", requestId, contextRequestId); if (llmRequest.getKvCacheTransferStart() == LlmRequest::TimePoint{}) { @@ -1736,6 +2011,7 @@ class CacheReceiver::Impl { if (perRequestCancel.load(std::memory_order_relaxed) || mTerminate.load(std::memory_order_relaxed)) { + cancellationObserved = true; TLLM_THROW("KV cache receive request %zu cancelled before request-info", requestId); } TLLM_CUDA_CHECK(cudaSetDevice(mDeviceId)); @@ -1753,8 +2029,12 @@ class CacheReceiver::Impl auto readyResult = receiveReadySignalDetailed(*session, perRequestCancel); TLLM_LOG_DEBUG("KV cache receive request %zu, context request %zu phase=%s end: result=%d.", requestId, contextRequestId, phase, static_cast(readyResult)); + readyTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + readyDiagnosticPending = readyResult == ReadySignalResult::kReady; if (readyResult == ReadySignalResult::kCancelled) { + logReceiverReadyDiagnostic(requestId, contextRequestId, readySignalResultName(readyResult), readyTime); + cancellationObserved = true; if (common::getEnvDisaggEnableInflightCancel()) { session->poisonReservedRecvBuffers(); @@ -1763,11 +2043,13 @@ class CacheReceiver::Impl } if (readyResult == ReadySignalResult::kNotReady) { + logReceiverReadyDiagnostic(requestId, contextRequestId, readySignalResultName(readyResult), readyTime); session->releaseReservedRecvBuffers(); TLLM_THROW("KV cache receive request %zu was rejected by the context peer", requestId); } if (readyResult == ReadySignalResult::kMixed) { + logReceiverReadyDiagnostic(requestId, contextRequestId, readySignalResultName(readyResult), readyTime); if (common::getEnvDisaggEnableInflightCancel()) { session->poisonReservedRecvBuffers(); @@ -1786,58 +2068,48 @@ class CacheReceiver::Impl receiveSync(*session); TLLM_LOG_DEBUG( "KV cache receive request %zu, context request %zu phase=%s end.", requestId, contextRequestId, phase); + auto const completedTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; llmRequest.setKvCacheTransferEnd(LlmRequest::getSteadyClockNow()); + logReceiverReadyDiagnostic(requestId, contextRequestId, "ready", readyTime); + readyDiagnosticPending = false; + logTransferDiagnostic("receiver", "completed", requestId, contextRequestId, phase, completedTime); } catch (std::exception const& err) { + auto const terminalTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + if (readyDiagnosticPending) + { + logReceiverReadyDiagnostic(requestId, contextRequestId, "ready", readyTime); + readyDiagnosticPending = false; + } if (common::getEnvDisaggEnableInflightCancel() && session.has_value()) { session->poisonReservedRecvBuffers(); } llmRequest.setKvCacheTransferEnd(LlmRequest::getSteadyClockNow()); - if (isDisaggTransferDiagnosticsEnabled()) - { - try - { - std::lock_guard lock(getDisaggDiagnosticsLogMutex()); - TLLM_LOG_INFO( - "[DISAGG_DIAG][receiver-transfer] t=%.9f rank=%d action=failed request=%zu " - "context_request=%zu phase=%s", - getSteadyClockTimeSeconds(), mpi::MpiComm::world().getRank(), requestId, contextRequestId, - phase); - } - catch (...) - { - // Preserve the original transfer exception if diagnostics fail. - } - } + auto const* action + = cancellationObserved || perRequestCancel.load(std::memory_order_relaxed) ? "cancelled" : "failed"; + logTransferDiagnostic("receiver", action, requestId, contextRequestId, phase, terminalTime); TLLM_LOG_ERROR("KV cache receive request %zu, context request %zu failed in phase=%s: %s", requestId, contextRequestId, phase, err.what()); throw; } catch (...) { + auto const terminalTime = isDisaggTransferDiagnosticsEnabled() ? getSteadyClockTimeSeconds() : 0.0; + if (readyDiagnosticPending) + { + logReceiverReadyDiagnostic(requestId, contextRequestId, "ready", readyTime); + readyDiagnosticPending = false; + } if (common::getEnvDisaggEnableInflightCancel() && session.has_value()) { session->poisonReservedRecvBuffers(); } llmRequest.setKvCacheTransferEnd(LlmRequest::getSteadyClockNow()); - if (isDisaggTransferDiagnosticsEnabled()) - { - try - { - std::lock_guard lock(getDisaggDiagnosticsLogMutex()); - TLLM_LOG_INFO( - "[DISAGG_DIAG][receiver-transfer] t=%.9f rank=%d action=failed request=%zu " - "context_request=%zu phase=%s", - getSteadyClockTimeSeconds(), mpi::MpiComm::world().getRank(), requestId, contextRequestId, - phase); - } - catch (...) - { - // Preserve the original transfer exception if diagnostics fail. - } - } + auto const* action + = cancellationObserved || perRequestCancel.load(std::memory_order_relaxed) ? "cancelled" : "failed"; + logTransferDiagnostic("receiver", action, requestId, contextRequestId, phase, terminalTime); TLLM_LOG_ERROR( "KV cache receive request %zu, context request %zu failed in phase=%s with an unknown " "exception", diff --git a/scripts/disagg_admission_telemetry.py b/scripts/disagg_admission_telemetry.py index ef5d1aa27036..b64287a7d96c 100644 --- a/scripts/disagg_admission_telemetry.py +++ b/scripts/disagg_admission_telemetry.py @@ -60,6 +60,9 @@ class Admission: candidate_requests: tuple[tuple[str, float], ...] admitted_requests: tuple[str, ...] deferred_requests: tuple[str, ...] + candidate_requests_omitted: int + admitted_requests_omitted: int + deferred_requests_omitted: int @dataclass(frozen=True) @@ -118,6 +121,118 @@ class ReleasePoint: source: str +@dataclass(frozen=True) +class _LifecycleMark: + """One request-scoped lifecycle point in a single clock domain.""" + + time_s: float + request: str + local_request: str | None + tag: str + category: str + action: str + log_source: str + emitter_source: str | None + host: str + role: str + rank: str + clock: str + + +_CTX_ACTIONS = { + "queued", + "send-queued", + "receiver-info-ready", + "credit-received", + "request-info-complete", + "first-write", + "first-write-submitted", + "service-start", +} +_GEN_ACTIONS = { + "capacity-prepared", + "request-info-sent", + "request-info-submitted", + "peer-ready", + "ready", + "local-ready", +} +_TERMINAL_ACTIONS = { + "completed", + "complete", + "reaped", + "failed", + "failure", + "cancelled", + "canceled", + "timeout", + "timed-out", +} +_SUCCESS_TERMINAL_ACTIONS = {"completed", "complete", "reaped"} +_LIFECYCLE_INTERVALS: dict[str, tuple[tuple[str, ...], tuple[str, ...]]] = { + "ctx_queued_to_credit": ( + ("ctx-queued", "sender-queued"), + ("sender-credit",), + ), + "ctx_credit_to_first_write": ( + ("sender-credit",), + ("sender-first-write",), + ), + "ctx_first_write_to_terminal": ( + ("sender-first-write",), + ("sender-terminal", "ctx-terminal"), + ), + "gen_arrival_to_first_gate2": ( + ("gen-arrival",), + ("gate2-seen",), + ), + "gen_arrival_to_activation": ( + ("gen-arrival",), + ("gen-activation",), + ), + "gen_activation_to_first_gate2": ( + ("gen-activation",), + ("gate2-seen",), + ), + "gate2_defer_to_admit": ( + ("gate2-deferred",), + ("gate2-admitted",), + ), + "gate2_admit_to_submit": ( + ("gate2-admitted",), + ("submit", "receiver-submitted"), + ), + "submit_to_request_info": ( + ("submit", "receiver-submitted"), + ("request-info",), + ), + "ready_to_reap": ( + ("local-ready", "receiver-completed"), + ("reap",), + ), + "gen_arrival_to_decode_start": ( + ("gen-arrival",), + ("decode-start",), + ), + "ready_to_decode_start": ( + ("local-ready", "receiver-completed"), + ("decode-start",), + ), + "reap_to_decode_start": ( + ("reap",), + ("decode-start",), + ), +} +_ACTIVE_AGE_BUCKETS = ( + (0.01, "lt_10ms"), + (0.1, "10ms_to_100ms"), + (1.0, "100ms_to_1s"), + (10.0, "1s_to_10s"), + (60.0, "10s_to_60s"), + (math.inf, "gte_60s"), +) + + def parse_diagnostic_line(line: str) -> DiagnosticEvent | None: """Parse one diagnostic line, returning ``None`` when it is unusable. @@ -321,14 +436,20 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: } for source, samples in sorted(aggregate_gaps_by_source.items()) } + lifecycle = _analyze_request_lifecycles(sorted_events) + remaining_work_ground_truth = _analyze_remaining_work_ground_truth(sorted_events) + known_backlog_release = _known_backlog_release_analysis(ranks) return { - "schema_version": 1, + "schema_version": 2, "rank_namespace": "source-path::rank" if namespace_by_source else "rank", "aggregate_scope": "all-input-sources" if namespace_by_source else "single-source", "parsed_event_count": len(sorted_events), "event_counts": dict(sorted(category_counts.items())), "ranks": ranks, + "lifecycle": lifecycle, + "remaining_work_ground_truth": remaining_work_ground_truth, + "known_backlog_release": known_backlog_release, "aggregate": { "completed_service_intervals": len(aggregate_service_intervals), "completed_blocks": aggregate_completed_blocks, @@ -687,6 +808,830 @@ def _analyze_rank( return analysis, service_intervals, selected_gaps +def _analyze_request_lifecycles(events: list[DiagnosticEvent]) -> dict[str, object]: + marks = _collect_lifecycle_marks(events) + timelines: dict[tuple[tuple[str, str, str, str, str], str], list[_LifecycleMark]] = defaultdict( + list + ) + tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]] = defaultdict(set) + for mark in marks: + domain = _lifecycle_domain(mark) + timelines[(domain, mark.request)].append(mark) + tag_domains[(mark.request, mark.tag)].add(domain) + + request_records: list[dict[str, object]] = [] + for (domain, request), request_marks in sorted( + timelines.items(), + key=lambda item: (item[0][0], item[0][1]), + ): + request_marks.sort(key=lambda mark: (mark.time_s, mark.tag)) + marks_by_tag: dict[str, list[_LifecycleMark]] = defaultdict(list) + for mark in request_marks: + marks_by_tag[mark.tag].append(mark) + + intervals = { + name: _evaluate_lifecycle_interval( + request, + domain, + marks_by_tag, + start_tags, + end_tags, + tag_domains, + ) + for name, (start_tags, end_tags) in _LIFECYCLE_INTERVALS.items() + } + first_timestamps = { + tag: tag_marks[0].time_s for tag, tag_marks in sorted(marks_by_tag.items()) if tag_marks + } + local_requests = sorted( + {mark.local_request for mark in request_marks if mark.local_request is not None} + ) + emitter_sources = sorted( + {mark.emitter_source for mark in request_marks if mark.emitter_source is not None} + ) + request_records.append( + { + "request": request, + "local_requests": local_requests, + "log_source": domain[0], + "host": domain[1], + "role": domain[2], + "rank": domain[3], + "clock": domain[4], + "clock_domain": _clock_domain_label(domain), + "emitter_sources": emitter_sources, + "first_timestamps": first_timestamps, + "intervals": intervals, + } + ) + + return { + "clock_domain_policy": ( + "Durations require the same input log source, host, role, rank, and clock. " + "Matching request IDs in another source or clock domain are correlated only " + "for censoring; their raw timestamps are never subtracted." + ), + "correlation_request_policy": ( + "Prefer a nonzero context_request/disaggregated request ID; otherwise use request." + ), + "request_count": len({mark.request for mark in marks}), + "clock_domain_request_count": len(request_records), + "requests": request_records, + "interval_coverage": _lifecycle_interval_coverage( + request_records, + marks, + ), + } + + +def _collect_lifecycle_marks(events: list[DiagnosticEvent]) -> list[_LifecycleMark]: + marks: list[_LifecycleMark] = [] + seen: set[tuple[object, ...]] = set() + for event in events: + category = _normalize_diag_token(event.category) + action = _normalize_diag_token(event.fields.get("action", "")) + role = _event_role(event, category, action) + direct_request = _event_correlation_request(event) + local_request = _event_local_request(event) + + def add( + tag: str, + request: str | None = direct_request, + time_s: float = event.time_s, + detail: str | None = None, + ) -> None: + if request is None or not math.isfinite(time_s) or time_s < 0.0: + return + mark = _LifecycleMark( + time_s=time_s, + request=request, + local_request=local_request if request == direct_request else None, + tag=tag, + category=category, + action=detail or action, + log_source=event.source or "", + emitter_source=event.fields.get("source"), + host=_event_host(event), + role=role, + rank=event.rank, + clock=_event_clock(event), + ) + identity = ( + _lifecycle_domain(mark), + mark.request, + mark.tag, + mark.time_s, + mark.category, + mark.action, + ) + if identity not in seen: + seen.add(identity) + marks.append(mark) + + if category == "gen-arrival": + add("gen-arrival") + if category == "gen-activation": + add("gen-activation") + + if category in {"gate1", "gate-1"}: + for field, tag in ( + ("waiting_requests", "gate1-waiting"), + ("fitting_requests", "gate1-fitting"), + ("blocked_requests", "gate1-blocked"), + ): + for request in _parse_request_ids(event.fields.get(field)): + add(tag, request=request, detail=field) + + if category in {"decision", "admission", "gate2", "gate-2"}: + for field, tag in ( + ("active_requests", "gate2-active"), + ("candidate_requests", "gate2-seen"), + ("waiting_requests", "gate2-seen"), + ("fitting_requests", "gate2-seen"), + ("admitted_requests", "gate2-admitted"), + ("deferred_requests", "gate2-deferred"), + ): + for request in _parse_request_ids(event.fields.get(field)): + add(tag, request=request, detail=field) + if tag in {"gate2-admitted", "gate2-deferred"}: + add("gate2-seen", request=request, detail=field) + if action in {"admit", "admitted"}: + add("gate2-admitted") + add("gate2-seen") + elif action in {"defer", "deferred"}: + add("gate2-deferred") + add("gate2-seen") + + if category == "submit": + submit_time = _as_float(event.fields.get("submit_start_t")) + if submit_time is None or submit_time < 0.0 or submit_time > event.time_s: + submit_time = event.time_s + add("submit", time_s=submit_time) + + if category == "reap": + add("reap") + ready_time = _as_float(event.fields.get("ready_t")) + ready_time_source = event.fields.get("ready_time_source") + if ( + ready_time_source != "cpp-global" + and ready_time is not None + and 0.0 <= ready_time <= event.time_s + ): + add("local-ready", time_s=ready_time, detail="ready_t") + + if category == "gen-service" and action == "decode-start-proxy": + add("decode-start") + + if category == "ctx-transfer": + if action in {"queued", "send-queued"}: + add("ctx-queued") + if action in _TERMINAL_ACTIONS: + add("ctx-terminal") + + if category == "sender-transfer": + if action in {"queued", "send-queued"}: + add("sender-queued") + add("ctx-queued") + if action in { + "credit-received", + "receiver-info-ready", + "request-info-complete", + }: + add("sender-credit") + if action in {"service-start", "first-write", "first-write-submitted"}: + add("sender-first-write") + if action in _TERMINAL_ACTIONS: + add("sender-terminal") + + if category == "receiver-transfer": + if action in {"submitted", "request-info-submitted"}: + add("receiver-submitted") + if action in { + "request-info-sent", + "request-info-submitted", + }: + add("request-info") + if action in {"peer-ready", "ready"}: + add("peer-ready") + if action == "local-ready": + add("local-ready") + if action in _TERMINAL_ACTIONS: + add("receiver-terminal") + if action in _SUCCESS_TERMINAL_ACTIONS: + add("receiver-completed") + + if category == "python-transfer": + if role == "ctx": + if action in {"queued", "send-queued"}: + add("ctx-queued") + add("sender-queued") + if action in { + "credit-received", + "receiver-info-ready", + "request-info-complete", + }: + add("sender-credit") + if action in { + "service-start", + "first-write", + "first-write-submitted", + }: + add("sender-first-write") + if action in _TERMINAL_ACTIONS: + add("sender-terminal") + add("ctx-terminal") + elif role == "gen": + if action in {"submitted", "capacity-prepared"}: + add("receiver-submitted") + if action in {"request-info-sent", "request-info-submitted"}: + add("request-info") + if action == "local-ready": + add("local-ready") + if action in _TERMINAL_ACTIONS: + add("receiver-terminal") + if action in _SUCCESS_TERMINAL_ACTIONS: + add("receiver-completed") + + if category == "receiver-slot" and action in {"release", "released"}: + add("physical-release") + + return sorted( + marks, + key=lambda mark: ( + _lifecycle_domain(mark), + mark.request, + mark.time_s, + mark.tag, + ), + ) + + +def _evaluate_lifecycle_interval( + request: str, + domain: tuple[str, str, str, str, str], + marks_by_tag: dict[str, list[_LifecycleMark]], + start_tags: tuple[str, ...], + end_tags: tuple[str, ...], + tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]], +) -> dict[str, object]: + starts = [mark for tag in start_tags for mark in marks_by_tag.get(tag, ())] + ends = [mark for tag in end_tags for mark in marks_by_tag.get(tag, ())] + starts.sort(key=lambda mark: mark.time_s) + ends.sort(key=lambda mark: mark.time_s) + if not starts and not ends: + return { + "status": "not-applicable", + "duration_s": None, + "censor_reason": None, + } + + start = starts[0] if starts else None + end = ( + next((candidate for candidate in ends if candidate.time_s >= start.time_s), None) + if start is not None + else None + ) + if start is not None and end is not None: + return { + "status": "observed", + "duration_s": end.time_s - start.time_s, + "start_t": start.time_s, + "start_kind": f"{start.category}:{start.action or start.tag}", + "end_t": end.time_s, + "end_kind": f"{end.category}:{end.action or end.tag}", + "censor_reason": None, + } + + if start is None: + reason = _cross_domain_endpoint_reason( + "start", + request, + domain, + start_tags, + tag_domains, + ) + if reason is None: + reason = "missing_start" + return { + "status": "censored", + "duration_s": None, + "start_t": None, + "end_t": ends[0].time_s if ends else None, + "censor_reason": reason, + } + + if ends: + reason = "end_precedes_start" + else: + reason = _cross_domain_endpoint_reason( + "end", + request, + domain, + end_tags, + tag_domains, + ) + if reason is None: + reason = "missing_end" + return { + "status": "censored", + "duration_s": None, + "start_t": start.time_s, + "end_t": None, + "censor_reason": reason, + } + + +def _cross_domain_endpoint_reason( + endpoint: str, + request: str, + domain: tuple[str, str, str, str, str], + tags: tuple[str, ...], + tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]], +) -> str | None: + other_domains = { + candidate + for tag in tags + for candidate in tag_domains.get((request, tag), ()) + if candidate != domain + } + if not other_domains: + return None + if any(candidate[0] != domain[0] for candidate in other_domains): + return f"{endpoint}_in_other_log_source" + return f"{endpoint}_in_other_clock_domain" + + +def _lifecycle_interval_coverage( + request_records: list[dict[str, object]], + marks: list[_LifecycleMark], +) -> dict[str, object]: + records_by_request: dict[str, list[dict[str, object]]] = defaultdict(list) + for record in request_records: + records_by_request[str(record["request"])].append(record) + mark_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]] = defaultdict(set) + for mark in marks: + mark_domains[(mark.request, mark.tag)].add(_lifecycle_domain(mark)) + + coverage: dict[str, object] = {} + for name, (start_tags, end_tags) in _LIFECYCLE_INTERVALS.items(): + durations: list[float] = [] + eligible_requests = 0 + observed_requests = 0 + reasons: Counter[str] = Counter() + for request, records in records_by_request.items(): + start_domains = { + domain for tag in start_tags for domain in mark_domains.get((request, tag), ()) + } + end_domains = { + domain for tag in end_tags for domain in mark_domains.get((request, tag), ()) + } + if not start_domains and not end_domains: + continue + eligible_requests += 1 + request_durations = [ + float(interval["duration_s"]) + for record in records + if isinstance((interval := record["intervals"][name]), dict) + and interval.get("status") == "observed" + and interval.get("duration_s") is not None + ] + if request_durations: + observed_requests += 1 + durations.extend(request_durations) + continue + if start_domains and end_domains: + common_domains = start_domains.intersection(end_domains) + if common_domains: + reasons["end_precedes_start"] += 1 + elif any( + start_domain[0] != end_domain[0] + for start_domain in start_domains + for end_domain in end_domains + ): + reasons["cross_log_source_only"] += 1 + else: + reasons["cross_clock_domain_only"] += 1 + elif start_domains: + reasons["missing_end"] += 1 + else: + reasons["missing_start"] += 1 + coverage[name] = { + "eligible_requests": eligible_requests, + "observed_requests": observed_requests, + "observed_samples": len(durations), + "censored_requests": eligible_requests - observed_requests, + "censor_reasons": dict(sorted(reasons.items())), + "duration_s": _summary(durations), + } + return coverage + + +def _analyze_remaining_work_ground_truth( + events: list[DiagnosticEvent], +) -> dict[str, object]: + marks = _collect_lifecycle_marks(events) + timelines: dict[tuple[tuple[str, str, str, str, str], str], dict[str, list[_LifecycleMark]]] = ( + defaultdict(lambda: defaultdict(list)) + ) + tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]] = defaultdict(set) + for mark in marks: + domain = _lifecycle_domain(mark) + timelines[(domain, mark.request)][mark.tag].append(mark) + tag_domains[(mark.request, mark.tag)].add(domain) + for tags in timelines.values(): + for tag_marks in tags.values(): + tag_marks.sort(key=lambda mark: mark.time_s) + + samples: list[dict[str, object]] = [] + seen: set[tuple[object, ...]] = set() + omission_seen: set[tuple[object, ...]] = set() + active_request_ids_omitted = 0 + for event in events: + category = _normalize_diag_token(event.category) + if category not in {"decision", "admission", "gate2", "gate-2"}: + continue + active_pairs = _parse_request_blocks(event.fields.get("active_requests")) + active_blocks = dict(active_pairs) + active_requests = _parse_request_ids(event.fields.get("active_requests")) + domain = _event_domain(event) + sequence = event.fields.get("sequence") + omitted = _as_int(event.fields.get("active_requests_omitted")) or 0 + omission_identity = (domain, sequence or event.time_s) + if omitted > 0 and omission_identity not in omission_seen: + active_request_ids_omitted += omitted + omission_seen.add(omission_identity) + if not active_requests: + continue + for request in active_requests: + identity = (domain, sequence or event.time_s, request) + if identity in seen: + continue + seen.add(identity) + tags = timelines.get((domain, request), {}) + ready = min( + ( + mark + for tag in ("local-ready", "receiver-completed") + for mark in tags.get(tag, ()) + if mark.time_s >= event.time_s + ), + key=lambda mark: mark.time_s, + default=None, + ) + reap = next( + (mark for mark in tags.get("reap", ()) if mark.time_s >= event.time_s), + None, + ) + prior_submits = [mark for mark in tags.get("submit", ()) if mark.time_s <= event.time_s] + active_age_s = event.time_s - prior_submits[-1].time_s if prior_submits else None + ready_reason = _remaining_endpoint_censor_reason( + request, + domain, + ("local-ready", "receiver-completed"), + event.time_s, + tags, + tag_domains, + ) + reap_reason = _remaining_endpoint_censor_reason( + request, + domain, + "reap", + event.time_s, + tags, + tag_domains, + ) + samples.append( + { + "request": request, + "blocks": active_blocks.get(request), + "sequence": sequence, + "decision_t": event.time_s, + "log_source": domain[0], + "host": domain[1], + "role": domain[2], + "rank": domain[3], + "clock": domain[4], + "clock_domain": _clock_domain_label(domain), + "active_age_s": active_age_s, + "active_age_bucket": _active_age_bucket(active_age_s), + "ready_t": ready.time_s if ready is not None else None, + "ready_kind": ( + f"{ready.category}:{ready.action or ready.tag}" + if ready is not None + else None + ), + "residual_ready_s": ( + ready.time_s - event.time_s if ready is not None else None + ), + "ready_censor_reason": (None if ready is not None else ready_reason), + "reap_t": reap.time_s if reap is not None else None, + "residual_reap_s": (reap.time_s - event.time_s if reap is not None else None), + "reap_censor_reason": (None if reap is not None else reap_reason), + } + ) + + samples.sort( + key=lambda sample: ( + str(sample["log_source"]), + str(sample["clock_domain"]), + float(sample["decision_t"]), + str(sample["request"]), + ) + ) + by_age_bucket: dict[str, object] = {} + bucket_names = [label for _, label in _ACTIVE_AGE_BUCKETS] + ["unknown"] + for bucket in bucket_names: + bucket_samples = [sample for sample in samples if sample["active_age_bucket"] == bucket] + if not bucket_samples: + continue + by_age_bucket[bucket] = { + "samples": len(bucket_samples), + "residual_ready_s": _summary( + [ + float(sample["residual_ready_s"]) + for sample in bucket_samples + if sample["residual_ready_s"] is not None + ] + ), + "residual_reap_s": _summary( + [ + float(sample["residual_reap_s"]) + for sample in bucket_samples + if sample["residual_reap_s"] is not None + ] + ), + "ready_coverage": _remaining_coverage(bucket_samples, "ready"), + "reap_coverage": _remaining_coverage(bucket_samples, "reap"), + } + + return { + "definition": ( + "For each Gate-2 admission snapshot and active request, residual_ready_s and " + "residual_reap_s use only later GEN events in the identical input source and " + "clock domain. CTX timestamps are never used." + ), + "active_decision_samples": len(samples), + "active_request_ids_omitted": active_request_ids_omitted, + "identity_coverage": { + "observed": len(samples), + "omitted": active_request_ids_omitted, + "fraction": _safe_ratio( + len(samples), + len(samples) + active_request_ids_omitted, + ), + }, + "unique_requests": len({str(sample["request"]) for sample in samples}), + "samples": samples, + "residual_ready_s": _summary( + [ + float(sample["residual_ready_s"]) + for sample in samples + if sample["residual_ready_s"] is not None + ] + ), + "residual_reap_s": _summary( + [ + float(sample["residual_reap_s"]) + for sample in samples + if sample["residual_reap_s"] is not None + ] + ), + "active_age_s": _summary( + [ + float(sample["active_age_s"]) + for sample in samples + if sample["active_age_s"] is not None + ] + ), + "ready_coverage": _remaining_coverage(samples, "ready"), + "reap_coverage": _remaining_coverage(samples, "reap"), + "by_active_age_bucket": by_age_bucket, + } + + +def _remaining_endpoint_censor_reason( + request: str, + domain: tuple[str, str, str, str, str], + tag: str | tuple[str, ...], + decision_time_s: float, + tags: dict[str, list[_LifecycleMark]], + tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]], +) -> str: + endpoint_tags = (tag,) if isinstance(tag, str) else tag + endpoint_name = tag if isinstance(tag, str) else "ready" + same_domain = [mark for endpoint_tag in endpoint_tags for mark in tags.get(endpoint_tag, ())] + if any(mark.time_s < decision_time_s for mark in same_domain): + return f"{endpoint_name}_precedes_decision" + other_domains = { + candidate + for endpoint_tag in endpoint_tags + for candidate in tag_domains.get((request, endpoint_tag), ()) + if candidate != domain + } + if any(candidate[0] != domain[0] for candidate in other_domains): + return f"{endpoint_name}_in_other_log_source" + if other_domains: + return f"{endpoint_name}_in_other_clock_domain" + return f"missing_{endpoint_name}" + + +def _remaining_coverage( + samples: list[dict[str, object]], + endpoint: str, +) -> dict[str, object]: + residual_field = f"residual_{endpoint}_s" + reason_field = f"{endpoint}_censor_reason" + observed = sum(sample.get(residual_field) is not None for sample in samples) + reasons = Counter( + str(sample[reason_field]) + for sample in samples + if sample.get(residual_field) is None and sample.get(reason_field) is not None + ) + return { + "eligible": len(samples), + "observed": observed, + "censored": len(samples) - observed, + "censor_reasons": dict(sorted(reasons.items())), + } + + +def _known_backlog_release_analysis(ranks: dict[str, object]) -> dict[str, object]: + samples: list[dict[str, object]] = [] + for rank, rank_analysis in ranks.items(): + if not isinstance(rank_analysis, dict): + continue + release_analysis = rank_analysis.get("release_to_admission") + if not isinstance(release_analysis, dict): + continue + by_source = release_analysis.get("by_source") + if not isinstance(by_source, dict): + continue + for release_source, source_analysis in by_source.items(): + if not isinstance(source_analysis, dict): + continue + source_samples = source_analysis.get("samples") + if not isinstance(source_samples, list): + continue + for sample in source_samples: + if isinstance(sample, dict) and not sample.get("backlog_identity_unknown", True): + samples.append( + { + "rank": rank, + "release_source": release_source, + **sample, + } + ) + + return { + "known_backlog_only": True, + "samples": samples, + "release_samples": len(samples), + "next_decision_coverage": { + "observed": sum(sample.get("decision_gap_s") is not None for sample in samples), + "censored": sum(sample.get("decision_gap_s") is None for sample in samples), + }, + "successful_admission_coverage": { + "observed": sum( + sample.get("successful_admission_gap_s") is not None for sample in samples + ), + "censored": sum(sample.get("successful_admission_gap_s") is None for sample in samples), + }, + "release_to_next_decision_s": _summary( + [ + float(sample["decision_gap_s"]) + for sample in samples + if sample.get("decision_gap_s") is not None + ] + ), + "release_to_successful_admission_s": _summary( + [ + float(sample["successful_admission_gap_s"]) + for sample in samples + if sample.get("successful_admission_gap_s") is not None + ] + ), + } + + +def _parse_request_ids(value: str | None) -> list[str]: + if value is None: + return [] + request_ids: list[str] = [] + for item in value.strip("[](){}").split(","): + token = item.strip() + if not token or token in {"-", "none", "None", "null"}: + continue + request = token.split(":", 1)[0].strip() + if request and request not in {"-", "none", "None", "null"}: + request_ids.append(request) + return request_ids + + +def _event_correlation_request(event: DiagnosticEvent) -> str | None: + for field in ( + "context_request", + "context_request_id", + "disagg_request", + "disagg_request_id", + ): + request = event.fields.get(field) + if request not in {None, "", "-", "-1", "0"}: + return request + for field in ("request", "request_id"): + request = event.fields.get(field) + if request not in {None, "", "-", "-1"}: + return request + return None + + +def _event_local_request(event: DiagnosticEvent) -> str | None: + for field in ("local_request", "local_request_id", "request", "request_id"): + request = event.fields.get(field) + if request not in {None, "", "-", "-1"}: + return request + return None + + +def _normalize_diag_token(value: str) -> str: + return value.strip().lower().replace("_", "-") + + +def _event_role(event: DiagnosticEvent, category: str, action: str) -> str: + explicit_role = _normalize_diag_token(event.fields.get("role", "")) + if explicit_role in {"ctx", "context", "sender"}: + return "ctx" + if explicit_role in {"gen", "generation", "receiver"}: + return "gen" + if category in {"ctx-transfer", "sender-transfer"}: + return "ctx" + if category in { + "gen-arrival", + "gen-activation", + "gate1", + "gate-1", + "decision", + "admission", + "gate2", + "gate-2", + "submit", + "receiver-transfer", + "receiver-slot", + "reap", + "gen-service", + "status-poll", + }: + return "gen" + if category == "python-transfer": + if action in _CTX_ACTIONS: + return "ctx" + if action in _GEN_ACTIONS or action in _TERMINAL_ACTIONS: + return "gen" + return "unknown" + + +def _event_host(event: DiagnosticEvent) -> str: + return ( + event.fields.get("host") + or event.fields.get("hostname") + or event.fields.get("node") + or "" + ) + + +def _event_clock(event: DiagnosticEvent) -> str: + return event.fields.get("clock_domain") or event.fields.get("clock") or "local_steady" + + +def _event_domain(event: DiagnosticEvent) -> tuple[str, str, str, str, str]: + category = _normalize_diag_token(event.category) + action = _normalize_diag_token(event.fields.get("action", "")) + return ( + event.source or "", + _event_host(event), + _event_role(event, category, action), + event.rank, + _event_clock(event), + ) + + +def _lifecycle_domain(mark: _LifecycleMark) -> tuple[str, str, str, str, str]: + return mark.log_source, mark.host, mark.role, mark.rank, mark.clock + + +def _clock_domain_label(domain: tuple[str, str, str, str, str]) -> str: + source, host, role, rank, clock = domain + return f"{source}::host={host}::role={role}::rank={rank}::clock={clock}" + + +def _active_age_bucket(active_age_s: float | None) -> str: + if active_age_s is None or active_age_s < 0.0: + return "unknown" + for upper_bound, label in _ACTIVE_AGE_BUCKETS: + if active_age_s < upper_bound: + return label + return "unknown" + + def _collect_admissions(events: list[DiagnosticEvent]) -> tuple[list[Admission], dict[str, float]]: admissions: list[Admission] = [] request_blocks: dict[str, float] = {} @@ -724,6 +1669,18 @@ def _collect_admissions(events: list[DiagnosticEvent]) -> tuple[list[Admission], candidate_requests=tuple(candidate_requests), admitted_requests=tuple(request for request, _ in admitted_request_blocks), deferred_requests=tuple(request for request, _ in deferred_request_blocks), + candidate_requests_omitted=max( + 0, + _as_int(event.fields.get("candidate_requests_omitted")) or 0, + ), + admitted_requests_omitted=max( + 0, + _as_int(event.fields.get("admitted_requests_omitted")) or 0, + ), + deferred_requests_omitted=max( + 0, + _as_int(event.fields.get("deferred_requests_omitted")) or 0, + ), ) ) admissions.sort(key=lambda admission: admission.time_s) @@ -839,7 +1796,8 @@ def _collect_unsuccessful_requests(events: list[DiagnosticEvent]) -> set[str]: return { request for event in events - if (request := event.fields.get("request")) is not None and _event_outcome(event) is False + if (request := _event_correlation_request(event)) is not None + and _event_outcome(event) is False } @@ -860,14 +1818,30 @@ def _event_outcome(event: DiagnosticEvent) -> bool | None: return False action = event.fields.get("action", "").lower() - if event.category == "receiver-transfer" and action in { + transfer_categories = { + "ctx-transfer", + "gen-transfer", + "sender-transfer", + "receiver-transfer", + "python-transfer", + } + if event.category in transfer_categories and action in { "failed", + "failure", "cancelled", "canceled", "aborted", "timeout", + "timed-out", }: return False + if event.category in transfer_categories and action in { + "completed", + "complete", + "success", + "succeeded", + }: + return True state = event.fields.get("state", "").upper() if any(marker in state for marker in ("ERROR", "FAIL", "CANCEL", "TIMEOUT")): @@ -978,7 +1952,7 @@ def _match_slot_intervals( if acquired.time_s > event.time_s: unmatched_releases += 1 continue - request = acquired.fields.get("request") or event.fields.get("request") + request = _event_correlation_request(acquired) or _event_correlation_request(event) if request is None: continue intervals.append( @@ -1302,7 +2276,7 @@ def _backlog_request_ids_at(admissions: list[Admission], time_s: float) -> set[s (candidate for candidate in reversed(admissions) if candidate.time_s <= time_s), None, ) - if admission is None or admission.deferred <= 0: + if admission is None or admission.deferred <= 0 or admission.deferred_requests_omitted > 0: return set() return set(admission.deferred_requests) @@ -1387,6 +2361,8 @@ def _fixed_multiplier_counterfactual( or admission.active_blocks is None or admission.budget_blocks is None or not admission.candidate_requests + or admission.candidate_requests_omitted > 0 + or admission.deferred_requests_omitted > 0 ): continue diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 98e273e702d8..5c4d533379bf 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -81,6 +81,31 @@ _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED = os.getenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS") == "1" +def _diagnostic_now_s() -> float: + return tensorrt_llm.bindings.steady_clock_now().total_seconds() + + +def _log_python_transfer_diagnostic( + *, + rank: int, + role: str, + action: str, + request: int, + timestamp_s: Optional[float] = None, + **fields: object, +) -> None: + if not _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + return + if timestamp_s is None: + timestamp_s = _diagnostic_now_s() + encoded_fields = " ".join(f"{key}={value}" for key, value in fields.items()) + logger.info( + "[DISAGG_DIAG][python-transfer] " + f"t={timestamp_s:.9f} clock=local_steady runtime=Python rank={rank} " + f"role={role} source=native action={action} request={request} {encoded_fields}" + ) + + @dataclass class RecvReqInfo: sender_req_id: int @@ -397,7 +422,18 @@ def setup_session(self, tx_session: "TxSession"): expected_count = len(self._registrar.get_peer_overlap(peer_ri, peer_ri.dp_rank).ranks) if self._is_req_ready(unique_rid, expected_count): with tx_session.lock: + became_ready = not tx_session.receiver_ready tx_session.receiver_ready = True + if became_ready: + _log_python_transfer_diagnostic( + rank=self._instance_rank, + role="ctx", + action="receiver-info-ready", + request=unique_rid, + local_request=tx_session.request_id, + received=expected_count, + expected=expected_count, + ) return def _get_session(self, unique_rid: Optional[int]) -> Optional["TxSession"]: @@ -595,6 +631,20 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): if timer: timer.record_transfer_start(write_meta.peer_rank) try: + if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + first_write_submit_time_s = _diagnostic_now_s() + if session.mark_first_write_submitted(first_write_submit_time_s): + _log_python_transfer_diagnostic( + rank=self._instance_rank, + role="ctx", + action="first-write-submitted", + request=write_meta.unique_rid, + timestamp_s=first_write_submit_time_s, + local_request=session.request_id, + peer_rank=write_meta.peer_rank, + slice=write_meta.slice_id, + bytes=int(write_meta.sizes.sum()), + ) status = self._agent.submit_transfer_requests(request) if not status.wait(): agent_result = AgentResult.FAILED @@ -1102,6 +1152,15 @@ def _save_peer_req_info(self, peer_transfer_req_info: RecvReqInfo): session = self._get_session(req_info.unique_rid) if session is not None and not session.receiver_ready: session.receiver_ready = True + _log_python_transfer_diagnostic( + rank=self._instance_rank, + role="ctx", + action="receiver-info-ready", + request=req_info.unique_rid, + local_request=session.request_id, + received=expected_transfers, + expected=expected_transfers, + ) def has_all_peer_req_infos(self, unique_rid: int) -> bool: req_info = self._get_first_req_info(unique_rid) @@ -1214,6 +1273,14 @@ def __init__( self._terminal_status: Optional[SessionStatus] = None self.transfer_start_time = None self.transfer_end_time = None + self._first_write_submit_time_s: Optional[float] = None + _log_python_transfer_diagnostic( + rank=self._sender._instance_rank, + role="ctx", + action="session-created", + request=self.disagg_request_id, + local_request=self.request_id, + ) # Must be last: makes session visible to listener thread, # so all attributes above must be initialized first. self._sender.setup_session(self) @@ -1245,6 +1312,21 @@ def status(self) -> SessionStatus: return SessionStatus.TRANSFERRING return SessionStatus.READY if self.receiver_ready else SessionStatus.INIT + @property + def first_write_submit_time_s(self) -> Optional[float]: + with self.lock: + return self._first_write_submit_time_s + + def mark_first_write_submitted(self, timestamp_s: float) -> bool: + """Record the first transfer-agent submission exactly once.""" + if not _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + return False + with self.lock: + if self._first_write_submit_time_s is not None: + return False + self._first_write_submit_time_s = timestamp_s + return True + def send(self, slice: KVSlice) -> None: if self.transfer_start_time is None: self.transfer_start_time = tensorrt_llm.bindings.global_steady_clock_now() @@ -1261,6 +1343,14 @@ def send(self, slice: KVSlice) -> None: task._unique_rid = self.disagg_request_id self.kv_tasks.append(task) req_info_snapshot = dict(self._sender._get_req_info(task._unique_rid) or {}) + _log_python_transfer_diagnostic( + rank=self._sender._instance_rank, + role="ctx", + action="send-queued", + request=self.disagg_request_id, + local_request=self.request_id, + slice=slice_id, + ) self._sender.dispatch_task(task, req_info_snapshot) def send_aux(self) -> AuxSendTask: @@ -1312,6 +1402,13 @@ def cancel(self) -> None: task.fail(exc) if self.aux_task is not None and self.aux_task.status == TaskStatus.INIT: self.aux_task.fail(exc) + _log_python_transfer_diagnostic( + rank=self._sender._instance_rank, + role="ctx", + action="cancelled", + request=self.disagg_request_id, + local_request=self.request_id, + ) # Send outside the lock to avoid holding it during I/O. self._sender.send_cancel_to_receivers(self.disagg_request_id) @@ -1644,6 +1741,20 @@ def dispatch_task(self, task: KVRecvTask): f"dispatch_task: RxSession {task._unique_rid} not found; " "session may have been closed before dispatch" ) + if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + capacity_prepared_time_s = _diagnostic_now_s() + if session.mark_capacity_prepared(capacity_prepared_time_s): + _log_python_transfer_diagnostic( + rank=self._registrar.self_rank_info.instance_rank, + role="gen", + action="capacity-prepared", + request=receiver_req.unique_rid, + timestamp_s=capacity_prepared_time_s, + local_request=session.request_id, + slice=task.slice_id, + expected=task.expected_transfers, + bounced=int(bounced), + ) session.mark_transferring(task.slice_id) # Cache sender endpoints so cancel() can send CANCEL_SESSION to them. session._sender_endpoints.update( @@ -1660,6 +1771,19 @@ def dispatch_task(self, task: KVRecvTask): receiver_req.bounce_dst_base = self._bounce.writer_base(key, i) receiver_req_bytes = receiver_req.to_bytes() self._request_sender_data(peer_infos.sender_endpoints[rank], receiver_req_bytes) + if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + request_info_sent_time_s = _diagnostic_now_s() + if session.mark_request_info_sent(request_info_sent_time_s): + _log_python_transfer_diagnostic( + rank=self._registrar.self_rank_info.instance_rank, + role="gen", + action="request-info-sent", + request=receiver_req.unique_rid, + timestamp_s=request_info_sent_time_s, + local_request=session.request_id, + slice=task.slice_id, + sent=len(peer_overlap.ranks), + ) return @staticmethod @@ -1842,11 +1966,20 @@ def __init__( self.transfer_start_time = None self.transfer_end_time = None self.kv_cache_size_bytes: int = 0 + self._capacity_prepared_time_s: Optional[float] = None + self._request_info_sent_time_s: Optional[float] = None self._kv_tasks: list[KVRecvTask] = [] self._aux_count = 0 self._aux_status: TaskStatus = TaskStatus.INIT self._sender_endpoints: set[str] = set() self.lock = threading.Lock() + _log_python_transfer_diagnostic( + rank=self._receiver._registrar.self_rank_info.instance_rank, + role="gen", + action="session-created", + request=self.disagg_request_id, + local_request=self.request_id, + ) self._receiver.setup_session(self) @property @@ -1881,6 +2014,29 @@ def mark_transferring(self, slice_id: int): with self.lock: self._kv_tasks[slice_id].mark_transferring() + def mark_capacity_prepared(self, timestamp_s: float) -> bool: + if not _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + return False + with self.lock: + if self._capacity_prepared_time_s is not None: + return False + self._capacity_prepared_time_s = timestamp_s + return True + + def mark_request_info_sent(self, timestamp_s: float) -> bool: + if not _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + return False + with self.lock: + if self._request_info_sent_time_s is not None: + return False + self._request_info_sent_time_s = timestamp_s + return True + + @property + def request_info_sent_time_s(self) -> Optional[float]: + with self.lock: + return self._request_info_sent_time_s + @property def kv_transfer_start_time_s(self) -> Optional[float]: start_times = [ @@ -2110,6 +2266,13 @@ def cancel(self) -> None: # A write may still be mid-flight, so quarantine the region rather than freeing # it; this keeps a cancelled transfer from leaking. No-op when bounce is off. self._receiver._bounce.orphan_reservation(rid_slice) + _log_python_transfer_diagnostic( + rank=self._receiver._registrar.self_rank_info.instance_rank, + role="gen", + action="cancelled", + request=self.disagg_request_id, + local_request=self.request_id, + ) # Send outside the lock to avoid holding it during I/O. self._receiver.send_cancel_to_senders(self.disagg_request_id, self._sender_endpoints) diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 4789d81258ff..426995d6671e 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -52,6 +52,10 @@ def _is_disagg_transfer_diagnostics_enabled() -> bool: return _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED +def _diagnostic_now_s() -> float: + return tensorrt_llm.bindings.steady_clock_now().total_seconds() + + def _find_consensus_request_ids(request_ids_all_ranks, sync_size): frequency_map = defaultdict(int) consensus = [] @@ -112,6 +116,7 @@ def __init__( self._send_reqs = {} self._recv_reqs = {} self._diagnostic_ready_rids = set() if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED else None + self._diagnostic_terminal_events = set() if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED else None self._wait_reqs = {} self._page_table = self._transfer_worker.page_table # _slice_num_bytes() is this rank's KV shard, so scale by tp_size to get the request total (kv_cache_size), @@ -123,6 +128,75 @@ def __init__( self._ever_had_send_session: bool = False self._ever_had_recv_session: bool = False + def _log_python_transfer_diagnostic( + self, + *, + role: str, + action: str, + request: int, + timestamp_s: Optional[float] = None, + **fields: object, + ) -> None: + if not _is_disagg_transfer_diagnostics_enabled(): + return + if timestamp_s is None: + timestamp_s = _diagnostic_now_s() + rank = getattr(getattr(self, "_dist", None), "rank", -1) + encoded_fields = " ".join(f"{key}={value}" for key, value in fields.items()) + logger.info( + "[DISAGG_DIAG][python-transfer] " + f"t={timestamp_s:.9f} clock=local_steady runtime=Python rank={rank} " + f"role={role} source=transceiver action={action} request={request} " + f"{encoded_fields}" + ) + + def _log_transfer_terminal( + self, + *, + role: str, + action: str, + rid: int, + session, + req, + ) -> None: + if not _is_disagg_transfer_diagnostics_enabled(): + return + terminal_events = getattr(self, "_diagnostic_terminal_events", None) + if terminal_events is None: + terminal_events = set() + self._diagnostic_terminal_events = terminal_events + event_key = (role, action, rid) + if event_key in terminal_events: + return + terminal_events.add(event_key) + + timestamp_s = _diagnostic_now_s() + fields: dict[str, object] = { + "local_request": getattr(req, "py_request_id", rid), + "status": getattr(getattr(session, "status", None), "value", "unknown"), + } + if role == "ctx": + first_write_time_s = getattr(session, "first_write_submit_time_s", None) + if first_write_time_s is not None: + fields["first_write_t"] = f"{first_write_time_s:.9f}" + fields["elapsed_from_first_write_ms"] = ( + f"{(timestamp_s - first_write_time_s) * 1000:.6f}" + ) + else: + request_info_sent_time_s = getattr(session, "request_info_sent_time_s", None) + ready_time_s = getattr(session, "kv_ready_time_s", None) + if request_info_sent_time_s is not None: + fields["request_info_sent_t"] = f"{request_info_sent_time_s:.9f}" + if ready_time_s is not None: + fields["ready_t"] = f"{ready_time_s:.9f}" + self._log_python_transfer_diagnostic( + role=role, + action=action, + request=rid, + timestamp_s=timestamp_s, + **fields, + ) + def _broadcast_instance_name(self) -> str: if self._dist.rank == 0: name = str(uuid.uuid4()) @@ -577,10 +651,32 @@ def request_and_receive_sync(self, req: LlmRequest): self._apply_aux(session, req) self._assert_disagg_history_declared(req) req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + self._log_transfer_terminal( + role="gen", + action="completed", + rid=rid, + session=session, + req=req, + ) else: req.state = LlmRequestState.DISAGG_TRANS_ERROR + self._log_transfer_terminal( + role="gen", + action="failed", + rid=rid, + session=session, + req=req, + ) except Exception: req.state = LlmRequestState.DISAGG_TRANS_ERROR + if session is not None: + self._log_transfer_terminal( + role="gen", + action="failed", + rid=rid, + session=session, + req=req, + ) raise finally: if session is not None: @@ -652,6 +748,32 @@ def check_context_transfer_status( to_process, cancelled, failed, completed, timed_out ) + if _is_disagg_transfer_diagnostics_enabled(): + for action, rids in ( + ("cancelled", cancelled), + ("failed", failed), + ("completed", completed), + ): + for rid in rids: + self._log_transfer_terminal( + role="ctx", + action=action, + rid=rid, + session=self._send_sessions[rid], + req=self._send_reqs[rid], + ) + for rid in timed_out: + session = self._send_sessions[rid] + req = self._send_reqs[rid] + self._log_python_transfer_diagnostic( + role="ctx", + action="wait-timeout", + request=rid, + local_request=req.py_request_id, + status=getattr(session.status, "value", "unknown"), + timeout_ms=self._sender_future_timeout_ms, + ) + for rid in cancelled: self._send_sessions[rid].close() del self._send_reqs[rid] @@ -688,36 +810,41 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): if rid in diagnostic_ready_rids: continue session = self._recv_sessions[rid] - service_start_time = getattr(session, "kv_transfer_start_time_s", None) + request_info_sent_time = getattr(session, "request_info_sent_time_s", None) ready_time = getattr(session, "kv_ready_time_s", None) if ready_time is None: diagnostic_ready_rids.add(rid) - logger.warning( - "[DISAGG_DIAG][python-transfer] " - f"rank={self._dist.rank} action=missing-native-timestamp " - f"request={self._recv_reqs[rid].py_request_id}" + req = self._recv_reqs[rid] + self._log_python_transfer_diagnostic( + role="gen", + action="missing-native-timestamp", + request=rid, + local_request=req.py_request_id, ) continue req = self._recv_reqs[rid] - if service_start_time is not None: - req.py_kv_transfer_service_start_time_s = service_start_time + if request_info_sent_time is not None: + req.py_kv_transfer_service_start_time_s = request_info_sent_time + req.py_kv_request_info_sent_time_s = request_info_sent_time req.py_kv_transfer_ready_time_s = ready_time diagnostic_ready_rids.add(rid) - service_ms = ( - (ready_time - service_start_time) * 1000 - if service_start_time is not None + receive_ms = ( + (ready_time - request_info_sent_time) * 1000 + if request_info_sent_time is not None else -1.0 ) - diagnostic_service_start_time = ( - service_start_time if service_start_time is not None else -1.0 + diagnostic_request_info_sent_time = ( + request_info_sent_time if request_info_sent_time is not None else -1.0 ) - logger.info( - "[DISAGG_DIAG][python-transfer] " - f"t={ready_time:.9f} rank={self._dist.rank} " - f"action=local-ready request={req.py_request_id} " - f"bytes={getattr(req, 'py_kv_cache_xfer_bytes', 0)} " - f"service_start_t={diagnostic_service_start_time:.9f} " - f"service_ms={service_ms:.3f}" + self._log_python_transfer_diagnostic( + role="gen", + action="local-ready", + request=rid, + timestamp_s=ready_time, + local_request=req.py_request_id, + bytes=getattr(req, "py_kv_cache_xfer_bytes", 0), + request_info_sent_t=f"{diagnostic_request_info_sent_time:.9f}", + receive_ms=f"{receive_ms:.3f}", ) to_process = self._build_to_process( self._recv_sessions, @@ -752,6 +879,21 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): to_process, cancelled, failed, completed ) + if _is_disagg_transfer_diagnostics_enabled(): + for action, rids in ( + ("cancelled", cancelled), + ("failed", failed), + ("completed", completed), + ): + for rid in rids: + self._log_transfer_terminal( + role="gen", + action=action, + rid=rid, + session=self._recv_sessions[rid], + req=self._recv_reqs[rid], + ) + cancelled_reqs = [] for rid in cancelled: cancelled_reqs.append(self._recv_reqs[rid]) diff --git a/tensorrt_llm/_torch/pyexecutor/llm_request.py b/tensorrt_llm/_torch/pyexecutor/llm_request.py index 88513bc326d6..e54f30b94f41 100644 --- a/tensorrt_llm/_torch/pyexecutor/llm_request.py +++ b/tensorrt_llm/_torch/pyexecutor/llm_request.py @@ -1207,6 +1207,11 @@ def executor_request_to_llm_request( llm_request.py_disaggregated_params = getattr(executor_request, "py_disaggregated_params", None) + disagg_executor_arrival_time = getattr( + executor_request, "py_disagg_gen_executor_arrival_time_s", None) + if disagg_executor_arrival_time is not None: + llm_request.py_disagg_gen_executor_arrival_time_s = ( + disagg_executor_arrival_time) llm_request.py_conversation_params = getattr(executor_request, "py_conversation_params", None) if child_req_ids: diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 7f4fe5b553f6..7edfdfd0fef1 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -10,6 +10,7 @@ import traceback from contextlib import contextmanager from enum import IntEnum +from itertools import islice from queue import Queue from typing import (TYPE_CHECKING, Callable, Dict, Iterable, List, Optional, Tuple, Union) @@ -29,6 +30,7 @@ is_trace_enabled, mpi_comm, mpi_disabled, nvtx_range, set_thread_local_mpi_comm, trace_func) +from tensorrt_llm.bindings import global_steady_clock_now from tensorrt_llm.bindings.executor import (DisServingRequestStats, FinishReason, InflightBatchingStats, IterationStats, KvCacheStats, @@ -118,16 +120,22 @@ def _stats_buffer_is_unbounded(max_stats_len: int) -> bool: _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED = ( os.getenv(_DISAGG_TRANSFER_DIAGNOSTICS_ENV) == "1") _DISAGG_STATUS_POLL_LOG_THRESHOLD_MS = 1.0 +_DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT = 64 def _is_disagg_transfer_diagnostics_enabled() -> bool: return _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED +def _get_global_steady_clock_now_in_seconds() -> float: + return global_steady_clock_now().total_seconds() + + def _format_disagg_diag_request_blocks( - request_blocks: Iterable[Tuple[int, int]]) -> str: + request_blocks: Iterable[Tuple[int, int]], + limit: int = _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT) -> str: encoded = ",".join(f"{request_id}:{blocks}" - for request_id, blocks in request_blocks) + for request_id, blocks in islice(request_blocks, limit)) return encoded or "-" @@ -3438,8 +3446,55 @@ def _log_disagg_transfer_diagnostic(self, category: str, **fields) -> None: @staticmethod def _disagg_diag_request_id(request: LlmRequest) -> int: + params = getattr(request, "py_disaggregated_params", None) + disagg_request_id = getattr(params, "disagg_request_id", + None) if params is not None else None + if isinstance(disagg_request_id, int): + return disagg_request_id return request.py_request_id + def _log_disagg_gen_ingress( + self, request_items: Iterable[RequestQueueItem]) -> None: + if not _is_disagg_transfer_diagnostics_enabled(): + return + + ingress_time = get_steady_clock_now_in_seconds() + for item in request_items: + request = item.request + if (request is None or request.request_type + != RequestType.REQUEST_TYPE_GENERATION_ONLY): + continue + request.py_disagg_gen_executor_arrival_time_s = ingress_time + self._log_disagg_transfer_diagnostic( + "gen-arrival", + t=ingress_time, + request=item.id, + boundary="executor-queue-to-waiting-queue") + + def _log_disagg_gen_activations(self, + requests: Iterable[LlmRequest]) -> None: + if not _is_disagg_transfer_diagnostics_enabled(): + return + + activation_time = get_steady_clock_now_in_seconds() + for request in requests: + if (getattr(request, "state", None) + != LlmRequestState.DISAGG_GENERATION_INIT): + continue + arrival_time = getattr(request, + "py_disagg_gen_executor_arrival_time_s", + None) + request.py_disagg_gen_executor_activation_time_s = activation_time + self._log_disagg_transfer_diagnostic( + "gen-activation", + t=activation_time, + request=self._disagg_diag_request_id(request), + local_request=request.py_request_id, + ingress_to_activation_ms=( + f"{(activation_time - arrival_time) * 1000:.6f}" + if isinstance(arrival_time, (int, float)) else "-1"), + state=getattr(request.state, "name", str(request.state))) + def _disagg_diag_request_blocks( self, requests: Iterable[LlmRequest], controller: DisaggTransferAdmissionController @@ -3459,12 +3514,22 @@ def _apply_disagg_transfer_admission( controller = self._get_disagg_transfer_admission_controller() if not (getattr(self, "kv_cache_transceiver", None) - and controller.enabled() and fitting_disagg_gen_init_requests): + and controller.enabled()): return fitting_disagg_gen_init_requests, False - admission_result = controller.select(self.active_requests, - fitting_disagg_gen_init_requests) - if _is_disagg_transfer_diagnostics_enabled(): + diagnostics_enabled = _is_disagg_transfer_diagnostics_enabled() + has_waiting_gen_init = diagnostics_enabled and any( + getattr(request, "state", None) == + LlmRequestState.DISAGG_GENERATION_INIT + for request in self.active_requests) + if not fitting_disagg_gen_init_requests and not has_waiting_gen_init: + return fitting_disagg_gen_init_requests, False + + decision_time = 0.0 + decision_sequence = 0 + candidate_request_blocks = [] + active_request_blocks = [] + if diagnostics_enabled: decision_time = get_steady_clock_now_in_seconds() decision_sequence = getattr( self, "_disagg_diag_admission_decision_sequence", 0) + 1 @@ -3473,8 +3538,81 @@ def _apply_disagg_transfer_admission( request for request in self.active_requests if request.is_disagg_generation_transmission_in_progress ] + waiting_requests = [ + request for request in self.active_requests + if (getattr(request, "state", None) == + LlmRequestState.DISAGG_GENERATION_INIT) + ] + waiting_request_ids = { + self._disagg_diag_request_id(request) + for request in waiting_requests + } + for request in fitting_disagg_gen_init_requests: + request_id = self._disagg_diag_request_id(request) + if request_id not in waiting_request_ids: + waiting_requests.append(request) + waiting_request_ids.add(request_id) + candidate_request_blocks = self._disagg_diag_request_blocks( fitting_disagg_gen_init_requests, controller) + active_request_blocks = self._disagg_diag_request_blocks( + active_requests, controller) + candidate_request_ids = { + request_id + for request_id, _ in candidate_request_blocks + } + blocked_count = sum( + self._disagg_diag_request_id(request) not in + candidate_request_ids for request in waiting_requests) + gate1_snapshot = (len(waiting_requests), + tuple(candidate_request_blocks), blocked_count) + last_gate1_snapshot = getattr(self, + "_last_disagg_diag_gate1_summary", + None) + if (gate1_snapshot != last_gate1_snapshot + or decision_sequence % 100 == 0): + waiting_request_blocks = self._disagg_diag_request_blocks( + waiting_requests, controller) + blocked_request_blocks = [ + request_block for request_block in waiting_request_blocks + if request_block[0] not in candidate_request_ids + ] + self._log_disagg_transfer_diagnostic( + "gate1", + t=decision_time, + sequence=decision_sequence, + waiting=len(waiting_request_blocks), + waiting_blocks=sum(blocks + for _, blocks in waiting_request_blocks), + waiting_requests=_format_disagg_diag_request_blocks( + waiting_request_blocks), + waiting_requests_omitted=max( + 0, + len(waiting_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), + fitting=len(candidate_request_blocks), + fitting_blocks=sum( + blocks for _, blocks in candidate_request_blocks), + fitting_requests=_format_disagg_diag_request_blocks( + candidate_request_blocks), + fitting_requests_omitted=max( + 0, + len(candidate_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), + blocked=len(blocked_request_blocks), + blocked_blocks=sum(blocks + for _, blocks in blocked_request_blocks), + blocked_requests=_format_disagg_diag_request_blocks( + blocked_request_blocks), + blocked_requests_omitted=max( + 0, + len(blocked_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT)) + self._last_disagg_diag_gate1_summary = gate1_snapshot + + admission_result = controller.select(self.active_requests, + fitting_disagg_gen_init_requests) + if diagnostics_enabled: admitted_request_ids = { self._disagg_diag_request_id(request) for request in admission_result.admitted_requests @@ -3487,8 +3625,6 @@ def _apply_disagg_transfer_admission( request_block for request_block in candidate_request_blocks if request_block[0] not in admitted_request_ids ] - active_request_blocks = self._disagg_diag_request_blocks( - active_requests, controller) candidate_transfer_blocks = sum( blocks for _, blocks in candidate_request_blocks) deferred_transfer_blocks = sum( @@ -3498,13 +3634,38 @@ def _apply_disagg_transfer_admission( t=decision_time, sequence=decision_sequence, runtime=type(self.kv_cache_transceiver).__name__, + active=len(active_request_blocks), active_blocks=admission_result.active_transfer_blocks, + active_requests=_format_disagg_diag_request_blocks( + active_request_blocks), + active_requests_omitted=max( + 0, + len(active_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), candidates=len(candidate_request_blocks), candidate_blocks=candidate_transfer_blocks, + candidate_requests=_format_disagg_diag_request_blocks( + candidate_request_blocks), + candidate_requests_omitted=max( + 0, + len(candidate_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), admitted=len(admitted_request_blocks), admitted_blocks=admission_result.admitted_transfer_blocks, + admitted_requests=_format_disagg_diag_request_blocks( + admitted_request_blocks), + admitted_requests_omitted=max( + 0, + len(admitted_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), deferred=admission_result.deferred_request_count, deferred_blocks=deferred_transfer_blocks, + deferred_requests=_format_disagg_diag_request_blocks( + deferred_request_blocks), + deferred_requests_omitted=max( + 0, + len(deferred_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), budget=controller.max_transfer_blocks) snapshot = (tuple(active_request_blocks), tuple(candidate_request_blocks), @@ -3520,18 +3681,34 @@ def _apply_disagg_transfer_admission( active_blocks=admission_result.active_transfer_blocks, active_requests=_format_disagg_diag_request_blocks( active_request_blocks), + active_requests_omitted=max( + 0, + len(active_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), candidates=len(candidate_request_blocks), candidate_blocks=candidate_transfer_blocks, candidate_requests=_format_disagg_diag_request_blocks( candidate_request_blocks), + candidate_requests_omitted=max( + 0, + len(candidate_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), admitted=len(admitted_request_blocks), admitted_blocks=(admission_result.admitted_transfer_blocks), admitted_requests=_format_disagg_diag_request_blocks( admitted_request_blocks), + admitted_requests_omitted=max( + 0, + len(admitted_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), deferred=admission_result.deferred_request_count, deferred_blocks=deferred_transfer_blocks, deferred_requests=_format_disagg_diag_request_blocks( deferred_request_blocks), + deferred_requests_omitted=max( + 0, + len(deferred_request_blocks) - + _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT), budget=controller.max_transfer_blocks) self._last_disagg_diag_admission = snapshot if admission_result.deferred_request_count > 0: @@ -5066,6 +5243,7 @@ def _fetch_and_enqueue_requests(self, waiting_queue: WaitingQueue, > 1) and self.dist.rank > 0: attach_py_objects_to_requests(new_requests, py_request_objects) + self._log_disagg_gen_ingress(new_requests) waiting_queue.add_requests(new_requests) def _pop_from_waiting_queue( @@ -5247,6 +5425,7 @@ def _respond_if_invalid(request: LlmRequest) -> bool: if not _respond_if_invalid(request) ] + self._log_disagg_gen_activations(validated_requests) self.active_requests.extend(validated_requests) return validated_requests @@ -5729,6 +5908,17 @@ def flag_if_kv_transfer_timed_out(req: LlmRequest, type: str) -> None: f"cache transfer timeout: elapsed {elapsed_time:.0f}ms > " f"kv_transfer_timeout_ms={timeout_ms}ms") req.py_kv_transfer_timed_out = True + if _is_disagg_transfer_diagnostics_enabled(): + category = ("ctx-transfer" + if type == "context" else "gen-transfer") + self._log_disagg_transfer_diagnostic( + category, + action="timeout", + request=self._disagg_diag_request_id(req), + local_request=req.py_request_id, + elapsed_ms=f"{elapsed_time:.3f}", + timeout_ms=timeout_ms, + state=getattr(req.state, "name", str(req.state))) for req in self.async_transfer_manager.requests_in_transfer().values(): flag_if_kv_transfer_timed_out(req, "context") @@ -5977,6 +6167,47 @@ def _prepare_disagg_gen_transmission_complete(self, scheduled_batch): for req in scheduled_batch.generation_requests: if req.is_disagg_generation_transmission_complete: req.state = LlmRequestState.GENERATION_IN_PROGRESS + if _is_disagg_transfer_diagnostics_enabled(): + decode_start_time = get_steady_clock_now_in_seconds() + arrival_time = getattr( + req, "py_disagg_gen_executor_arrival_time_s", None) + ready_time = getattr(req, "py_kv_transfer_ready_time_s", + None) + ready_time_source = "python-local" + ready_comparison_time = decode_start_time + if (not isinstance(ready_time, (int, float)) + or ready_time <= 0): + ready_time = None + if ready_time is None: + transfer_end = getattr(req, "kv_cache_transfer_end", + None) + total_seconds = getattr(transfer_end, "total_seconds", + None) + if callable(total_seconds): + transfer_end_time = total_seconds() + if (isinstance(transfer_end_time, (int, float)) + and transfer_end_time > 0): + ready_time = transfer_end_time + ready_time_source = "cpp-global" + ready_comparison_time = ( + _get_global_steady_clock_now_in_seconds()) + if ready_time is None: + ready_time_source = "unavailable" + self._log_disagg_transfer_diagnostic( + "gen-service", + t=decode_start_time, + action="decode-start-proxy", + boundary="trans-complete-to-generation", + request=self._disagg_diag_request_id(req), + local_request=req.py_request_id, + ready_time_source=ready_time_source, + arrival_to_decode_ms=( + f"{(decode_start_time - arrival_time) * 1000:.6f}" + if isinstance(arrival_time, + (int, float)) else "-1"), + ready_to_decode_ms=( + f"{(ready_comparison_time - ready_time) * 1000:.6f}" + if ready_time is not None else "-1")) req.context_current_position = req.prompt_len if self.kv_cache_transceiver is not None: self.kv_cache_transceiver.commit_blocks_for_reuse(req) @@ -6195,6 +6426,7 @@ def kv_connector_request_finished(req: LlmRequest): req, cache_block_ids): self.async_transfer_manager.start_transfer(req) + diagnostics_enabled = _is_disagg_transfer_diagnostics_enabled() if self.kv_cache_transceiver: for req in scheduled_requests: if req.is_context_only_request and ( @@ -6209,7 +6441,23 @@ def kv_connector_request_finished(req: LlmRequest): # Order is important here: we need to start the transfer before responding # to make sure the blocks are stored for reuse before they are sent. self.async_transfer_manager.start_transfer(req) + if diagnostics_enabled: + submit_start = get_steady_clock_now_in_seconds() self.kv_cache_transceiver.respond_and_send_async(req) + if diagnostics_enabled: + queued_time = get_steady_clock_now_in_seconds() + req.py_disagg_ctx_send_queued_time_s = queued_time + self._log_disagg_transfer_diagnostic( + "ctx-transfer", + t=queued_time, + action="queued", + runtime=type(self.kv_cache_transceiver).__name__, + request=self._disagg_diag_request_id(req), + local_request=req.py_request_id, + submit_start_t=f"{submit_start:.9f}", + submit_call_ms=( + f"{(queued_time - submit_start) * 1000:.6f}"), + state=getattr(req.state, "name", str(req.state))) if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None: req.py_kv_transfer_start_time = time.monotonic() @@ -6268,10 +6516,17 @@ def _check_cache_transfer_errors(self, error_msg_prefix: str): @nvtx_range("_check_disagg_ctx_cache_transfer_status") def _check_disagg_ctx_cache_transfer_status(self, atLeastNum: int = 0): + diagnostics_enabled = _is_disagg_transfer_diagnostics_enabled() + poll_start = (get_steady_clock_now_in_seconds() + if diagnostics_enabled else 0.0) finished_requests, error_requests = self.kv_cache_transceiver.check_context_transfer_status( atLeastNum) + poll_end = (get_steady_clock_now_in_seconds() + if diagnostics_enabled else 0.0) completed_req_ids = set(finished_requests + error_requests) + finished_req_ids = set( + finished_requests) if diagnostics_enabled else None requests_in_transfer = self.async_transfer_manager.requests_in_transfer( ) @@ -6286,6 +6541,32 @@ def _check_disagg_ctx_cache_transfer_status(self, atLeastNum: int = 0): request = requests_in_transfer[request_id] self._end_transfer_and_maybe_terminate(request) + if diagnostics_enabled: + reap_time = get_steady_clock_now_in_seconds() + queued_time = getattr(request, + "py_disagg_ctx_send_queued_time_s", None) + queued_to_reap_ms = ((reap_time - queued_time) * + 1000 if isinstance(queued_time, + (int, float)) else -1.0) + if finished_req_ids is not None and request_id in finished_req_ids: + outcome = "completed" + elif getattr(request, "py_kv_transfer_timed_out", False): + outcome = "timeout" + else: + outcome = "failed" + self._log_disagg_transfer_diagnostic( + "ctx-transfer", + t=reap_time, + action="reaped", + runtime=type(self.kv_cache_transceiver).__name__, + request=self._disagg_diag_request_id(request), + local_request=request.py_request_id, + outcome=outcome, + queued_t=(f"{queued_time:.9f}" if isinstance( + queued_time, (int, float)) else "-1"), + queued_to_reap_ms=f"{queued_to_reap_ms:.6f}", + poll_call_ms=f"{(poll_end - poll_start) * 1000:.6f}", + state=getattr(request.state, "name", str(request.state))) # The set of requests in transfer may have changed since we terminated some requests. requests_in_transfer = self.async_transfer_manager.requests_in_transfer( @@ -6395,20 +6676,22 @@ def _check_disagg_gen_cache_transfer_status(self, atLeastNum: int = 0): ready_time = getattr(request, "py_kv_transfer_ready_time_s", 0.0) ready_time_source = "python-local" + ready_comparison_time = reap_time if not ready_time: transfer_end = getattr(request, "kv_cache_transfer_end", None) if transfer_end is not None: - # This delay is measured entirely inside the local - # executor process. Applying the cross-rank clock - # offset would move the C++ timestamp into rank 0's - # domain and make it incomparable with reap_time. ready_time = max(0.0, transfer_end.total_seconds()) - ready_time_source = "cpp-local" + if ready_time: + ready_time_source = "cpp-global" + ready_comparison_time = ( + _get_global_steady_clock_now_in_seconds()) + else: + ready_time_source = "unavailable" else: ready_time_source = "unavailable" ready_to_reap_ms = (-1.0 if not ready_time else - (reap_time - ready_time) * 1000) + (ready_comparison_time - ready_time) * 1000) self._log_disagg_transfer_diagnostic( "reap", t=reap_time, diff --git a/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_8k1k_con4096_ctx1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml b/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_8k1k_con4096_ctx1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml index b37e2855173b..dd57bb89e4b6 100644 --- a/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_8k1k_con4096_ctx1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml +++ b/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_8k1k_con4096_ctx1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml @@ -35,7 +35,7 @@ environment: trtllm_repo: '' build_wheel: false work_dir: - worker_env_var: TLLM_LOG_LEVEL=INFO TRTLLM_SERVER_DISABLE_GC=1 TRTLLM_WORKER_DISABLE_GC=1 TRTLLM_ENABLE_PDL=1 ENROOT_ALLOW_DEV=yes TLLM_SPEC_DECODE_FORCE_NUM_ACCEPTED_TOKENS=1 + worker_env_var: TLLM_LOG_LEVEL=INFO TRTLLM_SERVER_DISABLE_GC=1 TRTLLM_WORKER_DISABLE_GC=1 TRTLLM_ENABLE_PDL=1 ENROOT_ALLOW_DEV=yes TLLM_SPEC_DECODE_FORCE_NUM_ACCEPTED_TOKENS=1 TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS=1 server_env_var: TRTLLM_SERVER_DISABLE_GC=1 profiling: nsys_on: false diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index 7e87745063fd..1faef2d3317c 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -547,7 +547,15 @@ def test_apply_emits_changed_admission_snapshot(self, monkeypatch): PyExecutor._apply_disagg_transfer_admission(executor, [candidate]) PyExecutor._apply_disagg_transfer_admission(executor, [candidate]) - assert log_info.call_count == 3 + assert log_info.call_count == 4 + gate1_messages = [ + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][gate1]" in call.args[0] + ] + assert len(gate1_messages) == 1 + assert "sequence=1" in gate1_messages[0] + assert "fitting_requests=2:1" in gate1_messages[0] decision_messages = [ call.args[0] for call in log_info.call_args_list @@ -556,6 +564,8 @@ def test_apply_emits_changed_admission_snapshot(self, monkeypatch): assert len(decision_messages) == 2 assert "sequence=1" in decision_messages[0] assert "sequence=2" in decision_messages[1] + assert "active_requests=1:1" in decision_messages[0] + assert "deferred_requests=2:1" in decision_messages[0] message = next( call.args[0] for call in log_info.call_args_list @@ -569,6 +579,34 @@ def test_apply_emits_changed_admission_snapshot(self, monkeypatch): assert "deferred_requests=2:1" in message assert "budget=1" in message + def test_apply_logs_gate1_when_no_request_fits(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor.kv_cache_transceiver = Mock() + waiting = _make_disagg_transfer_request(4, 64) + waiting.state = LlmRequestState.DISAGG_GENERATION_INIT + executor.active_requests = [waiting] + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=64, tokens_per_block=32 + ) + + admitted, wait_for_progress = PyExecutor._apply_disagg_transfer_admission(executor, []) + + assert admitted == [] + assert not wait_for_progress + gate1_message = next( + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][gate1]" in call.args[0] + ) + assert "waiting_requests=4:2" in gate1_message + assert "fitting_requests=-" in gate1_message + assert "blocked_requests=4:2" in gate1_message + def test_apply_missing_controller_preserves_candidates(self): executor = object.__new__(PyExecutor) executor.kv_cache_transceiver = Mock() @@ -649,6 +687,113 @@ def test_gen_only_no_context_bypasses_transfer_budget(self, monkeypatch): @pytest.mark.usefixtures("_clear_disagg_transfer_mode_env") class TestDisaggTransferIdleProgress: + def test_generation_ingress_diagnostics_off_is_noop(self, monkeypatch): + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", False) + now = Mock() + monkeypatch.setattr(py_executor_module, "get_steady_clock_now_in_seconds", now) + executor = object.__new__(PyExecutor) + request = Mock( + spec=["request_type"], + request_type=py_executor_module.RequestType.REQUEST_TYPE_GENERATION_ONLY, + ) + + PyExecutor._log_disagg_gen_ingress(executor, [RequestQueueItem(701, request)]) + + now.assert_not_called() + assert not hasattr(request, "py_disagg_gen_executor_arrival_time_s") + + def test_logs_generation_executor_ingress(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr( + py_executor_module, + "get_steady_clock_now_in_seconds", + lambda: 12.0, + ) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=2) + request = Mock(request_type=py_executor_module.RequestType.REQUEST_TYPE_GENERATION_ONLY) + item = RequestQueueItem(701, request) + + PyExecutor._log_disagg_gen_ingress(executor, [item]) + + assert request.py_disagg_gen_executor_arrival_time_s == 12.0 + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][gen-arrival]" in message + assert "rank=2" in message + assert "request=701" in message + assert "boundary=executor-queue-to-waiting-queue" in message + + def test_logs_generation_executor_activation_with_common_id(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr( + py_executor_module, + "get_steady_clock_now_in_seconds", + lambda: 12.5, + ) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=2) + request = _make_disagg_transfer_request(7, 64) + request.state = LlmRequestState.DISAGG_GENERATION_INIT + request.py_disaggregated_params = Mock(disagg_request_id=701) + request.py_disagg_gen_executor_arrival_time_s = 12.0 + + PyExecutor._log_disagg_gen_activations(executor, [request]) + + assert request.py_disagg_gen_executor_activation_time_s == 12.5 + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][gen-activation]" in message + assert "rank=2" in message + assert "request=701" in message + assert "local_request=7" in message + assert "ingress_to_activation_ms=500.000000" in message + assert "state=DISAGG_GENERATION_INIT" in message + + def test_decode_proxy_uses_global_clock_for_cpp_completion(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr( + py_executor_module, + "get_steady_clock_now_in_seconds", + lambda: 10.0, + ) + monkeypatch.setattr( + py_executor_module, + "_get_global_steady_clock_now_in_seconds", + lambda: 100.003, + ) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=2) + executor.resource_manager = Mock() + executor.resource_manager.resource_managers = {ResourceManagerType.SEQ_SLOT_MANAGER: Mock()} + executor._setup_sampler_step = Mock() + executor.model_engine = Mock(enable_spec_decode=False) + executor.kv_cache_transceiver = None + request = _make_disagg_transfer_request(7, 64) + request.is_disagg_generation_transmission_complete = True + request.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + request.py_disagg_gen_executor_arrival_time_s = 9.0 + request.py_kv_transfer_ready_time_s = None + request.kv_cache_transfer_end = Mock(total_seconds=Mock(return_value=100.0)) + request.context_phase_params = Mock(first_gen_tokens=[], draft_tokens=[]) + request.prompt_len = 64 + scheduled_batch = Mock(generation_requests=[request]) + + PyExecutor._prepare_disagg_gen_transmission_complete(executor, scheduled_batch) + + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][gen-service]" in message + assert "ready_time_source=cpp-global" in message + assert "arrival_to_decode_ms=1000.000000" in message + assert "ready_to_decode_ms=3.000000" in message + def test_gen_transfer_status_polls_active_transfers(self): executor = object.__new__(PyExecutor) executor.active_requests = [_make_disagg_transfer_request(1, 32, in_progress=True)] @@ -692,11 +837,109 @@ def mark_in_progress(req): assert "bytes=4096" in message assert "submit_call_ms=2.000000" in message + def test_context_send_emits_queued_boundary(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + timestamps = iter((10.0, 10.002)) + monkeypatch.setattr( + py_executor_module, + "get_steady_clock_now_in_seconds", + lambda: next(timestamps), + ) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor.kv_cache_manager = Mock() + executor.async_transfer_manager = Mock() + executor.kv_cache_transceiver = Mock(kv_transfer_timeout_ms=None) + executor.kv_connector_manager = None + executor._check_disagg_ctx_cache_transfer_status = Mock() + request = _make_disagg_transfer_request(17, 64) + request.is_context_only_request = True + request.is_context_finished = True + request.is_finished_due_to_length = False + request.is_finished_due_to_cancellation = False + request.state = LlmRequestState.CONTEXT_INIT + + PyExecutor._send_kv_async(executor, [request]) + + executor.kv_cache_transceiver.respond_and_send_async.assert_called_once_with(request) + assert request.py_disagg_ctx_send_queued_time_s == 10.002 + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][ctx-transfer]" in message + assert "action=queued" in message + assert "request=17" in message + assert "submit_call_ms=2.000000" in message + + def test_context_status_emits_completed_reap(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + timestamps = iter((10.0, 10.002, 10.003)) + monkeypatch.setattr( + py_executor_module, + "get_steady_clock_now_in_seconds", + lambda: next(timestamps), + ) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + request = _make_disagg_transfer_request(17, 64) + request.py_disagg_ctx_send_queued_time_s = 9.9 + request.py_kv_transfer_timed_out = False + executor.kv_cache_transceiver = Mock() + executor.kv_cache_transceiver.check_context_transfer_status.return_value = ( + [17], + [], + ) + executor.async_transfer_manager = Mock() + executor.async_transfer_manager.requests_in_transfer.return_value = {17: request} + executor._end_transfer_and_maybe_terminate = Mock() + executor._check_cache_transfer_errors = Mock() + executor._disagg_timed_out_ctx_cancelled_ids = set() + + PyExecutor._check_disagg_ctx_cache_transfer_status(executor, 0) + + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][ctx-transfer]" in message + assert "action=reaped" in message + assert "outcome=completed" in message + assert "queued_to_reap_ms=103.000000" in message + assert "poll_call_ms=2.000000" in message + + def test_transfer_timeout_emits_terminal_boundary(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(py_executor_module.time, "monotonic", lambda: 20.050) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor.kv_cache_transceiver = Mock(kv_transfer_timeout_ms=10) + executor.async_transfer_manager = Mock() + executor.async_transfer_manager.requests_in_transfer.return_value = {} + executor._is_disagg_inflight_cancel_active = Mock(return_value=False) + request = _make_disagg_transfer_request(19, 64, in_progress=True) + request.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + request.py_kv_transfer_start_time = 20.0 + request.py_kv_transfer_timed_out = False + executor.active_requests = [request] + + PyExecutor._check_kv_transfer_timeout(executor) + + assert request.py_kv_transfer_timed_out + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][gen-transfer]" in message + assert "action=timeout" in message + assert "request=19" in message + assert "timeout_ms=10" in message + @pytest.mark.parametrize( ("python_ready_time", "cpp_ready_time", "ready_time_source"), [ (9.5, None, "python-local"), - (0.0, 9.5, "cpp-local"), + (0.0, 9.5, "cpp-global"), ], ) def test_transfer_status_emits_ready_to_reap_delay( @@ -714,6 +957,11 @@ def test_transfer_status_emits_ready_to_reap_delay( monkeypatch.setattr( py_executor_module, "get_steady_clock_now_in_seconds", lambda: next(timestamps) ) + monkeypatch.setattr( + py_executor_module, + "_get_global_steady_clock_now_in_seconds", + lambda: 10.1, + ) executor = object.__new__(PyExecutor) executor.dist = Mock(rank=0) executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( diff --git a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py index c8a75f6d7f21..b9f8ad1423ff 100644 --- a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py +++ b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py @@ -16,6 +16,7 @@ from __future__ import annotations +import threading from dataclasses import dataclass from datetime import timedelta from typing import Optional @@ -39,6 +40,7 @@ @dataclass class _FakeRequest: state: Optional[LlmRequestState] = None + py_request_id: Optional[int] = None class _FakeTransferWorker: @@ -59,6 +61,7 @@ def __init__( is_completed: bool = False, has_failed: bool = False, kv_transfer_start_time_s: Optional[float] = None, + request_info_sent_time_s: Optional[float] = None, kv_ready_time_s: Optional[float] = None, ) -> None: self._rid = rid @@ -67,6 +70,7 @@ def __init__( self._is_completed = is_completed self._has_failed = has_failed self.kv_transfer_start_time_s = kv_transfer_start_time_s + self.request_info_sent_time_s = request_info_sent_time_s self.kv_ready_time_s = kv_ready_time_s self.blocking_calls: list[bool] = [] self.closed = False @@ -110,7 +114,7 @@ def _make_transceiver( ) -> KvCacheTransceiverV2: transceiver = object.__new__(KvCacheTransceiverV2) transceiver._send_sessions = sessions - transceiver._send_reqs = reqs or {rid: _FakeRequest() for rid in sessions} + transceiver._send_reqs = reqs or {rid: _FakeRequest(py_request_id=rid) for rid in sessions} transceiver._sender_future_timeout_ms = 123 # Attributes read by check_context_transfer_status before it processes sessions. transceiver._ever_had_send_session = True @@ -164,6 +168,32 @@ def test_context_transfer_status_bounded_poll_keeps_not_ready_session_queued() - assert transceiver._transfer_worker.sweep_count == 1 +def test_context_wait_timeout_is_nonterminal_diagnostic( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(transceiver_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + session = _FakeSession(rid=14, wait_result=WaitResult.TIMEOUT) + transceiver = _make_transceiver({14: session}) + transceiver._log_python_transfer_diagnostic = Mock() + transceiver._log_transfer_terminal = Mock() + + completed, failed = transceiver.check_context_transfer_status(at_least_request_num=1) + + assert completed == [] + assert failed == [] + assert not session.closed + assert 14 in transceiver._send_sessions + transceiver._log_transfer_terminal.assert_not_called() + transceiver._log_python_transfer_diagnostic.assert_called_once_with( + role="ctx", + action="wait-timeout", + request=14, + local_request=14, + status=SessionStatus.READY.value, + timeout_ms=123, + ) + + def test_context_transfer_status_block_all_uses_blocking_wait() -> None: session = _FakeSession(rid=12, wait_result=WaitResult.COMPLETED) req = _FakeRequest() @@ -271,6 +301,7 @@ def test_gen_transfer_status_stamps_first_local_ready_time( wait_result=WaitResult.COMPLETED, is_completed=True, kv_transfer_start_time_s=10.25, + request_info_sent_time_s=10.25, kv_ready_time_s=12.5, ) request = Mock( @@ -299,14 +330,17 @@ def test_gen_transfer_status_stamps_first_local_ready_time( assert failed == [] assert cancelled == [] assert request.py_kv_transfer_service_start_time_s == 10.25 + assert request.py_kv_request_info_sent_time_s == 10.25 assert request.py_kv_transfer_ready_time_s == 12.5 - message = log_info.call_args.args[0] + message = next( + call.args[0] for call in log_info.call_args_list if "action=local-ready" in call.args[0] + ) assert "[DISAGG_DIAG][python-transfer]" in message assert "rank=2" in message assert "request=21" in message assert "bytes=8192" in message - assert "service_start_t=10.250000000" in message - assert "service_ms=2250.000" in message + assert "request_info_sent_t=10.250000000" in message + assert "receive_ms=2250.000" in message def test_kv_recv_task_records_native_transfer_boundaries( @@ -418,6 +452,25 @@ def test_tx_session_wait_complete_defaults_to_blocking() -> None: assert task.wait_calls == [0.25] +def test_native_session_boundaries_are_recorded_once( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(native_transfer_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + tx_session = object.__new__(TxSession) + tx_session.lock = threading.Lock() + tx_session._first_write_submit_time_s = None + rx_session = object.__new__(RxSession) + rx_session.lock = threading.Lock() + rx_session._request_info_sent_time_s = None + + assert tx_session.mark_first_write_submitted(3.0) + assert not tx_session.mark_first_write_submitted(4.0) + assert tx_session.first_write_submit_time_s == 3.0 + assert rx_session.mark_request_info_sent(5.0) + assert not rx_session.mark_request_info_sent(6.0) + assert rx_session.request_info_sent_time_s == 5.0 + + def test_tx_session_wait_complete_nonblocking_returns_none_without_waiting() -> None: task = _FakeTask(TaskStatus.TRANSFERRING) session = _make_tx_session([task]) diff --git a/tests/unittest/tools/test_disagg_admission_telemetry.py b/tests/unittest/tools/test_disagg_admission_telemetry.py index d22652508d73..2b307d5ad5f6 100644 --- a/tests/unittest/tools/test_disagg_admission_telemetry.py +++ b/tests/unittest/tools/test_disagg_admission_telemetry.py @@ -31,12 +31,25 @@ analyze_events = _TELEMETRY.analyze_events main = _TELEMETRY.main parse_diagnostic_line = _TELEMETRY.parse_diagnostic_line +DiagnosticEvent = _TELEMETRY.DiagnosticEvent def _parse_lines(lines: list[str]): return [event for line in lines if (event := parse_diagnostic_line(line)) is not None] +def _parse_source_line(line: str, source: str): + event = parse_diagnostic_line(line) + assert event is not None + return DiagnosticEvent( + event.category, + event.time_s, + event.rank, + event.fields, + source, + ) + + def test_parse_diagnostic_line_accepts_rank_prefix_and_ignores_malformed_lines(): event = parse_diagnostic_line( "INFO [RANK 3] [DISAGG_DIAG][admission] t=12.5 active_blocks=8 " @@ -258,7 +271,8 @@ def test_failed_transfer_contributes_no_service_or_release_samples(): "manager=0xabc buffer=1", "[DISAGG_DIAG][reap] t=0.7 rank=0 request=9 blocks=4 outcome=failed " "state=DISAGG_TRANS_ERROR", - "[DISAGG_DIAG][receiver-transfer] t=0.8 rank=0 action=failed request=9", + "[DISAGG_DIAG][receiver-transfer] t=0.8 rank=0 action=failed " + "request=7 context_request=9", ] ) @@ -367,3 +381,95 @@ def test_multi_log_block_lookup_is_scoped_by_source(tmp_path, capsys): assert output["ranks"][f"{ctx_log}::rank=1"]["service"]["completed_blocks"] == 4 assert output["ranks"][f"{gen_log}::rank=1"]["service"]["completed_blocks"] == 2 + + +def test_cpp_lifecycle_and_remaining_work_use_same_domain_completion(): + source = "gen-worker.log" + events = [ + _parse_source_line(line, source) + for line in [ + "[DISAGG_DIAG][gen-arrival] t=0.0 rank=0 request=42", + "[DISAGG_DIAG][gen-activation] t=0.1 rank=0 request=42", + "[DISAGG_DIAG][decision] t=0.2 rank=0 sequence=1 " + "active_requests=- candidate_requests=42:8 admitted_requests=42:8 " + "deferred_requests=- admitted=1 deferred=0 budget=8", + "[DISAGG_DIAG][submit] t=0.25 rank=0 request=42 blocks=8", + "[DISAGG_DIAG][receiver-transfer] t=0.3 rank=0 " + "action=request-info-submitted request=7 context_request=42", + "[DISAGG_DIAG][decision] t=0.4 rank=0 sequence=2 " + "active_requests=42:8 candidate_requests=- admitted_requests=- " + "deferred_requests=- admitted=0 deferred=0 budget=8", + "[DISAGG_DIAG][receiver-transfer] t=0.9 rank=0 " + "action=completed request=7 context_request=42", + "[DISAGG_DIAG][reap] t=1.0 rank=0 request=42 blocks=8 outcome=completed", + "[DISAGG_DIAG][gen-service] t=1.05 rank=0 action=decode-start-proxy request=42", + ] + ] + + result = analyze_events(events) + remaining = result["remaining_work_ground_truth"] + sample = remaining["samples"][0] + + assert sample["request"] == "42" + assert sample["active_age_s"] == pytest.approx(0.15) + assert sample["residual_ready_s"] == pytest.approx(0.5) + assert sample["ready_kind"] == "receiver-transfer:completed" + assert sample["residual_reap_s"] == pytest.approx(0.6) + assert remaining["ready_coverage"] == { + "eligible": 1, + "observed": 1, + "censored": 0, + "censor_reasons": {}, + } + lifecycle = result["lifecycle"]["interval_coverage"] + assert lifecycle["gen_arrival_to_activation"]["duration_s"]["p50"] == pytest.approx(0.1) + assert lifecycle["gen_activation_to_first_gate2"]["duration_s"]["p50"] == pytest.approx(0.1) + assert lifecycle["gate2_admit_to_submit"]["duration_s"]["p50"] == pytest.approx(0.05) + assert lifecycle["submit_to_request_info"]["duration_s"]["p50"] == pytest.approx(0.05) + assert lifecycle["ready_to_reap"]["duration_s"]["p50"] == pytest.approx(0.1) + assert lifecycle["gen_arrival_to_decode_start"]["duration_s"]["p50"] == pytest.approx(1.05) + assert lifecycle["ready_to_decode_start"]["duration_s"]["p50"] == pytest.approx(0.15) + assert lifecycle["reap_to_decode_start"]["duration_s"]["p50"] == pytest.approx(0.05) + + +def test_remaining_work_censors_cross_source_endpoints(): + events = [ + _parse_source_line( + "[DISAGG_DIAG][decision] t=1.0 rank=0 sequence=1 " + "active_requests=42:8 admitted=0 deferred=1 budget=8", + "scheduler.log", + ), + _parse_source_line( + "[DISAGG_DIAG][receiver-transfer] t=1.1 rank=0 " + "action=completed request=7 context_request=42", + "receiver.log", + ), + _parse_source_line( + "[DISAGG_DIAG][reap] t=1.2 rank=0 request=42 outcome=completed", + "receiver.log", + ), + ] + + sample = analyze_events(events)["remaining_work_ground_truth"]["samples"][0] + + assert sample["residual_ready_s"] is None + assert sample["ready_censor_reason"] == "ready_in_other_log_source" + assert sample["residual_reap_s"] is None + assert sample["reap_censor_reason"] == "reap_in_other_log_source" + + +def test_cpp_global_ready_timestamp_is_not_subtracted_from_local_reap_clock(): + events = _parse_lines( + [ + "[DISAGG_DIAG][decision] t=0.5 rank=0 sequence=1 " + "active_requests=42:8 admitted=0 deferred=1 budget=8", + "[DISAGG_DIAG][reap] t=1.0 rank=0 request=42 ready_t=100.0 " + "ready_time_source=cpp-global ready_to_reap_ms=2.0 outcome=completed", + ] + ) + + sample = analyze_events(events)["remaining_work_ground_truth"]["samples"][0] + + assert sample["residual_ready_s"] is None + assert sample["ready_censor_reason"] == "missing_ready" + assert sample["residual_reap_s"] == pytest.approx(0.5) From 4a123112bbac9460f23b74c64af67071cc463f0f Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Fri, 24 Jul 2026 13:59:50 -0700 Subject: [PATCH 5/7] [NVBUG 6312828][test] extend cross-role transfer diagnostics Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- scripts/disagg_admission_telemetry.py | 902 +++++++++++++++++- .../_torch/disaggregation/native/transfer.py | 193 +++- .../_torch/disaggregation/transceiver.py | 46 +- tensorrt_llm/_torch/pyexecutor/py_executor.py | 403 +++++++- ...ctx1_tp1_gen1_tp4_eplb0_mtp0_ccb-NIXL.yaml | 2 +- ...1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml | 2 +- .../test_disagg_inflight_cancel_gate.py | 50 + .../_torch/executor/test_py_executor.py | 209 +++- .../test_transceiver_bounded_polling.py | 43 + .../tools/test_disagg_admission_telemetry.py | 365 +++++++ 10 files changed, 2100 insertions(+), 115 deletions(-) diff --git a/scripts/disagg_admission_telemetry.py b/scripts/disagg_admission_telemetry.py index b64287a7d96c..4ee00555f6e5 100644 --- a/scripts/disagg_admission_telemetry.py +++ b/scripts/disagg_admission_telemetry.py @@ -142,14 +142,25 @@ class _LifecycleMark: _CTX_ACTIONS = { "queued", "send-queued", + "timer-start", + "deadline-observed", + "cancel-requested", + "cancel-result", "receiver-info-ready", "credit-received", "request-info-complete", "first-write", "first-write-submitted", "service-start", + "local-complete", + "physical-complete", + "kv-physical-complete", } _GEN_ACTIONS = { + "timer-start", + "deadline-observed", + "cancel-requested", + "cancel-result", "capacity-prepared", "request-info-sent", "request-info-submitted", @@ -165,11 +176,29 @@ class _LifecycleMark: "failure", "cancelled", "canceled", - "timeout", - "timed-out", } _SUCCESS_TERMINAL_ACTIONS = {"completed", "complete", "reaped"} _LIFECYCLE_INTERVALS: dict[str, tuple[tuple[str, ...], tuple[str, ...]]] = { + "ctx_timer_to_credit": ( + ("ctx-timer-start",), + ("sender-credit",), + ), + "ctx_timer_to_deadline": ( + ("ctx-timer-start",), + ("ctx-deadline",), + ), + "ctx_deadline_to_credit": ( + ("ctx-deadline",), + ("sender-credit",), + ), + "ctx_deadline_to_cancel_result": ( + ("ctx-deadline",), + ("ctx-cancel-result",), + ), + "ctx_deadline_to_terminal": ( + ("ctx-deadline",), + ("sender-terminal", "ctx-terminal"), + ), "ctx_queued_to_credit": ( ("ctx-queued", "sender-queued"), ("sender-credit",), @@ -182,6 +211,38 @@ class _LifecycleMark: ("sender-first-write",), ("sender-terminal", "ctx-terminal"), ), + "ctx_first_write_to_kv_physical_complete": ( + ("sender-first-write",), + ("sender-kv-physical-complete",), + ), + "ctx_kv_physical_complete_to_terminal": ( + ("sender-kv-physical-complete",), + ("sender-terminal", "ctx-terminal"), + ), + "ctx_deadline_to_kv_physical_complete": ( + ("ctx-deadline",), + ("sender-kv-physical-complete",), + ), + "gate1_blocked_to_fitting": ( + ("gate1-blocked",), + ("gate1-fitting",), + ), + "gen_arrival_to_first_gate1": ( + ("gen-arrival",), + ("gate1-fitting", "gate1-blocked"), + ), + "gen_timer_to_deadline": ( + ("gen-timer-start",), + ("gen-deadline",), + ), + "gen_deadline_to_cancel_result": ( + ("gen-deadline",), + ("gen-cancel-result",), + ), + "gen_deadline_to_terminal": ( + ("gen-deadline",), + ("receiver-terminal",), + ), "gen_arrival_to_first_gate2": ( ("gen-arrival",), ("gate2-seen",), @@ -309,18 +370,73 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: sources = {event.source for event in sorted_events if event.source is not None} namespace_by_source = len(sources) > 1 category_counts = Counter(event.category for event in sorted_events) + known_instances: dict[tuple[str | None, str, str, str], set[str]] = defaultdict(set) + for event in sorted_events: + instance = event.fields.get("instance", "-") + if instance in {"-", "unknown"}: + continue + category = _normalize_diag_token(event.category) + action = _normalize_diag_token(event.fields.get("action", "")) + known_instances[ + ( + event.source, + _event_host(event), + _event_role(event, category, action), + event.rank, + ) + ].add(instance) + + def resolved_instance(event: DiagnosticEvent) -> str: + instance = event.fields.get("instance", "-") + if instance not in {"-", "unknown"}: + return instance + category = _normalize_diag_token(event.category) + action = _normalize_diag_token(event.fields.get("action", "")) + candidates = known_instances.get( + ( + event.source, + _event_host(event), + _event_role(event, category, action), + event.rank, + ), + set(), + ) + return next(iter(candidates)) if len(candidates) == 1 else instance + + emitters_by_legacy_rank: dict[str, set[tuple[str, str, str]]] = defaultdict(set) + for event in sorted_events: + legacy_rank = f"{event.source}::rank={event.rank}" if namespace_by_source else event.rank + category = _normalize_diag_token(event.category) + action = _normalize_diag_token(event.fields.get("action", "")) + emitters_by_legacy_rank[legacy_rank].add( + ( + _event_host(event), + resolved_instance(event), + _event_role(event, category, action), + ) + ) + split_legacy_ranks = { + legacy_rank + for legacy_rank, emitters in emitters_by_legacy_rank.items() + if len(emitters) > 1 + and any( + host != "" or instance not in {"-", "unknown"} + for host, instance, _ in emitters + ) + } events_by_rank: dict[str, list[DiagnosticEvent]] = defaultdict(list) for event in sorted_events: rank_key = f"{event.source}::rank={event.rank}" if namespace_by_source else event.rank + if rank_key in split_legacy_ranks: + category = _normalize_diag_token(event.category) + action = _normalize_diag_token(event.fields.get("action", "")) + rank_key = ( + f"{rank_key}::host={_event_host(event)}" + f"::instance={resolved_instance(event)}" + f"::role={_event_role(event, category, action)}" + ) events_by_rank[rank_key].append(event) - global_blocks = _collect_global_request_blocks(sorted_events) - blocks_by_source = { - source: _collect_global_request_blocks( - [event for event in sorted_events if event.source == source] - ) - for source in sources - } ranks: dict[str, object] = {} aggregate_service_intervals: list[ServiceInterval] = [] aggregate_selected_gaps: list[dict[str, object]] = [] @@ -339,9 +455,32 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: for rank in sorted(events_by_rank, key=_rank_sort_key): rank_events = events_by_rank[rank] - request_blocks = global_blocks - if namespace_by_source and rank_events: - request_blocks = blocks_by_source.get(rank_events[0].source, {}) + block_scope = sorted_events + if rank_events: + representative = rank_events[0] + instance = resolved_instance(representative) + if instance not in {"-", "unknown"}: + category = _normalize_diag_token(representative.category) + action = _normalize_diag_token(representative.fields.get("action", "")) + role = _event_role(representative, category, action) + block_scope = [ + event + for event in sorted_events + if event.source == representative.source + and _event_host(event) == _event_host(representative) + and resolved_instance(event) == instance + and _event_role( + event, + _normalize_diag_token(event.category), + _normalize_diag_token(event.fields.get("action", "")), + ) + == role + ] + elif namespace_by_source: + block_scope = [ + event for event in sorted_events if event.source == representative.source + ] + request_blocks = _collect_global_request_blocks(block_scope) rank_analysis, rank_intervals, selected_gaps = _analyze_rank(rank_events, request_blocks) ranks[rank] = rank_analysis aggregate_service_intervals.extend(rank_intervals) @@ -437,17 +576,23 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: for source, samples in sorted(aggregate_gaps_by_source.items()) } lifecycle = _analyze_request_lifecycles(sorted_events) + cross_host_correlation = _analyze_cross_host_correlation(sorted_events) remaining_work_ground_truth = _analyze_remaining_work_ground_truth(sorted_events) known_backlog_release = _known_backlog_release_analysis(ranks) return { - "schema_version": 2, - "rank_namespace": "source-path::rank" if namespace_by_source else "rank", + "schema_version": 3, + "rank_namespace": ( + "source-path::rank[::host::instance::role]" + if split_legacy_ranks + else ("source-path::rank" if namespace_by_source else "rank") + ), "aggregate_scope": "all-input-sources" if namespace_by_source else "single-source", "parsed_event_count": len(sorted_events), "event_counts": dict(sorted(category_counts.items())), "ranks": ranks, "lifecycle": lifecycle, + "cross_host_correlation": cross_host_correlation, "remaining_work_ground_truth": remaining_work_ground_truth, "known_backlog_release": known_backlog_release, "aggregate": { @@ -861,10 +1006,16 @@ def _analyze_request_lifecycles(events: list[DiagnosticEvent]) -> dict[str, obje "clock_domain": _clock_domain_label(domain), "emitter_sources": emitter_sources, "first_timestamps": first_timestamps, + "deadline_phase": _classify_deadline_phase(marks_by_tag), "intervals": intervals, } ) + deadline_phase_counts = Counter( + str(record["deadline_phase"]) + for record in request_records + if record["deadline_phase"] is not None + ) return { "clock_domain_policy": ( "Durations require the same input log source, host, role, rank, and clock. " @@ -876,6 +1027,7 @@ def _analyze_request_lifecycles(events: list[DiagnosticEvent]) -> dict[str, obje ), "request_count": len({mark.request for mark in marks}), "clock_domain_request_count": len(request_records), + "deadline_phase_counts": dict(sorted(deadline_phase_counts.items())), "requests": request_records, "interval_coverage": _lifecycle_interval_coverage( request_records, @@ -941,6 +1093,10 @@ def add( ): for request in _parse_request_ids(event.fields.get(field)): add(tag, request=request, detail=field) + if action in {"fitting", "fit"}: + add("gate1-fitting") + elif action in {"blocked", "deferred"}: + add("gate1-blocked") if category in {"decision", "admission", "gate2", "gate-2"}: for field, tag in ( @@ -985,9 +1141,27 @@ def add( if category == "ctx-transfer": if action in {"queued", "send-queued"}: add("ctx-queued") + if action == "timer-start": + add("ctx-timer-start") + elif action in {"deadline-observed", "timeout", "timed-out"}: + add("ctx-deadline") + elif action == "cancel-requested": + add("ctx-cancel-requested") + elif action == "cancel-result": + add("ctx-cancel-result") if action in _TERMINAL_ACTIONS: add("ctx-terminal") + if category == "gen-transfer": + if action == "timer-start": + add("gen-timer-start") + elif action in {"deadline-observed", "timeout", "timed-out"}: + add("gen-deadline") + elif action == "cancel-requested": + add("gen-cancel-requested") + elif action == "cancel-result": + add("gen-cancel-result") + if category == "sender-transfer": if action in {"queued", "send-queued"}: add("sender-queued") @@ -1000,6 +1174,12 @@ def add( add("sender-credit") if action in {"service-start", "first-write", "first-write-submitted"}: add("sender-first-write") + if action in { + "local-complete", + "physical-complete", + "kv-physical-complete", + }: + add("sender-kv-physical-complete") if action in _TERMINAL_ACTIONS: add("sender-terminal") @@ -1037,6 +1217,12 @@ def add( "first-write-submitted", }: add("sender-first-write") + if action in { + "local-complete", + "physical-complete", + "kv-physical-complete", + }: + add("sender-kv-physical-complete") if action in _TERMINAL_ACTIONS: add("sender-terminal") add("ctx-terminal") @@ -1066,6 +1252,429 @@ def add( ) +def _analyze_cross_host_correlation( + events: list[DiagnosticEvent], +) -> dict[str, object]: + """Join CTX and GEN causal points using diagnostic Unix timestamps.""" + points: dict[str, dict[str, list[dict[str, object]]]] = defaultdict(lambda: defaultdict(list)) + + def add( + event: DiagnosticEvent, + tag: str, + request: str | None = None, + ) -> None: + wall_time_s = _event_wall_time(event) + request = request or _event_correlation_request(event) + if wall_time_s is None or request is None: + return + points[request][tag].append( + { + "wall_t": wall_time_s, + "wall_semantics": event.fields.get("wall_semantics", "unknown"), + "local_t": event.time_s, + "local_clock": _event_clock(event), + "host": _event_host(event), + "instance": event.fields.get("instance", "-"), + "role": _event_role( + event, + _normalize_diag_token(event.category), + _normalize_diag_token(event.fields.get("action", "")), + ), + "rank": event.rank, + "log_source": event.source or "", + "category": _normalize_diag_token(event.category), + "action": _normalize_diag_token(event.fields.get("action", "")), + "sequence": event.fields.get("sequence"), + "previous": event.fields.get("previous"), + } + ) + + for event in events: + category = _normalize_diag_token(event.category) + action = _normalize_diag_token(event.fields.get("action", "")) + role = _event_role(event, category, action) + + if category == "gen-arrival": + add(event, "gen-arrival") + elif category == "gen-activation": + add(event, "gen-activation") + elif category in {"gate1", "gate-1"}: + if action in {"fitting", "fit"}: + add(event, "gate1-fitting") + elif action in {"blocked", "deferred"}: + add(event, "gate1-blocked") + for field, tag in ( + ("fitting_requests", "gate1-fitting"), + ("blocked_requests", "gate1-blocked"), + ): + for request in _parse_request_ids(event.fields.get(field)): + add(event, tag, request) + elif category in {"decision", "admission", "gate2", "gate-2"}: + if action in {"admit", "admitted"}: + add(event, "gate2-admitted") + elif action in {"defer", "deferred"}: + add(event, "gate2-deferred") + elif action == "ineligible": + add(event, "gate2-ineligible") + for field, tag in ( + ("admitted_requests", "gate2-admitted"), + ("deferred_requests", "gate2-deferred"), + ): + for request in _parse_request_ids(event.fields.get(field)): + add(event, tag, request) + elif category == "submit": + add(event, "gen-submit") + elif category == "gen-service" and action == "decode-start-proxy": + add(event, "gen-service-start") + elif category == "ctx-transfer": + if action in {"queued", "send-queued"}: + add(event, "ctx-queued") + elif action == "timer-start": + add(event, "ctx-timer-start") + elif action in {"deadline-observed", "timeout", "timed-out"}: + add(event, "ctx-deadline") + elif action == "cancel-result": + add(event, "ctx-cancel-result") + elif action in _TERMINAL_ACTIONS: + add(event, "ctx-terminal") + elif category == "gen-transfer": + if action == "timer-start": + add(event, "gen-timer-start") + elif action in {"deadline-observed", "timeout", "timed-out"}: + add(event, "gen-deadline") + elif action == "cancel-result": + add(event, "gen-cancel-result") + elif category in {"sender-transfer", "python-transfer"} and role == "ctx": + if action in { + "credit-received", + "receiver-info-ready", + "request-info-complete", + }: + add(event, "ctx-receiver-credit") + elif action in { + "service-start", + "first-write", + "first-write-submitted", + }: + add(event, "ctx-first-write") + elif action in { + "local-complete", + "physical-complete", + "kv-physical-complete", + }: + add(event, "ctx-kv-physical-complete") + elif action in _TERMINAL_ACTIONS: + add(event, "ctx-terminal") + elif category in {"receiver-transfer", "python-transfer"} and role == "gen": + if action in {"request-info-sent", "request-info-submitted"}: + add(event, "gen-request-info") + elif action == "local-ready": + add(event, "gen-local-ready") + elif action in _TERMINAL_ACTIONS: + add(event, "gen-terminal") + + records: list[dict[str, object]] = [] + relationship_counts: Counter[str] = Counter() + joined_request_count = 0 + for request, tags in sorted(points.items()): + selected_points = { + tag: _select_cross_host_point(tag, tag_points) + for tag, tag_points in sorted(tags.items()) + if tag_points + } + first_gate1_points = tags.get("gate1-fitting", []) + tags.get("gate1-blocked", []) + if first_gate1_points: + selected_points["gate1-first"] = _select_cross_host_point( + "gate1-first", first_gate1_points + ) + deadline_point = selected_points.get("ctx-deadline") + if deadline_point is not None: + gate2_state = _gate2_state_at_deadline(tags, float(deadline_point["wall_t"])) + if gate2_state is not None: + selected_points["gate2-state-at-ctx-deadline"] = gate2_state + has_ctx = any(tag.startswith("ctx-") for tag in selected_points) + has_gen = any(tag.startswith("gen-") or tag.startswith("gate") for tag in selected_points) + if has_ctx and has_gen: + joined_request_count += 1 + + deadline_relationship = _classify_cross_host_deadline(selected_points) + if deadline_relationship is not None: + relationship_counts[deadline_relationship] += 1 + + records.append( + { + "request": request, + "joined_ctx_gen": has_ctx and has_gen, + "points": selected_points, + "wall_intervals_s": { + name: _wall_interval(selected_points, start_tag, end_tag) + for name, start_tag, end_tag in ( + ("ctx_timer_to_gen_arrival", "ctx-timer-start", "gen-arrival"), + ("ctx_timer_to_first_gate1", "ctx-timer-start", "gate1-first"), + ("ctx_timer_to_first_gate2_defer", "ctx-timer-start", "gate2-deferred"), + ("ctx_timer_to_gate2_admit", "ctx-timer-start", "gate2-admitted"), + ("ctx_timer_to_gen_submit", "ctx-timer-start", "gen-submit"), + ("ctx_timer_to_deadline", "ctx-timer-start", "ctx-deadline"), + ("gate2_defer_to_ctx_deadline", "gate2-deferred", "ctx-deadline"), + ) + }, + "ctx_deadline_relationship": deadline_relationship, + "ctx_deadline_phase_coverage": _cross_host_phase_coverage(selected_points), + "wall_clock_anomalies": _cross_host_wall_anomalies(selected_points), + } + ) + + return { + "clock_policy": ( + "Unix wall timestamps permit CTX/GEN request correlation across " + "hosts but depend on cluster clock synchronization and may step. " + "Boundary-sampled points are preferred; emission points include " + "logger and polling delay. Cross-rank progress boundaries select " + "the latest observed emitter while deadline/defer triggers select " + "the earliest. Emitter coverage is observed, not proof that every " + "required rank logged. Use this view for causal classification and " + "long queue delays, not sub-millisecond service or throughput " + "estimation." + ), + "request_count": len(records), + "joined_ctx_gen_request_count": joined_request_count, + "ctx_deadline_relationship_counts": dict(sorted(relationship_counts.items())), + "requests": records, + } + + +def _select_cross_host_point( + tag: str, + tag_points: list[dict[str, object]], +) -> dict[str, object]: + """Select an earliest trigger or conservative latest progress boundary.""" + points_by_emitter: dict[tuple[object, ...], list[dict[str, object]]] = defaultdict(list) + select_earliest = tag in { + "ctx-queued", + "ctx-timer-start", + "ctx-deadline", + "gate2-deferred", + "gate1-first", + } + for point in tag_points: + emitter = ( + point["log_source"], + point["host"], + point["instance"], + point["role"], + point["rank"], + ) + points_by_emitter[emitter].append(point) + + selector = min if select_earliest else max + emitter_points = [] + for emitter_candidates in points_by_emitter.values(): + boundary_candidates = [ + point for point in emitter_candidates if point["wall_semantics"] == "boundary-sampled" + ] + preferred = boundary_candidates or emitter_candidates + emitter_points.append(selector(preferred, key=lambda point: float(point["wall_t"]))) + selected = dict(selector(emitter_points, key=lambda point: float(point["wall_t"]))) + selected["selection"] = "earliest" if select_earliest else "latest" + selected["observed_emitter_count"] = len(emitter_points) + selected["observed_hosts"] = sorted({str(point["host"]) for point in emitter_points}) + selected["observed_ranks"] = sorted( + {str(point["rank"]) for point in emitter_points}, + key=_rank_sort_key, + ) + selected["wall_semantics_seen"] = sorted({str(point["wall_semantics"]) for point in tag_points}) + return selected + + +def _event_wall_time(event: DiagnosticEvent) -> float | None: + if _normalize_diag_token(event.fields.get("wall_clock", "")) != "unix": + return None + wall_time_s = _as_float(event.fields.get("wall_t")) + if wall_time_s is None or not math.isfinite(wall_time_s): + return None + return wall_time_s + + +def _gate2_state_at_deadline( + tags: dict[str, list[dict[str, object]]], + deadline_t: float, +) -> dict[str, object] | None: + transitions_by_emitter: dict[tuple[object, ...], list[tuple[str, dict[str, object]]]] = ( + defaultdict(list) + ) + for state, tag in ( + ("deferred", "gate2-deferred"), + ("ineligible", "gate2-ineligible"), + ("admitted", "gate2-admitted"), + ): + for point in tags.get(tag, ()): + if float(point["wall_t"]) > deadline_t: + continue + emitter = ( + point["log_source"], + point["host"], + point["instance"], + point["role"], + point["rank"], + ) + transitions_by_emitter[emitter].append((state, point)) + if not transitions_by_emitter: + return None + + emitter_states = [ + max( + transitions, + key=lambda item: ( + float(item[1]["wall_t"]), + {"deferred": 0, "ineligible": 1, "admitted": 2}[item[0]], + ), + ) + for transitions in transitions_by_emitter.values() + ] + observed_states = {state for state, _ in emitter_states} + if len(observed_states) == 1: + state = next(iter(observed_states)) + elif "admitted" in observed_states: + state = "partial-admission" + elif "ineligible" in observed_states: + state = "ineligible" + else: + state = "deferred" + _, point = max(emitter_states, key=lambda item: float(item[1]["wall_t"])) + result = dict(point) + result["state"] = state + result["observed_emitter_count"] = len(emitter_states) + result["observed_emitter_state_counts"] = dict( + sorted(Counter(emitter_state for emitter_state, _ in emitter_states).items()) + ) + result["selection"] = "latest-transition-at-or-before-deadline" + return result + + +def _wall_interval( + points: dict[str, dict[str, object]], + start_tag: str, + end_tag: str, +) -> float | None: + start = points.get(start_tag) + end = points.get(end_tag) + if start is None or end is None: + return None + interval_s = float(end["wall_t"]) - float(start["wall_t"]) + return interval_s if interval_s >= 0.0 else None + + +def _classify_cross_host_deadline( + points: dict[str, dict[str, object]], +) -> str | None: + deadline = points.get("ctx-deadline") + if deadline is None: + return None + deadline_t = float(deadline["wall_t"]) + timer_start = points.get("ctx-timer-start") + if timer_start is not None and deadline_t < float(timer_start["wall_t"]): + return "wall-clock-order-uncertain" + phase_coverage = _cross_host_phase_coverage(points) + if phase_coverage is not None and phase_coverage["timestamp_inversions"]: + return "wall-clock-order-uncertain" + + gate2_admitted = points.get("gate2-admitted") + if gate2_admitted is None or float(gate2_admitted["wall_t"]) > deadline_t: + gate2_state = points.get("gate2-state-at-ctx-deadline") + if gate2_state is not None and gate2_state.get("state") in { + "admitted", + "partial-admission", + }: + return "partial-gate2-admission-before-global-admission" + if gate2_state is not None and gate2_state.get("state") == "deferred": + return "during-gate2-deferral" + if gate2_state is not None and gate2_state.get("state") == "ineligible": + return "gate2-ineligible-at-deadline" + gate1 = points.get("gate1-first") + if gate1 is None or float(gate1["wall_t"]) > deadline_t: + return "before-gen-gate1" + return "before-gen-gate2-admission" + + ordered_milestones = ( + ("gate2-admitted", "after-gen-admission-before-submit"), + ("gen-submit", "after-gen-submit-before-ctx-credit"), + ("ctx-receiver-credit", "after-ctx-credit-before-first-write"), + ("ctx-first-write", "during-kv-physical-transfer"), + ( + "ctx-kv-physical-complete", + "after-kv-physical-complete-before-terminal", + ), + ("ctx-terminal", "after-terminal"), + ) + furthest_phase = "after-gen-admission-before-submit" + for tag, phase in ordered_milestones: + milestone = points.get(tag) + if milestone is not None and float(milestone["wall_t"]) <= deadline_t: + furthest_phase = phase + return furthest_phase + + +def _cross_host_phase_coverage( + points: dict[str, dict[str, object]], +) -> dict[str, object] | None: + deadline = points.get("ctx-deadline") + if deadline is None: + return None + deadline_t = float(deadline["wall_t"]) + ordered_tags = ( + "gate2-admitted", + "gen-submit", + "ctx-receiver-credit", + "ctx-first-write", + "ctx-kv-physical-complete", + "ctx-terminal", + ) + observed = [ + tag for tag in ordered_tags if tag in points and float(points[tag]["wall_t"]) <= deadline_t + ] + if not observed: + return { + "observed_at_or_before_deadline": [], + "missing_before_furthest": [], + "timestamp_inversions": [], + } + furthest_index = max(ordered_tags.index(tag) for tag in observed) + missing = [tag for tag in ordered_tags[:furthest_index] if tag not in observed] + inversions = [] + prior_tag = None + prior_time = None + for tag in observed: + tag_time = float(points[tag]["wall_t"]) + if prior_time is not None and tag_time < prior_time: + inversions.append(f"{prior_tag}->{tag}") + prior_tag = tag + prior_time = tag_time + return { + "observed_at_or_before_deadline": observed, + "missing_before_furthest": missing, + "timestamp_inversions": inversions, + } + + +def _cross_host_wall_anomalies( + points: dict[str, dict[str, object]], +) -> list[str]: + anomalies: list[str] = [] + for start_tag, end_tag in ( + ("ctx-timer-start", "ctx-deadline"), + ("gate2-admitted", "gen-submit"), + ("gen-submit", "ctx-receiver-credit"), + ("ctx-receiver-credit", "ctx-first-write"), + ("ctx-first-write", "ctx-kv-physical-complete"), + ("ctx-kv-physical-complete", "ctx-terminal"), + ): + start = points.get(start_tag) + end = points.get(end_tag) + if start is not None and end is not None and float(end["wall_t"]) < float(start["wall_t"]): + anomalies.append(f"negative:{start_tag}->{end_tag}") + return anomalies + + def _evaluate_lifecycle_interval( request: str, domain: tuple[str, str, str, str, str], @@ -1141,6 +1750,28 @@ def _evaluate_lifecycle_interval( } +def _classify_deadline_phase( + marks_by_tag: dict[str, list[_LifecycleMark]], +) -> str | None: + """Locate a CTX deadline in the sender-side transfer lifecycle.""" + deadlines = marks_by_tag.get("ctx-deadline", ()) + if not deadlines: + return None + deadline_t = deadlines[0].time_s + + phase_boundaries = ( + ("pre-credit", "sender-credit"), + ("sender-queue", "sender-first-write"), + ("kv-transfer-service", "sender-kv-physical-complete"), + ("completion-visibility", "ctx-terminal"), + ) + for phase, boundary_tag in phase_boundaries: + boundaries = marks_by_tag.get(boundary_tag, ()) + if not boundaries or boundaries[0].time_s > deadline_t: + return phase + return "after-terminal" + + def _cross_domain_endpoint_reason( endpoint: str, request: str, @@ -1245,20 +1876,31 @@ def _analyze_remaining_work_ground_truth( samples: list[dict[str, object]] = [] seen: set[tuple[object, ...]] = set() omission_seen: set[tuple[object, ...]] = set() + membership_overflow = _collect_membership_overflow(events) active_request_ids_omitted = 0 + active_request_ids_recovered_from_overflow = 0 for event in events: category = _normalize_diag_token(event.category) - if category not in {"decision", "admission", "gate2", "gate-2"}: + if category not in {"decision", "admission"}: + continue + if "active_requests" not in event.fields and "active_requests_omitted" not in event.fields: continue - active_pairs = _parse_request_blocks(event.fields.get("active_requests")) - active_blocks = dict(active_pairs) - active_requests = _parse_request_ids(event.fields.get("active_requests")) - domain = _event_domain(event) sequence = event.fields.get("sequence") + domain = _event_domain(event) + active_overflow = _membership_overflow_for_event(event, "active", membership_overflow) + active_pairs = _parse_request_blocks(event.fields.get("active_requests")) + active_overflow + active_blocks = dict(active_pairs) + active_requests = [request for request, _ in active_pairs] omitted = _as_int(event.fields.get("active_requests_omitted")) or 0 - omission_identity = (domain, sequence or event.time_s) - if omitted > 0 and omission_identity not in omission_seen: - active_request_ids_omitted += omitted + omission_identity = ( + domain, + event.fields.get("instance", "-"), + sequence or event.time_s, + ) + if omission_identity not in omission_seen: + recovered = min(omitted, len(active_overflow)) + active_request_ids_recovered_from_overflow += recovered + active_request_ids_omitted += max(0, omitted - recovered) omission_seen.add(omission_identity) if not active_requests: continue @@ -1372,6 +2014,7 @@ def _analyze_remaining_work_ground_truth( ), "active_decision_samples": len(samples), "active_request_ids_omitted": active_request_ids_omitted, + "active_request_ids_recovered_from_overflow": (active_request_ids_recovered_from_overflow), "identity_coverage": { "observed": len(samples), "omitted": active_request_ids_omitted, @@ -1567,6 +2210,7 @@ def _event_role(event: DiagnosticEvent, category: str, action: str) -> str: if category in { "gen-arrival", "gen-activation", + "gen-transfer", "gate1", "gate-1", "decision", @@ -1574,6 +2218,7 @@ def _event_role(event: DiagnosticEvent, category: str, action: str) -> str: "gate2", "gate-2", "submit", + "decision-members", "receiver-transfer", "receiver-slot", "reap", @@ -1632,20 +2277,95 @@ def _active_age_bucket(active_age_s: float | None) -> str: return "unknown" +def _membership_snapshot_key( + event: DiagnosticEvent, membership: str, *, definition: bool = False +) -> tuple[tuple[str, str, str, str, str], str, str, str] | None: + reference_field = "snapshot_version" if definition else f"{membership}_snapshot" + reference = event.fields.get(reference_field) or event.fields.get("sequence") + if reference is None: + return None + return ( + _event_domain(event), + event.fields.get("instance", "-"), + reference, + membership, + ) + + +def _membership_overflow_for_event( + event: DiagnosticEvent, + membership: str, + overflow: dict[ + tuple[tuple[str, str, str, str, str], str, str, str], + list[tuple[str, float]], + ], +) -> list[tuple[str, float]]: + key = _membership_snapshot_key(event, membership) + if key is None: + return [] + tail = overflow.get(key, []) + omitted = _as_int(event.fields.get(f"{membership}_requests_omitted")) + if omitted is not None and omitted != len(tail): + return [] + return tail + + def _collect_admissions(events: list[DiagnosticEvent]) -> tuple[list[Admission], dict[str, float]]: admissions: list[Admission] = [] request_blocks: dict[str, float] = {} + membership_overflow = _collect_membership_overflow(events) + decision_keys = { + ( + _event_domain(event), + event.fields.get("instance", "-"), + event.fields.get("sequence"), + ) + for event in events + if _normalize_diag_token(event.category) == "decision" + and event.fields.get("sequence") is not None + } + compatibility_admissions = { + ( + _event_domain(event), + event.fields.get("instance", "-"), + event.fields.get("sequence"), + ): event + for event in events + if _normalize_diag_token(event.category) == "admission" + and event.fields.get("sequence") is not None + } for event in events: - if event.category != "admission": + category = _normalize_diag_token(event.category) + if category not in {"decision", "admission"}: continue - candidate_requests = _parse_request_blocks(event.fields.get("candidate_requests")) - admitted_request_blocks = _parse_request_blocks(event.fields.get("admitted_requests")) - deferred_request_blocks = _parse_request_blocks(event.fields.get("deferred_requests")) - for request, blocks in ( - candidate_requests + admitted_request_blocks + deferred_request_blocks - ): - request_blocks[request] = blocks - + event_key = ( + _event_domain(event), + event.fields.get("instance", "-"), + event.fields.get("sequence"), + ) + if category == "admission" and event_key in decision_keys: + continue + if category == "decision" and event_key in compatibility_admissions: + compatibility_event = compatibility_admissions[event_key] + event = DiagnosticEvent( + category=event.category, + time_s=event.time_s, + rank=event.rank, + fields={**compatibility_event.fields, **event.fields}, + source=event.source, + ) + candidate_overflow = _membership_overflow_for_event(event, "candidate", membership_overflow) + admitted_overflow = _membership_overflow_for_event(event, "admitted", membership_overflow) + deferred_overflow = _membership_overflow_for_event(event, "deferred", membership_overflow) + candidate_requests = ( + _parse_request_blocks(event.fields.get("candidate_requests")) + candidate_overflow + ) + admitted_request_blocks = ( + _parse_request_blocks(event.fields.get("admitted_requests")) + admitted_overflow + ) + deferred_request_blocks = ( + _parse_request_blocks(event.fields.get("deferred_requests")) + deferred_overflow + ) admitted = _as_int(event.fields.get("admitted")) deferred = _as_int(event.fields.get("deferred")) if admitted is None: @@ -1654,6 +2374,35 @@ def _collect_admissions(events: list[DiagnosticEvent]) -> tuple[list[Admission], deferred = len(deferred_request_blocks) if admitted < 0 or deferred < 0: continue + candidate_omitted = max( + 0, + (_as_int(event.fields.get("candidate_requests_omitted")) or 0) + - len(candidate_overflow), + ) + if ( + event.fields.get("candidate_snapshot") is not None + and candidate_omitted == 0 + and len(candidate_requests) >= admitted + deferred + ): + admitted_request_blocks = candidate_requests[:admitted] + deferred_request_blocks = candidate_requests[admitted : admitted + deferred] + admitted_omitted = 0 + deferred_omitted = 0 + else: + admitted_omitted = max( + 0, + (_as_int(event.fields.get("admitted_requests_omitted")) or 0) + - len(admitted_overflow), + ) + deferred_omitted = max( + 0, + (_as_int(event.fields.get("deferred_requests_omitted")) or 0) + - len(deferred_overflow), + ) + for request, blocks in ( + candidate_requests + admitted_request_blocks + deferred_request_blocks + ): + request_blocks[request] = blocks budget = _as_float(event.fields.get("budget")) if budget is not None and budget <= 0.0: budget = None @@ -1669,24 +2418,77 @@ def _collect_admissions(events: list[DiagnosticEvent]) -> tuple[list[Admission], candidate_requests=tuple(candidate_requests), admitted_requests=tuple(request for request, _ in admitted_request_blocks), deferred_requests=tuple(request for request, _ in deferred_request_blocks), - candidate_requests_omitted=max( - 0, - _as_int(event.fields.get("candidate_requests_omitted")) or 0, - ), - admitted_requests_omitted=max( - 0, - _as_int(event.fields.get("admitted_requests_omitted")) or 0, - ), - deferred_requests_omitted=max( - 0, - _as_int(event.fields.get("deferred_requests_omitted")) or 0, - ), + candidate_requests_omitted=candidate_omitted, + admitted_requests_omitted=admitted_omitted, + deferred_requests_omitted=deferred_omitted, ) ) admissions.sort(key=lambda admission: admission.time_s) return admissions, request_blocks +def _collect_membership_overflow( + events: list[DiagnosticEvent], +) -> dict[ + tuple[tuple[str, str, str, str, str], str, str, str], + list[tuple[str, float]], +]: + chunks: dict[ + tuple[tuple[str, str, str, str, str], str, str, str], + dict[int, list[tuple[str, float]]], + ] = defaultdict(dict) + expected_chunk_counts: dict[tuple[tuple[str, str, str, str, str], str, str, str], int] = {} + conflicts: set[tuple[tuple[str, str, str, str, str], str, str, str]] = set() + for event in events: + if _normalize_diag_token(event.category) != "decision-members": + continue + membership = _normalize_diag_token(event.fields.get("membership", "")) + if membership not in {"active", "candidate", "admitted", "deferred"}: + continue + chunk_index = _as_int(event.fields.get("chunk_index")) + chunk_count = _as_int(event.fields.get("chunk_count")) + if ( + chunk_index is None + or chunk_index <= 0 + or chunk_count is None + or chunk_count <= 0 + or chunk_index > chunk_count + ): + continue + key = _membership_snapshot_key(event, membership, definition=True) + if key is None: + continue + previous_count = expected_chunk_counts.setdefault(key, chunk_count) + if previous_count != chunk_count: + conflicts.add(key) + continue + request_blocks = _parse_request_blocks(event.fields.get("requests")) + previous_chunk = chunks[key].get(chunk_index) + if previous_chunk is not None and previous_chunk != request_blocks: + conflicts.add(key) + continue + chunks[key][chunk_index] = request_blocks + + overflow: dict[ + tuple[tuple[str, str, str, str, str], str, str, str], + list[tuple[str, float]], + ] = {} + for key, key_chunks in chunks.items(): + expected = expected_chunk_counts[key] + if key in conflicts or set(key_chunks) != set(range(1, expected + 1)): + continue + request_blocks = [ + request_block + for chunk_index in range(1, expected + 1) + for request_block in key_chunks[chunk_index] + ] + request_ids = [request for request, _ in request_blocks] + if len(request_ids) != len(set(request_ids)): + continue + overflow[key] = request_blocks + return overflow + + def _collect_decisions( events: list[DiagnosticEvent], admissions: list[Admission] ) -> list[Decision]: @@ -1729,10 +2531,18 @@ def _collect_global_request_blocks(events: list[DiagnosticEvent]) -> dict[str, f conflicts: set[str] = set() for event in events: pairs: list[tuple[str, float]] = [] - if event.category == "admission": + category = _normalize_diag_token(event.category) + if category in {"decision", "admission"}: for field in ("candidate_requests", "admitted_requests", "deferred_requests"): pairs.extend(_parse_request_blocks(event.fields.get(field))) - elif event.category in {"submit", "reap"}: + elif category == "decision-members": + pairs.extend(_parse_request_blocks(event.fields.get("requests"))) + elif category in {"gate1", "gate-1", "gate2", "gate-2"}: + request = _event_correlation_request(event) + blocks = _as_float(event.fields.get("blocks")) + if request is not None and blocks is not None and blocks >= 0.0: + pairs.append((request, blocks)) + elif category in {"submit", "reap"}: request = event.fields.get("request") blocks = _as_float(event.fields.get("blocks")) if request is not None and blocks is not None and blocks >= 0.0: @@ -1831,8 +2641,6 @@ def _event_outcome(event: DiagnosticEvent) -> bool | None: "cancelled", "canceled", "aborted", - "timeout", - "timed-out", }: return False if event.category in transfer_categories and action in { diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 5c4d533379bf..478d7b4963ac 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -91,18 +91,29 @@ def _log_python_transfer_diagnostic( role: str, action: str, request: int, + instance: str = "-", timestamp_s: Optional[float] = None, + wall_timestamp_s: Optional[float] = None, + wall_semantics: str = "emission", **fields: object, ) -> None: if not _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: return if timestamp_s is None: timestamp_s = _diagnostic_now_s() + if wall_timestamp_s is None: + wall_timestamp_s = time.time() + host = os.getenv("HOSTNAME", "unknown").replace(" ", "_") + instance = instance.replace(" ", "_") encoded_fields = " ".join(f"{key}={value}" for key, value in fields.items()) logger.info( "[DISAGG_DIAG][python-transfer] " - f"t={timestamp_s:.9f} clock=local_steady runtime=Python rank={rank} " - f"role={role} source=native action={action} request={request} {encoded_fields}" + f"t={timestamp_s:.9f} clock=local_steady " + f"wall_t={wall_timestamp_s:.9f} wall_clock=unix " + f"wall_semantics={wall_semantics} runtime=Python " + f"host={host} instance={instance} rank={rank} role={role} " + f"source=native action={action} " + f"request={request} {encoded_fields}" ) @@ -291,6 +302,8 @@ def __init__( super().__init__(params) self.slice_id = slice_id self.transferred_count = 0 + self.kv_physical_complete_time_s: Optional[float] = None + self.kv_physical_complete_wall_time_s: Optional[float] = None self._slice = kv_slice self._prompt_len = prompt_len self._beam_width = beam_width @@ -424,12 +437,18 @@ def setup_session(self, tx_session: "TxSession"): with tx_session.lock: became_ready = not tx_session.receiver_ready tx_session.receiver_ready = True - if became_ready: + if became_ready and _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + receiver_ready_time_s = _diagnostic_now_s() + receiver_ready_wall_time_s = time.time() _log_python_transfer_diagnostic( rank=self._instance_rank, role="ctx", action="receiver-info-ready", request=unique_rid, + instance=self._registrar.self_rank_info.instance_name, + timestamp_s=receiver_ready_time_s, + wall_timestamp_s=receiver_ready_wall_time_s, + wall_semantics="boundary-sampled", local_request=tx_session.request_id, received=expected_count, expected=expected_count, @@ -603,6 +622,9 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): agent_result = AgentResult.SUCCESS send_slot_id = None + first_write_boundary = None + agent_wait_complete_time_s = None + agent_wait_complete_wall_time_s = None if write_meta.src_ptrs.size > 0: try: request, send_slot_id = build_send_request( @@ -631,21 +653,18 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): if timer: timer.record_transfer_start(write_meta.peer_rank) try: + status = self._agent.submit_transfer_requests(request) if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: first_write_submit_time_s = _diagnostic_now_s() - if session.mark_first_write_submitted(first_write_submit_time_s): - _log_python_transfer_diagnostic( - rank=self._instance_rank, - role="ctx", - action="first-write-submitted", - request=write_meta.unique_rid, - timestamp_s=first_write_submit_time_s, - local_request=session.request_id, - peer_rank=write_meta.peer_rank, - slice=write_meta.slice_id, - bytes=int(write_meta.sizes.sum()), + first_write_submit_wall_time_s = time.time() + if session.mark_first_write_submitted( + first_write_submit_time_s, + first_write_submit_wall_time_s, + ): + first_write_boundary = ( + first_write_submit_time_s, + first_write_submit_wall_time_s, ) - status = self._agent.submit_transfer_requests(request) if not status.wait(): agent_result = AgentResult.FAILED last_status = getattr(status, "last_status_str", lambda: "")() @@ -664,9 +683,31 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): ) logger.error(detail) task.fail(RuntimeError(detail)) + elif _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + agent_wait_complete_time_s = _diagnostic_now_s() + agent_wait_complete_wall_time_s = time.time() finally: if send_slot_id is not None: self._bounce.release_send(send_slot_id) + elif _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + agent_wait_complete_time_s = _diagnostic_now_s() + agent_wait_complete_wall_time_s = time.time() + + if first_write_boundary is not None: + _log_python_transfer_diagnostic( + rank=self._instance_rank, + role="ctx", + action="first-write-submitted", + request=write_meta.unique_rid, + instance=self._registrar.self_rank_info.instance_name, + timestamp_s=first_write_boundary[0], + wall_timestamp_s=first_write_boundary[1], + wall_semantics="boundary-sampled", + local_request=session.request_id, + peer_rank=write_meta.peer_rank, + slice=write_meta.slice_id, + bytes=int(write_meta.sizes.sum()), + ) if timer: timer.record_transfer_end(write_meta.peer_rank) @@ -694,6 +735,16 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): task.print_perf_info(write_meta.peer_rank, ri.instance_name, ri.instance_rank) with task.lock: + if ( + agent_result == AgentResult.SUCCESS + and agent_wait_complete_time_s is not None + and ( + task.kv_physical_complete_time_s is None + or agent_wait_complete_time_s > task.kv_physical_complete_time_s + ) + ): + task.kv_physical_complete_time_s = agent_wait_complete_time_s + task.kv_physical_complete_wall_time_s = agent_wait_complete_wall_time_s task.transferred_count += 1 count = task.transferred_count @@ -711,6 +762,54 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): task.complete() if all(t.status == TaskStatus.TRANSFERRED for t in session.kv_tasks): session.transfer_end_time = tensorrt_llm.bindings.global_steady_clock_now() + if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + task_boundaries = [] + for kv_task in session.kv_tasks: + with kv_task.lock: + task_boundaries.append( + ( + kv_task.kv_physical_complete_time_s, + kv_task.kv_physical_complete_wall_time_s, + ) + ) + if all( + steady_time is not None and wall_time is not None + for steady_time, wall_time in task_boundaries + ): + kv_physical_complete_time_s, (kv_physical_complete_wall_time_s) = max( + task_boundaries, key=lambda boundary: boundary[0] + ) + else: + kv_physical_complete_time_s = None + kv_physical_complete_wall_time_s = None + if ( + kv_physical_complete_time_s is not None + and session.mark_kv_physical_complete(kv_physical_complete_time_s) + ): + first_write_time_s = session.first_write_submit_time_s + kv_physical_service_ms = ( + (kv_physical_complete_time_s - first_write_time_s) * 1000 + if first_write_time_s is not None + else -1.0 + ) + _log_python_transfer_diagnostic( + rank=self._instance_rank, + role="ctx", + action="kv-physical-complete", + request=write_meta.unique_rid, + instance=self._registrar.self_rank_info.instance_name, + timestamp_s=kv_physical_complete_time_s, + wall_timestamp_s=kv_physical_complete_wall_time_s, + wall_semantics="boundary-sampled", + local_request=session.request_id, + slices=len(session.kv_tasks), + first_write_t=( + f"{first_write_time_s:.9f}" + if first_write_time_s is not None + else "-1" + ), + kv_physical_service_ms=(f"{kv_physical_service_ms:.6f}"), + ) logger.debug( f"deliver_kv_to_agent completed: unique_rid={write_meta.unique_rid}, " @@ -1152,15 +1251,22 @@ def _save_peer_req_info(self, peer_transfer_req_info: RecvReqInfo): session = self._get_session(req_info.unique_rid) if session is not None and not session.receiver_ready: session.receiver_ready = True - _log_python_transfer_diagnostic( - rank=self._instance_rank, - role="ctx", - action="receiver-info-ready", - request=req_info.unique_rid, - local_request=session.request_id, - received=expected_transfers, - expected=expected_transfers, - ) + if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + receiver_ready_time_s = _diagnostic_now_s() + receiver_ready_wall_time_s = time.time() + _log_python_transfer_diagnostic( + rank=self._instance_rank, + role="ctx", + action="receiver-info-ready", + request=req_info.unique_rid, + instance=self._registrar.self_rank_info.instance_name, + timestamp_s=receiver_ready_time_s, + wall_timestamp_s=receiver_ready_wall_time_s, + wall_semantics="boundary-sampled", + local_request=session.request_id, + received=expected_transfers, + expected=expected_transfers, + ) def has_all_peer_req_infos(self, unique_rid: int) -> bool: req_info = self._get_first_req_info(unique_rid) @@ -1274,11 +1380,14 @@ def __init__( self.transfer_start_time = None self.transfer_end_time = None self._first_write_submit_time_s: Optional[float] = None + self._first_write_submit_wall_time_s: Optional[float] = None + self._kv_physical_complete_time_s: Optional[float] = None _log_python_transfer_diagnostic( rank=self._sender._instance_rank, role="ctx", action="session-created", request=self.disagg_request_id, + instance=self._sender._registrar.self_rank_info.instance_name, local_request=self.request_id, ) # Must be last: makes session visible to listener thread, @@ -1317,7 +1426,14 @@ def first_write_submit_time_s(self) -> Optional[float]: with self.lock: return self._first_write_submit_time_s - def mark_first_write_submitted(self, timestamp_s: float) -> bool: + @property + def kv_physical_complete_time_s(self) -> Optional[float]: + with self.lock: + return self._kv_physical_complete_time_s + + def mark_first_write_submitted( + self, timestamp_s: float, wall_timestamp_s: Optional[float] = None + ) -> bool: """Record the first transfer-agent submission exactly once.""" if not _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: return False @@ -1325,6 +1441,17 @@ def mark_first_write_submitted(self, timestamp_s: float) -> bool: if self._first_write_submit_time_s is not None: return False self._first_write_submit_time_s = timestamp_s + self._first_write_submit_wall_time_s = wall_timestamp_s + return True + + def mark_kv_physical_complete(self, timestamp_s: float) -> bool: + """Record sender-side KV physical completion exactly once.""" + if not _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: + return False + with self.lock: + if self._kv_physical_complete_time_s is not None: + return False + self._kv_physical_complete_time_s = timestamp_s return True def send(self, slice: KVSlice) -> None: @@ -1348,6 +1475,7 @@ def send(self, slice: KVSlice) -> None: role="ctx", action="send-queued", request=self.disagg_request_id, + instance=self._sender._registrar.self_rank_info.instance_name, local_request=self.request_id, slice=slice_id, ) @@ -1405,8 +1533,9 @@ def cancel(self) -> None: _log_python_transfer_diagnostic( rank=self._sender._instance_rank, role="ctx", - action="cancelled", + action="cancel-requested", request=self.disagg_request_id, + instance=self._sender._registrar.self_rank_info.instance_name, local_request=self.request_id, ) # Send outside the lock to avoid holding it during I/O. @@ -1743,13 +1872,17 @@ def dispatch_task(self, task: KVRecvTask): ) if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: capacity_prepared_time_s = _diagnostic_now_s() + capacity_prepared_wall_time_s = time.time() if session.mark_capacity_prepared(capacity_prepared_time_s): _log_python_transfer_diagnostic( rank=self._registrar.self_rank_info.instance_rank, role="gen", action="capacity-prepared", request=receiver_req.unique_rid, + instance=self._registrar.self_rank_info.instance_name, timestamp_s=capacity_prepared_time_s, + wall_timestamp_s=capacity_prepared_wall_time_s, + wall_semantics="boundary-sampled", local_request=session.request_id, slice=task.slice_id, expected=task.expected_transfers, @@ -1773,13 +1906,17 @@ def dispatch_task(self, task: KVRecvTask): self._request_sender_data(peer_infos.sender_endpoints[rank], receiver_req_bytes) if _DISAGG_TRANSFER_DIAGNOSTICS_ENABLED: request_info_sent_time_s = _diagnostic_now_s() + request_info_sent_wall_time_s = time.time() if session.mark_request_info_sent(request_info_sent_time_s): _log_python_transfer_diagnostic( rank=self._registrar.self_rank_info.instance_rank, role="gen", action="request-info-sent", request=receiver_req.unique_rid, + instance=self._registrar.self_rank_info.instance_name, timestamp_s=request_info_sent_time_s, + wall_timestamp_s=request_info_sent_wall_time_s, + wall_semantics="boundary-sampled", local_request=session.request_id, slice=task.slice_id, sent=len(peer_overlap.ranks), @@ -1978,6 +2115,7 @@ def __init__( role="gen", action="session-created", request=self.disagg_request_id, + instance=self._receiver._registrar.self_rank_info.instance_name, local_request=self.request_id, ) self._receiver.setup_session(self) @@ -2269,8 +2407,9 @@ def cancel(self) -> None: _log_python_transfer_diagnostic( rank=self._receiver._registrar.self_rank_info.instance_rank, role="gen", - action="cancelled", + action="cancel-requested", request=self.disagg_request_id, + instance=self._receiver._registrar.self_rank_info.instance_name, local_request=self.request_id, ) # Send outside the lock to avoid holding it during I/O. diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 426995d6671e..d5b47385d20e 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -141,12 +141,18 @@ def _log_python_transfer_diagnostic( return if timestamp_s is None: timestamp_s = _diagnostic_now_s() + wall_timestamp_s = time.time() rank = getattr(getattr(self, "_dist", None), "rank", -1) + host = os.getenv("HOSTNAME", "unknown").replace(" ", "_") + instance = getattr(self, "_instance_name", "-") encoded_fields = " ".join(f"{key}={value}" for key, value in fields.items()) logger.info( "[DISAGG_DIAG][python-transfer] " - f"t={timestamp_s:.9f} clock=local_steady runtime=Python rank={rank} " - f"role={role} source=transceiver action={action} request={request} " + f"t={timestamp_s:.9f} clock=local_steady " + f"wall_t={wall_timestamp_s:.9f} wall_clock=unix " + f"wall_semantics=emission runtime=Python " + f"host={host} instance={instance} rank={rank} role={role} " + f"source=transceiver action={action} request={request} " f"{encoded_fields}" ) @@ -177,11 +183,17 @@ def _log_transfer_terminal( } if role == "ctx": first_write_time_s = getattr(session, "first_write_submit_time_s", None) + kv_physical_complete_time_s = getattr(session, "kv_physical_complete_time_s", None) if first_write_time_s is not None: fields["first_write_t"] = f"{first_write_time_s:.9f}" fields["elapsed_from_first_write_ms"] = ( f"{(timestamp_s - first_write_time_s) * 1000:.6f}" ) + if kv_physical_complete_time_s is not None: + fields["kv_physical_complete_t"] = f"{kv_physical_complete_time_s:.9f}" + fields["kv_physical_complete_to_terminal_ms"] = ( + f"{(timestamp_s - kv_physical_complete_time_s) * 1000:.6f}" + ) else: request_info_sent_time_s = getattr(session, "request_info_sent_time_s", None) ready_time_s = getattr(session, "kv_ready_time_s", None) @@ -1008,20 +1020,38 @@ def cancel_request(self, req: LlmRequest) -> bool: has_transferring = False if rid in self._send_sessions: - self._send_sessions[rid].cancel() - if self._send_sessions[rid].has_transferring_tasks(): + session = self._send_sessions[rid] + req = self._send_reqs[rid] + session.cancel() + if session.has_transferring_tasks(): has_transferring = True else: - self._send_sessions[rid].close() + self._log_transfer_terminal( + role="ctx", + action="cancelled", + rid=rid, + session=session, + req=req, + ) + session.close() del self._send_reqs[rid] del self._send_sessions[rid] if rid in self._recv_sessions: - self._recv_sessions[rid].cancel() - if self._recv_sessions[rid].has_transferring_tasks(): + session = self._recv_sessions[rid] + req = self._recv_reqs[rid] + session.cancel() + if session.has_transferring_tasks(): has_transferring = True else: - self._recv_sessions[rid].close() + self._log_transfer_terminal( + role="gen", + action="cancelled", + rid=rid, + session=session, + req=req, + ) + session.close() del self._recv_reqs[rid] del self._recv_sessions[rid] diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 7edfdfd0fef1..02aed1573adf 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -3438,11 +3438,23 @@ def _log_disagg_transfer_diagnostic(self, category: str, **fields) -> None: timestamp = fields.pop("t", None) if timestamp is None: timestamp = get_steady_clock_now_in_seconds() + wall_timestamp = fields.pop("wall_t", None) + wall_semantics = fields.pop("wall_semantics", "emission") + if wall_timestamp is None: + wall_timestamp = time.time() rank = getattr(getattr(self, "dist", None), "rank", -1) + host = os.getenv("HOSTNAME", "unknown").replace(" ", "_") + transceiver = getattr(self, "kv_cache_transceiver", None) + instance = getattr(transceiver, "_instance_name", "-") + if not isinstance(instance, str): + instance = "-" encoded_fields = " ".join(f"{key}={value}" for key, value in fields.items()) logger.info(f"[DISAGG_DIAG][{category}] t={timestamp:.9f} " - f"rank={rank} {encoded_fields}") + f"clock=local_steady wall_t={wall_timestamp:.9f} " + f"wall_clock=unix wall_semantics={wall_semantics} " + f"source=pyexecutor host={host} " + f"instance={instance} rank={rank} {encoded_fields}") @staticmethod def _disagg_diag_request_id(request: LlmRequest) -> int: @@ -3451,14 +3463,204 @@ def _disagg_diag_request_id(request: LlmRequest) -> int: None) if params is not None else None if isinstance(disagg_request_id, int): return disagg_request_id + ctx_request_id = getattr(params, "ctx_request_id", + None) if params is not None else None + if isinstance(ctx_request_id, int): + return ctx_request_id return request.py_request_id + @staticmethod + def _disagg_diag_schedule_style(request: LlmRequest) -> str: + params = getattr(request, "py_disaggregated_params", None) + schedule_style = getattr(params, "schedule_style", + None) if params is not None else None + if schedule_style is None: + return "unknown" + return getattr(schedule_style, "name", str(schedule_style)) + + def _log_disagg_effective_config_once(self, request: LlmRequest, + role: str) -> None: + if not _is_disagg_transfer_diagnostics_enabled(): + return + logged_roles = getattr(self, "_disagg_diag_config_logged_roles", None) + if logged_roles is None: + logged_roles = set() + self._disagg_diag_config_logged_roles = logged_roles + if role in logged_roles: + return + logged_roles.add(role) + + transceiver = getattr(self, "kv_cache_transceiver", None) + controller = self._get_disagg_transfer_admission_controller() + cache_config = getattr(getattr(self, "llm_args", None), + "cache_transceiver_config", None) + transfer_worker = getattr(transceiver, "_transfer_worker", None) + sender = getattr(transfer_worker, "_sender", None) + transfer_threads = getattr(sender, "_num_threads", None) + if not isinstance(transfer_threads, int): + transfer_threads = -1 + timeout_ms = getattr(transceiver, "kv_transfer_timeout_ms", None) + poll_interval_ms = getattr(transceiver, "kv_transfer_poll_interval_ms", + None) + sender_future_timeout_ms = getattr( + cache_config, "kv_transfer_sender_future_timeout_ms", None) + dist = getattr(self, "dist", None) + + def topology_size(name: str) -> int: + value = getattr(dist, name, None) + return value if isinstance(value, int) else -1 + + self._log_disagg_transfer_diagnostic( + "config", + action="effective", + role=role, + runtime=type(transceiver).__name__, + requested_runtime=getattr(cache_config, "transceiver_runtime", "-"), + backend=getattr(cache_config, "backend", "-"), + schedule_style=self._disagg_diag_schedule_style(request), + max_tokens_in_buffer=getattr(cache_config, "max_tokens_in_buffer", + None), + gate2_enabled=int(controller.enabled()), + budget_blocks=controller.max_transfer_blocks, + tokens_per_block=controller.tokens_per_block, + timeout_ms=timeout_ms, + sender_future_timeout_ms=sender_future_timeout_ms, + poll_interval_ms=poll_interval_ms, + transfer_threads=transfer_threads, + world_size=topology_size("world_size"), + tp_size=topology_size("tp_size"), + pp_size=topology_size("pp_size"), + cp_size=topology_size("cp_size"), + attention_dp=int(getattr(self, "enable_attention_dp", False)), + participant_coverage="observed-only", + async_transfer=int(self._uses_async_disagg_gen_transfer()), + inflight_cancel_requested=int(is_disagg_inflight_cancel_enabled()), + overlap_scheduler=int( + not getattr(self, "disable_overlap_scheduler", False))) + + def _capture_disagg_transfer_timer( + self, request: LlmRequest) -> Optional[Tuple[float, float, float]]: + transceiver = getattr(self, "kv_cache_transceiver", None) + timeout_ms = getattr(transceiver, "kv_transfer_timeout_ms", None) + if timeout_ms is None: + return None + + timer_start = time.monotonic() + timer_wall_time = time.time() + request.py_kv_transfer_start_time = timer_start + return timer_start, timer_wall_time, float(timeout_ms) + + def _log_disagg_transfer_timer_start( + self, request: LlmRequest, role: str, + timer_capture: Optional[Tuple[float, float, float]]) -> None: + if not _is_disagg_transfer_diagnostics_enabled(): + return + + self._log_disagg_effective_config_once(request, role) + if timer_capture is None: + return + timer_start, timer_wall_time, timeout_ms = timer_capture + category = "ctx-transfer" if role == "ctx" else "gen-transfer" + self._log_disagg_transfer_diagnostic( + category, + t=timer_start, + wall_t=timer_wall_time, + wall_semantics="boundary-sampled", + action="timer-start", + request=self._disagg_diag_request_id(request), + local_request=request.py_request_id, + timer_start_t=f"{timer_start:.9f}", + deadline_t=f"{timer_start + timeout_ms / 1000.0:.9f}", + timer_clock="python_monotonic", + timeout_ms=timeout_ms, + anchor=("after-respond-and-send-async" if role == "ctx" else + "after-request-and-receive-async-batch"), + schedule_style=self._disagg_diag_schedule_style(request), + state=getattr(request.state, "name", str(request.state))) + + def _start_disagg_transfer_timer(self, request: LlmRequest, + role: str) -> None: + timer_capture = self._capture_disagg_transfer_timer(request) + self._log_disagg_transfer_timer_start(request, role, timer_capture) + + def _log_disagg_gate_transition(self, request: LlmRequest, gate: int, + state: str, blocks: int, + decision_time: float, + decision_wall_time: float, + decision_sequence: int) -> None: + state_attr = f"py_disagg_diag_gate{gate}_state" + previous = getattr(request, state_attr, None) + if previous == state: + return + setattr(request, state_attr, state) + + fields = { + "action": + state, + "request": + self._disagg_diag_request_id(request), + "local_request": + request.py_request_id, + "blocks": + blocks, + "sequence": + decision_sequence, + "previous": + previous or "-", + "reason": + (f"capacity-scheduler-{'selected' if state == 'fitting' else 'not-selected'}" + if gate == 1 else "transport-window"), + } + activation_time = getattr(request, + "py_disagg_gen_executor_activation_time_s", + None) + if isinstance(activation_time, (int, float)): + fields["activation_to_transition_ms"] = ( + f"{(decision_time - activation_time) * 1000:.6f}") + + if gate == 2 and state == "deferred": + first_defer_time = getattr( + request, "py_disagg_diag_gate2_first_defer_time_s", None) + if not isinstance(first_defer_time, (int, float)): + request.py_disagg_diag_gate2_first_defer_time_s = decision_time + fields["first_defer"] = 1 + else: + fields["first_defer"] = 0 + elif gate == 2 and state == "admitted": + first_defer_time = getattr( + request, "py_disagg_diag_gate2_first_defer_time_s", None) + fields["deferral_ms"] = ( + f"{(decision_time - first_defer_time) * 1000:.6f}" + if isinstance(first_defer_time, (int, float)) else "0.000000") + + self._log_disagg_transfer_diagnostic(f"gate{gate}", + t=decision_time, + wall_t=decision_wall_time, + wall_semantics="boundary-sampled", + **fields) + if gate == 1 and state == "blocked": + previous_gate2 = getattr(request, "py_disagg_diag_gate2_state", + None) + if previous_gate2 not in {None, "ineligible"}: + self._log_disagg_gate_transition( + request, + gate=2, + state="ineligible", + blocks=blocks, + decision_time=decision_time, + decision_wall_time=decision_wall_time, + decision_sequence=decision_sequence) + else: + request.py_disagg_diag_gate2_state = "ineligible" + request.py_disagg_diag_gate2_first_defer_time_s = None + def _log_disagg_gen_ingress( self, request_items: Iterable[RequestQueueItem]) -> None: if not _is_disagg_transfer_diagnostics_enabled(): return ingress_time = get_steady_clock_now_in_seconds() + ingress_wall_time = time.time() for item in request_items: request = item.request if (request is None or request.request_type @@ -3468,6 +3670,8 @@ def _log_disagg_gen_ingress( self._log_disagg_transfer_diagnostic( "gen-arrival", t=ingress_time, + wall_t=ingress_wall_time, + wall_semantics="boundary-sampled", request=item.id, boundary="executor-queue-to-waiting-queue") @@ -3477,6 +3681,7 @@ def _log_disagg_gen_activations(self, return activation_time = get_steady_clock_now_in_seconds() + activation_wall_time = time.time() for request in requests: if (getattr(request, "state", None) != LlmRequestState.DISAGG_GENERATION_INIT): @@ -3488,6 +3693,8 @@ def _log_disagg_gen_activations(self, self._log_disagg_transfer_diagnostic( "gen-activation", t=activation_time, + wall_t=activation_wall_time, + wall_semantics="boundary-sampled", request=self._disagg_diag_request_id(request), local_request=request.py_request_id, ingress_to_activation_ms=( @@ -3503,6 +3710,50 @@ def _disagg_diag_request_blocks( controller._estimate_request_blocks(request)) for request in requests] + def _log_disagg_membership_snapshot(self, membership: str, + request_blocks: List[Tuple[int, int]], + decision_time: float, + decision_wall_time: float, + decision_sequence: int) -> int: + snapshots = getattr(self, "_disagg_diag_membership_snapshots", None) + if snapshots is None: + snapshots = {} + self._disagg_diag_membership_snapshots = snapshots + request_snapshot = tuple(request_blocks) + previous = snapshots.get(membership) + changed = previous is None or previous[1] != request_snapshot + if changed: + snapshot_version = 1 if previous is None else previous[0] + 1 + snapshots[membership] = (snapshot_version, request_snapshot) + else: + snapshot_version = previous[0] + + limit = _DISAGG_DIAGNOSTIC_REQUEST_LIST_LIMIT + overflow = request_blocks[limit:] + if not overflow or (not changed and decision_sequence % 100 != 0): + return snapshot_version + chunk_count = (len(overflow) + limit - 1) // limit + for chunk_index, offset in enumerate(range(0, len(overflow), limit), + start=1): + chunk = overflow[offset:offset + limit] + self._log_disagg_transfer_diagnostic( + "decision-members", + t=decision_time, + wall_t=decision_wall_time, + wall_semantics="boundary-sampled", + role="gen", + sequence=decision_sequence, + snapshot_version=snapshot_version, + membership=membership, + chunk_index=chunk_index, + chunk_count=chunk_count, + total_requests=len(request_blocks), + total_blocks=sum(blocks for _, blocks in request_blocks), + requests=_format_disagg_diag_request_blocks(chunk), + request_count=len(chunk), + request_blocks=sum(blocks for _, blocks in chunk)) + return snapshot_version + def _apply_disagg_transfer_admission( self, fitting_disagg_gen_init_requests: List[LlmRequest] ) -> Tuple[List[LlmRequest], bool]: @@ -3526,11 +3777,13 @@ def _apply_disagg_transfer_admission( return fitting_disagg_gen_init_requests, False decision_time = 0.0 + decision_wall_time = 0.0 decision_sequence = 0 candidate_request_blocks = [] active_request_blocks = [] if diagnostics_enabled: decision_time = get_steady_clock_now_in_seconds() + decision_wall_time = time.time() decision_sequence = getattr( self, "_disagg_diag_admission_decision_sequence", 0) + 1 self._disagg_diag_admission_decision_sequence = decision_sequence @@ -3561,6 +3814,18 @@ def _apply_disagg_transfer_admission( request_id for request_id, _ in candidate_request_blocks } + for request in waiting_requests: + request_id = self._disagg_diag_request_id(request) + state = ("fitting" + if request_id in candidate_request_ids else "blocked") + self._log_disagg_gate_transition( + request, + gate=1, + state=state, + blocks=controller._estimate_request_blocks(request), + decision_time=decision_time, + decision_wall_time=decision_wall_time, + decision_sequence=decision_sequence) blocked_count = sum( self._disagg_diag_request_id(request) not in candidate_request_ids for request in waiting_requests) @@ -3580,6 +3845,8 @@ def _apply_disagg_transfer_admission( self._log_disagg_transfer_diagnostic( "gate1", t=decision_time, + wall_t=decision_wall_time, + wall_semantics="boundary-sampled", sequence=decision_sequence, waiting=len(waiting_request_blocks), waiting_blocks=sum(blocks @@ -3625,14 +3892,36 @@ def _apply_disagg_transfer_admission( request_block for request_block in candidate_request_blocks if request_block[0] not in admitted_request_ids ] + for request in fitting_disagg_gen_init_requests: + request_id = self._disagg_diag_request_id(request) + state = ("admitted" + if request_id in admitted_request_ids else "deferred") + self._log_disagg_gate_transition( + request, + gate=2, + state=state, + blocks=controller._estimate_request_blocks(request), + decision_time=decision_time, + decision_wall_time=decision_wall_time, + decision_sequence=decision_sequence) candidate_transfer_blocks = sum( blocks for _, blocks in candidate_request_blocks) deferred_transfer_blocks = sum( blocks for _, blocks in deferred_request_blocks) + active_snapshot = self._log_disagg_membership_snapshot( + "active", active_request_blocks, decision_time, + decision_wall_time, decision_sequence) + candidate_snapshot = self._log_disagg_membership_snapshot( + "candidate", candidate_request_blocks, decision_time, + decision_wall_time, decision_sequence) self._log_disagg_transfer_diagnostic( "decision", t=decision_time, + wall_t=decision_wall_time, + wall_semantics="boundary-sampled", sequence=decision_sequence, + active_snapshot=active_snapshot, + candidate_snapshot=candidate_snapshot, runtime=type(self.kv_cache_transceiver).__name__, active=len(active_request_blocks), active_blocks=admission_result.active_transfer_blocks, @@ -3675,7 +3964,11 @@ def _apply_disagg_transfer_admission( self._log_disagg_transfer_diagnostic( "admission", t=decision_time, + wall_t=decision_wall_time, + wall_semantics="boundary-sampled", sequence=decision_sequence, + active_snapshot=active_snapshot, + candidate_snapshot=candidate_snapshot, runtime=type(self.kv_cache_transceiver).__name__, active=len(active_request_blocks), active_blocks=admission_result.active_transfer_blocks, @@ -5780,14 +6073,39 @@ def _is_disagg_inflight_cancel_active(self) -> bool: self._disagg_inflight_cancel_unsupported_logged = True return False - def _request_kv_transfer_cancellation(self, request: LlmRequest) -> bool: + def _request_kv_transfer_cancellation(self, + request: LlmRequest, + reason: str = "user") -> bool: """Best-effort cancellation that leaves ownership intact on errors.""" + role = ("ctx" if getattr(request, "is_context_only_request", False) + is True else "gen") + category = "ctx-transfer" if role == "ctx" else "gen-transfer" try: - return self.kv_cache_transceiver.cancel_request(request) + is_cancelled = self.kv_cache_transceiver.cancel_request(request) except Exception as error: logger.error(f"KV transfer cancellation failed for request " f"{request.py_request_id}; will retry: {error}") + if _is_disagg_transfer_diagnostics_enabled(): + self._log_disagg_transfer_diagnostic( + category, + action="cancel-result", + request=self._disagg_diag_request_id(request), + local_request=request.py_request_id, + reason=reason, + result="exception", + exception=type(error).__name__, + state=getattr(request.state, "name", str(request.state))) return False + if _is_disagg_transfer_diagnostics_enabled(): + self._log_disagg_transfer_diagnostic( + category, + action="cancel-result", + request=self._disagg_diag_request_id(request), + local_request=request.py_request_id, + reason=reason, + result="accepted" if is_cancelled else "retry", + state=getattr(request.state, "name", str(request.state))) + return is_cancelled @nvtx_range("_cancel_timed_out_gen_transfers") def _cancel_timed_out_gen_transfers(self) -> None: @@ -5809,6 +6127,7 @@ def _cancel_timed_out_gen_transfers(self) -> None: if req.is_disagg_generation_transmission_in_progress } current_time = time.monotonic() + current_wall_time = time.time() for request in requests_in_transfer.values(): if request.py_kv_transfer_start_time is None: continue @@ -5820,6 +6139,23 @@ def _cancel_timed_out_gen_transfers(self) -> None: f"Requesting cancellation for generation request " f"{request.py_request_id} due to KV cache transfer timeout") request.py_kv_transfer_timed_out = True + if _is_disagg_transfer_diagnostics_enabled(): + self._log_disagg_transfer_diagnostic( + "gen-transfer", + t=current_time, + wall_t=current_wall_time, + wall_semantics="boundary-sampled", + action="deadline-observed", + request=self._disagg_diag_request_id(request), + local_request=request.py_request_id, + timer_start_t=( + f"{request.py_kv_transfer_start_time:.9f}"), + timer_clock="python_monotonic", + elapsed_ms=f"{elapsed_time:.3f}", + timeout_ms=timeout_ms, + cancel_mode="inflight", + state=getattr(request.state, "name", + str(request.state))) user_canceled_ids = set(self.canceled_req_ids) local_timed_out_ids = sorted( @@ -5856,7 +6192,8 @@ def _cancel_timed_out_gen_transfers(self) -> None: if request_id in self._disagg_timed_out_gen_cancelled_ids: continue - is_cancelled = self._request_kv_transfer_cancellation(request) + is_cancelled = self._request_kv_transfer_cancellation( + request, reason="deadline") if is_cancelled: self._disagg_timed_out_gen_cancelled_ids.add(request_id) logger.warning( @@ -5896,13 +6233,15 @@ def _check_kv_transfer_timeout(self): def flag_if_kv_transfer_timed_out(req: LlmRequest, type: str) -> None: current_time = time.monotonic() + current_wall_time = time.time() if req.py_kv_transfer_start_time is None: return elapsed_time = (current_time - req.py_kv_transfer_start_time) * 1000 if elapsed_time > timeout_ms and not req.py_kv_transfer_timed_out: + inflight_cancel_active = ( + self._is_disagg_inflight_cancel_active()) verb = ("Requesting cancellation for" - if self._is_disagg_inflight_cancel_active() else - "Observed timeout on") + if inflight_cancel_active else "Observed timeout on") logger.warning( f"{verb} {type} request {req.py_request_id} due to KV " f"cache transfer timeout: elapsed {elapsed_time:.0f}ms > " @@ -5913,11 +6252,18 @@ def flag_if_kv_transfer_timed_out(req: LlmRequest, type: str) -> None: if type == "context" else "gen-transfer") self._log_disagg_transfer_diagnostic( category, - action="timeout", + t=current_time, + wall_t=current_wall_time, + wall_semantics="boundary-sampled", + action="deadline-observed", request=self._disagg_diag_request_id(req), local_request=req.py_request_id, + timer_start_t=(f"{req.py_kv_transfer_start_time:.9f}"), + timer_clock="python_monotonic", elapsed_ms=f"{elapsed_time:.3f}", timeout_ms=timeout_ms, + cancel_mode=("inflight" + if inflight_cancel_active else "legacy"), state=getattr(req.state, "name", str(req.state))) for req in self.async_transfer_manager.requests_in_transfer().values(): @@ -6383,16 +6729,32 @@ def _recv_disagg_gen_cache(self, new_gen_reqs): diagnostics_enabled = _is_disagg_transfer_diagnostics_enabled() controller = (self._get_disagg_transfer_admission_controller() if diagnostics_enabled else None) + submit_diagnostics = [] for req in new_gen_reqs: if diagnostics_enabled: submit_start = get_steady_clock_now_in_seconds() self.kv_cache_transceiver.request_and_receive_async(req) if diagnostics_enabled: submit_end = get_steady_clock_now_in_seconds() - assert controller is not None + submit_wall_time = time.time() + submit_diagnostics.append( + (req, submit_start, submit_end, submit_wall_time)) + + timer_captures = [] + for req in new_gen_reqs: + if req.state == LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS: + timer_captures.append( + (req, self._capture_disagg_transfer_timer(req))) + + if diagnostics_enabled: + assert controller is not None + for (req, submit_start, submit_end, + submit_wall_time) in submit_diagnostics: self._log_disagg_transfer_diagnostic( "submit", t=submit_end, + wall_t=submit_wall_time, + wall_semantics="boundary-sampled", runtime=type(self.kv_cache_transceiver).__name__, request=self._disagg_diag_request_id(req), blocks=controller._estimate_request_blocks(req), @@ -6401,11 +6763,8 @@ def _recv_disagg_gen_cache(self, new_gen_reqs): submit_call_ms=( f"{(submit_end - submit_start) * 1000:.6f}"), state=getattr(req.state, "name", str(req.state))) - - if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None: - for req in new_gen_reqs: - if req.state == LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS: - req.py_kv_transfer_start_time = time.monotonic() + for req, timer_capture in timer_captures: + self._log_disagg_transfer_timer_start(req, "gen", timer_capture) self._check_disagg_gen_cache_transfer_status(0) @@ -6446,10 +6805,15 @@ def kv_connector_request_finished(req: LlmRequest): self.kv_cache_transceiver.respond_and_send_async(req) if diagnostics_enabled: queued_time = get_steady_clock_now_in_seconds() + queued_wall_time = time.time() + timer_capture = self._capture_disagg_transfer_timer(req) + if diagnostics_enabled: req.py_disagg_ctx_send_queued_time_s = queued_time self._log_disagg_transfer_diagnostic( "ctx-transfer", t=queued_time, + wall_t=queued_wall_time, + wall_semantics="boundary-sampled", action="queued", runtime=type(self.kv_cache_transceiver).__name__, request=self._disagg_diag_request_id(req), @@ -6458,9 +6822,8 @@ def kv_connector_request_finished(req: LlmRequest): submit_call_ms=( f"{(queued_time - submit_start) * 1000:.6f}"), state=getattr(req.state, "name", str(req.state))) - - if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None: - req.py_kv_transfer_start_time = time.monotonic() + self._log_disagg_transfer_timer_start( + req, "ctx", timer_capture) if self.kv_connector_manager: if not self.disable_overlap_scheduler: @@ -6562,6 +6925,8 @@ def _check_disagg_ctx_cache_transfer_status(self, atLeastNum: int = 0): request=self._disagg_diag_request_id(request), local_request=request.py_request_id, outcome=outcome, + deadline_observed=int( + getattr(request, "py_kv_transfer_timed_out", False)), queued_t=(f"{queued_time:.9f}" if isinstance( queued_time, (int, float)) else "-1"), queued_to_reap_ms=f"{queued_to_reap_ms:.6f}", @@ -6579,7 +6944,8 @@ def _check_disagg_ctx_cache_transfer_status(self, atLeastNum: int = 0): or request_id in self._disagg_timed_out_ctx_cancelled_ids): continue - is_cancelled = self._request_kv_transfer_cancellation(request) + is_cancelled = self._request_kv_transfer_cancellation( + request, reason="deadline") if not is_cancelled: continue @@ -7312,7 +7678,8 @@ def _handle_responses(self, emit_first_iter: bool = True): timed_out_requests.append(request) continue - is_cancelled = self._request_kv_transfer_cancellation(request) + is_cancelled = self._request_kv_transfer_cancellation( + request, reason="deadline") if is_cancelled: # _handle_errors enters response collectives under ADP. # Defer it until the rank-uniform vote below. diff --git a/tests/scripts/perf-sanity/disaggregated/gb200_gpt-oss-120b-fp4_8k1k_con1024_ctx1_tp1_gen1_tp4_eplb0_mtp0_ccb-NIXL.yaml b/tests/scripts/perf-sanity/disaggregated/gb200_gpt-oss-120b-fp4_8k1k_con1024_ctx1_tp1_gen1_tp4_eplb0_mtp0_ccb-NIXL.yaml index 9f4b7086060d..f5e62e307149 100644 --- a/tests/scripts/perf-sanity/disaggregated/gb200_gpt-oss-120b-fp4_8k1k_con1024_ctx1_tp1_gen1_tp4_eplb0_mtp0_ccb-NIXL.yaml +++ b/tests/scripts/perf-sanity/disaggregated/gb200_gpt-oss-120b-fp4_8k1k_con1024_ctx1_tp1_gen1_tp4_eplb0_mtp0_ccb-NIXL.yaml @@ -42,7 +42,7 @@ environment: # existing gpt-oss perf-sanity configs; PDL / no-parallel-weight-load / # disabled NCCL graph register / disabled UCC collective are preserved so # the same load & runtime paths are exercised). - worker_env_var: "TLLM_LOG_LEVEL=INFO TRTLLM_SERVER_DISABLE_GC=1 TRTLLM_WORKER_DISABLE_GC=1 TRTLLM_ENABLE_PDL=1 ENROOT_ALLOW_DEV=yes OVERRIDE_QUANT_ALGO=W4A8_MXFP4_MXFP8 TRT_LLM_DISABLE_LOAD_WEIGHTS_IN_PARALLEL=True NCCL_GRAPH_REGISTER=0 OMPI_MCA_coll_ucc_enable=0" + worker_env_var: "TLLM_LOG_LEVEL=INFO TRTLLM_SERVER_DISABLE_GC=1 TRTLLM_WORKER_DISABLE_GC=1 TRTLLM_ENABLE_PDL=1 ENROOT_ALLOW_DEV=yes OVERRIDE_QUANT_ALGO=W4A8_MXFP4_MXFP8 TRT_LLM_DISABLE_LOAD_WEIGHTS_IN_PARALLEL=True NCCL_GRAPH_REGISTER=0 OMPI_MCA_coll_ucc_enable=0 TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS=1" server_env_var: "TRTLLM_SERVER_DISABLE_GC=1" profiling: nsys_on: false diff --git a/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_8k1k_con4096_ctx1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml b/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_8k1k_con4096_ctx1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml index dd57bb89e4b6..b37e2855173b 100644 --- a/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_8k1k_con4096_ctx1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml +++ b/tests/scripts/perf-sanity/disaggregated/gb300_deepseek-r1-fp4_8k1k_con4096_ctx1_dep4_gen1_dep16_eplb0_mtp1_ccb-NIXL.yaml @@ -35,7 +35,7 @@ environment: trtllm_repo: '' build_wheel: false work_dir: - worker_env_var: TLLM_LOG_LEVEL=INFO TRTLLM_SERVER_DISABLE_GC=1 TRTLLM_WORKER_DISABLE_GC=1 TRTLLM_ENABLE_PDL=1 ENROOT_ALLOW_DEV=yes TLLM_SPEC_DECODE_FORCE_NUM_ACCEPTED_TOKENS=1 TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS=1 + worker_env_var: TLLM_LOG_LEVEL=INFO TRTLLM_SERVER_DISABLE_GC=1 TRTLLM_WORKER_DISABLE_GC=1 TRTLLM_ENABLE_PDL=1 ENROOT_ALLOW_DEV=yes TLLM_SPEC_DECODE_FORCE_NUM_ACCEPTED_TOKENS=1 server_env_var: TRTLLM_SERVER_DISABLE_GC=1 profiling: nsys_on: false diff --git a/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py b/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py index 3d9a3f83dcc3..af5958883a69 100644 --- a/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py +++ b/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py @@ -99,6 +99,56 @@ def test_unsupported_transceiver_warns_once(monkeypatch): warning.assert_called_once() +@pytest.mark.parametrize( + ("cancelled", "expected_result"), + [(True, "accepted"), (False, "retry")], +) +def test_cancel_result_diagnostic_distinguishes_accept_and_retry( + monkeypatch, cancelled, expected_result +): + monkeypatch.setattr(executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(executor_module.logger, "info", log_info) + request = _make_timeout_request() + request.is_context_only_request = True + request.py_disaggregated_params = None + executor = object.__new__(PyExecutor) + executor.dist = SimpleNamespace(rank=0) + executor.kv_cache_transceiver = Mock() + executor.kv_cache_transceiver.cancel_request.return_value = cancelled + + result = PyExecutor._request_kv_transfer_cancellation(executor, request, reason="deadline") + + assert result is cancelled + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][ctx-transfer]" in message + assert "action=cancel-result" in message + assert "reason=deadline" in message + assert f"result={expected_result}" in message + + +def test_cancel_result_diagnostic_records_exception(monkeypatch): + monkeypatch.setattr(executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(executor_module.logger, "info", log_info) + request = _make_timeout_request() + request.is_context_only_request = False + request.py_disaggregated_params = None + executor = object.__new__(PyExecutor) + executor.dist = SimpleNamespace(rank=0) + executor.kv_cache_transceiver = Mock() + executor.kv_cache_transceiver.cancel_request.side_effect = RuntimeError("cancel failed") + + result = PyExecutor._request_kv_transfer_cancellation(executor, request, reason="deadline") + + assert not result + message = log_info.call_args.args[0] + assert "[DISAGG_DIAG][gen-transfer]" in message + assert "action=cancel-result" in message + assert "result=exception" in message + assert "exception=RuntimeError" in message + + def test_flag_unset_generation_timeout_uses_rank_uniform_cleanup(): request = _make_timeout_request() executor = _make_response_handler_stub([request], [True, False]) diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index 1faef2d3317c..208eef1ad4c2 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -434,6 +434,10 @@ def _make_disagg_transfer_request( req.py_prompt_len = prompt_len req.total_input_len_cp = prompt_len if total_input_len_cp is None else total_input_len_cp req.is_disagg_generation_transmission_in_progress = in_progress + req.py_disaggregated_params = None + req.py_disagg_diag_gate1_state = None + req.py_disagg_diag_gate2_state = None + req.py_disagg_diag_gate2_first_defer_time_s = None return req @@ -447,6 +451,15 @@ def _clear_disagg_transfer_mode_env(monkeypatch: pytest.MonkeyPatch) -> None: @pytest.mark.usefixtures("_clear_disagg_transfer_mode_env") class TestDisaggTransferAdmissionController: + def test_diagnostic_request_id_falls_back_to_context_request_id(self): + request = _make_disagg_transfer_request(7, 64) + request.py_disaggregated_params = types.SimpleNamespace( + disagg_request_id=None, + ctx_request_id=701, + ) + + assert PyExecutor._disagg_diag_request_id(request) == 701 + def test_disabled_preserves_candidates(self): controller = DisaggTransferAdmissionController( max_tokens_in_buffer=None, tokens_per_block=32 @@ -547,15 +560,30 @@ def test_apply_emits_changed_admission_snapshot(self, monkeypatch): PyExecutor._apply_disagg_transfer_admission(executor, [candidate]) PyExecutor._apply_disagg_transfer_admission(executor, [candidate]) - assert log_info.call_count == 4 - gate1_messages = [ + assert log_info.call_count == 6 + gate1_summary_messages = [ call.args[0] for call in log_info.call_args_list - if "[DISAGG_DIAG][gate1]" in call.args[0] + if "[DISAGG_DIAG][gate1]" in call.args[0] and "waiting_requests=" in call.args[0] ] - assert len(gate1_messages) == 1 - assert "sequence=1" in gate1_messages[0] - assert "fitting_requests=2:1" in gate1_messages[0] + assert len(gate1_summary_messages) == 1 + assert "sequence=1" in gate1_summary_messages[0] + assert "fitting_requests=2:1" in gate1_summary_messages[0] + gate1_transition = next( + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][gate1]" in call.args[0] and "action=fitting" in call.args[0] + ) + assert "request=2" in gate1_transition + assert "previous=-" in gate1_transition + gate2_transition = next( + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][gate2]" in call.args[0] + ) + assert "action=deferred" in gate2_transition + assert "request=2" in gate2_transition + assert "first_defer=1" in gate2_transition decision_messages = [ call.args[0] for call in log_info.call_args_list @@ -601,11 +629,92 @@ def test_apply_logs_gate1_when_no_request_fits(self, monkeypatch): gate1_message = next( call.args[0] for call in log_info.call_args_list - if "[DISAGG_DIAG][gate1]" in call.args[0] + if "[DISAGG_DIAG][gate1]" in call.args[0] and "waiting_requests=" in call.args[0] ) assert "waiting_requests=4:2" in gate1_message assert "fitting_requests=-" in gate1_message assert "blocked_requests=4:2" in gate1_message + transition_message = next( + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][gate1]" in call.args[0] and "action=blocked" in call.args[0] + ) + assert "request=4" in transition_message + + def test_gate2_transitions_are_not_truncated(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor.kv_cache_transceiver = Mock() + executor._is_kv_manager_v2 = False + executor.active_requests = [_make_disagg_transfer_request(1, 32, in_progress=True)] + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + candidates = [_make_disagg_transfer_request(request_id, 32) for request_id in range(2, 72)] + + PyExecutor._apply_disagg_transfer_admission(executor, candidates) + PyExecutor._apply_disagg_transfer_admission(executor, candidates) + + transition_messages = [ + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][gate2]" in call.args[0] and "action=deferred" in call.args[0] + ] + assert len(transition_messages) == 70 + assert any("request=71 " in message for message in transition_messages) + membership_messages = [ + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][decision-members]" in call.args[0] + ] + assert len(membership_messages) == 1 + assert "membership=candidate" in membership_messages[0] + assert "snapshot_version=1" in membership_messages[0] + assert "71:1" in membership_messages[0] + + def test_gate2_deferral_episode_resets_when_gate1_blocks(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + decision_times = iter((1.0, 2.0, 3.0)) + monkeypatch.setattr( + py_executor_module, + "get_steady_clock_now_in_seconds", + lambda: next(decision_times), + ) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor.kv_cache_transceiver = Mock() + executor._is_kv_manager_v2 = False + active = _make_disagg_transfer_request(1, 32, in_progress=True) + candidate = _make_disagg_transfer_request(2, 32) + candidate.state = LlmRequestState.DISAGG_GENERATION_INIT + executor.active_requests = [active] + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=32, tokens_per_block=32 + ) + + PyExecutor._apply_disagg_transfer_admission(executor, [candidate]) + executor.active_requests = [active, candidate] + PyExecutor._apply_disagg_transfer_admission(executor, []) + executor.active_requests = [] + PyExecutor._apply_disagg_transfer_admission(executor, [candidate]) + + gate2_messages = [ + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][gate2]" in call.args[0] + ] + assert any("action=deferred" in message for message in gate2_messages) + assert any("action=ineligible" in message for message in gate2_messages) + readmitted = next(message for message in gate2_messages if "action=admitted" in message) + assert "previous=ineligible" in readmitted + assert "deferral_ms=0.000000" in readmitted def test_apply_missing_controller_preserves_candidates(self): executor = object.__new__(PyExecutor) @@ -808,7 +917,7 @@ def test_async_receive_emits_submit_work(self, monkeypatch): monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) log_info = Mock() monkeypatch.setattr(py_executor_module.logger, "info", log_info) - timestamps = iter((10.0, 10.002)) + timestamps = iter((10.0, 10.002, 10.003)) monkeypatch.setattr( py_executor_module, "get_steady_clock_now_in_seconds", lambda: next(timestamps) ) @@ -830,17 +939,87 @@ def mark_in_progress(req): PyExecutor._recv_disagg_gen_cache(executor, [request]) - message = log_info.call_args.args[0] + message = next( + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][submit]" in call.args[0] + ) assert "[DISAGG_DIAG][submit]" in message assert "request=7" in message assert "blocks=2" in message assert "bytes=4096" in message assert "submit_call_ms=2.000000" in message + @pytest.mark.parametrize( + ("role", "category"), + [("ctx", "ctx-transfer"), ("gen", "gen-transfer")], + ) + def test_transfer_timer_records_exact_start_and_effective_config( + self, monkeypatch, role, category + ): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(py_executor_module.time, "monotonic", lambda: 12.345) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=0) + executor.disable_overlap_scheduler = False + executor.kv_cache_transceiver = Mock( + kv_transfer_timeout_ms=60_000, + kv_transfer_poll_interval_ms=5, + ) + executor.llm_args = types.SimpleNamespace( + cache_transceiver_config=types.SimpleNamespace( + transceiver_runtime="Python", + backend="NIXL", + max_tokens_in_buffer=9216, + kv_transfer_sender_future_timeout_ms=60_000, + ) + ) + executor._disagg_transfer_admission_controller = DisaggTransferAdmissionController( + max_tokens_in_buffer=9216, tokens_per_block=32 + ) + request = _make_disagg_transfer_request(7, 8192) + request.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + request.py_disaggregated_params = types.SimpleNamespace( + disagg_request_id=701, + schedule_style=types.SimpleNamespace(name="CONTEXT_FIRST"), + ) + + PyExecutor._start_disagg_transfer_timer(executor, request, role) + + assert request.py_kv_transfer_start_time == 12.345 + config_message = next( + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][config]" in call.args[0] + ) + assert "gate2_enabled=1" in config_message + assert "budget_blocks=288" in config_message + assert "timeout_ms=60000" in config_message + assert "schedule_style=CONTEXT_FIRST" in config_message + timer_message = next( + call.args[0] + for call in log_info.call_args_list + if f"[DISAGG_DIAG][{category}]" in call.args[0] and "action=timer-start" in call.args[0] + ) + assert "t=12.345000000" in timer_message + assert "request=701" in timer_message + assert "timer_start_t=12.345000000" in timer_message + assert "deadline_t=72.345000000" in timer_message + assert "timer_clock=python_monotonic" in timer_message + expected_anchor = ( + "after-respond-and-send-async" + if role == "ctx" + else "after-request-and-receive-async-batch" + ) + assert f"anchor={expected_anchor}" in timer_message + def test_context_send_emits_queued_boundary(self, monkeypatch): monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) - timestamps = iter((10.0, 10.002)) + timestamps = iter((10.0, 10.002, 10.003)) monkeypatch.setattr( py_executor_module, "get_steady_clock_now_in_seconds", @@ -866,7 +1045,11 @@ def test_context_send_emits_queued_boundary(self, monkeypatch): executor.kv_cache_transceiver.respond_and_send_async.assert_called_once_with(request) assert request.py_disagg_ctx_send_queued_time_s == 10.002 - message = log_info.call_args.args[0] + message = next( + call.args[0] + for call in log_info.call_args_list + if "[DISAGG_DIAG][ctx-transfer]" in call.args[0] and "action=queued" in call.args[0] + ) assert "[DISAGG_DIAG][ctx-transfer]" in message assert "action=queued" in message assert "request=17" in message @@ -908,7 +1091,7 @@ def test_context_status_emits_completed_reap(self, monkeypatch): assert "queued_to_reap_ms=103.000000" in message assert "poll_call_ms=2.000000" in message - def test_transfer_timeout_emits_terminal_boundary(self, monkeypatch): + def test_transfer_timeout_emits_nonterminal_deadline_boundary(self, monkeypatch): monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) monkeypatch.setattr(py_executor_module.time, "monotonic", lambda: 20.050) @@ -931,7 +1114,7 @@ def test_transfer_timeout_emits_terminal_boundary(self, monkeypatch): assert request.py_kv_transfer_timed_out message = log_info.call_args.args[0] assert "[DISAGG_DIAG][gen-transfer]" in message - assert "action=timeout" in message + assert "action=deadline-observed" in message assert "request=19" in message assert "timeout_ms=10" in message diff --git a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py index b9f8ad1423ff..ea4195b9ec92 100644 --- a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py +++ b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py @@ -459,6 +459,7 @@ def test_native_session_boundaries_are_recorded_once( tx_session = object.__new__(TxSession) tx_session.lock = threading.Lock() tx_session._first_write_submit_time_s = None + tx_session._kv_physical_complete_time_s = None rx_session = object.__new__(RxSession) rx_session.lock = threading.Lock() rx_session._request_info_sent_time_s = None @@ -466,11 +467,53 @@ def test_native_session_boundaries_are_recorded_once( assert tx_session.mark_first_write_submitted(3.0) assert not tx_session.mark_first_write_submitted(4.0) assert tx_session.first_write_submit_time_s == 3.0 + assert tx_session.mark_kv_physical_complete(4.0) + assert not tx_session.mark_kv_physical_complete(5.0) + assert tx_session.kv_physical_complete_time_s == 4.0 assert rx_session.mark_request_info_sent(5.0) assert not rx_session.mark_request_info_sent(6.0) assert rx_session.request_info_sent_time_s == 5.0 +@pytest.mark.parametrize("has_transferring_tasks", [False, True]) +def test_cancel_request_separates_request_from_terminal_diagnostic( + has_transferring_tasks: bool, +) -> None: + request = Mock( + request_id=7, + py_request_id=7, + py_disaggregated_params=None, + ) + session = Mock(status=SessionStatus.CANCELLED) + session.has_transferring_tasks.return_value = has_transferring_tasks + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._wait_reqs = {} + transceiver._send_sessions = {7: session} + transceiver._send_reqs = {7: request} + transceiver._recv_sessions = {} + transceiver._recv_reqs = {} + transceiver._log_transfer_terminal = Mock() + + result = transceiver.cancel_request(request) + + session.cancel.assert_called_once() + if has_transferring_tasks: + assert not result + assert 7 in transceiver._send_sessions + transceiver._log_transfer_terminal.assert_not_called() + else: + assert result + assert 7 not in transceiver._send_sessions + transceiver._log_transfer_terminal.assert_called_once_with( + role="ctx", + action="cancelled", + rid=7, + session=session, + req=request, + ) + session.close.assert_called_once() + + def test_tx_session_wait_complete_nonblocking_returns_none_without_waiting() -> None: task = _FakeTask(TaskStatus.TRANSFERRING) session = _make_tx_session([task]) diff --git a/tests/unittest/tools/test_disagg_admission_telemetry.py b/tests/unittest/tools/test_disagg_admission_telemetry.py index 2b307d5ad5f6..43450dacc502 100644 --- a/tests/unittest/tools/test_disagg_admission_telemetry.py +++ b/tests/unittest/tools/test_disagg_admission_telemetry.py @@ -432,6 +432,371 @@ def test_cpp_lifecycle_and_remaining_work_use_same_domain_completion(): assert lifecycle["reap_to_decode_start"]["duration_s"]["p50"] == pytest.approx(0.05) +def test_deadline_is_nonterminal_and_classified_by_sender_phase(): + events = _parse_lines( + [ + "[DISAGG_DIAG][ctx-transfer] t=0.00 clock=local_steady host=ctx " + "rank=0 action=timer-start request=42", + "[DISAGG_DIAG][ctx-transfer] t=0.01 clock=local_steady host=ctx " + "rank=0 action=queued request=42", + "[DISAGG_DIAG][sender-transfer] t=0.20 clock=local_steady host=ctx " + "rank=0 action=credit-received request=42", + "[DISAGG_DIAG][python-transfer] t=0.30 clock=local_steady host=ctx " + "rank=0 role=ctx source=native action=first-write-submitted request=42", + "[DISAGG_DIAG][ctx-transfer] t=0.60 clock=local_steady host=ctx " + "rank=0 action=deadline-observed request=42", + "[DISAGG_DIAG][ctx-transfer] t=0.65 clock=local_steady host=ctx " + "rank=0 action=cancel-result request=42 result=retry", + "[DISAGG_DIAG][python-transfer] t=0.80 clock=local_steady host=ctx " + "rank=0 role=ctx source=native action=kv-physical-complete request=42", + "[DISAGG_DIAG][ctx-transfer] t=0.90 clock=local_steady host=ctx " + "rank=0 action=reaped request=42 outcome=completed", + ] + ) + + result = analyze_events(events) + lifecycle = result["lifecycle"] + record = next( + request + for request in lifecycle["requests"] + if request["request"] == "42" and request["role"] == "ctx" + ) + + assert result["schema_version"] == 3 + assert lifecycle["deadline_phase_counts"] == {"kv-transfer-service": 1} + assert record["deadline_phase"] == "kv-transfer-service" + assert record["intervals"]["ctx_timer_to_deadline"]["duration_s"] == pytest.approx(0.6) + assert record["intervals"]["ctx_deadline_to_kv_physical_complete"][ + "duration_s" + ] == pytest.approx(0.2) + assert record["intervals"]["ctx_deadline_to_cancel_result"]["duration_s"] == pytest.approx(0.05) + assert _TELEMETRY._collect_unsuccessful_requests(events) == set() + + +def test_cross_host_wall_clock_joins_ctx_deadline_to_gen_deferral(): + events = _parse_lines( + [ + "[DISAGG_DIAG][ctx-transfer] t=100.0 clock=local_steady " + "wall_t=1000.0 wall_clock=unix wall_semantics=boundary-sampled " + "host=ctx rank=0 action=timer-start request=42", + "[DISAGG_DIAG][gate1] t=1.0 clock=local_steady " + "wall_t=1005.0 wall_clock=unix wall_semantics=boundary-sampled " + "host=gen rank=0 action=fitting request=42", + "[DISAGG_DIAG][gate2] t=2.0 clock=local_steady " + "wall_t=1006.0 wall_clock=unix wall_semantics=boundary-sampled " + "host=gen rank=0 action=deferred request=42", + "[DISAGG_DIAG][ctx-transfer] t=160.0 clock=local_steady " + "wall_t=1060.0 wall_clock=unix wall_semantics=boundary-sampled " + "host=ctx rank=0 action=deadline-observed request=42", + "[DISAGG_DIAG][gate2] t=62.0 clock=local_steady " + "wall_t=1062.0 wall_clock=unix wall_semantics=boundary-sampled " + "host=gen rank=0 action=admitted request=42", + ] + ) + + correlation = analyze_events(events)["cross_host_correlation"] + record = correlation["requests"][0] + + assert correlation["joined_ctx_gen_request_count"] == 1 + assert correlation["ctx_deadline_relationship_counts"] == {"during-gate2-deferral": 1} + assert record["ctx_deadline_relationship"] == "during-gate2-deferral" + assert record["wall_intervals_s"]["ctx_timer_to_gate2_admit"] == pytest.approx(62.0) + assert record["wall_intervals_s"]["gate2_defer_to_ctx_deadline"] == pytest.approx(54.0) + + +def test_cross_host_partial_gate2_admission_is_not_labeled_before_gate1(): + events = _parse_lines( + [ + "[DISAGG_DIAG][ctx-transfer] t=0 wall_t=0 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx rank=0 " + "action=timer-start request=42", + "[DISAGG_DIAG][gate2] t=2 wall_t=50 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 " + "action=admitted request=42", + "[DISAGG_DIAG][gate2] t=3 wall_t=70 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=1 " + "action=admitted request=42", + "[DISAGG_DIAG][ctx-transfer] t=60 wall_t=60 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx rank=0 " + "action=deadline-observed request=42", + ] + ) + + record = analyze_events(events)["cross_host_correlation"]["requests"][0] + + assert record["points"]["gate2-state-at-ctx-deadline"]["state"] == "admitted" + assert record["ctx_deadline_relationship"] == "partial-gate2-admission-before-global-admission" + + +def test_cross_host_progress_uses_latest_observed_rank(): + events = _parse_lines( + [ + "[DISAGG_DIAG][ctx-transfer] t=0 wall_t=0 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx rank=0 " + "action=timer-start request=42", + "[DISAGG_DIAG][gate2] t=1 wall_t=10 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 " + "action=deferred request=42", + "[DISAGG_DIAG][gate2] t=2 wall_t=50 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 " + "action=admitted request=42", + "[DISAGG_DIAG][gate2] t=3 wall_t=70 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=1 " + "action=admitted request=42", + "[DISAGG_DIAG][ctx-transfer] t=60 wall_t=60 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx rank=0 " + "action=deadline-observed request=42", + ] + ) + + record = analyze_events(events)["cross_host_correlation"]["requests"][0] + admission = record["points"]["gate2-admitted"] + + assert admission["wall_t"] == pytest.approx(70.0) + assert admission["selection"] == "latest" + assert admission["observed_emitter_count"] == 2 + assert record["ctx_deadline_relationship"] == "partial-gate2-admission-before-global-admission" + + +def test_cross_host_gate2_ineligible_closes_deferral_episode(): + events = _parse_lines( + [ + "[DISAGG_DIAG][ctx-transfer] t=0 wall_t=0 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx rank=0 " + "action=timer-start request=42", + "[DISAGG_DIAG][gate1] t=1 wall_t=5 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 " + "action=fitting request=42", + "[DISAGG_DIAG][gate2] t=2 wall_t=10 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 " + "action=deferred request=42", + "[DISAGG_DIAG][gate1] t=3 wall_t=20 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 " + "action=blocked request=42", + "[DISAGG_DIAG][gate2] t=3 wall_t=20 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 " + "action=ineligible request=42", + "[DISAGG_DIAG][ctx-transfer] t=60 wall_t=60 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx rank=0 " + "action=deadline-observed request=42", + ] + ) + + record = analyze_events(events)["cross_host_correlation"]["requests"][0] + + assert record["points"]["gate2-state-at-ctx-deadline"]["state"] == "ineligible" + assert record["ctx_deadline_relationship"] == "gate2-ineligible-at-deadline" + + +def test_cross_host_phase_uses_furthest_milestone_with_missing_credit(): + events = _parse_lines( + [ + "[DISAGG_DIAG][ctx-transfer] t=0 wall_t=0 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx rank=0 " + "action=timer-start request=42", + "[DISAGG_DIAG][gate2] t=1 wall_t=10 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 " + "action=admitted request=42", + "[DISAGG_DIAG][submit] t=2 wall_t=20 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 request=42", + "[DISAGG_DIAG][python-transfer] t=3 wall_t=30 wall_clock=unix " + "wall_semantics=emission host=ctx rank=0 role=ctx " + "action=first-write-submitted request=42", + "[DISAGG_DIAG][ctx-transfer] t=40 wall_t=40 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx rank=0 " + "action=deadline-observed request=42", + ] + ) + + record = analyze_events(events)["cross_host_correlation"]["requests"][0] + + assert record["ctx_deadline_relationship"] == "during-kv-physical-transfer" + assert record["ctx_deadline_phase_coverage"]["missing_before_furthest"] == [ + "ctx-receiver-credit" + ] + + +def test_uncapped_gate_transitions_cover_every_request(): + events = _parse_lines( + [ + event + for request_id in range(1, 71) + for event in ( + f"[DISAGG_DIAG][gen-arrival] t=0.0 rank=0 request={request_id}", + f"[DISAGG_DIAG][gate1] t=0.1 rank=0 action=fitting " + f"request={request_id} blocks=1 sequence=1", + f"[DISAGG_DIAG][gate2] t=0.2 rank=0 action=deferred " + f"request={request_id} blocks=1 sequence=1", + ) + ] + ) + + lifecycle = analyze_events(events)["lifecycle"] + + assert lifecycle["request_count"] == 70 + assert lifecycle["interval_coverage"]["gen_arrival_to_first_gate1"]["observed_requests"] == 70 + assert lifecycle["interval_coverage"]["gen_arrival_to_first_gate2"]["observed_requests"] == 70 + assert _TELEMETRY._collect_global_request_blocks(events)["70"] == 1 + + +def test_versioned_candidate_snapshot_reconstructs_unchanged_decisions(): + prefix = ",".join(f"{request}:1" for request in range(1, 65)) + tail = ",".join(f"{request}:1" for request in range(65, 71)) + events = _parse_lines( + [ + "[DISAGG_DIAG][decision-members] t=1 rank=0 role=gen " + "instance=gen sequence=1 snapshot_version=1 membership=candidate " + f"chunk_index=1 chunk_count=1 requests={tail}", + "[DISAGG_DIAG][decision] t=1 rank=0 instance=gen sequence=1 " + "candidate_snapshot=1 active_snapshot=1 active_requests=- " + "active_requests_omitted=0 " + f"candidate_requests={prefix} candidate_requests_omitted=6 " + "admitted=0 admitted_requests=- admitted_requests_omitted=0 " + f"deferred=70 deferred_requests={prefix} " + "deferred_requests_omitted=6 budget=1", + "[DISAGG_DIAG][decision] t=2 rank=0 instance=gen sequence=2 " + "candidate_snapshot=1 active_snapshot=1 active_requests=- " + "active_requests_omitted=0 " + f"candidate_requests={prefix} candidate_requests_omitted=6 " + "admitted=0 admitted_requests=- admitted_requests_omitted=0 " + f"deferred=70 deferred_requests={prefix} " + "deferred_requests_omitted=6 budget=1", + ] + ) + + admissions, _ = _TELEMETRY._collect_admissions(events) + + assert len(admissions) == 2 + assert all(len(admission.candidate_requests) == 70 for admission in admissions) + assert all(len(admission.deferred_requests) == 70 for admission in admissions) + assert all(admission.candidate_requests_omitted == 0 for admission in admissions) + assert admissions[1].deferred_requests[-1] == "70" + + +def test_remaining_work_ignores_gate_transition_and_reports_incomplete_tail(): + prefix = ",".join(f"{request}:1" for request in range(1, 65)) + partial_tail = ",".join(f"{request}:1" for request in range(65, 68)) + events = _parse_lines( + [ + "[DISAGG_DIAG][gate2] t=0.9 rank=0 instance=gen sequence=1 " + "action=deferred request=99 blocks=1", + "[DISAGG_DIAG][decision-members] t=1 rank=0 role=gen " + "instance=gen sequence=1 snapshot_version=1 membership=active " + f"chunk_index=1 chunk_count=2 requests={partial_tail}", + "[DISAGG_DIAG][decision] t=1 rank=0 instance=gen sequence=1 " + "active_snapshot=1 candidate_snapshot=1 " + f"active_requests={prefix} active_requests_omitted=6 " + "candidate_requests=- candidate_requests_omitted=0 " + "admitted=0 deferred=0 budget=1", + ] + ) + + remaining = _TELEMETRY._analyze_remaining_work_ground_truth(events) + + assert remaining["active_decision_samples"] == 64 + assert remaining["active_request_ids_recovered_from_overflow"] == 0 + assert remaining["active_request_ids_omitted"] == 6 + assert remaining["identity_coverage"]["fraction"] == pytest.approx(64 / 70) + + +def test_membership_snapshots_do_not_cross_instances_or_ranks(): + prefix = ",".join(f"{request}:1" for request in range(1, 65)) + events = _parse_lines( + [ + "[DISAGG_DIAG][decision-members] t=1 rank=0 role=gen " + "instance=gen-a snapshot_version=1 membership=candidate " + "chunk_index=1 chunk_count=1 requests=65:1", + "[DISAGG_DIAG][decision] t=1 rank=0 instance=gen-a sequence=1 " + "candidate_snapshot=1 active_snapshot=1 active_requests=- " + "active_requests_omitted=0 " + f"candidate_requests={prefix} candidate_requests_omitted=1 " + "admitted=0 deferred=65 budget=1", + "[DISAGG_DIAG][decision-members] t=1 rank=1 role=gen " + "instance=gen-b snapshot_version=1 membership=candidate " + "chunk_index=1 chunk_count=1 requests=165:1", + "[DISAGG_DIAG][decision] t=1 rank=1 instance=gen-b sequence=1 " + "candidate_snapshot=1 active_snapshot=1 active_requests=- " + "active_requests_omitted=0 " + f"candidate_requests={prefix} candidate_requests_omitted=1 " + "admitted=0 deferred=65 budget=1", + ] + ) + + admissions, _ = _TELEMETRY._collect_admissions(events) + + assert {admission.candidate_requests[-1][0] for admission in admissions} == { + "65", + "165", + } + + +def test_single_log_same_rank_is_namespaced_by_instance(): + events = [ + _parse_source_line( + "[DISAGG_DIAG][decision] t=1 host=node instance=gen-a rank=0 " + "sequence=1 active_blocks=0 candidate_requests=1:1 admitted=1 " + "admitted_requests=1:1 deferred=0 budget=1", + "combined.log", + ), + _parse_source_line( + "[DISAGG_DIAG][decision] t=1 host=node instance=gen-b rank=0 " + "sequence=1 active_blocks=0 candidate_requests=101:1 admitted=1 " + "admitted_requests=101:1 deferred=0 budget=1", + "combined.log", + ), + ] + + result = analyze_events(events) + + assert sorted(result["ranks"]) == [ + "0::host=node::instance=gen-a::role=gen", + "0::host=node::instance=gen-b::role=gen", + ] + assert result["rank_namespace"] == "source-path::rank[::host::instance::role]" + + +def test_missing_native_instance_aliases_to_single_known_instance(): + events = [ + _parse_source_line( + "[DISAGG_DIAG][submit] t=1 host=node instance=gen rank=0 request=1 blocks=4", + "worker.log", + ), + _parse_source_line( + "[DISAGG_DIAG][python-transfer] t=2 host=node rank=0 role=gen " + "action=local-ready request=1 outcome=completed", + "worker.log", + ), + _parse_source_line( + "[DISAGG_DIAG][reap] t=3 host=node instance=gen rank=0 " + "request=1 blocks=4 outcome=completed", + "worker.log", + ), + ] + + result = analyze_events(events) + + assert list(result["ranks"]) == ["0"] + assert result["ranks"]["0"]["service"]["latency_s"]["p50"] == pytest.approx(1.0) + + +@pytest.mark.parametrize("deadline_action", ["timeout", "timed-out"]) +def test_legacy_timeout_action_is_an_observation_not_a_terminal_outcome( + deadline_action, +): + events = _parse_lines( + [ + "[DISAGG_DIAG][ctx-transfer] t=0.0 rank=0 action=timer-start request=42", + f"[DISAGG_DIAG][ctx-transfer] t=1.0 rank=0 action={deadline_action} request=42", + "[DISAGG_DIAG][ctx-transfer] t=2.0 rank=0 action=reaped request=42 outcome=completed", + ] + ) + + result = analyze_events(events) + + assert result["lifecycle"]["deadline_phase_counts"] == {"pre-credit": 1} + assert _TELEMETRY._collect_unsuccessful_requests(events) == set() + + def test_remaining_work_censors_cross_source_endpoints(): events = [ _parse_source_line( From 4ca5af1fdb15c7133a58238e639f35ffa38f63d3 Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Fri, 24 Jul 2026 14:30:42 -0700 Subject: [PATCH 6/7] [NVBUG 6312828][test] fix telemetry correlation domains Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- scripts/disagg_admission_telemetry.py | 108 +++++++---- tensorrt_llm/_torch/pyexecutor/py_executor.py | 13 +- .../_torch/executor/test_py_executor.py | 32 +++- .../tools/test_disagg_admission_telemetry.py | 176 +++++++++++++++++- 4 files changed, 289 insertions(+), 40 deletions(-) diff --git a/scripts/disagg_admission_telemetry.py b/scripts/disagg_admission_telemetry.py index 4ee00555f6e5..d6be97878d32 100644 --- a/scripts/disagg_admission_telemetry.py +++ b/scripts/disagg_admission_telemetry.py @@ -35,6 +35,8 @@ _FIELD_PATTERN = re.compile(r"([A-Za-z_][A-Za-z0-9_]*)=([^\s]+)") _RANK_PATTERN = re.compile(r"\[RANK\s+(\d+)]") +_ClockDomain = tuple[str, str, str, str, str, str] + @dataclass(frozen=True) class DiagnosticEvent: @@ -134,6 +136,7 @@ class _LifecycleMark: log_source: str emitter_source: str | None host: str + instance: str role: str rank: str clock: str @@ -403,6 +406,22 @@ def resolved_instance(event: DiagnosticEvent) -> str: ) return next(iter(candidates)) if len(candidates) == 1 else instance + sorted_events = [ + event + if resolved_instance(event) == event.fields.get("instance", "-") + else DiagnosticEvent( + event.category, + event.time_s, + event.rank, + { + **event.fields, + "instance": resolved_instance(event), + }, + event.source, + ) + for event in sorted_events + ] + emitters_by_legacy_rank: dict[str, set[tuple[str, str, str]]] = defaultdict(set) for event in sorted_events: legacy_rank = f"{event.source}::rank={event.rank}" if namespace_by_source else event.rank @@ -955,10 +974,8 @@ def _analyze_rank( def _analyze_request_lifecycles(events: list[DiagnosticEvent]) -> dict[str, object]: marks = _collect_lifecycle_marks(events) - timelines: dict[tuple[tuple[str, str, str, str, str], str], list[_LifecycleMark]] = defaultdict( - list - ) - tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]] = defaultdict(set) + timelines: dict[tuple[_ClockDomain, str], list[_LifecycleMark]] = defaultdict(list) + tag_domains: dict[tuple[str, str], set[_ClockDomain]] = defaultdict(set) for mark in marks: domain = _lifecycle_domain(mark) timelines[(domain, mark.request)].append(mark) @@ -1003,6 +1020,7 @@ def _analyze_request_lifecycles(events: list[DiagnosticEvent]) -> dict[str, obje "role": domain[2], "rank": domain[3], "clock": domain[4], + "instance": domain[5], "clock_domain": _clock_domain_label(domain), "emitter_sources": emitter_sources, "first_timestamps": first_timestamps, @@ -1018,9 +1036,9 @@ def _analyze_request_lifecycles(events: list[DiagnosticEvent]) -> dict[str, obje ) return { "clock_domain_policy": ( - "Durations require the same input log source, host, role, rank, and clock. " - "Matching request IDs in another source or clock domain are correlated only " - "for censoring; their raw timestamps are never subtracted." + "Durations require the same input log source, host, instance, role, rank, " + "and clock. Matching request IDs in another source or clock domain are " + "correlated only for censoring; their raw timestamps are never subtracted." ), "correlation_request_policy": ( "Prefer a nonzero context_request/disaggregated request ID; otherwise use request." @@ -1064,6 +1082,7 @@ def add( log_source=event.source or "", emitter_source=event.fields.get("source"), host=_event_host(event), + instance=event.fields.get("instance", "-"), role=role, rank=event.rank, clock=_event_clock(event), @@ -1671,17 +1690,37 @@ def _cross_host_wall_anomalies( start = points.get(start_tag) end = points.get(end_tag) if start is not None and end is not None and float(end["wall_t"]) < float(start["wall_t"]): - anomalies.append(f"negative:{start_tag}->{end_tag}") + used_multi_emitter_selection = ( + int(start.get("observed_emitter_count", 1)) > 1 + or int(end.get("observed_emitter_count", 1)) > 1 + ) + selected_different_emitters = _cross_host_emitter(start) != _cross_host_emitter(end) + prefix = ( + "cross-emitter-selection" + if used_multi_emitter_selection and selected_different_emitters + else "negative" + ) + anomalies.append(f"{prefix}:{start_tag}->{end_tag}") return anomalies +def _cross_host_emitter(point: dict[str, object]) -> tuple[object, ...]: + return ( + point.get("log_source"), + point.get("host"), + point.get("instance"), + point.get("role"), + point.get("rank"), + ) + + def _evaluate_lifecycle_interval( request: str, - domain: tuple[str, str, str, str, str], + domain: _ClockDomain, marks_by_tag: dict[str, list[_LifecycleMark]], start_tags: tuple[str, ...], end_tags: tuple[str, ...], - tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]], + tag_domains: dict[tuple[str, str], set[_ClockDomain]], ) -> dict[str, object]: starts = [mark for tag in start_tags for mark in marks_by_tag.get(tag, ())] ends = [mark for tag in end_tags for mark in marks_by_tag.get(tag, ())] @@ -1775,9 +1814,9 @@ def _classify_deadline_phase( def _cross_domain_endpoint_reason( endpoint: str, request: str, - domain: tuple[str, str, str, str, str], + domain: _ClockDomain, tags: tuple[str, ...], - tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]], + tag_domains: dict[tuple[str, str], set[_ClockDomain]], ) -> str | None: other_domains = { candidate @@ -1799,7 +1838,7 @@ def _lifecycle_interval_coverage( records_by_request: dict[str, list[dict[str, object]]] = defaultdict(list) for record in request_records: records_by_request[str(record["request"])].append(record) - mark_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]] = defaultdict(set) + mark_domains: dict[tuple[str, str], set[_ClockDomain]] = defaultdict(set) for mark in marks: mark_domains[(mark.request, mark.tag)].add(_lifecycle_domain(mark)) @@ -1861,10 +1900,10 @@ def _analyze_remaining_work_ground_truth( events: list[DiagnosticEvent], ) -> dict[str, object]: marks = _collect_lifecycle_marks(events) - timelines: dict[tuple[tuple[str, str, str, str, str], str], dict[str, list[_LifecycleMark]]] = ( - defaultdict(lambda: defaultdict(list)) + timelines: dict[tuple[_ClockDomain, str], dict[str, list[_LifecycleMark]]] = defaultdict( + lambda: defaultdict(list) ) - tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]] = defaultdict(set) + tag_domains: dict[tuple[str, str], set[_ClockDomain]] = defaultdict(set) for mark in marks: domain = _lifecycle_domain(mark) timelines[(domain, mark.request)][mark.tag].append(mark) @@ -1894,7 +1933,6 @@ def _analyze_remaining_work_ground_truth( omitted = _as_int(event.fields.get("active_requests_omitted")) or 0 omission_identity = ( domain, - event.fields.get("instance", "-"), sequence or event.time_s, ) if omission_identity not in omission_seen: @@ -1953,6 +1991,7 @@ def _analyze_remaining_work_ground_truth( "role": domain[2], "rank": domain[3], "clock": domain[4], + "instance": domain[5], "clock_domain": _clock_domain_label(domain), "active_age_s": active_age_s, "active_age_bucket": _active_age_bucket(active_age_s), @@ -2009,8 +2048,8 @@ def _analyze_remaining_work_ground_truth( return { "definition": ( "For each Gate-2 admission snapshot and active request, residual_ready_s and " - "residual_reap_s use only later GEN events in the identical input source and " - "clock domain. CTX timestamps are never used." + "residual_reap_s use only later GEN events in the identical input source, " + "instance, and clock domain. CTX timestamps are never used." ), "active_decision_samples": len(samples), "active_request_ids_omitted": active_request_ids_omitted, @@ -2054,11 +2093,11 @@ def _analyze_remaining_work_ground_truth( def _remaining_endpoint_censor_reason( request: str, - domain: tuple[str, str, str, str, str], + domain: _ClockDomain, tag: str | tuple[str, ...], decision_time_s: float, tags: dict[str, list[_LifecycleMark]], - tag_domains: dict[tuple[str, str], set[tuple[str, str, str, str, str]]], + tag_domains: dict[tuple[str, str], set[_ClockDomain]], ) -> str: endpoint_tags = (tag,) if isinstance(tag, str) else tag endpoint_name = tag if isinstance(tag, str) else "ready" @@ -2247,7 +2286,7 @@ def _event_clock(event: DiagnosticEvent) -> str: return event.fields.get("clock_domain") or event.fields.get("clock") or "local_steady" -def _event_domain(event: DiagnosticEvent) -> tuple[str, str, str, str, str]: +def _event_domain(event: DiagnosticEvent) -> _ClockDomain: category = _normalize_diag_token(event.category) action = _normalize_diag_token(event.fields.get("action", "")) return ( @@ -2256,16 +2295,17 @@ def _event_domain(event: DiagnosticEvent) -> tuple[str, str, str, str, str]: _event_role(event, category, action), event.rank, _event_clock(event), + event.fields.get("instance", "-"), ) -def _lifecycle_domain(mark: _LifecycleMark) -> tuple[str, str, str, str, str]: - return mark.log_source, mark.host, mark.role, mark.rank, mark.clock +def _lifecycle_domain(mark: _LifecycleMark) -> _ClockDomain: + return mark.log_source, mark.host, mark.role, mark.rank, mark.clock, mark.instance -def _clock_domain_label(domain: tuple[str, str, str, str, str]) -> str: - source, host, role, rank, clock = domain - return f"{source}::host={host}::role={role}::rank={rank}::clock={clock}" +def _clock_domain_label(domain: _ClockDomain) -> str: + source, host, role, rank, clock, instance = domain + return f"{source}::host={host}::role={role}::rank={rank}::clock={clock}::instance={instance}" def _active_age_bucket(active_age_s: float | None) -> str: @@ -2279,7 +2319,7 @@ def _active_age_bucket(active_age_s: float | None) -> str: def _membership_snapshot_key( event: DiagnosticEvent, membership: str, *, definition: bool = False -) -> tuple[tuple[str, str, str, str, str], str, str, str] | None: +) -> tuple[_ClockDomain, str, str, str] | None: reference_field = "snapshot_version" if definition else f"{membership}_snapshot" reference = event.fields.get(reference_field) or event.fields.get("sequence") if reference is None: @@ -2296,7 +2336,7 @@ def _membership_overflow_for_event( event: DiagnosticEvent, membership: str, overflow: dict[ - tuple[tuple[str, str, str, str, str], str, str, str], + tuple[_ClockDomain, str, str, str], list[tuple[str, float]], ], ) -> list[tuple[str, float]]: @@ -2430,15 +2470,15 @@ def _collect_admissions(events: list[DiagnosticEvent]) -> tuple[list[Admission], def _collect_membership_overflow( events: list[DiagnosticEvent], ) -> dict[ - tuple[tuple[str, str, str, str, str], str, str, str], + tuple[_ClockDomain, str, str, str], list[tuple[str, float]], ]: chunks: dict[ - tuple[tuple[str, str, str, str, str], str, str, str], + tuple[_ClockDomain, str, str, str], dict[int, list[tuple[str, float]]], ] = defaultdict(dict) - expected_chunk_counts: dict[tuple[tuple[str, str, str, str, str], str, str, str], int] = {} - conflicts: set[tuple[tuple[str, str, str, str, str], str, str, str]] = set() + expected_chunk_counts: dict[tuple[_ClockDomain, str, str, str], int] = {} + conflicts: set[tuple[_ClockDomain, str, str, str]] = set() for event in events: if _normalize_diag_token(event.category) != "decision-members": continue @@ -2470,7 +2510,7 @@ def _collect_membership_overflow( chunks[key][chunk_index] = request_blocks overflow: dict[ - tuple[tuple[str, str, str, str, str], str, str, str], + tuple[_ClockDomain, str, str, str], list[tuple[str, float]], ] = {} for key, key_chunks in chunks.items(): diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 02aed1573adf..1758fd9cabb5 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -3457,7 +3457,8 @@ def _log_disagg_transfer_diagnostic(self, category: str, **fields) -> None: f"instance={instance} rank={rank} {encoded_fields}") @staticmethod - def _disagg_diag_request_id(request: LlmRequest) -> int: + def _disagg_diag_request_id(request: LlmRequest | ExecutorRequest, + fallback_request_id: int | None = None) -> int: params = getattr(request, "py_disaggregated_params", None) disagg_request_id = getattr(params, "disagg_request_id", None) if params is not None else None @@ -3467,7 +3468,12 @@ def _disagg_diag_request_id(request: LlmRequest) -> int: None) if params is not None else None if isinstance(ctx_request_id, int): return ctx_request_id - return request.py_request_id + py_request_id = getattr(request, "py_request_id", None) + if isinstance(py_request_id, int): + return py_request_id + if isinstance(fallback_request_id, int): + return fallback_request_id + raise ValueError("A diagnostic request ID is required") @staticmethod def _disagg_diag_schedule_style(request: LlmRequest) -> str: @@ -3672,7 +3678,8 @@ def _log_disagg_gen_ingress( t=ingress_time, wall_t=ingress_wall_time, wall_semantics="boundary-sampled", - request=item.id, + request=self._disagg_diag_request_id(request, item.id), + local_request=item.id, boundary="executor-queue-to-waiting-queue") def _log_disagg_gen_activations(self, diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index 208eef1ad4c2..16b7984c0f8f 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -823,7 +823,10 @@ def test_logs_generation_executor_ingress(self, monkeypatch): monkeypatch.setattr(py_executor_module.logger, "info", log_info) executor = object.__new__(PyExecutor) executor.dist = Mock(rank=2) - request = Mock(request_type=py_executor_module.RequestType.REQUEST_TYPE_GENERATION_ONLY) + request = types.SimpleNamespace( + request_type=py_executor_module.RequestType.REQUEST_TYPE_GENERATION_ONLY, + py_disaggregated_params=None, + ) item = RequestQueueItem(701, request) PyExecutor._log_disagg_gen_ingress(executor, [item]) @@ -833,8 +836,35 @@ def test_logs_generation_executor_ingress(self, monkeypatch): assert "[DISAGG_DIAG][gen-arrival]" in message assert "rank=2" in message assert "request=701" in message + assert "local_request=701" in message assert "boundary=executor-queue-to-waiting-queue" in message + def test_generation_ingress_uses_context_request_id_for_correlation(self, monkeypatch): + monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") + monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr( + py_executor_module, + "get_steady_clock_now_in_seconds", + lambda: 12.0, + ) + log_info = Mock() + monkeypatch.setattr(py_executor_module.logger, "info", log_info) + executor = object.__new__(PyExecutor) + executor.dist = Mock(rank=2) + request = types.SimpleNamespace( + request_type=py_executor_module.RequestType.REQUEST_TYPE_GENERATION_ONLY, + py_disaggregated_params=types.SimpleNamespace( + disagg_request_id=None, + ctx_request_id=1701, + ), + ) + + PyExecutor._log_disagg_gen_ingress(executor, [RequestQueueItem(701, request)]) + + message = log_info.call_args.args[0] + assert "request=1701" in message + assert "local_request=701" in message + def test_logs_generation_executor_activation_with_common_id(self, monkeypatch): monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) diff --git a/tests/unittest/tools/test_disagg_admission_telemetry.py b/tests/unittest/tools/test_disagg_admission_telemetry.py index 43450dacc502..11d748af2d7e 100644 --- a/tests/unittest/tools/test_disagg_admission_telemetry.py +++ b/tests/unittest/tools/test_disagg_admission_telemetry.py @@ -516,6 +516,8 @@ def test_cross_host_partial_gate2_admission_is_not_labeled_before_gate1(): "[DISAGG_DIAG][gate2] t=3 wall_t=70 wall_clock=unix " "wall_semantics=boundary-sampled host=gen rank=1 " "action=admitted request=42", + "[DISAGG_DIAG][submit] t=4 wall_t=55 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen rank=0 request=42", "[DISAGG_DIAG][ctx-transfer] t=60 wall_t=60 wall_clock=unix " "wall_semantics=boundary-sampled host=ctx rank=0 " "action=deadline-observed request=42", @@ -526,6 +528,65 @@ def test_cross_host_partial_gate2_admission_is_not_labeled_before_gate1(): assert record["points"]["gate2-state-at-ctx-deadline"]["state"] == "admitted" assert record["ctx_deadline_relationship"] == "partial-gate2-admission-before-global-admission" + assert "negative:gate2-admitted->gen-submit" not in record["wall_clock_anomalies"] + assert "cross-emitter-selection:gate2-admitted->gen-submit" in record["wall_clock_anomalies"] + + +def test_cross_host_wall_clock_preserves_same_emitter_negative_order(): + events = _parse_lines( + [ + "[DISAGG_DIAG][ctx-transfer] t=10 wall_t=10 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx instance=ctx rank=0 " + "action=timer-start request=42", + "[DISAGG_DIAG][ctx-transfer] t=5 wall_t=5 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx instance=ctx rank=0 " + "action=deadline-observed request=42", + ] + ) + + record = analyze_events(events)["cross_host_correlation"]["requests"][0] + + assert "negative:ctx-timer-start->ctx-deadline" in record["wall_clock_anomalies"] + + +def test_cross_host_wall_clock_preserves_single_cross_emitter_inversion(): + events = _parse_lines( + [ + "[DISAGG_DIAG][submit] t=10 wall_t=10 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen instance=gen rank=0 " + "request=42", + "[DISAGG_DIAG][sender-transfer] t=9 wall_t=9 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx instance=ctx rank=0 " + "role=ctx action=receiver-info-ready request=42", + ] + ) + + record = analyze_events(events)["cross_host_correlation"]["requests"][0] + + assert "negative:gen-submit->ctx-receiver-credit" in record["wall_clock_anomalies"] + + +def test_cross_host_wall_clock_preserves_multi_rank_same_emitter_inversion(): + events = _parse_lines( + [ + "[DISAGG_DIAG][gate2] t=70 wall_t=70 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen instance=gen rank=0 " + "action=admitted request=42", + "[DISAGG_DIAG][gate2] t=60 wall_t=60 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen instance=gen rank=1 " + "action=admitted request=42", + "[DISAGG_DIAG][submit] t=55 wall_t=55 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen instance=gen rank=0 " + "request=42", + "[DISAGG_DIAG][submit] t=50 wall_t=50 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen instance=gen rank=1 " + "request=42", + ] + ) + + record = analyze_events(events)["cross_host_correlation"]["requests"][0] + + assert "negative:gate2-admitted->gen-submit" in record["wall_clock_anomalies"] def test_cross_host_progress_uses_latest_observed_rank(): @@ -648,14 +709,14 @@ def test_versioned_candidate_snapshot_reconstructs_unchanged_decisions(): "instance=gen sequence=1 snapshot_version=1 membership=candidate " f"chunk_index=1 chunk_count=1 requests={tail}", "[DISAGG_DIAG][decision] t=1 rank=0 instance=gen sequence=1 " - "candidate_snapshot=1 active_snapshot=1 active_requests=- " + "candidate_snapshot=1 active_snapshot=1 active_blocks=0 active_requests=- " "active_requests_omitted=0 " f"candidate_requests={prefix} candidate_requests_omitted=6 " "admitted=0 admitted_requests=- admitted_requests_omitted=0 " f"deferred=70 deferred_requests={prefix} " "deferred_requests_omitted=6 budget=1", "[DISAGG_DIAG][decision] t=2 rank=0 instance=gen sequence=2 " - "candidate_snapshot=1 active_snapshot=1 active_requests=- " + "candidate_snapshot=1 active_snapshot=1 active_blocks=0 active_requests=- " "active_requests_omitted=0 " f"candidate_requests={prefix} candidate_requests_omitted=6 " "admitted=0 admitted_requests=- admitted_requests_omitted=0 " @@ -671,6 +732,10 @@ def test_versioned_candidate_snapshot_reconstructs_unchanged_decisions(): assert all(len(admission.deferred_requests) == 70 for admission in admissions) assert all(admission.candidate_requests_omitted == 0 for admission in admissions) assert admissions[1].deferred_requests[-1] == "70" + fixed = analyze_events(events)["ranks"]["0"]["fixed_multiplier_counterfactual"] + assert len(fixed["samples"]) == 2 + assert fixed["next_deferred_required_multiplier"]["count"] == 2 + assert [sample["next_deferred_request"] for sample in fixed["samples"]] == ["1", "1"] def test_remaining_work_ignores_gate_transition_and_reports_incomplete_tail(): @@ -699,6 +764,43 @@ def test_remaining_work_ignores_gate_transition_and_reports_incomplete_tail(): assert remaining["identity_coverage"]["fraction"] == pytest.approx(64 / 70) +def test_remaining_work_recovers_complete_active_snapshot_overflow(): + prefix = ",".join(f"{request}:1" for request in range(1, 65)) + tail = ",".join(f"{request}:1" for request in range(65, 71)) + events = _parse_lines( + [ + "[DISAGG_DIAG][decision-members] t=1 rank=0 role=gen " + "instance=gen sequence=1 snapshot_version=1 membership=active " + f"chunk_index=1 chunk_count=1 requests={tail}", + "[DISAGG_DIAG][decision] t=1 rank=0 instance=gen sequence=1 " + "active_snapshot=1 candidate_snapshot=1 " + f"active_requests={prefix} active_requests_omitted=6 " + "candidate_requests=- candidate_requests_omitted=0 " + "admitted=0 deferred=0 budget=1", + *[ + "[DISAGG_DIAG][python-transfer] t=2 rank=0 role=gen " + f"instance=gen action=local-ready request={request} outcome=completed" + for request in range(1, 71) + ], + ] + ) + + remaining = _TELEMETRY._analyze_remaining_work_ground_truth(events) + + assert remaining["active_decision_samples"] == 70 + assert remaining["active_request_ids_recovered_from_overflow"] == 6 + assert remaining["active_request_ids_omitted"] == 0 + assert remaining["identity_coverage"]["fraction"] == 1.0 + assert remaining["ready_coverage"] == { + "eligible": 70, + "observed": 70, + "censored": 0, + "censor_reasons": {}, + } + assert remaining["residual_ready_s"]["count"] == 70 + assert remaining["residual_ready_s"]["p50"] == pytest.approx(1.0) + + def test_membership_snapshots_do_not_cross_instances_or_ranks(): prefix = ",".join(f"{request}:1" for request in range(1, 65)) events = _parse_lines( @@ -755,12 +857,74 @@ def test_single_log_same_rank_is_namespaced_by_instance(): assert result["rank_namespace"] == "source-path::rank[::host::instance::role]" +def test_lifecycle_and_remaining_work_are_namespaced_by_instance(): + events = [ + _parse_source_line(line, "combined.log") + for line in [ + "[DISAGG_DIAG][gen-arrival] t=0 host=node instance=gen-a rank=0 request=42", + "[DISAGG_DIAG][gen-activation] t=0.5 host=node instance=gen-a rank=0 request=42", + "[DISAGG_DIAG][decision] t=1 host=node instance=gen-a rank=0 " + "sequence=1 active_requests=42:1 active_requests_omitted=0 " + "candidate_requests=- admitted=0 deferred=0 budget=1", + "[DISAGG_DIAG][python-transfer] t=2 host=node instance=gen-a " + "rank=0 role=gen action=local-ready request=42 outcome=completed", + "[DISAGG_DIAG][gen-arrival] t=10 host=node instance=gen-b rank=0 request=42", + "[DISAGG_DIAG][gen-activation] t=11 host=node instance=gen-b rank=0 request=42", + "[DISAGG_DIAG][decision] t=11 host=node instance=gen-b rank=0 " + "sequence=1 active_requests=42:1 active_requests_omitted=0 " + "candidate_requests=- admitted=0 deferred=0 budget=1", + "[DISAGG_DIAG][python-transfer] t=20 host=node instance=gen-b " + "rank=0 role=gen action=local-ready request=42 outcome=completed", + ] + ] + + result = analyze_events(events) + lifecycle = result["lifecycle"] + remaining = result["remaining_work_ground_truth"] + + assert lifecycle["clock_domain_request_count"] == 2 + assert {record["instance"] for record in lifecycle["requests"]} == { + "gen-a", + "gen-b", + } + activation_coverage = lifecycle["interval_coverage"]["gen_arrival_to_activation"] + assert activation_coverage["observed_samples"] == 2 + assert sorted( + record["intervals"]["gen_arrival_to_activation"]["duration_s"] + for record in lifecycle["requests"] + if record["intervals"]["gen_arrival_to_activation"]["status"] == "observed" + ) == [0.5, 1.0] + assert remaining["active_decision_samples"] == 2 + assert {sample["instance"] for sample in remaining["samples"]} == { + "gen-a", + "gen-b", + } + assert sorted(sample["residual_ready_s"] for sample in remaining["samples"]) == [ + 1.0, + 9.0, + ] + + def test_missing_native_instance_aliases_to_single_known_instance(): events = [ + _parse_source_line( + "[DISAGG_DIAG][gen-arrival] t=0 host=node instance=gen rank=0 request=1", + "worker.log", + ), + _parse_source_line( + "[DISAGG_DIAG][gen-activation] t=0.5 host=node rank=0 request=1", + "worker.log", + ), _parse_source_line( "[DISAGG_DIAG][submit] t=1 host=node instance=gen rank=0 request=1 blocks=4", "worker.log", ), + _parse_source_line( + "[DISAGG_DIAG][decision] t=1.5 host=node instance=gen rank=0 " + "sequence=1 active_requests=1:4 active_requests_omitted=0 " + "candidate_requests=- admitted=0 deferred=0 budget=4", + "worker.log", + ), _parse_source_line( "[DISAGG_DIAG][python-transfer] t=2 host=node rank=0 role=gen " "action=local-ready request=1 outcome=completed", @@ -777,6 +941,14 @@ def test_missing_native_instance_aliases_to_single_known_instance(): assert list(result["ranks"]) == ["0"] assert result["ranks"]["0"]["service"]["latency_s"]["p50"] == pytest.approx(1.0) + lifecycle_interval = result["lifecycle"]["requests"][0]["intervals"][ + "gen_arrival_to_activation" + ] + assert lifecycle_interval["status"] == "observed" + assert lifecycle_interval["duration_s"] == pytest.approx(0.5) + remaining = result["remaining_work_ground_truth"] + assert remaining["active_decision_samples"] == 1 + assert remaining["residual_ready_s"]["p50"] == pytest.approx(0.5) @pytest.mark.parametrize("deadline_action", ["timeout", "timed-out"]) From ced9d0b7a0b9a89306026a14621607c2ebcaa0d0 Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Fri, 24 Jul 2026 17:01:43 -0700 Subject: [PATCH 7/7] [NVBUG 6312828][test] harden disaggregated admission telemetry Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- scripts/disagg_admission_telemetry.py | 915 +++++++++++++----- .../_torch/disaggregation/diagnostics.py | 33 + .../_torch/disaggregation/native/transfer.py | 6 +- .../_torch/disaggregation/transceiver.py | 45 +- tensorrt_llm/_torch/pyexecutor/py_executor.py | 32 +- .../_torch/executor/test_py_executor.py | 30 +- .../disaggregated/test_disagg_diagnostics.py | 32 + .../test_transceiver_bounded_polling.py | 73 +- .../tools/test_disagg_admission_telemetry.py | 310 ++++++ 9 files changed, 1226 insertions(+), 250 deletions(-) create mode 100644 tensorrt_llm/_torch/disaggregation/diagnostics.py create mode 100644 tests/unittest/disaggregated/test_disagg_diagnostics.py diff --git a/scripts/disagg_admission_telemetry.py b/scripts/disagg_admission_telemetry.py index d6be97878d32..24bc3d5b8b85 100644 --- a/scripts/disagg_admission_telemetry.py +++ b/scripts/disagg_admission_telemetry.py @@ -25,8 +25,10 @@ import json import math import re +from bisect import bisect_left, bisect_right from collections import Counter, defaultdict, deque from dataclasses import dataclass +from heapq import heappop, heappush from pathlib import Path from statistics import fmean from typing import Iterable, Sequence @@ -142,6 +144,26 @@ class _LifecycleMark: clock: str +@dataclass(frozen=True) +class _ReleaseIndexes: + """Chronological indexes shared by release/refill analyses on one rank.""" + + decision_times: tuple[float, ...] + admission_times: tuple[float, ...] + admission_by_sequence: dict[str, Admission] + admission_by_decision: tuple[Admission | None, ...] + signal_keys: tuple[tuple[float, int], ...] + signal_values: tuple[tuple[int, float | None], ...] + successful_decision_indices: tuple[int, ...] + successful_decision_times: tuple[float, ...] + successful_positions_without_detail: tuple[int, ...] + successful_positions_by_request: dict[str, tuple[int, ...]] + submits: tuple[PointEvent, ...] + submit_times: tuple[float, ...] + submit_positions_by_request: dict[str, tuple[int, ...]] + submits_are_chronological: bool + + _CTX_ACTIONS = { "queued", "send-queued", @@ -170,6 +192,7 @@ class _LifecycleMark: "peer-ready", "ready", "local-ready", + "global-ready", } _TERMINAL_ACTIONS = { "completed", @@ -274,6 +297,10 @@ class _LifecycleMark: ("local-ready", "receiver-completed"), ("reap",), ), + "global_ready_to_reap": ( + ("global-ready",), + ("reap",), + ), "gen_arrival_to_decode_start": ( ("gen-arrival",), ("decode-start",), @@ -282,6 +309,10 @@ class _LifecycleMark: ("local-ready", "receiver-completed"), ("decode-start",), ), + "global_ready_to_decode_start": ( + ("global-ready",), + ("decode-start",), + ), "reap_to_decode_start": ( ("reap",), ("decode-start",), @@ -352,6 +383,67 @@ def read_diagnostic_events(paths: Iterable[str | Path]) -> list[DiagnosticEvent] return events +def _resolve_missing_emitter_identity( + events: list[DiagnosticEvent], +) -> list[DiagnosticEvent]: + """Fill missing native host/instance fields only when one emitter fits. + + Native receiver-slot records do not currently carry the Python host and + worker instance. A log source can still identify them safely when every + explicit event for the same rank and role names one emitter. Ambiguous + combined logs remain unresolved instead of being guessed. + """ + emitters: dict[tuple[str | None, str, str], set[tuple[str, str]]] = defaultdict(set) + for event in events: + category = _normalize_diag_token(event.category) + action = _normalize_diag_token(event.fields.get("action", "")) + role = _event_role(event, category, action) + host = _event_host(event) + instance = event.fields.get("instance", "-") + if host == "" or instance in {"-", "unknown"}: + continue + emitters[(event.source, event.rank, role)].add((host, instance)) + + resolved: list[DiagnosticEvent] = [] + for event in events: + category = _normalize_diag_token(event.category) + action = _normalize_diag_token(event.fields.get("action", "")) + role = _event_role(event, category, action) + host = _event_host(event) + instance = event.fields.get("instance", "-") + candidates = emitters.get((event.source, event.rank, role), set()) + if host != "": + candidates = {candidate for candidate in candidates if candidate[0] == host} + if instance not in {"-", "unknown"}: + candidates = {candidate for candidate in candidates if candidate[1] == instance} + if len(candidates) != 1: + resolved.append(event) + continue + resolved_host, resolved_instance = next(iter(candidates)) + inferred_fields = [] + fields = dict(event.fields) + if host == "": + fields["host"] = resolved_host + inferred_fields.append("host") + if instance in {"-", "unknown"}: + fields["instance"] = resolved_instance + inferred_fields.append("instance") + if not inferred_fields: + resolved.append(event) + continue + fields["identity_inferred"] = ",".join(inferred_fields) + resolved.append( + DiagnosticEvent( + event.category, + event.time_s, + event.rank, + fields, + event.source, + ) + ) + return resolved + + def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: """Calculate admission-window measurements from parsed events. @@ -367,7 +459,7 @@ def analyze_events(events: Iterable[DiagnosticEvent]) -> dict[str, object]: A JSON-serializable analysis dictionary. """ sorted_events = sorted( - events, + _resolve_missing_emitter_identity(list(events)), key=lambda event: (event.source or "", event.rank, event.time_s), ) sources = {event.source for event in sorted_events if event.source is not None} @@ -457,21 +549,6 @@ def resolved_instance(event: DiagnosticEvent) -> str: events_by_rank[rank_key].append(event) ranks: dict[str, object] = {} - aggregate_service_intervals: list[ServiceInterval] = [] - aggregate_selected_gaps: list[dict[str, object]] = [] - aggregate_gaps_by_source: dict[str, list[dict[str, object]]] = defaultdict(list) - aggregate_slot_refill_gaps: list[float] = [] - aggregate_progress_credits: list[float] = [] - aggregate_fixed_multipliers: list[float] = [] - aggregate_poll_durations_ms: list[float] = [] - aggregate_progress_poll_durations_ms: list[float] = [] - aggregate_no_progress_poll_durations_ms: list[float] = [] - aggregate_reported_ready_to_reap_ms: list[float] = [] - aggregate_physical_release_to_reap_s: list[float] = [] - aggregate_invalid_ready_to_reap_samples = 0 - aggregate_busy_s = 0.0 - aggregate_completed_blocks = 0.0 - for rank in sorted(events_by_rank, key=_rank_sort_key): rank_events = events_by_rank[rank] block_scope = sorted_events @@ -500,76 +577,168 @@ def resolved_instance(event: DiagnosticEvent) -> str: event for event in sorted_events if event.source == representative.source ] request_blocks = _collect_global_request_blocks(block_scope) - rank_analysis, rank_intervals, selected_gaps = _analyze_rank(rank_events, request_blocks) + rank_analysis, _, _ = _analyze_rank(rank_events, request_blocks) ranks[rank] = rank_analysis - aggregate_service_intervals.extend(rank_intervals) - aggregate_selected_gaps.extend(selected_gaps) - release_analysis = rank_analysis["release_to_admission"] - if isinstance(release_analysis, dict): - by_source = release_analysis["by_source"] - if isinstance(by_source, dict): - for source, source_analysis in by_source.items(): - if isinstance(source_analysis, dict): - aggregate_gaps_by_source[source].extend(source_analysis["samples"]) - - service = rank_analysis["service"] - if isinstance(service, dict): - aggregate_busy_s += float(service["busy_s"]) - aggregate_completed_blocks += float(service["completed_blocks"]) - receiver_slots = rank_analysis["receiver_slots"] - if isinstance(receiver_slots, dict): - aggregate_slot_refill_gaps.extend(receiver_slots["backlog_refill_gap_samples_s"]) - progress = rank_analysis["linear_progress_credit"] - if isinstance(progress, dict): - aggregate_progress_credits.extend(progress["credit_samples_blocks"]) - counterfactual = rank_analysis["fixed_multiplier_counterfactual"] - if isinstance(counterfactual, dict): - aggregate_fixed_multipliers.extend( - counterfactual["next_deferred_required_multiplier_samples"] - ) - status_poll = rank_analysis["status_poll"] - if isinstance(status_poll, dict): - aggregate_poll_durations_ms.extend(status_poll["duration_samples_ms"]) - aggregate_progress_poll_durations_ms.extend(status_poll["progress_duration_samples_ms"]) - aggregate_no_progress_poll_durations_ms.extend( - status_poll["no_progress_duration_samples_ms"] - ) - scheduler_visibility = rank_analysis["scheduler_visibility"] - if isinstance(scheduler_visibility, dict): - aggregate_reported_ready_to_reap_ms.extend( - scheduler_visibility["reported_ready_to_reap_samples_ms"] - ) - aggregate_physical_release_to_reap_s.extend( - scheduler_visibility["physical_release_to_reap_samples_s"] - ) - aggregate_invalid_ready_to_reap_samples += int( - scheduler_visibility["invalid_reported_ready_to_reap_samples"] - ) + lifecycle = _analyze_request_lifecycles(sorted_events) + cross_host_correlation = _analyze_cross_host_correlation(sorted_events) + remaining_work_ground_truth = _analyze_remaining_work_ground_truth(sorted_events) + known_backlog_release = _known_backlog_release_analysis(ranks) + + return { + "schema_version": 3, + "rank_namespace": ( + "source-path::rank[::host::instance::role]" + if split_legacy_ranks + else ("source-path::rank" if namespace_by_source else "rank") + ), + "aggregate_scope": "all-input-sources" if namespace_by_source else "single-source", + "parsed_event_count": len(sorted_events), + "inferred_emitter_identity_events": sum( + "identity_inferred" in event.fields for event in sorted_events + ), + "event_counts": dict(sorted(category_counts.items())), + "ranks": ranks, + "lifecycle": lifecycle, + "cross_host_correlation": cross_host_correlation, + "remaining_work_ground_truth": remaining_work_ground_truth, + "known_backlog_release": known_backlog_release, + "aggregate": _aggregate_rank_reports( + rank_report for rank_report in ranks.values() if isinstance(rank_report, dict) + ), + "model": { + "shadow_multiplier": "1 + throughput_blocks_per_s * refill_gap_s / budget_blocks", + "fixed_multiplier_counterfactual": ( + "max(1, (active_blocks + FCFS_prefix_blocks) / budget_blocks)" + ), + "linear_progress_credit": ( + "sum(request_blocks * elapsed_service_s / realized_service_s)" + ), + "caveat": ( + "Retrospective service and progress use completed intervals; they are validation " + "estimates, not online remaining-work measurements. Python global-ready is " + "sampled after GEN consensus when available; local-ready is only a rank-local " + "bound, and reap is scheduler-visible." + ), + }, + } + + +def _aggregate_rank_reports( + rank_reports: Iterable[dict[str, object]], +) -> dict[str, object]: + """Reconstruct the public aggregate from already-analyzed rank reports.""" + service_latencies: list[float] = [] + selected_samples: list[dict[str, object]] = [] + gaps_by_source: dict[str, list[dict[str, object]]] = defaultdict(list) + slot_refill_gaps: list[float] = [] + progress_credits: list[float] = [] + fixed_multipliers: list[float] = [] + poll_durations_ms: list[float] = [] + progress_poll_durations_ms: list[float] = [] + no_progress_poll_durations_ms: list[float] = [] + reported_ready_to_reap_ms: list[float] = [] + global_ready_to_reap_s: list[float] = [] + physical_release_to_reap_s: list[float] = [] + invalid_ready_to_reap_samples = 0 + completed_service_intervals = 0 + busy_s = 0.0 + completed_blocks = 0.0 + + for rank_report in rank_reports: + service = rank_report["service"] + assert isinstance(service, dict) + intervals = service["intervals"] + assert isinstance(intervals, list) + completed_service_intervals += len(intervals) + service_latencies.extend(float(interval["latency_s"]) for interval in intervals) + busy_s += float(service["busy_s"]) + completed_blocks += float(service["completed_blocks"]) + + release_analysis = rank_report["release_to_admission"] + assert isinstance(release_analysis, dict) + rank_selected_samples = release_analysis["selected_samples"] + assert isinstance(rank_selected_samples, list) + selected_samples.extend(rank_selected_samples) + rank_gaps_by_source = release_analysis["by_source"] + assert isinstance(rank_gaps_by_source, dict) + for source, source_report in rank_gaps_by_source.items(): + assert isinstance(source_report, dict) + source_samples = source_report["samples"] + assert isinstance(source_samples, list) + gaps_by_source[str(source)].extend(source_samples) + + receiver_slots = rank_report["receiver_slots"] + assert isinstance(receiver_slots, dict) + slot_refill_gaps.extend( + float(value) for value in receiver_slots["backlog_refill_gap_samples_s"] + ) + + progress = rank_report["linear_progress_credit"] + assert isinstance(progress, dict) + progress_credits.extend(float(value) for value in progress["credit_samples_blocks"]) + + counterfactual = rank_report["fixed_multiplier_counterfactual"] + assert isinstance(counterfactual, dict) + fixed_multipliers.extend( + float(value) for value in counterfactual["next_deferred_required_multiplier_samples"] + ) + + status_poll = rank_report["status_poll"] + assert isinstance(status_poll, dict) + poll_durations_ms.extend(float(value) for value in status_poll["duration_samples_ms"]) + progress_poll_durations_ms.extend( + float(value) for value in status_poll["progress_duration_samples_ms"] + ) + no_progress_poll_durations_ms.extend( + float(value) for value in status_poll["no_progress_duration_samples_ms"] + ) + + scheduler_visibility = rank_report["scheduler_visibility"] + assert isinstance(scheduler_visibility, dict) + python_transfer = rank_report["python_transfer"] + assert isinstance(python_transfer, dict) + global_ready_to_reap_s.extend( + float(value) for value in python_transfer["global_ready_to_reap_samples_s"] + ) + reported_ready_to_reap_ms.extend( + float(value) for value in scheduler_visibility["reported_ready_to_reap_samples_ms"] + ) + physical_release_to_reap_s.extend( + float(value) for value in scheduler_visibility["physical_release_to_reap_samples_s"] + ) + invalid_ready_to_reap_samples += int( + scheduler_visibility["invalid_reported_ready_to_reap_samples"] + ) - aggregate_throughput = _safe_ratio(aggregate_completed_blocks, aggregate_busy_s) selected_decision_gaps = [ float(sample["decision_gap_s"]) - for sample in aggregate_selected_gaps + for sample in selected_samples if sample.get("decision_gap_s") is not None ] selected_successful_admission_gaps = [ float(sample["successful_admission_gap_s"]) - for sample in aggregate_selected_gaps + for sample in selected_samples if sample.get("successful_admission_gap_s") is not None ] selected_refill_gaps = [ float(sample["refill_gap_s"]) - for sample in aggregate_selected_gaps + for sample in selected_samples if sample.get("refill_gap_s") is not None ] shadow_samples = [ float(sample["shadow_multiplier"]) - for sample in aggregate_selected_gaps + for sample in selected_samples if sample.get("shadow_multiplier") is not None ] - aggregate_release_bounds = { + release_bounds = { source: { - "decision_gap_s": _summary([float(sample["decision_gap_s"]) for sample in samples]), + "decision_gap_s": _summary( + [ + float(sample["decision_gap_s"]) + for sample in samples + if sample.get("decision_gap_s") is not None + ] + ), "successful_admission_gap_s": _summary( [ float(sample["successful_admission_gap_s"]) @@ -592,71 +761,34 @@ def resolved_instance(event: DiagnosticEvent) -> str: ] ), } - for source, samples in sorted(aggregate_gaps_by_source.items()) + for source, samples in sorted(gaps_by_source.items()) } - lifecycle = _analyze_request_lifecycles(sorted_events) - cross_host_correlation = _analyze_cross_host_correlation(sorted_events) - remaining_work_ground_truth = _analyze_remaining_work_ground_truth(sorted_events) - known_backlog_release = _known_backlog_release_analysis(ranks) - return { - "schema_version": 3, - "rank_namespace": ( - "source-path::rank[::host::instance::role]" - if split_legacy_ranks - else ("source-path::rank" if namespace_by_source else "rank") + "completed_service_intervals": completed_service_intervals, + "completed_blocks": completed_blocks, + "busy_rank_seconds": busy_s, + "throughput_blocks_per_s": _safe_ratio(completed_blocks, busy_s), + "service_latency_s": _summary(service_latencies), + "selected_physical_release_to_next_decision_gap_s": _summary(selected_decision_gaps), + "selected_physical_release_to_successful_admission_gap_s": _summary( + selected_successful_admission_gaps ), - "aggregate_scope": "all-input-sources" if namespace_by_source else "single-source", - "parsed_event_count": len(sorted_events), - "event_counts": dict(sorted(category_counts.items())), - "ranks": ranks, - "lifecycle": lifecycle, - "cross_host_correlation": cross_host_correlation, - "remaining_work_ground_truth": remaining_work_ground_truth, - "known_backlog_release": known_backlog_release, - "aggregate": { - "completed_service_intervals": len(aggregate_service_intervals), - "completed_blocks": aggregate_completed_blocks, - "busy_rank_seconds": aggregate_busy_s, - "throughput_blocks_per_s": aggregate_throughput, - "service_latency_s": _summary( - [interval.end_s - interval.start_s for interval in aggregate_service_intervals] - ), - "selected_physical_release_to_next_decision_gap_s": _summary(selected_decision_gaps), - "selected_physical_release_to_successful_admission_gap_s": _summary( - selected_successful_admission_gaps - ), - "selected_physical_release_to_refill_gap_s": _summary(selected_refill_gaps), - "release_bounds_by_source": aggregate_release_bounds, - "receiver_slot_refill_gap_s": _summary(aggregate_slot_refill_gaps), - "selected_physical_shadow_multiplier": _summary(shadow_samples), - "next_deferred_required_fixed_multiplier": _summary(aggregate_fixed_multipliers), - "linear_progress_credit_blocks": _summary(aggregate_progress_credits), - "status_poll": { - "duration_ms": _summary(aggregate_poll_durations_ms), - "progress_duration_ms": _summary(aggregate_progress_poll_durations_ms), - "no_progress_duration_ms": _summary(aggregate_no_progress_poll_durations_ms), - }, - "scheduler_visibility": { - "reported_ready_to_reap_ms": _summary(aggregate_reported_ready_to_reap_ms), - "physical_release_to_reap_s": _summary(aggregate_physical_release_to_reap_s), - "invalid_reported_ready_to_reap_samples": (aggregate_invalid_ready_to_reap_samples), - }, + "selected_physical_release_to_refill_gap_s": _summary(selected_refill_gaps), + "release_bounds_by_source": release_bounds, + "receiver_slot_refill_gap_s": _summary(slot_refill_gaps), + "selected_physical_shadow_multiplier": _summary(shadow_samples), + "next_deferred_required_fixed_multiplier": _summary(fixed_multipliers), + "linear_progress_credit_blocks": _summary(progress_credits), + "status_poll": { + "duration_ms": _summary(poll_durations_ms), + "progress_duration_ms": _summary(progress_poll_durations_ms), + "no_progress_duration_ms": _summary(no_progress_poll_durations_ms), }, - "model": { - "shadow_multiplier": "1 + throughput_blocks_per_s * refill_gap_s / budget_blocks", - "fixed_multiplier_counterfactual": ( - "max(1, (active_blocks + FCFS_prefix_blocks) / budget_blocks)" - ), - "linear_progress_credit": ( - "sum(request_blocks * elapsed_service_s / realized_service_s)" - ), - "caveat": ( - "Retrospective service and progress use completed intervals; they are validation " - "estimates, not online remaining-work measurements. Python local-ready is a " - "rank-local bound and reap is scheduler-visible; runtime control requires " - "conservative cross-rank aggregation or global-ready semantics." - ), + "scheduler_visibility": { + "global_ready_to_reap_s": _summary(global_ready_to_reap_s), + "reported_ready_to_reap_ms": _summary(reported_ready_to_reap_ms), + "physical_release_to_reap_s": _summary(physical_release_to_reap_s), + "invalid_reported_ready_to_reap_samples": invalid_ready_to_reap_samples, }, } @@ -667,14 +799,29 @@ def analyze_log_paths(paths: Iterable[str | Path]) -> dict[str, object]: events = read_diagnostic_events(path_list) result = analyze_events(events) if len(path_list) > 1: + events_by_source: dict[str, list[DiagnosticEvent]] = defaultdict(list) + for event in events: + if event.source is not None: + events_by_source[event.source].append(event) + + ranks = result["ranks"] + assert isinstance(ranks, dict) source_aggregates: dict[str, object] = {} for path_like in path_list: source = str(path_like) - source_result = analyze_events(event for event in events if event.source == source) + source_events = events_by_source.get(source, []) + source_rank_prefix = f"{source}::rank=" + source_rank_reports = [ + rank_report + for rank, rank_report in ranks.items() + if str(rank).startswith(source_rank_prefix) and isinstance(rank_report, dict) + ] source_aggregates[source] = { - "parsed_event_count": source_result["parsed_event_count"], - "event_counts": source_result["event_counts"], - "aggregate": source_result["aggregate"], + "parsed_event_count": len(source_events), + "event_counts": dict( + sorted(Counter(event.category for event in source_events).items()) + ), + "aggregate": _aggregate_rank_reports(source_rank_reports), } result["source_aggregates"] = source_aggregates return result @@ -697,6 +844,13 @@ def _analyze_rank( excluded_requests=unsuccessful_requests, completed_only=True, ) + global_ready = _collect_points( + events, + "python-transfer", + action="global-ready", + excluded_requests=unsuccessful_requests, + completed_only=True, + ) reaps = _collect_points( events, "reap", @@ -732,14 +886,25 @@ def _analyze_rank( "local-ready": [ ReleasePoint(point.time_s, point.request, "local-ready") for point in local_ready ], + "global-ready": [ + ReleasePoint(point.time_s, point.request, "global-ready") for point in global_ready + ], "reap": [ReleasePoint(point.time_s, point.request, "reap") for point in reaps], "receiver-slot": [ ReleasePoint(interval.end_s, interval.request, "receiver-slot") for interval in physical_service_intervals ], } + release_indexes = _build_release_indexes(decisions, admissions, submits) gaps_by_source = { - source: _match_release_gaps(points, decisions, admissions, submits, throughput) + source: _match_release_gaps( + points, + decisions, + admissions, + submits, + throughput, + indexes=release_indexes, + ) for source, points in release_points.items() } selected_source = _select_release_source(release_points) @@ -750,10 +915,12 @@ def _analyze_rank( decisions, admissions, unsuccessful_requests, + indexes=release_indexes, ) progress_samples = _linear_progress_credit(admissions, service_intervals) fixed_multiplier_samples = _fixed_multiplier_counterfactual(admissions) ready_to_reap_samples = _point_pair_gaps(local_ready, reaps) + global_ready_to_reap_samples = _point_pair_gaps(global_ready, reaps) physical_release_to_reap_samples = _point_pair_gaps( [PointEvent(interval.end_s, interval.request) for interval in physical_service_intervals], reaps, @@ -811,6 +978,13 @@ def _analyze_rank( [float(sample["gap_s"]) for sample in ready_to_reap_samples] ), "pairs": ready_to_reap_samples, + "global_ready_to_reap_samples_s": [ + float(sample["gap_s"]) for sample in global_ready_to_reap_samples + ], + "global_ready_to_reap_s": _summary( + [float(sample["gap_s"]) for sample in global_ready_to_reap_samples] + ), + "global_ready_pairs": global_ready_to_reap_samples, }, "status_poll": { "samples": status_poll_samples, @@ -931,9 +1105,9 @@ def _analyze_rank( for source, samples in gaps_by_source.items() }, "policy_note": ( - "Python local-ready is a rank-local idle-opportunity bound; reap is a " - "conservative scheduler-visible bound. An adaptive policy must aggregate " - "conservatively across ranks or use a global-ready signal." + "Python global-ready is the conservative all-rank logical release when " + "available. Local-ready is only a rank-local idle-opportunity bound; reap " + "is the later scheduler-visible bound." ), }, "fixed_multiplier_counterfactual": { @@ -1152,7 +1326,12 @@ def add( and ready_time is not None and 0.0 <= ready_time <= event.time_s ): - add("local-ready", time_s=ready_time, detail="ready_t") + ready_tag = ( + "global-ready" + if ready_time_source == "python-global-consensus" + else "local-ready" + ) + add(ready_tag, time_s=ready_time, detail="ready_t") if category == "gen-service" and action == "decode-start-proxy": add("decode-start") @@ -1214,6 +1393,8 @@ def add( add("peer-ready") if action == "local-ready": add("local-ready") + elif action == "global-ready": + add("global-ready") if action in _TERMINAL_ACTIONS: add("receiver-terminal") if action in _SUCCESS_TERMINAL_ACTIONS: @@ -1252,6 +1433,8 @@ def add( add("request-info") if action == "local-ready": add("local-ready") + elif action == "global-ready": + add("global-ready") if action in _TERMINAL_ACTIONS: add("receiver-terminal") if action in _SUCCESS_TERMINAL_ACTIONS: @@ -1305,6 +1488,8 @@ def add( "action": _normalize_diag_token(event.fields.get("action", "")), "sequence": event.fields.get("sequence"), "previous": event.fields.get("previous"), + "completion_scope": event.fields.get("completion_scope"), + "expected_participants": _as_int(event.fields.get("expected_participants")), } ) @@ -1389,6 +1574,8 @@ def add( add(event, "gen-request-info") elif action == "local-ready": add(event, "gen-local-ready") + elif action == "global-ready": + add(event, "gen-global-ready") elif action in _TERMINAL_ACTIONS: add(event, "gen-terminal") @@ -1433,6 +1620,16 @@ def add( ("ctx_timer_to_first_gate2_defer", "ctx-timer-start", "gate2-deferred"), ("ctx_timer_to_gate2_admit", "ctx-timer-start", "gate2-admitted"), ("ctx_timer_to_gen_submit", "ctx-timer-start", "gen-submit"), + ( + "ctx_timer_to_gen_global_ready", + "ctx-timer-start", + "gen-global-ready", + ), + ( + "gen_submit_to_global_ready", + "gen-submit", + "gen-global-ready", + ), ("ctx_timer_to_deadline", "ctx-timer-start", "ctx-deadline"), ("gate2_defer_to_ctx_deadline", "gate2-deferred", "ctx-deadline"), ) @@ -1502,6 +1699,16 @@ def _select_cross_host_point( key=_rank_sort_key, ) selected["wall_semantics_seen"] = sorted({str(point["wall_semantics"]) for point in tag_points}) + expected_participants = [ + int(point["expected_participants"]) + for point in tag_points + if point.get("expected_participants") is not None + ] + if expected_participants: + expected = max(expected_participants) + selected["expected_participants"] = expected + selected["observed_participants"] = len(emitter_points) + selected["participant_coverage_complete"] = len(emitter_points) >= expected return selected @@ -1948,10 +2155,15 @@ def _analyze_remaining_work_ground_truth( continue seen.add(identity) tags = timelines.get((domain, request), {}) + ready_tags = ( + ("global-ready",) + if tags.get("global-ready") + else ("local-ready", "receiver-completed") + ) ready = min( ( mark - for tag in ("local-ready", "receiver-completed") + for tag in ready_tags for mark in tags.get(tag, ()) if mark.time_s >= event.time_s ), @@ -1967,7 +2179,7 @@ def _analyze_remaining_work_ground_truth( ready_reason = _remaining_endpoint_censor_reason( request, domain, - ("local-ready", "receiver-completed"), + ready_tags, event.time_s, tags, tag_domains, @@ -2001,6 +2213,11 @@ def _analyze_remaining_work_ground_truth( if ready is not None else None ), + "ready_scope": ( + "global-consensus" + if ready_tags == ("global-ready",) + else "rank-local-or-terminal" + ), "residual_ready_s": ( ready.time_s - event.time_s if ready is not None else None ), @@ -2049,7 +2266,9 @@ def _analyze_remaining_work_ground_truth( "definition": ( "For each Gate-2 admission snapshot and active request, residual_ready_s and " "residual_reap_s use only later GEN events in the identical input source, " - "instance, and clock domain. CTX timestamps are never used." + "instance, and clock domain. Global-ready is preferred per request; older " + "logs fall back to rank-local or terminal completion. CTX timestamps are " + "never used." ), "active_decision_samples": len(samples), "active_request_ids_omitted": active_request_ids_omitted, @@ -2998,50 +3217,196 @@ def _submit_to_interval_start_gaps( return samples +def _build_release_indexes( + decisions: list[Decision], + admissions: list[Admission], + submits: list[PointEvent], +) -> _ReleaseIndexes: + """Build immutable indexes reused by release/refill analyses.""" + decision_times = tuple(decision.time_s for decision in decisions) + admission_times = tuple(admission.time_s for admission in admissions) + + admission_by_sequence: dict[str, Admission] = {} + for admission in admissions: + if admission.sequence is not None: + admission_by_sequence.setdefault(admission.sequence, admission) + + def matching_admission(decision: Decision) -> Admission | None: + if decision.sequence is not None: + match = admission_by_sequence.get(decision.sequence) + if match is not None: + return match + start = bisect_left(admission_times, decision.time_s - 1e-9) + for position in range(start, len(admissions)): + admission = admissions[position] + if admission.time_s > decision.time_s + 1e-9: + break + if math.isclose( + admission.time_s, + decision.time_s, + rel_tol=0.0, + abs_tol=1e-9, + ): + return admission + return None + + admission_by_decision = tuple(matching_admission(decision) for decision in decisions) + + signal_by_key: dict[tuple[float, int], tuple[int, float | None]] = {} + for decision in decisions: + signal_by_key.setdefault( + (decision.time_s, 1), + (decision.deferred, decision.budget_blocks), + ) + for admission in admissions: + signal_by_key.setdefault( + (admission.time_s, 0), + (admission.deferred, admission.budget_blocks), + ) + signal_keys = tuple(sorted(signal_by_key)) + signal_values = tuple(signal_by_key[key] for key in signal_keys) + + successful_decision_indices = tuple( + index for index, decision in enumerate(decisions) if decision.admitted > 0 + ) + successful_decision_times = tuple( + decisions[index].time_s for index in successful_decision_indices + ) + successful_positions_without_detail: list[int] = [] + successful_positions_by_request: dict[str, list[int]] = defaultdict(list) + for successful_position, decision_index in enumerate(successful_decision_indices): + admission = admission_by_decision[decision_index] + if admission is None: + successful_positions_without_detail.append(successful_position) + continue + for request in set(admission.admitted_requests): + successful_positions_by_request[request].append(successful_position) + + ordered_submits = tuple(submits) + submit_times = tuple(submit.time_s for submit in ordered_submits) + submit_positions_by_request: dict[str, list[int]] = defaultdict(list) + for position, submit in enumerate(ordered_submits): + submit_positions_by_request[submit.request].append(position) + + return _ReleaseIndexes( + decision_times=decision_times, + admission_times=admission_times, + admission_by_sequence=admission_by_sequence, + admission_by_decision=admission_by_decision, + signal_keys=signal_keys, + signal_values=signal_values, + successful_decision_indices=successful_decision_indices, + successful_decision_times=successful_decision_times, + successful_positions_without_detail=tuple(successful_positions_without_detail), + successful_positions_by_request={ + request: tuple(positions) + for request, positions in successful_positions_by_request.items() + }, + submits=ordered_submits, + submit_times=submit_times, + submit_positions_by_request={ + request: tuple(positions) for request, positions in submit_positions_by_request.items() + }, + submits_are_chronological=all( + previous <= current for previous, current in zip(submit_times, submit_times[1:]) + ), + ) + + +def _first_position_at_or_after( + positions: tuple[int, ...], + minimum_position: int, +) -> int | None: + offset = bisect_left(positions, minimum_position) + return positions[offset] if offset < len(positions) else None + + +def _next_successful_decision( + release_time_s: float, + backlog_requests: set[str] | None, + decisions: list[Decision], + indexes: _ReleaseIndexes, +) -> tuple[Decision | None, Admission | None]: + minimum_position = bisect_right(indexes.successful_decision_times, release_time_s) + if minimum_position >= len(indexes.successful_decision_indices): + return None, None + + if backlog_requests is None: + successful_position = minimum_position + else: + candidates: list[int] = [] + missing_detail_position = _first_position_at_or_after( + indexes.successful_positions_without_detail, + minimum_position, + ) + if missing_detail_position is not None: + candidates.append(missing_detail_position) + for request in backlog_requests: + request_position = _first_position_at_or_after( + indexes.successful_positions_by_request.get(request, ()), + minimum_position, + ) + if request_position is not None: + candidates.append(request_position) + if not candidates: + return None, None + successful_position = min(candidates) + + decision_index = indexes.successful_decision_indices[successful_position] + return decisions[decision_index], indexes.admission_by_decision[decision_index] + + def _match_release_gaps( releases: list[ReleasePoint], decisions: list[Decision], admissions: list[Admission], submits: list[PointEvent], throughput_blocks_per_s: float | None, + *, + indexes: _ReleaseIndexes | None = None, ) -> list[dict[str, object]]: + indexes = indexes or _build_release_indexes(decisions, admissions, submits) samples: list[dict[str, object]] = [] for release in sorted(releases, key=lambda point: point.time_s): - prior = _latest_backlog_signal(decisions, admissions, release.time_s) + prior = _latest_backlog_signal( + decisions, + admissions, + release.time_s, + indexes=indexes, + ) if prior is None or prior[0] <= 0: continue - next_decision = next( - (decision for decision in decisions if decision.time_s > release.time_s), - None, - ) - if next_decision is None: + next_decision_position = bisect_right(indexes.decision_times, release.time_s) + if next_decision_position >= len(decisions): continue - backlog_requests = _backlog_request_ids_at(admissions, release.time_s) + next_decision = decisions[next_decision_position] + backlog_requests = _backlog_request_ids_at( + admissions, + release.time_s, + indexes=indexes, + ) backlog_identity_unknown = not backlog_requests - successful_admission = None + successful_admission, detailed_admission = _next_successful_decision( + release.time_s, + None if backlog_identity_unknown else backlog_requests, + decisions, + indexes, + ) matched_backlog_requests: set[str] = set() - for decision in decisions: - if decision.time_s <= release.time_s or decision.admitted <= 0: - continue - if backlog_identity_unknown: - successful_admission = decision - break - detailed_admission = _matching_admission(decision, admissions) + if successful_admission is not None and not backlog_identity_unknown: if detailed_admission is None: backlog_identity_unknown = True - successful_admission = decision - break - matched = set(detailed_admission.admitted_requests).intersection(backlog_requests) - if matched: - successful_admission = decision - matched_backlog_requests = matched - break + else: + matched_backlog_requests = set(detailed_admission.admitted_requests).intersection( + backlog_requests + ) refill = ( _find_refill_submit( submits, successful_admission, admissions, matched_backlog_requests or None, + indexes=indexes, ) if successful_admission is not None else None @@ -3105,31 +3470,87 @@ def _find_refill_submit( decision: Decision, admissions: list[Admission], required_requests: set[str] | None = None, + *, + indexes: _ReleaseIndexes | None = None, ) -> PointEvent | None: - candidates = [submit for submit in submits if submit.time_s >= decision.time_s] + indexes = indexes or _build_release_indexes([decision], admissions, submits) + if not indexes.submits_are_chronological: + candidates = [submit for submit in indexes.submits if submit.time_s >= decision.time_s] + if required_requests: + return next( + (submit for submit in candidates if submit.request in required_requests), + None, + ) + admission = _matching_admission(decision, admissions, indexes=indexes) + if admission is not None and admission.admitted_requests: + admitted = set(admission.admitted_requests) + return next( + (submit for submit in candidates if submit.request in admitted), + None, + ) + return candidates[0] if candidates else None + + minimum_position = bisect_left(indexes.submit_times, decision.time_s) + if minimum_position >= len(indexes.submits): + return None + + def first_submit_for(requests: set[str]) -> PointEvent | None: + positions = [] + for request in requests: + position = _first_position_at_or_after( + indexes.submit_positions_by_request.get(request, ()), + minimum_position, + ) + if position is not None: + positions.append(position) + return indexes.submits[min(positions)] if positions else None + if required_requests: - return next( - (submit for submit in candidates if submit.request in required_requests), - None, - ) - admission = _matching_admission(decision, admissions) + return first_submit_for(required_requests) + admission = _matching_admission(decision, admissions, indexes=indexes) if admission is not None and admission.admitted_requests: - admitted = set(admission.admitted_requests) - return next((submit for submit in candidates if submit.request in admitted), None) - return candidates[0] if candidates else None + return first_submit_for(set(admission.admitted_requests)) + return indexes.submits[minimum_position] -def _backlog_request_ids_at(admissions: list[Admission], time_s: float) -> set[str]: - admission = next( - (candidate for candidate in reversed(admissions) if candidate.time_s <= time_s), - None, - ) +def _backlog_request_ids_at( + admissions: list[Admission], + time_s: float, + *, + indexes: _ReleaseIndexes | None = None, +) -> set[str]: + indexes = indexes or _build_release_indexes([], admissions, []) + position = bisect_right(indexes.admission_times, time_s) - 1 + admission = admissions[position] if position >= 0 else None if admission is None or admission.deferred <= 0 or admission.deferred_requests_omitted > 0: return set() return set(admission.deferred_requests) -def _matching_admission(decision: Decision, admissions: list[Admission]) -> Admission | None: +def _matching_admission( + decision: Decision, + admissions: list[Admission], + *, + indexes: _ReleaseIndexes | None = None, +) -> Admission | None: + if indexes is not None: + if decision.sequence is not None: + match = indexes.admission_by_sequence.get(decision.sequence) + if match is not None: + return match + start = bisect_left(indexes.admission_times, decision.time_s - 1e-9) + for position in range(start, len(admissions)): + admission = admissions[position] + if admission.time_s > decision.time_s + 1e-9: + break + if math.isclose( + admission.time_s, + decision.time_s, + rel_tol=0.0, + abs_tol=1e-9, + ): + return admission + return None if decision.sequence is not None: match = next( (admission for admission in admissions if admission.sequence == decision.sequence), @@ -3153,29 +3574,26 @@ def _matching_admission(decision: Decision, admissions: list[Admission]) -> Admi def _latest_backlog_signal( - decisions: list[Decision], admissions: list[Admission], time_s: float + decisions: list[Decision], + admissions: list[Admission], + time_s: float, + *, + indexes: _ReleaseIndexes | None = None, ) -> tuple[int, float | None] | None: - signals = [ - (decision.time_s, 1, decision.deferred, decision.budget_blocks) - for decision in decisions - if decision.time_s <= time_s - ] - signals.extend( - (admission.time_s, 0, admission.deferred, admission.budget_blocks) - for admission in admissions - if admission.time_s <= time_s - ) - if not signals: + indexes = indexes or _build_release_indexes(decisions, admissions, []) + position = bisect_right(indexes.signal_keys, (time_s, 2)) - 1 + if position < 0: return None - _, _, deferred, budget = max(signals, key=lambda signal: (signal[0], signal[1])) - return deferred, budget + return indexes.signal_values[position] def _select_release_source(release_points: dict[str, list[ReleasePoint]]) -> str | None: - # Only the C++ path has a directly observed physical release signal. - # Python local-ready and consensus reap are complementary bounds, so the - # report intentionally does not collapse them into one selected source. - return "receiver-slot" if release_points["receiver-slot"] else None + # Prefer the directly observed physical release. Global-ready is a safe + # logical fallback because it is sampled only after GEN consensus. A + # rank-local ready event is not sufficient to declare reusable capacity. + if release_points["receiver-slot"]: + return "receiver-slot" + return "global-ready" if release_points["global-ready"] else None def _slot_refill_gaps( @@ -3183,7 +3601,10 @@ def _slot_refill_gaps( decisions: list[Decision], admissions: list[Admission], excluded_requests: set[str], + *, + indexes: _ReleaseIndexes | None = None, ) -> list[float]: + indexes = indexes or _build_release_indexes(decisions, admissions, []) by_slot: dict[tuple[str, str], list[SlotInterval]] = defaultdict(list) for interval in intervals: by_slot[(interval.manager, interval.buffer)].append(interval) @@ -3193,7 +3614,12 @@ def _slot_refill_gaps( for current, following in zip(slot_intervals, slot_intervals[1:]): if current.request in excluded_requests or following.request in excluded_requests: continue - prior = _latest_backlog_signal(decisions, admissions, current.end_s) + prior = _latest_backlog_signal( + decisions, + admissions, + current.end_s, + indexes=indexes, + ) if prior is not None and prior[0] > 0 and following.start_s >= current.end_s: gaps.append(following.start_s - current.end_s) return gaps @@ -3272,39 +3698,74 @@ def _fixed_multiplier_counterfactual( def _linear_progress_credit( admissions: list[Admission], intervals: list[ServiceInterval] ) -> list[dict[str, object]]: - samples: list[dict[str, object]] = [] - for admission in admissions: - if admission.deferred <= 0: - continue - in_progress = [ - interval - for interval in intervals - if interval.blocks is not None - and interval.start_s <= admission.time_s < interval.end_s - and interval.end_s > interval.start_s - ] - if not in_progress: - continue - original_blocks = sum(interval.blocks or 0.0 for interval in in_progress) - credit = sum( - (interval.blocks or 0.0) - * (admission.time_s - interval.start_s) - / (interval.end_s - interval.start_s) - for interval in in_progress + deferred_admissions = sorted( + ( + (admission.time_s, original_position, admission) + for original_position, admission in enumerate(admissions) + if admission.deferred > 0 + ), + key=lambda item: (item[0], item[1]), + ) + service_starts = sorted( + ( + interval.start_s, + interval_index, + interval, ) + for interval_index, interval in enumerate(intervals) + if interval.blocks is not None and interval.end_s > interval.start_s + ) + + samples_by_position: dict[int, dict[str, object]] = {} + active_ends: list[tuple[float, int, float, float, float]] = [] + next_start = 0 + active_count = 0 + original_blocks = 0.0 + progress_rate = 0.0 + progress_offset = 0.0 + for decision_time_s, original_position, admission in deferred_admissions: + while next_start < len(service_starts) and service_starts[next_start][0] <= decision_time_s: + _, interval_index, interval = service_starts[next_start] + blocks = interval.blocks or 0.0 + duration_s = interval.end_s - interval.start_s + rate = blocks / duration_s + offset = rate * interval.start_s + heappush( + active_ends, + (interval.end_s, interval_index, blocks, rate, offset), + ) + active_count += 1 + original_blocks += blocks + progress_rate += rate + progress_offset += offset + next_start += 1 + + while active_ends and active_ends[0][0] <= decision_time_s: + _, _, blocks, rate, offset = heappop(active_ends) + active_count -= 1 + original_blocks -= blocks + progress_rate -= rate + progress_offset -= offset + + if active_count <= 0: + continue + + credit = decision_time_s * progress_rate - progress_offset fraction = _safe_ratio(credit, original_blocks) or 0.0 - samples.append( - { - "decision_t": admission.time_s, - "in_progress_requests": len(in_progress), - "logged_active_blocks": admission.active_blocks, - "original_in_progress_blocks": original_blocks, - "estimated_progress_credit_blocks": credit, - "estimated_remaining_blocks": original_blocks - credit, - "estimated_progress_fraction": fraction, - } - ) - return samples + samples_by_position[original_position] = { + "decision_t": admission.time_s, + "in_progress_requests": active_count, + "logged_active_blocks": admission.active_blocks, + "original_in_progress_blocks": original_blocks, + "estimated_progress_credit_blocks": credit, + "estimated_remaining_blocks": original_blocks - credit, + "estimated_progress_fraction": fraction, + } + return [ + samples_by_position[position] + for position in range(len(admissions)) + if position in samples_by_position + ] def _union_duration(intervals: list[ServiceInterval]) -> float: diff --git a/tensorrt_llm/_torch/disaggregation/diagnostics.py b/tensorrt_llm/_torch/disaggregation/diagnostics.py new file mode 100644 index 000000000000..cf78734ebd2f --- /dev/null +++ b/tensorrt_llm/_torch/disaggregation/diagnostics.py @@ -0,0 +1,33 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import os +import socket + + +def get_diagnostic_host_identity() -> tuple[str, str]: + """Return a stable node identity and the source used to derive it. + + ``HOSTNAME`` can be inherited from the container launch node in Slurm + jobs, so prefer Slurm's execution-node identity and then the kernel + hostname. + """ + slurm_node = os.getenv("SLURMD_NODENAME") + if slurm_node: + return _sanitize_diagnostic_value(slurm_node), "slurm" + + try: + hostname = socket.gethostname() + except OSError: + hostname = "" + if hostname: + return _sanitize_diagnostic_value(hostname), "socket" + + environment_hostname = os.getenv("HOSTNAME") + if environment_hostname: + return _sanitize_diagnostic_value(environment_hostname), "environment" + return "unknown", "fallback" + + +def _sanitize_diagnostic_value(value: str) -> str: + return "_".join(value.split()) diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 478d7b4963ac..0d8af6ddfa5b 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -54,6 +54,7 @@ TxSessionBase, WaitResult, ) +from tensorrt_llm._torch.disaggregation.diagnostics import get_diagnostic_host_identity from tensorrt_llm._torch.disaggregation.native.auxiliary import AuxBuffer from tensorrt_llm._torch.disaggregation.native.messenger import ZMQMessenger, decode_message from tensorrt_llm._torch.disaggregation.native.mixers.ssm.peer import MambaPolicy @@ -103,7 +104,7 @@ def _log_python_transfer_diagnostic( timestamp_s = _diagnostic_now_s() if wall_timestamp_s is None: wall_timestamp_s = time.time() - host = os.getenv("HOSTNAME", "unknown").replace(" ", "_") + host, host_source = get_diagnostic_host_identity() instance = instance.replace(" ", "_") encoded_fields = " ".join(f"{key}={value}" for key, value in fields.items()) logger.info( @@ -111,7 +112,8 @@ def _log_python_transfer_diagnostic( f"t={timestamp_s:.9f} clock=local_steady " f"wall_t={wall_timestamp_s:.9f} wall_clock=unix " f"wall_semantics={wall_semantics} runtime=Python " - f"host={host} instance={instance} rank={rank} role={role} " + f"host={host} host_source={host_source} " + f"instance={instance} rank={rank} role={role} " f"source=native action={action} " f"request={request} {encoded_fields}" ) diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index d5b47385d20e..021a3b907a02 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -22,6 +22,7 @@ WaitResult, get_unique_rid, ) +from tensorrt_llm._torch.disaggregation.diagnostics import get_diagnostic_host_identity from tensorrt_llm._torch.disaggregation.native.bounce import ( config_from_size as bounce_config_from_size, ) @@ -135,23 +136,27 @@ def _log_python_transfer_diagnostic( action: str, request: int, timestamp_s: Optional[float] = None, + wall_timestamp_s: Optional[float] = None, + wall_semantics: str = "emission", **fields: object, ) -> None: if not _is_disagg_transfer_diagnostics_enabled(): return if timestamp_s is None: timestamp_s = _diagnostic_now_s() - wall_timestamp_s = time.time() + if wall_timestamp_s is None: + wall_timestamp_s = time.time() rank = getattr(getattr(self, "_dist", None), "rank", -1) - host = os.getenv("HOSTNAME", "unknown").replace(" ", "_") + host, host_source = get_diagnostic_host_identity() instance = getattr(self, "_instance_name", "-") encoded_fields = " ".join(f"{key}={value}" for key, value in fields.items()) logger.info( "[DISAGG_DIAG][python-transfer] " f"t={timestamp_s:.9f} clock=local_steady " f"wall_t={wall_timestamp_s:.9f} wall_clock=unix " - f"wall_semantics=emission runtime=Python " - f"host={host} instance={instance} rank={rank} role={role} " + f"wall_semantics={wall_semantics} runtime=Python " + f"host={host} host_source={host_source} " + f"instance={instance} rank={rank} role={role} " f"source=transceiver action={action} request={request} " f"{encoded_fields}" ) @@ -892,6 +897,38 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): ) if _is_disagg_transfer_diagnostics_enabled(): + global_ready_time = _diagnostic_now_s() if completed else None + global_ready_wall_time = time.time() if completed else None + expected_participants = 1 + if self._gen_need_sync: + expected_participants = ( + self._mapping.pp_size + if self._mapping.enable_attention_dp + else self._mapping.world_size + ) + for rid in completed: + assert global_ready_time is not None + assert global_ready_wall_time is not None + session = self._recv_sessions[rid] + req = self._recv_reqs[rid] + req.py_kv_transfer_global_ready_time_s = global_ready_time + local_ready_time = getattr(session, "kv_ready_time_s", None) + self._log_python_transfer_diagnostic( + role="gen", + action="global-ready", + request=rid, + timestamp_s=global_ready_time, + wall_timestamp_s=global_ready_wall_time, + wall_semantics="boundary-sampled", + local_request=req.py_request_id, + completion_scope=("global-consensus" if self._gen_need_sync else "rank-local"), + expected_participants=expected_participants, + local_ready_t=( + f"{local_ready_time:.9f}" + if isinstance(local_ready_time, (int, float)) + else "-1" + ), + ) for action, rids in ( ("cancelled", cancelled), ("failed", failed), diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 1758fd9cabb5..ccc3915e9a22 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -26,6 +26,8 @@ except ImportError: from cuda import cudart +from tensorrt_llm._torch.disaggregation.diagnostics import \ + get_diagnostic_host_identity from tensorrt_llm._utils import (CUASSERT, customized_gc_thresholds, is_trace_enabled, mpi_comm, mpi_disabled, nvtx_range, set_thread_local_mpi_comm, @@ -3443,7 +3445,7 @@ def _log_disagg_transfer_diagnostic(self, category: str, **fields) -> None: if wall_timestamp is None: wall_timestamp = time.time() rank = getattr(getattr(self, "dist", None), "rank", -1) - host = os.getenv("HOSTNAME", "unknown").replace(" ", "_") + host, host_source = get_diagnostic_host_identity() transceiver = getattr(self, "kv_cache_transceiver", None) instance = getattr(transceiver, "_instance_name", "-") if not isinstance(instance, str): @@ -3453,7 +3455,7 @@ def _log_disagg_transfer_diagnostic(self, category: str, **fields) -> None: logger.info(f"[DISAGG_DIAG][{category}] t={timestamp:.9f} " f"clock=local_steady wall_t={wall_timestamp:.9f} " f"wall_clock=unix wall_semantics={wall_semantics} " - f"source=pyexecutor host={host} " + f"source=pyexecutor host={host} host_source={host_source} " f"instance={instance} rank={rank} {encoded_fields}") @staticmethod @@ -6524,13 +6526,21 @@ def _prepare_disagg_gen_transmission_complete(self, scheduled_batch): decode_start_time = get_steady_clock_now_in_seconds() arrival_time = getattr( req, "py_disagg_gen_executor_arrival_time_s", None) - ready_time = getattr(req, "py_kv_transfer_ready_time_s", + ready_time = getattr(req, + "py_kv_transfer_global_ready_time_s", None) - ready_time_source = "python-local" + ready_time_source = "python-global-consensus" ready_comparison_time = decode_start_time if (not isinstance(ready_time, (int, float)) or ready_time <= 0): ready_time = None + if ready_time is None: + ready_time = getattr(req, "py_kv_transfer_ready_time_s", + None) + ready_time_source = "python-local" + if (not isinstance(ready_time, (int, float)) + or ready_time <= 0): + ready_time = None if ready_time is None: transfer_end = getattr(req, "kv_cache_transfer_end", None) @@ -7046,10 +7056,18 @@ def _check_disagg_gen_cache_transfer_status(self, atLeastNum: int = 0): outcome = "failed" else: outcome = "cancelled" - ready_time = getattr(request, "py_kv_transfer_ready_time_s", - 0.0) - ready_time_source = "python-local" + ready_time = getattr(request, + "py_kv_transfer_global_ready_time_s", 0.0) + ready_time_source = "python-global-consensus" ready_comparison_time = reap_time + if (not isinstance(ready_time, (int, float)) + or ready_time <= 0): + ready_time = getattr(request, "py_kv_transfer_ready_time_s", + 0.0) + ready_time_source = "python-local" + if (not isinstance(ready_time, (int, float)) + or ready_time <= 0): + ready_time = 0.0 if not ready_time: transfer_end = getattr(request, "kv_cache_transfer_end", None) diff --git a/tests/unittest/_torch/executor/test_py_executor.py b/tests/unittest/_torch/executor/test_py_executor.py index 16b7984c0f8f..3c53ae75a767 100644 --- a/tests/unittest/_torch/executor/test_py_executor.py +++ b/tests/unittest/_torch/executor/test_py_executor.py @@ -1149,18 +1149,29 @@ def test_transfer_timeout_emits_nonterminal_deadline_boundary(self, monkeypatch) assert "timeout_ms=10" in message @pytest.mark.parametrize( - ("python_ready_time", "cpp_ready_time", "ready_time_source"), + ( + "global_ready_time", + "python_ready_time", + "cpp_ready_time", + "ready_time_source", + "expected_ready_time", + "expected_ready_to_reap_ms", + ), [ - (9.5, None, "python-local"), - (0.0, 9.5, "cpp-global"), + (9.7, 9.5, None, "python-global-consensus", 9.7, 400.0), + (None, 9.5, None, "python-local", 9.5, 600.0), + (None, 0.0, 9.5, "cpp-global", 9.5, 600.0), ], ) def test_transfer_status_emits_ready_to_reap_delay( self, monkeypatch, + global_ready_time, python_ready_time, cpp_ready_time, ready_time_source, + expected_ready_time, + expected_ready_to_reap_ms, ): monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") monkeypatch.setattr(py_executor_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) @@ -1185,6 +1196,7 @@ def test_transfer_status_emits_ready_to_reap_delay( executor.canceled_req_ids = [] request = _make_disagg_transfer_request(8, 32, in_progress=True) request.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + request.py_kv_transfer_global_ready_time_s = global_ready_time request.py_kv_transfer_ready_time_s = python_ready_time request.kv_cache_transfer_end = ( None @@ -1207,9 +1219,9 @@ def complete(_at_least_num): message = log_info.call_args.args[0] assert "[DISAGG_DIAG][reap]" in message assert "request=8" in message - assert "ready_t=9.500000000" in message + assert f"ready_t={expected_ready_time:.9f}" in message assert f"ready_time_source={ready_time_source}" in message - assert "ready_to_reap_ms=600.000000" in message + assert f"ready_to_reap_ms={expected_ready_to_reap_ms:.6f}" in message assert "poll_call_ms=100.000000" in message assert "outcome=completed" in message @@ -1591,8 +1603,8 @@ class StopLocalSchedule(RuntimeError): executor.scheduler = Mock() calls = [] - executor.kv_cache_manager.prepare_expect_snapshot_points.side_effect = ( - lambda requests: calls.append(("prepare", requests)) + executor.kv_cache_manager.prepare_expect_snapshot_points.side_effect = lambda requests: ( + calls.append(("prepare", requests)) ) def stop_after_schedule(requests, inflight_req_ids): @@ -1626,8 +1638,8 @@ class StopSchedule(RuntimeError): executor.scheduler = Mock() calls = [] - executor.kv_cache_manager.prepare_expect_snapshot_points.side_effect = ( - lambda requests: calls.append(("prepare", requests)) + executor.kv_cache_manager.prepare_expect_snapshot_points.side_effect = lambda requests: ( + calls.append(("prepare", requests)) ) def stop_after_schedule(requests, inflight_req_ids): diff --git a/tests/unittest/disaggregated/test_disagg_diagnostics.py b/tests/unittest/disaggregated/test_disagg_diagnostics.py new file mode 100644 index 000000000000..c115b2f9ed1b --- /dev/null +++ b/tests/unittest/disaggregated/test_disagg_diagnostics.py @@ -0,0 +1,32 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from tensorrt_llm._torch.disaggregation.diagnostics import get_diagnostic_host_identity + + +def test_diagnostic_host_prefers_slurm_execution_node(monkeypatch): + monkeypatch.setenv("SLURMD_NODENAME", "node 11") + monkeypatch.setenv("HOSTNAME", "stale-launch-node") + monkeypatch.setattr("socket.gethostname", lambda: "container-node") + + assert get_diagnostic_host_identity() == ("node_11", "slurm") + + +def test_diagnostic_host_falls_back_to_kernel_hostname(monkeypatch): + monkeypatch.delenv("SLURMD_NODENAME", raising=False) + monkeypatch.setenv("HOSTNAME", "stale-launch-node") + monkeypatch.setattr("socket.gethostname", lambda: "execution-node") + + assert get_diagnostic_host_identity() == ("execution-node", "socket") + + +def test_diagnostic_host_falls_back_to_environment(monkeypatch): + monkeypatch.delenv("SLURMD_NODENAME", raising=False) + monkeypatch.setenv("HOSTNAME", "environment node") + + def fail_to_read_hostname(): + raise OSError("hostname unavailable") + + monkeypatch.setattr("socket.gethostname", fail_to_read_hostname) + + assert get_diagnostic_host_identity() == ("environment_node", "environment") diff --git a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py index ea4195b9ec92..5ad57a59894f 100644 --- a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py +++ b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py @@ -72,6 +72,8 @@ def __init__( self.kv_transfer_start_time_s = kv_transfer_start_time_s self.request_info_sent_time_s = request_info_sent_time_s self.kv_ready_time_s = kv_ready_time_s + self.transfer_end_time = None + self.kv_cache_size_bytes = 0 self.blocking_calls: list[bool] = [] self.closed = False @@ -294,6 +296,7 @@ def test_gen_transfer_status_stamps_first_local_ready_time( ) -> None: monkeypatch.setenv("TRTLLM_DISAGG_TRANSFER_DIAGNOSTICS", "1") monkeypatch.setattr(transceiver_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + monkeypatch.setattr(transceiver_module, "_diagnostic_now_s", lambda: 12.6) log_info = Mock() monkeypatch.setattr(transceiver_module.logger, "info", log_info) session = _FakeSession( @@ -311,7 +314,8 @@ def test_gen_transfer_status_stamps_first_local_ready_time( ) transceiver = object.__new__(KvCacheTransceiverV2) transceiver._ever_had_recv_session = True - transceiver._gen_need_sync = False + transceiver._gen_need_sync = True + transceiver._mapping = Mock(enable_attention_dp=False, world_size=4) transceiver._recv_sessions = {21: session} transceiver._recv_reqs = {21: request} transceiver._diagnostic_ready_rids = set() @@ -332,6 +336,7 @@ def test_gen_transfer_status_stamps_first_local_ready_time( assert request.py_kv_transfer_service_start_time_s == 10.25 assert request.py_kv_request_info_sent_time_s == 10.25 assert request.py_kv_transfer_ready_time_s == 12.5 + assert request.py_kv_transfer_global_ready_time_s == 12.6 message = next( call.args[0] for call in log_info.call_args_list if "action=local-ready" in call.args[0] ) @@ -341,6 +346,72 @@ def test_gen_transfer_status_stamps_first_local_ready_time( assert "bytes=8192" in message assert "request_info_sent_t=10.250000000" in message assert "receive_ms=2250.000" in message + global_ready_message = next( + call.args[0] for call in log_info.call_args_list if "action=global-ready" in call.args[0] + ) + assert "t=12.600000000" in global_ready_message + assert "completion_scope=global-consensus" in global_ready_message + assert "expected_participants=4" in global_ready_message + assert "local_ready_t=12.500000000" in global_ready_message + + +def test_gen_global_ready_uses_one_consensus_boundary_for_batch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(transceiver_module, "_DISAGG_TRANSFER_DIAGNOSTICS_ENABLED", True) + diagnostic_clock = Mock(side_effect=(12.6, 12.7, 12.8)) + monkeypatch.setattr(transceiver_module, "_diagnostic_now_s", diagnostic_clock) + wall_clock = Mock(side_effect=(100.1, 100.2, 100.3, 100.4, 100.5)) + monkeypatch.setattr(transceiver_module.time, "time", wall_clock) + log_info = Mock() + monkeypatch.setattr(transceiver_module.logger, "info", log_info) + sessions = { + rid: _FakeSession( + rid=rid, + wait_result=WaitResult.COMPLETED, + is_completed=True, + request_info_sent_time_s=10.0, + kv_ready_time_s=12.5, + ) + for rid in (21, 22) + } + requests = { + rid: Mock( + py_request_id=rid, + py_kv_cache_xfer_bytes=8192, + state=LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS, + ) + for rid in sessions + } + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._ever_had_recv_session = True + transceiver._gen_need_sync = True + transceiver._mapping = Mock(enable_attention_dp=False, world_size=4) + transceiver._recv_sessions = sessions + transceiver._recv_reqs = requests + transceiver._diagnostic_ready_rids = set(sessions) + transceiver._dist = Mock(rank=2) + transceiver._collect_done = Mock(return_value=(list(sessions), [])) + transceiver._gen_consensus = Mock(return_value=list(sessions)) + transceiver._build_to_process = Mock(return_value=list(sessions)) + transceiver._gen_consensus_outcome = Mock(return_value=([], [], list(sessions))) + transceiver._need_aux_transfer = Mock(return_value=False) + transceiver._assert_disagg_history_declared = Mock() + transceiver._close_failed_sessions = Mock() + + completed, failed, cancelled = transceiver.check_gen_transfer_status(at_least_request_num=0) + + assert completed == [21, 22] + assert failed == [] + assert cancelled == [] + assert {request.py_kv_transfer_global_ready_time_s for request in requests.values()} == {12.6} + assert diagnostic_clock.call_count == 3 + global_ready_messages = [ + call.args[0] for call in log_info.call_args_list if "action=global-ready" in call.args[0] + ] + assert len(global_ready_messages) == 2 + assert all("wall_t=100.100000000" in message for message in global_ready_messages) + assert all("wall_semantics=boundary-sampled" in message for message in global_ready_messages) def test_kv_recv_task_records_native_transfer_boundaries( diff --git a/tests/unittest/tools/test_disagg_admission_telemetry.py b/tests/unittest/tools/test_disagg_admission_telemetry.py index 11d748af2d7e..09b753235ac0 100644 --- a/tests/unittest/tools/test_disagg_admission_telemetry.py +++ b/tests/unittest/tools/test_disagg_admission_telemetry.py @@ -29,6 +29,7 @@ _SPEC.loader.exec_module(_TELEMETRY) analyze_events = _TELEMETRY.analyze_events +analyze_log_paths = _TELEMETRY.analyze_log_paths main = _TELEMETRY.main parse_diagnostic_line = _TELEMETRY.parse_diagnostic_line DiagnosticEvent = _TELEMETRY.DiagnosticEvent @@ -383,6 +384,84 @@ def test_multi_log_block_lookup_is_scoped_by_source(tmp_path, capsys): assert output["ranks"][f"{gen_log}::rank=1"]["service"]["completed_blocks"] == 2 +def test_native_slot_identity_is_inferred_from_one_matching_python_emitter(): + source = "gen-worker.log" + events = [ + _parse_source_line(line, source) + for line in [ + "[DISAGG_DIAG][decision] t=0 rank=0 host=gen-node instance=GEN_0 " + "role=gen active_blocks=0 candidate_requests=42:8 admitted=1 " + "admitted_requests=42:8 deferred=0 deferred_requests=- budget=8", + "[DISAGG_DIAG][receiver-slot] t=0.1 rank=0 action=acquired " + "request=42 manager=m buffer=0", + "[DISAGG_DIAG][receiver-slot] t=0.9 rank=0 action=released " + "request=42 manager=m buffer=0", + ] + ] + + result = analyze_events(events) + + assert result["inferred_emitter_identity_events"] == 2 + assert result["rank_namespace"] == "rank" + assert result["ranks"]["0"]["receiver_slots"]["service_latency_s"]["p50"] == pytest.approx(0.8) + + +def test_multi_log_source_aggregates_reuse_top_level_rank_analysis(tmp_path, monkeypatch): + first_log = tmp_path / "first.log" + second_log = tmp_path / "second.log" + first_log.write_text( + "\n".join( + [ + "[DISAGG_DIAG][decision] t=0 rank=0 sequence=1 " + "active_blocks=0 candidate_requests=1:4 admitted=1 " + "admitted_requests=1:4 deferred=0 deferred_requests=- budget=4", + "[DISAGG_DIAG][submit] t=0.1 rank=0 request=1 blocks=4", + "[DISAGG_DIAG][python-transfer] t=1.1 rank=0 " + "action=local-ready request=1 outcome=completed", + "[DISAGG_DIAG][reap] t=1.2 rank=0 request=1 blocks=4 outcome=completed", + ] + ), + encoding="utf-8", + ) + second_log.write_text( + "\n".join( + [ + "[DISAGG_DIAG][decision] t=2 rank=0 sequence=1 " + "active_blocks=0 candidate_requests=2:8 admitted=1 " + "admitted_requests=2:8 deferred=0 deferred_requests=- budget=8", + "[DISAGG_DIAG][submit] t=2.1 rank=0 request=2 blocks=8", + "[DISAGG_DIAG][python-transfer] t=4.1 rank=0 " + "action=local-ready request=2 outcome=completed", + "[DISAGG_DIAG][reap] t=4.2 rank=0 request=2 blocks=8 outcome=completed", + ] + ), + encoding="utf-8", + ) + paths = [str(first_log), str(second_log)] + events = _TELEMETRY.read_diagnostic_events(paths) + expected_aggregates = { + path: analyze_events(event for event in events if event.source == path)["aggregate"] + for path in paths + } + + original_analyze_events = _TELEMETRY.analyze_events + analyze_calls = 0 + + def counted_analyze_events(events): + nonlocal analyze_calls + analyze_calls += 1 + return original_analyze_events(events) + + monkeypatch.setattr(_TELEMETRY, "analyze_events", counted_analyze_events) + + result = analyze_log_paths(paths) + + assert analyze_calls == 1 + assert { + path: result["source_aggregates"][path]["aggregate"] for path in paths + } == expected_aggregates + + def test_cpp_lifecycle_and_remaining_work_use_same_domain_completion(): source = "gen-worker.log" events = [ @@ -432,6 +511,63 @@ def test_cpp_lifecycle_and_remaining_work_use_same_domain_completion(): assert lifecycle["reap_to_decode_start"]["duration_s"]["p50"] == pytest.approx(0.05) +def test_global_ready_drives_logical_release_and_remaining_work(): + source = "gen-worker.log" + events = [ + _parse_source_line(line, source) + for line in [ + "[DISAGG_DIAG][submit] t=0.1 rank=0 host=gen instance=GEN_0 request=42 blocks=8", + "[DISAGG_DIAG][decision] t=0.5 rank=0 host=gen instance=GEN_0 " + "active_requests=42:8 candidate_requests=- admitted_requests=- " + "deferred_requests=- admitted=0 deferred=0 budget=8", + "[DISAGG_DIAG][python-transfer] t=1.0 rank=0 host=gen instance=GEN_0 " + "role=gen action=local-ready request=42", + "[DISAGG_DIAG][python-transfer] t=1.2 rank=0 host=gen instance=GEN_0 " + "role=gen action=global-ready request=42 completion_scope=global-consensus " + "expected_participants=2", + "[DISAGG_DIAG][reap] t=1.3 rank=0 host=gen instance=GEN_0 " + "request=42 blocks=8 outcome=completed", + "[DISAGG_DIAG][gen-service] t=1.4 rank=0 host=gen instance=GEN_0 " + "action=decode-start-proxy request=42", + ] + ] + + result = analyze_events(events) + rank = result["ranks"]["0"] + remaining = result["remaining_work_ground_truth"]["samples"][0] + lifecycle = result["lifecycle"]["interval_coverage"] + + assert rank["release_to_admission"]["selected_release_source"] == "global-ready" + assert rank["python_transfer"]["global_ready_to_reap_s"]["p50"] == pytest.approx(0.1) + assert remaining["ready_t"] == pytest.approx(1.2) + assert remaining["ready_scope"] == "global-consensus" + assert remaining["residual_ready_s"] == pytest.approx(0.7) + assert lifecycle["global_ready_to_reap"]["duration_s"]["p50"] == pytest.approx(0.1) + assert lifecycle["global_ready_to_decode_start"]["duration_s"]["p50"] == pytest.approx(0.2) + + +def test_reap_preserves_python_global_ready_scope_when_transfer_log_is_missing(): + events = _parse_lines( + [ + "[DISAGG_DIAG][decision] t=0.5 rank=0 instance=GEN_0 " + "active_requests=42:8 candidate_requests=- admitted_requests=- " + "deferred_requests=- admitted=0 deferred=0 budget=8", + "[DISAGG_DIAG][reap] t=1.3 rank=0 instance=GEN_0 request=42 " + "ready_t=1.2 ready_time_source=python-global-consensus " + "blocks=8 outcome=completed", + ] + ) + + result = analyze_events(events) + sample = result["remaining_work_ground_truth"]["samples"][0] + marks = _TELEMETRY._collect_lifecycle_marks(events) + + assert sample["ready_t"] == pytest.approx(1.2) + assert sample["ready_scope"] == "global-consensus" + assert any(mark.tag == "global-ready" for mark in marks) + assert not any(mark.tag == "local-ready" for mark in marks) + + def test_deadline_is_nonterminal_and_classified_by_sender_phase(): events = _parse_lines( [ @@ -504,6 +640,37 @@ def test_cross_host_wall_clock_joins_ctx_deadline_to_gen_deferral(): assert record["wall_intervals_s"]["gate2_defer_to_ctx_deadline"] == pytest.approx(54.0) +def test_cross_host_global_ready_uses_latest_participant_boundary(): + events = _parse_lines( + [ + "[DISAGG_DIAG][ctx-transfer] t=1 wall_t=1000 wall_clock=unix " + "wall_semantics=boundary-sampled host=ctx instance=CTX_0 rank=0 " + "action=timer-start request=42", + "[DISAGG_DIAG][submit] t=2 wall_t=1005 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen instance=GEN_0 rank=0 " + "request=42", + "[DISAGG_DIAG][python-transfer] t=3 wall_t=1010 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen instance=GEN_0 rank=0 " + "role=gen action=global-ready request=42 completion_scope=global-consensus " + "expected_participants=2", + "[DISAGG_DIAG][python-transfer] t=4 wall_t=1011 wall_clock=unix " + "wall_semantics=boundary-sampled host=gen instance=GEN_0 rank=1 " + "role=gen action=global-ready request=42 completion_scope=global-consensus " + "expected_participants=2", + ] + ) + + record = analyze_events(events)["cross_host_correlation"]["requests"][0] + global_ready = record["points"]["gen-global-ready"] + + assert global_ready["wall_t"] == pytest.approx(1011.0) + assert global_ready["observed_participants"] == 2 + assert global_ready["expected_participants"] == 2 + assert global_ready["participant_coverage_complete"] is True + assert record["wall_intervals_s"]["ctx_timer_to_gen_global_ready"] == pytest.approx(11.0) + assert record["wall_intervals_s"]["gen_submit_to_global_ready"] == pytest.approx(6.0) + + def test_cross_host_partial_gate2_admission_is_not_labeled_before_gate1(): events = _parse_lines( [ @@ -1010,3 +1177,146 @@ def test_cpp_global_ready_timestamp_is_not_subtracted_from_local_reap_clock(): assert sample["residual_ready_s"] is None assert sample["ready_censor_reason"] == "missing_ready" assert sample["residual_reap_s"] == pytest.approx(0.5) + + +def test_indexed_release_matching_avoids_repeated_admission_scans(monkeypatch): + sample_count = 2_000 + decisions = [] + admissions = [] + submits = [] + releases = [] + for index in range(sample_count): + decision_time = float(index) + admitted_request = str(index) + deferred_request = str(index + 1) + decisions.append( + _TELEMETRY.Decision( + decision_time, + str(index), + admitted=1, + deferred=1, + budget_blocks=4.0, + ) + ) + admissions.append( + _TELEMETRY.Admission( + decision_time, + str(index), + admitted=1, + deferred=1, + budget_blocks=4.0, + active_blocks=0.0, + candidate_requests=( + (admitted_request, 4.0), + (deferred_request, 4.0), + ), + admitted_requests=(admitted_request,), + deferred_requests=(deferred_request,), + candidate_requests_omitted=0, + admitted_requests_omitted=0, + deferred_requests_omitted=0, + ) + ) + submits.append( + _TELEMETRY.PointEvent( + decision_time + 0.01, + admitted_request, + ) + ) + if index + 1 < sample_count: + releases.append( + _TELEMETRY.ReleasePoint( + decision_time + 0.5, + admitted_request, + "reap", + ) + ) + + indexes = _TELEMETRY._build_release_indexes(decisions, admissions, submits) + monkeypatch.setattr( + _TELEMETRY, + "_matching_admission", + lambda *_args, **_kwargs: pytest.fail( + "indexed release matching fell back to a linear admission scan" + ), + ) + + samples = _TELEMETRY._match_release_gaps( + releases, + decisions, + admissions, + submits, + throughput_blocks_per_s=4.0, + indexes=indexes, + ) + + assert len(samples) == sample_count - 1 + assert samples[0]["matched_backlog_request_ids"] == ["1"] + assert samples[-1]["matched_backlog_request_ids"] == [str(sample_count - 1)] + + +def test_linear_progress_sweep_matches_interval_scan(): + admissions = [ + _TELEMETRY.Admission( + time_s=time_s, + sequence=str(index), + admitted=0, + deferred=1, + budget_blocks=8.0, + active_blocks=6.0, + candidate_requests=(), + admitted_requests=(), + deferred_requests=(), + candidate_requests_omitted=0, + admitted_requests_omitted=0, + deferred_requests_omitted=0, + ) + for index, time_s in enumerate((2.5, 0.5, 4.0, 1.5)) + ] + intervals = [ + _TELEMETRY.ServiceInterval("a", 0.0, 2.0, 4.0, "submit", "ready"), + _TELEMETRY.ServiceInterval("b", 1.0, 3.0, 2.0, "submit", "ready"), + _TELEMETRY.ServiceInterval("c", 2.5, 5.0, 5.0, "submit", "ready"), + ] + + expected = [] + for admission in admissions: + active = [ + interval + for interval in intervals + if interval.blocks is not None + and interval.start_s <= admission.time_s < interval.end_s + and interval.end_s > interval.start_s + ] + original_blocks = sum(interval.blocks or 0.0 for interval in active) + credit = sum( + (interval.blocks or 0.0) + * (admission.time_s - interval.start_s) + / (interval.end_s - interval.start_s) + for interval in active + ) + expected.append( + { + "decision_t": admission.time_s, + "in_progress_requests": len(active), + "original_in_progress_blocks": original_blocks, + "estimated_progress_credit_blocks": credit, + "estimated_remaining_blocks": original_blocks - credit, + "estimated_progress_fraction": credit / original_blocks, + } + ) + + actual = _TELEMETRY._linear_progress_credit(admissions, intervals) + + assert [sample["decision_t"] for sample in actual] == [ + sample["decision_t"] for sample in expected + ] + for actual_sample, expected_sample in zip(actual, expected): + assert actual_sample["in_progress_requests"] == expected_sample["in_progress_requests"] + for field in ( + "original_in_progress_blocks", + "estimated_progress_credit_blocks", + "estimated_remaining_blocks", + "estimated_progress_fraction", + ): + assert actual_sample[field] == pytest.approx(expected_sample[field])