Skip to content
Open
20 changes: 16 additions & 4 deletions tensorrt_llm/_torch/pyexecutor/py_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,8 +94,7 @@
from .kv_cache.mamba_cache_manager import (BaseMambaCacheManager,
MixedMambaHybridCacheManager)
from .kv_cache_stats import append_kv_cache_iteration_stats
from .llm_request import (ATTENTION_DP_DUMMY_REQUEST_ID,
MAX_SPEC_DECODE_POSITIONS, ExecutorRequest,
from .llm_request import (ATTENTION_DP_DUMMY_REQUEST_ID, ExecutorRequest,
LlmRequest, LlmRequestState, LlmResponse,
MultimodalEncoderRequestError, get_draft_token_length,
initialize_multimodal_encoder_request,
Expand Down Expand Up @@ -7769,9 +7768,22 @@ def _accumulate_spec_dec_stats(self, sample_state: SampleState) -> None:
if self.max_draft_len > 0 else drafted_step)
request.py_total_draft_tokens += drafted
request.py_total_accepted_draft_tokens += py_num_accepted
for pos in range(min(drafted_step, MAX_SPEC_DECODE_POSITIONS)):
# Grow rather than clamp. max_draft_len can exceed the arrays'
# initial capacity under tree drafting, and truncating here would
# silently drop every position past it -- leaving the arrays
# unable to reconcile with py_total_accepted_draft_tokens, which
# is exact. The Prometheus label-cardinality bound is enforced in
# MetricsCollector, not here, so per-request accuracy and metrics
# cardinality no longer share one constant.
if drafted_step > len(request.py_per_pos_drafted):
request.py_per_pos_drafted.extend(
[0] * (drafted_step - len(request.py_per_pos_drafted)))
if py_num_accepted > len(request.py_per_pos_accepted):
request.py_per_pos_accepted.extend(
[0] * (py_num_accepted - len(request.py_per_pos_accepted)))
for pos in range(drafted_step):
request.py_per_pos_drafted[pos] += 1
for pos in range(min(py_num_accepted, MAX_SPEC_DECODE_POSITIONS)):
for pos in range(py_num_accepted):
request.py_per_pos_accepted[pos] += 1

def _handle_errors(self,
Expand Down
9 changes: 9 additions & 0 deletions tensorrt_llm/executor/postproc_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,15 @@ class PostprocArgs:
num_prompt_tokens_offset: int = 0
tokenizer: Optional[TransformersTokenizer] = None
ctx_usage: Optional[Any] = None
# Per-request speculative-decoding acceptance stats. Set on the base so
# every endpoint's args subclass inherits one opt-in path. Mirrors the
# server's per_request_spec_decode_stats setting; there is no per-request
# opt-in, so clients need send nothing.
return_spec_decode_stats: bool = False
# Fixed per-step draft bound, or None when draft_len_schedule makes it vary
# by batch size. Sizes the emitted acceptance histogram so its length is a
# function of configuration rather than of what a request happened to hit.
spec_decode_num_spec_tokens: Optional[int] = None


@dataclass(kw_only=True)
Expand Down
23 changes: 23 additions & 0 deletions tensorrt_llm/executor/result.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,16 @@ class LogProbsResult(NamedTuple):
generation: Optional[TokenLogprobs | SimpleTokenLogprobs] = None


class _SpecDecCounters(NamedTuple):
"""Speculative-decoding counters reported by the request behind one sequence.

Field names match the executor response so a formatter can read either.
"""
per_pos_drafted: Optional[List[int]] = None
per_pos_accepted: Optional[List[int]] = None
spec_dec_totals: Optional[tuple[int, int]] = None


class ResponseWrapper:
"""1. Wrapper of runtime response with optional outputs computed post runtime.
2. A workaround to pass around RequestPerfMetrics.
Expand Down Expand Up @@ -157,6 +167,10 @@ class CompletionOutput:
repr=False)
# the result of result_handler passed to postprocess workers
_postprocess_result: Any = None
# this sequence's own spec-decode counters; see _handle_sequence
_spec_dec_counters: Optional[_SpecDecCounters] = field(default=None,
init=False,
repr=False)

def __getstate__(self) -> dict:
# _incremental_states holds a tokenizers.DecodeStream (a Rust object,
Expand Down Expand Up @@ -367,6 +381,15 @@ def _handle_sequence(self,
output.disaggregated_params = self.disaggregated_params
output._last_token_ids_len = len(output.token_ids)
output._last_logprobs_len = len(output.logprobs)
# With n > 1 each candidate is its own child request with its own
# counters, while the request-level copies on self are overwritten by
# every candidate's response. Capture this sequence's here so a choice
# never reports whichever candidate happened to respond last.
output._spec_dec_counters = _SpecDecCounters(
per_pos_drafted=getattr(response_tensors, 'per_pos_drafted', None),
per_pos_accepted=getattr(response_tensors, 'per_pos_accepted',
None),
spec_dec_totals=getattr(response_tensors, 'spec_dec_totals', None))
decoder_output_prefix = ()
if (self.sampling_params.exclude_input_from_output
or getattr(self, "_streaming", False)):
Expand Down
11 changes: 11 additions & 0 deletions tensorrt_llm/llmapi/llm_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -5148,6 +5148,17 @@ class BaseLlmArgs(StrictBaseModel):
"the request sets X-TRTLLM-return-metrics: 1.",
status="prototype")

per_request_spec_decode_stats: bool = Field(
default=False,
description=
"Include per-request speculative-decoding acceptance statistics on each "
"response choice. Server-side opt-in only: unlike return_perf_metrics "
"this needs no per-request header, so benchmarking clients that "
"discover the payload by shape do not have to know they are talking to "
"TensorRT-LLM. Deliberately independent of return_perf_metrics, which "
"also mounts the Prometheus endpoint. PyTorch backend only.",
status="prototype")

perf_metrics_output_dir: Optional[str] = Field(
default=None,
description="Directory for per-process performance metrics JSONL "
Expand Down
11 changes: 10 additions & 1 deletion tensorrt_llm/metrics/collector.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,14 @@

from .enums import MetricNames

# Upper bound on distinct ``token_position`` label values for the per-position
# speculative-decoding counters. Each position is one Prometheus time series per
# model, so this caps label cardinality no matter how deep a request drafts.
# This bound used to be implicit in the per-request arrays' fixed length; those
# now grow with max_draft_len so per-request accuracy is not truncated, which
# makes the metrics-side limit its own explicit concern.
MAX_SPEC_DECODE_POSITION_LABELS = 16


# Adapted from https://github.com/vllm-project/vllm/blob/v0.10.0rc1/vllm/engine/metrics.py#L30
class MetricsCollector:
Expand Down Expand Up @@ -696,7 +704,8 @@ def log_request_metrics_dict(self, metrics_dict: dict) -> None:
if per_pos_drafted[i] > 0:
last_nonzero = i
break
for pos in range(last_nonzero + 1):
for pos in range(
min(last_nonzero + 1, MAX_SPEC_DECODE_POSITION_LABELS)):
labels_with_pos = {
**self.labels, self.labelname_token_pos: pos
}
Expand Down
85 changes: 80 additions & 5 deletions tensorrt_llm/serve/openai_protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@
from openai_harmony import ReasoningEffort
from pydantic import (AliasChoices, BaseModel, ConfigDict, Field,
NonNegativeInt, PositiveInt, field_validator,
model_validator)
model_serializer, model_validator)
from typing_extensions import Annotated, Required, TypeAlias, TypedDict

from tensorrt_llm.executor.request import LoRARequest
Expand Down Expand Up @@ -151,6 +151,73 @@ class OpenAIBaseModel(BaseModel):
model_config = ConfigDict(extra="forbid", populate_by_name=True)


class SpeculativeDecodingStats(OpenAIBaseModel):
"""Per-request speculative-decoding acceptance for one generated sequence.

Opt-in: emitted only when the server sets per_request_spec_decode_stats.
Absent entirely when the request drafted nothing. PyTorch backend only.

``mean acceptance length`` is deliberately not a field here: it is already
reported per choice as ``avg_decoded_tokens_per_iter``, and duplicating it
would let the two drift. Consumers derive it as
``1 + total_accepted_draft_tokens / num_spec_steps``.

Three identities hold on every emitted record, and a consumer may rely on
them:

* ``sum(acceptance_histogram) == num_spec_steps``
* ``sum(j * acceptance_histogram[j]) == total_accepted_draft_tokens``
* ``total_accepted_draft_tokens <= total_draft_tokens``
"""

acceptance_rate: float = Field(
description="Accepted draft tokens divided by proposed draft tokens.")
total_accepted_draft_tokens: int = Field(
description="Draft tokens accepted across the request, excluding the "
"always-accepted bonus token.")
total_draft_tokens: int = Field(
description="Draft tokens proposed across the request. For tree "
"drafting this counts paths, not tree nodes, mirroring the "
"getMaxDraftPathLen clamp in updateNumTokensPerIteration.")
num_spec_steps: int = Field(
description="Verify steps performed for the request. Equals the sum of "
"acceptance_histogram.")
acceptance_histogram: List[int] = Field(
description="Dense histogram indexed by accepted-draft count: entry j "
"is the number of verify steps that accepted exactly j draft tokens. "
"Tree-agnostic -- it records output lengths per step and encodes no "
"parent/child structure.")
num_spec_tokens: Optional[int] = Field(
default=None,
description="Maximum draft length per step, when the run has a fixed "
"bound. None under draft_len_schedule, where the bound varies by batch "
"size.")


class _OmitsAbsentSpecDecodeStats(OpenAIBaseModel):
"""Drops ``speculative_decoding`` from serialized output when it is absent.

Per-request spec-decode stats are off by default, so without this every
response on the paths that serialize with a plain ``model_dump()`` -- the
non-streaming chat and completions responses, and the completions stream
(``exclude_unset=False``) -- would gain ``"speculative_decoding": null`` for
every user, whether or not they enabled ``per_request_spec_decode_stats``.

Scoped to this one field on purpose. Blanket ``exclude_none`` would also
strip unrelated optional fields that clients may rely on being present, and
the dump calls are spread across the serving layer rather than funnelled
through one place where an ``exclude=`` argument could be applied.
"""

@model_serializer(mode="wrap")
def _omit_absent_spec_decode_stats(self, handler: Any) -> Any:
data = handler(self)
if isinstance(data, dict) and data.get("speculative_decoding",
...) is None:
data.pop("speculative_decoding", None)
return data


class StreamOptions(OpenAIBaseModel):
include_usage: Optional[bool] = True
continuous_usage_stats: Optional[bool] = False
Expand Down Expand Up @@ -295,7 +362,7 @@ class CompletionLogProbs(OpenAIBaseModel):
top_logprobs: List[Optional[Dict[str, float]]] = Field(default_factory=list)


class CompletionResponseChoice(OpenAIBaseModel):
class CompletionResponseChoice(_OmitsAbsentSpecDecodeStats):
index: int
text: str
token_ids: Optional[List[int]] = None
Expand All @@ -312,6 +379,8 @@ class CompletionResponseChoice(OpenAIBaseModel):
)
disaggregated_params: Optional[DisaggregatedParams] = Field(default=None)
avg_decoded_tokens_per_iter: Optional[float] = Field(default=None)
speculative_decoding: Optional[SpeculativeDecodingStats] = Field(
default=None)


class CompletionResponse(OpenAIBaseModel):
Expand All @@ -326,7 +395,7 @@ class CompletionResponse(OpenAIBaseModel):
prompt_token_ids: Optional[Union[List[List[int]], List[int]]] = None


class CompletionResponseStreamChoice(OpenAIBaseModel):
class CompletionResponseStreamChoice(_OmitsAbsentSpecDecodeStats):
index: int
text: str
token_ids: Optional[List[int]] = None
Expand All @@ -340,6 +409,8 @@ class CompletionResponseStreamChoice(OpenAIBaseModel):
"including encountering the EOS token"),
)
avg_decoded_tokens_per_iter: Optional[float] = Field(default=None)
speculative_decoding: Optional[SpeculativeDecodingStats] = Field(
default=None)


class CompletionStreamResponse(OpenAIBaseModel):
Expand Down Expand Up @@ -856,7 +927,7 @@ class ChatCompletionLogProbs(OpenAIBaseModel):
content: Optional[List[ChatCompletionLogProbsContent]] = None


class ChatCompletionResponseChoice(OpenAIBaseModel):
class ChatCompletionResponseChoice(_OmitsAbsentSpecDecodeStats):
index: int
message: ChatMessage
logprobs: Optional[ChatCompletionLogProbs] = None
Expand All @@ -868,6 +939,8 @@ class ChatCompletionResponseChoice(OpenAIBaseModel):

disaggregated_params: Optional[DisaggregatedParams] = Field(default=None)
avg_decoded_tokens_per_iter: Optional[float] = Field(default=None)
speculative_decoding: Optional[SpeculativeDecodingStats] = Field(
default=None)


class ChatCompletionResponse(OpenAIBaseModel):
Expand All @@ -894,13 +967,15 @@ class DeltaMessage(OpenAIBaseModel):
tool_calls: Optional[List[DeltaToolCall]] = None


class ChatCompletionResponseStreamChoice(OpenAIBaseModel):
class ChatCompletionResponseStreamChoice(_OmitsAbsentSpecDecodeStats):
index: int
delta: DeltaMessage
logprobs: Optional[ChatCompletionLogProbs] = None
finish_reason: Optional[str] = None
stop_reason: Optional[Union[int, str]] = None
avg_decoded_tokens_per_iter: Optional[float] = Field(default=None)
speculative_decoding: Optional[SpeculativeDecodingStats] = Field(
default=None)


class ChatCompletionStreamResponse(OpenAIBaseModel):
Expand Down
56 changes: 55 additions & 1 deletion tensorrt_llm/serve/openai_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,7 @@
from tensorrt_llm._utils import EnergyMonitor
# yapf: disable
from tensorrt_llm.executor import CppExecutorError
from tensorrt_llm.executor.postproc_worker import PostprocParams
from tensorrt_llm.executor.postproc_worker import PostprocArgs, PostprocParams
from tensorrt_llm.executor.request import DEFAULT_REQUEST_PRIORITY
from tensorrt_llm.inputs import prompt_inputs
from tensorrt_llm.inputs.data import TokensPrompt
Expand Down Expand Up @@ -691,6 +691,21 @@ def _image_output_size(image) -> Optional[str]:
return f"{width}x{height}"


def resolve_spec_decode_num_spec_tokens(args: Any) -> Optional[int]:
"""Fixed per-step draft bound to report, or None when there is not one.

Emitted as ``num_spec_tokens`` and used to size the acceptance histogram, so
a wrong answer here makes every histogram the wrong width. Returns None both
when speculative decoding is off and when ``draft_len_schedule`` is set --
the bound genuinely varies by batch size there, and None is the honest
answer rather than reporting whichever value happened to be configured.
"""
spec_config = getattr(args, "speculative_config", None) if args else None
if spec_config is None or getattr(spec_config, "draft_len_schedule", None):
return None
return getattr(spec_config, "max_draft_len", None)


class OpenAIServer(_VideoRoutesMixin):

@staticmethod
Expand Down Expand Up @@ -786,6 +801,13 @@ def __init__(
None) if args else None)
self._collect_perf_metrics = (self._expose_perf_metrics
or perf_metrics_output_dir is not None)
# Per-request spec-decode acceptance stats. Deliberately independent of
# return_perf_metrics: coupling them would mean asking for acceptance
# numbers silently mounts the Prometheus endpoint too.
self._per_request_spec_decode_stats = bool(
args and getattr(args, "per_request_spec_decode_stats", False))
self._spec_decode_num_spec_tokens = resolve_spec_decode_num_spec_tokens(
args)
# AsyncLLM uses this flag to request engine-level snapshots. Preserve the
# original value separately because only it controls public headers.
if self._collect_perf_metrics and args is not None:
Expand Down Expand Up @@ -1978,6 +2000,28 @@ async def _iteration_stats_collector_loop(self):
logger.info("Iteration stats collector loop cancelled")
raise

def _apply_spec_decode_stats_opt_in(self,
postproc_args: PostprocArgs) -> None:
"""Enable per-request spec-decode stats when the server opted in.

Server-side only, deliberately: per_request_spec_decode_stats in the
YAML config is the entire opt-in, and a client sends nothing extra.
This differs from return_perf_metrics, which additionally requires a
per-request X-TRTLLM-return-metrics header.

The reason is that benchmarking clients discover this payload by shape
rather than being told which engine they are talking to -- requiring a
vendor-specific request header would mean the client has to know it is
talking to TensorRT-LLM before it can find out, which it does not. The
cost stays opt-in because an operator who does not set the YAML field
pays nothing, and one who does has asked for exactly this.
"""
if not self._per_request_spec_decode_stats:
return
postproc_args.return_spec_decode_stats = True
postproc_args.spec_decode_num_spec_tokens = (
self._spec_decode_num_spec_tokens)
Comment thread
debermudez marked this conversation as resolved.

async def openai_chat(self, request: ChatCompletionRequest,
raw_request: Request) -> Response:

Expand Down Expand Up @@ -2212,6 +2256,7 @@ async def chat_stream_generator(
err_type="BadRequestError",
status_code=HTTPStatus.BAD_REQUEST)
postproc_args = ChatPostprocArgs.from_request(request)
self._apply_spec_decode_stats_opt_in(postproc_args)
if (is_kimi_k3 and request.add_generation_prompt
and request.prompt_token_ids is None
and request.prompt_token_ids_b64 is None
Expand Down Expand Up @@ -2993,6 +3038,7 @@ async def generator_wrapper(generator: AsyncIterator[Any]):
request.conversation_params)
for idx, prompt in enumerate(prompts):
postproc_args = CompletionPostprocArgs.from_request(request)
self._apply_spec_decode_stats_opt_in(postproc_args)
postproc_args.prompt_idx = idx
postproc_args.stream_response_id = stream_response_id
postproc_args.stream_created = stream_created
Expand Down Expand Up @@ -3165,6 +3211,14 @@ async def create_streaming_generator(promise: RequestOutput,
tracing.extract_trace_headers(raw_request.headers))

postproc_args = ChatCompletionPostprocArgs.from_request(request)
# No spec-decode opt-in here on purpose. The Harmony handlers build
# their choices in harmony_adapter, which carries no per-request
# spec-decode data at all -- avg_decoded_tokens_per_iter is absent
# from that path too -- so setting the flag would configure
# something nothing reads. Extending Harmony should cover both
# fields together; handle_non_streaming_response would need the
# GenerationResult threaded through, as it currently receives only
# the outputs.
postproc_params = PostprocParams(
post_processor=chat_harmony_streaming_post_processor
if request.stream else chat_harmony_post_processor,
Expand Down
Loading