From 42369e8c401a8719ac1b17219de3bf4a2d799191 Mon Sep 17 00:00:00 2001 From: Elias Bermudez Date: Thu, 3 Sep 2026 19:47:09 -0700 Subject: [PATCH 1/6] None: fix per-position spec-decode arrays truncating at 16 positions py_per_pos_drafted / py_per_pos_accepted were clamped to MAX_SPEC_DECODE_POSITIONS (16) in PyExecutor._accumulate_spec_dec_stats, silently dropping every position beyond it. max_draft_len can exceed 16 under tree drafting (EAGLE3 dynamic-tree, Medusa), so in those configurations the arrays stopped reconciling with py_total_accepted_draft_tokens -- which is exact -- and the trtllm_spec_decode_{drafted,accepted}_tokens_total per-position counters under-reported with no indication anything was lost. The constant was bounding two unrelated things: per-request buffer size, an accuracy concern, and Prometheus token_position label cardinality, an operational one. The second was emergent rather than stated -- MetricsCollector iterates range(len(per_pos_drafted)) with no clamp of its own, so the cardinality bound held only because the arrays happened to be length 16. Separate the two: grow the arrays on demand in the accumulator, and add an explicit MAX_SPEC_DECODE_POSITION_LABELS clamp in MetricsCollector. MAX_SPEC_DECODE_POSITIONS is retained as the initial allocation size, so the common case (max_draft_len <= 16) never reallocates, and the two limits can now move independently. Signed-off-by: Elias Bermudez --- tensorrt_llm/_torch/pyexecutor/py_executor.py | 20 ++++++++++++--- tensorrt_llm/metrics/collector.py | 11 +++++++- .../executor/test_spec_dec_stats_pairing.py | 23 +++++++++++++++++ tests/unittest/metrics/test_collector.py | 25 ++++++++++++++++++- 4 files changed, 73 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 62ccc9abeb35..d0b0c8ad85ef 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -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, @@ -7775,9 +7774,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, diff --git a/tensorrt_llm/metrics/collector.py b/tensorrt_llm/metrics/collector.py index d6f30806498f..6f3726b54c1b 100644 --- a/tensorrt_llm/metrics/collector.py +++ b/tensorrt_llm/metrics/collector.py @@ -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: @@ -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 } diff --git a/tests/unittest/executor/test_spec_dec_stats_pairing.py b/tests/unittest/executor/test_spec_dec_stats_pairing.py index 10c09afd5317..8d7997c4a323 100644 --- a/tests/unittest/executor/test_spec_dec_stats_pairing.py +++ b/tests/unittest/executor/test_spec_dec_stats_pairing.py @@ -114,6 +114,29 @@ def test_tree_drafting_clamps_totals_to_max_path_len(self): assert request.py_total_accepted_draft_tokens == 3 assert request.py_per_pos_drafted[:13] == [1] * 12 + [0] + def test_positions_beyond_initial_capacity_are_not_truncated(self): + # max_draft_len can exceed MAX_SPEC_DECODE_POSITIONS with tree drafting + # (EAGLE3 dynamic-tree, Medusa). The per-pos arrays start at that size + # but must grow rather than clamp, otherwise every position past the + # initial capacity is silently dropped and the arrays stop reconciling + # with py_total_accepted_draft_tokens -- which is exact. + deep = MAX_SPEC_DECODE_POSITIONS + 4 + request = _fake_request(verified=deep, accepted=deep - 2, draft_buffer_len=deep) + _accumulate([request], max_draft_len=deep) + assert request.py_per_pos_drafted[:deep] == [1] * deep + assert request.py_per_pos_accepted[:deep - 2] == [1] * (deep - 2) + # The survival array must sum to the exact accepted total: that identity + # is what a consumer derives the acceptance histogram from. + assert sum(request.py_per_pos_accepted) == request.py_total_accepted_draft_tokens + + def test_capacity_not_grown_when_within_initial_size(self): + # The common case (max_draft_len <= MAX_SPEC_DECODE_POSITIONS) must not + # reallocate: growth is a tail path, not per-step overhead. + request = _fake_request(verified=4, accepted=3, draft_buffer_len=4) + _accumulate([request], max_draft_len=4) + assert len(request.py_per_pos_drafted) == MAX_SPEC_DECODE_POSITIONS + assert len(request.py_per_pos_accepted) == MAX_SPEC_DECODE_POSITIONS + def _make_llm_request(request_id, seq_slot): return LlmRequest( diff --git a/tests/unittest/metrics/test_collector.py b/tests/unittest/metrics/test_collector.py index 7b1719abaef3..163b4bf0bd89 100644 --- a/tests/unittest/metrics/test_collector.py +++ b/tests/unittest/metrics/test_collector.py @@ -19,7 +19,7 @@ import pytest from prometheus_client import REGISTRY -from tensorrt_llm.metrics.collector import MetricsCollector +from tensorrt_llm.metrics.collector import MAX_SPEC_DECODE_POSITION_LABELS, MetricsCollector from tensorrt_llm.metrics.enums import MetricNames, RequestEventTiming from tensorrt_llm.metrics.perf_utils import process_req_perf_metrics @@ -1171,6 +1171,29 @@ def test_absent_arrays_no_error(self, collector): metrics = {MetricsCollector.labelname_finish_reason: "end_id"} collector.log_request_metrics_dict(metrics) # must not raise + def test_label_cardinality_is_capped(self, collector): + """Deep drafting must not create unbounded token_position series. + + The per-request arrays grow with max_draft_len so per-request acceptance + is not truncated, but each position is one Prometheus time series per + model, so the metrics side caps them independently. + """ + deep = MAX_SPEC_DECODE_POSITION_LABELS + 5 + metrics = { + MetricsCollector.labelname_finish_reason: "end_id", + MetricNames.SPEC_DEC_DRAFTED_PER_POS: [1] * deep, + MetricNames.SPEC_DEC_ACCEPTED_PER_POS: [1] * deep, + } + collector.log_request_metrics_dict(metrics) + existing_pos = { + sample.labels.get("token_position") + for metric in REGISTRY.collect() + if metric.name == collector.counter_tokens_drafted_per_position._name + for sample in metric.samples + if sample.name.endswith("_total") + } + assert existing_pos == {str(p) for p in range(MAX_SPEC_DECODE_POSITION_LABELS)} + def test_no_observation_without_finish_reason(self, collector): """No per-position counter updates when finish_reason is missing.""" metrics = { From 59de05beb55adf0516828e2d4541de987eab2e18 Mon Sep 17 00:00:00 2001 From: Elias Bermudez Date: Thu, 3 Sep 2026 19:56:28 -0700 Subject: [PATCH 2/6] None: feat: per-request speculative-decoding acceptance stats on the response TRT-LLM already computes per-request speculative-decoding acceptance on the PyTorch backend, but discards the attribution at the HTTP boundary: the exact (accepted, drafted) totals reach the serve layer and are written only to the server-side perf-metrics JSONL, and the per-position vectors are consumed into Prometheus counters, which sum across every request. A benchmarking client therefore cannot obtain per-request acceptance at all. Emit it per choice as `speculative_decoding`, alongside the existing `avg_decoded_tokens_per_iter`, which already establishes a per-request spec-decode field in the response body. No new instrumentation: every value is already an attribute on GenerationResultBase, which is what the postprocess handlers already receive. `per_pos_accepted` is prefix-cumulative and therefore a survival function, so the acceptance histogram is its negative first difference. Totals come from `spec_dec_totals`, which the executor accumulates exactly, rather than from summing the vectors. Mean acceptance length is deliberately not a field: it would duplicate `avg_decoded_tokens_per_iter` on the same choice and the two could drift; consumers derive it as 1 + accepted/steps. Enablement is server-side only, via the `per_request_spec_decode_stats` TorchLlmArgs field. This deliberately differs from `return_perf_metrics`, which additionally requires a per-request `X-TRTLLM-return-metrics` header: benchmarking clients discover this payload by shape rather than being told which engine they are talking to, so requiring a vendor-specific request header would mean the client must already know it is talking to TensorRT-LLM in order to find out. The cost stays opt-in because an operator who does not set the field pays nothing. Kept independent of `return_perf_metrics`, which also mounts the Prometheus endpoint -- coupling them would mean asking for acceptance numbers silently starts a metrics server. Off by default. Absent for requests that never drafted, on non-terminal stream chunks, and on the C++/TRT backend, which has no per-position vectors. Not wired for /v1/responses, whose schema has no choices array to attach to. Signed-off-by: Elias Bermudez --- tensorrt_llm/executor/postproc_worker.py | 9 ++ tensorrt_llm/llmapi/llm_args.py | 11 ++ tensorrt_llm/serve/openai_protocol.py | 51 ++++++ tensorrt_llm/serve/openai_server.py | 42 ++++- tensorrt_llm/serve/postprocess_handlers.py | 74 ++++++++- .../usage/llm_args_golden_manifest.json | 5 + .../integration/test_lists/test-db/l0_a10.yml | 1 + .../api_stability/references/llm.yaml | 4 + .../test_spec_decode_stats_payload.py | 152 ++++++++++++++++++ 9 files changed, 346 insertions(+), 3 deletions(-) create mode 100644 tests/unittest/executor/test_spec_decode_stats_payload.py diff --git a/tensorrt_llm/executor/postproc_worker.py b/tensorrt_llm/executor/postproc_worker.py index 184b9f924c8e..2f038d698f67 100644 --- a/tensorrt_llm/executor/postproc_worker.py +++ b/tensorrt_llm/executor/postproc_worker.py @@ -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) diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 40412b882a1c..9623cda104b4 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -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 " diff --git a/tensorrt_llm/serve/openai_protocol.py b/tensorrt_llm/serve/openai_protocol.py index 0d85fd554229..d03870ff47e4 100644 --- a/tensorrt_llm/serve/openai_protocol.py +++ b/tensorrt_llm/serve/openai_protocol.py @@ -151,6 +151,49 @@ 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 StreamOptions(OpenAIBaseModel): include_usage: Optional[bool] = True continuous_usage_stats: Optional[bool] = False @@ -312,6 +355,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): @@ -340,6 +385,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): @@ -868,6 +915,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): @@ -901,6 +950,8 @@ class ChatCompletionResponseStreamChoice(OpenAIBaseModel): 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): diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index 6efa29a49279..b5dbbf86f769 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -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 @@ -786,6 +786,21 @@ 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)) + spec_config = getattr(args, "speculative_config", + None) if args else None + if spec_config is None or getattr(spec_config, "draft_len_schedule", + None): + # No spec decode, or draft_len_schedule makes the per-step bound + # vary by batch size; None is the honest answer for both. + self._spec_decode_num_spec_tokens = None + else: + self._spec_decode_num_spec_tokens = getattr(spec_config, + "max_draft_len", None) # 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: @@ -1979,6 +1994,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) + async def openai_chat(self, request: ChatCompletionRequest, raw_request: Request) -> Response: @@ -2213,6 +2250,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 @@ -2994,6 +3032,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 @@ -3166,6 +3205,7 @@ async def create_streaming_generator(promise: RequestOutput, tracing.extract_trace_headers(raw_request.headers)) postproc_args = ChatCompletionPostprocArgs.from_request(request) + self._apply_spec_decode_stats_opt_in(postproc_args) postproc_params = PostprocParams( post_processor=chat_harmony_streaming_post_processor if request.stream else chat_harmony_post_processor, diff --git a/tensorrt_llm/serve/postprocess_handlers.py b/tensorrt_llm/serve/postprocess_handlers.py index 3e440543e67d..39eb163d5779 100644 --- a/tensorrt_llm/serve/postprocess_handlers.py +++ b/tensorrt_llm/serve/postprocess_handlers.py @@ -53,8 +53,9 @@ CompletionStreamResponse, DeltaFunctionCall, DeltaMessage, DeltaToolCall, FunctionCall, PromptTokensDetails, ResponsesRequest, - ResponsesResponse, StreamOptions, ToolCall, - UsageInfo, to_disaggregated_params) + ResponsesResponse, SpeculativeDecodingStats, + StreamOptions, ToolCall, UsageInfo, + to_disaggregated_params) from .tool_parser.base_tool_parser import (BaseToolParser, warn_if_tool_call_unparsed) from .tool_parser.core_types import StreamingParseResult, ToolCallItem @@ -146,6 +147,67 @@ def from_request(cls, request: ChatCompletionRequest): ) +def _build_spec_decode_stats( + rsp: GenerationResultBase, args: PostprocArgs, + finish_reason: Optional[str]) -> Optional[SpeculativeDecodingStats]: + """Derive per-request speculative-decoding acceptance for one sequence. + + Returns None -- meaning the field is omitted entirely -- when the caller has + not opted in, when the sequence has not finished (streaming carries this + only on the terminal chunk), when the request never drafted, or when the + executor attached no per-position vectors, as on the non-PyTorch backend. + + per_pos_accepted is prefix-cumulative and therefore a survival function: + entry k counts the steps that accepted *at least* k+1 draft tokens. The + acceptance histogram is its negative first difference. Totals come from + spec_dec_totals, which the executor accumulates exactly, rather than from + summing the vectors. + """ + if not args.return_spec_decode_stats or finish_reason is None: + return None + totals = getattr(rsp, 'spec_dec_totals', None) + per_pos_accepted = getattr(rsp, 'per_pos_accepted', None) + per_pos_drafted = getattr(rsp, 'per_pos_drafted', None) + if not totals or not per_pos_drafted or not per_pos_accepted: + return None + accepted, drafted = totals + # Position 0 is incremented once for every step that drafted at all, so it + # is the verify-step count. + num_spec_steps = per_pos_drafted[0] + if drafted <= 0 or num_spec_steps <= 0: + return None + + # Survival is non-increasing, so its positive entries form a prefix whose + # length is the deepest acceptance any single step reached. + deepest = 0 + for count in per_pos_accepted: + if count <= 0: + break + deepest += 1 + + # Size to the configured draft budget when there is one, so the histogram's + # length describes the configuration rather than what this request happened + # to reach. Under draft_len_schedule there is no fixed bound, so it sizes to + # the observed depth instead. + num_spec_tokens = args.spec_decode_num_spec_tokens + width = deepest if num_spec_tokens is None else max(deepest, + num_spec_tokens) + histogram = [0] * (width + 1) + histogram[0] = num_spec_steps - per_pos_accepted[0] + for j in range(1, deepest + 1): + following = per_pos_accepted[j] if j < len(per_pos_accepted) else 0 + histogram[j] = per_pos_accepted[j - 1] - following + + return SpeculativeDecodingStats( + acceptance_rate=accepted / drafted, + total_accepted_draft_tokens=accepted, + total_draft_tokens=drafted, + num_spec_steps=num_spec_steps, + acceptance_histogram=histogram, + num_spec_tokens=num_spec_tokens, + ) + + def _ensure_stream_metadata(args: Any, rsp: GenerationResultBase, prefix: str) -> Tuple[str, int]: if args.stream_response_id is None: @@ -521,6 +583,8 @@ def yield_first_chat(num_tokens: int, avg_decoded_tokens_per_iter=getattr(rsp, 'avg_decoded_tokens_per_iter', None), + speculative_decoding=_build_spec_decode_stats( + rsp, args, output.finish_reason), stop_reason=output.stop_reason, ) if args.return_logprobs: @@ -686,6 +750,8 @@ def chat_response_post_processor( avg_decoded_tokens_per_iter=getattr(rsp, 'avg_decoded_tokens_per_iter', None), + speculative_decoding=_build_spec_decode_stats( + rsp, args, output.finish_reason), ) if output.finish_reason == "stop" and args.has_tool_call.get( output.index, False): @@ -814,6 +880,8 @@ def completion_stream_post_processor(rsp: DetokenizedGenerationResultBase, avg_decoded_tokens_per_iter=getattr(rsp, 'avg_decoded_tokens_per_iter', None), + speculative_decoding=_build_spec_decode_stats( + rsp, args, output.finish_reason), ) if args.return_logprobs: logprobs = output.logprobs_diff @@ -882,6 +950,8 @@ def completion_response_post_processor( avg_decoded_tokens_per_iter=getattr(rsp, 'avg_decoded_tokens_per_iter', None), + speculative_decoding=_build_spec_decode_stats( + rsp, args, output.finish_reason), ) if args.return_logprobs: logprobs = output.logprobs diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json index 257d1ba70c11..7bb5d55949c0 100644 --- a/tensorrt_llm/usage/llm_args_golden_manifest.json +++ b/tensorrt_llm/usage/llm_args_golden_manifest.json @@ -1069,6 +1069,11 @@ "kind": "value", "path": "peft_cache_config.optimal_adapter_size" }, + { + "capture_policy": "bool", + "kind": "value", + "path": "per_request_spec_decode_stats" + }, { "capture_policy": "int", "kind": "value", diff --git a/tests/integration/test_lists/test-db/l0_a10.yml b/tests/integration/test_lists/test-db/l0_a10.yml index 8415c70c1dab..b8f2cef13d9b 100644 --- a/tests/integration/test_lists/test-db/l0_a10.yml +++ b/tests/integration/test_lists/test-db/l0_a10.yml @@ -145,6 +145,7 @@ l0_a10: - unittest/llmapi/test_additional_model_outputs.py -m "gpu1" # executor - unittest/executor/test_spec_dec_stats_pairing.py + - unittest/executor/test_spec_decode_stats_payload.py - unittest/executor/test_postprocessor_hook.py - unittest/executor/test_proxy_postproc_terminate.py - unittest/executor/test_proxy_fast_death.py diff --git a/tests/unittest/api_stability/references/llm.yaml b/tests/unittest/api_stability/references/llm.yaml index 8728888a2868..447c0a1e329f 100644 --- a/tests/unittest/api_stability/references/llm.yaml +++ b/tests/unittest/api_stability/references/llm.yaml @@ -35,6 +35,10 @@ methods: annotation: bool default: False status: prototype + per_request_spec_decode_stats: + annotation: bool + default: False + status: prototype perf_metrics_output_dir: annotation: Optional[str] default: null diff --git a/tests/unittest/executor/test_spec_decode_stats_payload.py b/tests/unittest/executor/test_spec_decode_stats_payload.py new file mode 100644 index 000000000000..824b805486a7 --- /dev/null +++ b/tests/unittest/executor/test_spec_decode_stats_payload.py @@ -0,0 +1,152 @@ +# 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. +"""Derivation tests for the per-request speculative_decoding response field. + +py_per_pos_accepted is prefix-cumulative -- entry k counts steps that accepted +at least k+1 drafts -- so it is a survival function, and the acceptance +histogram is its negative first difference. These tests pin that derivation and +the three identities a consumer relies on, since a payload that violates any of +them is indistinguishable from corruption on the client side: + +* sum(acceptance_histogram) == num_spec_steps +* sum(j * acceptance_histogram[j]) == total_accepted_draft_tokens +* total_accepted_draft_tokens <= total_draft_tokens + +Totals come from spec_dec_totals (exact) while the histogram comes from the +per-position vectors, so the second identity is the one that would break first +if the two ever stopped agreeing. + +GPU-free logic, but importing postprocess_handlers is unproven on the CPU-only +CI stage, so this file is wired into a GPU list (l0_a10.yml) like its sibling +test_spec_dec_stats_pairing.py. +""" + +from types import SimpleNamespace + +import pytest +from pytest import param + +from tensorrt_llm._torch.pyexecutor.llm_request import MAX_SPEC_DECODE_POSITIONS +from tensorrt_llm.serve.postprocess_handlers import _build_spec_decode_stats + + +def _rsp(survival, num_spec_steps, totals): + """Result stub. `survival` is the leading non-zero part of per_pos_accepted.""" + accepted = list(survival) + [0] * (MAX_SPEC_DECODE_POSITIONS - len(survival)) + drafted = [num_spec_steps] + [0] * (MAX_SPEC_DECODE_POSITIONS - 1) + return SimpleNamespace(per_pos_accepted=accepted, + per_pos_drafted=drafted, + spec_dec_totals=totals) + + +def _args(*, enabled=True, num_spec_tokens=None): + return SimpleNamespace(return_spec_decode_stats=enabled, + spec_decode_num_spec_tokens=num_spec_tokens) + + +def _assert_identities(stats): + histogram = stats.acceptance_histogram + assert sum(histogram) == stats.num_spec_steps + assert sum(j * c for j, c in enumerate( + histogram)) == stats.total_accepted_draft_tokens + assert stats.total_accepted_draft_tokens <= stats.total_draft_tokens + + +class TestHistogramDerivation: + + @pytest.mark.parametrize( + "survival, steps, totals, num_spec_tokens, expected", + [ + # 20 steps: 8 accepted none, 6 accepted 2, 6 accepted all 3. + param([12, 12, 6], 20, (30, 60), 3, [8, 0, 6, 6], id="mixed"), + param([], 10, (0, 30), 3, [10, 0, 0, 0], id="all_rejected"), + param([5, 5, 5], 5, (15, 15), 3, [0, 0, 0, 5], id="all_accepted"), + ], + ) # fmt: skip + def test_histogram(self, survival, steps, totals, num_spec_tokens, + expected): + stats = _build_spec_decode_stats(_rsp(survival, steps, totals), + _args(num_spec_tokens=num_spec_tokens), + "stop") + assert stats.acceptance_histogram == expected + _assert_identities(stats) + + def test_mean_acceptance_length_is_derivable(self): + # Not a field: consumers derive it, and it must match the counts. The + # field would otherwise duplicate avg_decoded_tokens_per_iter on the + # same choice and the two could drift. + stats = _build_spec_decode_stats(_rsp([12, 12, 6], 20, (30, 60)), + _args(num_spec_tokens=3), "stop") + assert 1 + (stats.total_accepted_draft_tokens / + stats.num_spec_steps) == 2.5 + + def test_histogram_padded_to_configured_budget(self): + # Length must describe the draft budget, not the depth this particular + # request happened to reach, so it is stable across requests. + stats = _build_spec_decode_stats(_rsp([3], 5, (3, 25)), + _args(num_spec_tokens=5), "stop") + assert len(stats.acceptance_histogram) == 6 + _assert_identities(stats) + + def test_adaptive_drafting_reports_no_fixed_bound(self): + # Under draft_len_schedule there is no fixed per-step bound, so + # num_spec_tokens is None and the histogram sizes to observed depth. + stats = _build_spec_decode_stats(_rsp([12, 12, 6], 20, (30, 60)), + _args(num_spec_tokens=None), "stop") + assert stats.num_spec_tokens is None + assert stats.acceptance_histogram == [8, 0, 6, 6] + _assert_identities(stats) + + def test_deep_drafting_beyond_initial_capacity(self): + # Tree drafting can exceed MAX_SPEC_DECODE_POSITIONS; the executor grows + # the vectors rather than truncating, so the identities must still hold. + depth = MAX_SPEC_DECODE_POSITIONS + 4 + survival = [1] * depth + stats = _build_spec_decode_stats( + SimpleNamespace(per_pos_accepted=survival, + per_pos_drafted=[1] + [0] * (depth - 1), + spec_dec_totals=(depth, depth)), + _args(num_spec_tokens=depth), "stop") + _assert_identities(stats) + + +class TestOmission: + """The field is absent, not null-filled, whenever it cannot be trusted.""" + + def test_absent_when_not_opted_in(self): + assert _build_spec_decode_stats(_rsp([12], 20, (30, 60)), + _args(enabled=False, + num_spec_tokens=3), + "stop") is None + + def test_absent_on_non_terminal_stream_chunk(self): + # Streaming carries this only on the chunk bearing finish_reason; + # intermediate chunks would report a partial request as if complete. + assert _build_spec_decode_stats(_rsp([12], 20, (30, 60)), + _args(num_spec_tokens=3), None) is None + + def test_absent_when_nothing_drafted(self): + assert _build_spec_decode_stats(_rsp([], 0, (0, 0)), + _args(num_spec_tokens=3), + "stop") is None + + def test_absent_without_per_position_vectors(self): + # The C++/TRT backend populates spec metrics via + # updateNumTokensPerIteration and has no per-position vectors. + assert _build_spec_decode_stats( + SimpleNamespace(per_pos_accepted=None, + per_pos_drafted=None, + spec_dec_totals=None), _args(num_spec_tokens=3), + "stop") is None From 0caf26375bad36ab14b5938b1b22a81da0a29568 Mon Sep 17 00:00:00 2001 From: Elias Bermudez Date: Thu, 17 Sep 2026 09:46:23 -0700 Subject: [PATCH 3/6] None: fix: omit absent spec-decode stats instead of serializing null Per-request spec-decode stats are off by default, so nearly every response carries none. The field must then be absent from the wire rather than present as null: the serving layer dumps responses several ways -- plain model_dump() for non-streaming chat and completions, exclude_unset=False for the completions stream, exclude_none=True for the chat stream -- and only the last would have dropped the null on its own. Every other user would have gained a "speculative_decoding": null key whether or not they enabled per_request_spec_decode_stats. Add a field-scoped serializer on the four choice models that drops the key when it holds nothing. Deliberately not blanket exclude_none, which would also strip unrelated optional fields clients may rely on being present -- avg_decoded_tokens_per_iter sits on these same models and must keep serializing as null. Regression tests cover every dump style the serving layer uses, in both directions, plus that scoping. Also drop the spec-decode opt-in from the Harmony path, where it was dead configuration. 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 the flag configured something nothing reads. Extending Harmony should cover both fields together and needs handle_non_streaming_response to receive the GenerationResult, which today it does not. Signed-off-by: Elias Bermudez --- tensorrt_llm/serve/openai_protocol.py | 34 ++++- tensorrt_llm/serve/openai_server.py | 9 +- .../integration/test_lists/test-db/l0_a10.yml | 1 + .../test_spec_decode_serialization.py | 129 ++++++++++++++++++ 4 files changed, 167 insertions(+), 6 deletions(-) create mode 100644 tests/unittest/executor/test_spec_decode_serialization.py diff --git a/tensorrt_llm/serve/openai_protocol.py b/tensorrt_llm/serve/openai_protocol.py index d03870ff47e4..230a1ab85d7d 100644 --- a/tensorrt_llm/serve/openai_protocol.py +++ b/tensorrt_llm/serve/openai_protocol.py @@ -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 @@ -194,6 +194,30 @@ class SpeculativeDecodingStats(OpenAIBaseModel): "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 @@ -338,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 @@ -371,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 @@ -903,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 @@ -943,7 +967,7 @@ class DeltaMessage(OpenAIBaseModel): tool_calls: Optional[List[DeltaToolCall]] = None -class ChatCompletionResponseStreamChoice(OpenAIBaseModel): +class ChatCompletionResponseStreamChoice(_OmitsAbsentSpecDecodeStats): index: int delta: DeltaMessage logprobs: Optional[ChatCompletionLogProbs] = None diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index b5dbbf86f769..6be9e1809cd0 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -3205,7 +3205,14 @@ async def create_streaming_generator(promise: RequestOutput, tracing.extract_trace_headers(raw_request.headers)) postproc_args = ChatCompletionPostprocArgs.from_request(request) - self._apply_spec_decode_stats_opt_in(postproc_args) + # 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, diff --git a/tests/integration/test_lists/test-db/l0_a10.yml b/tests/integration/test_lists/test-db/l0_a10.yml index b8f2cef13d9b..3d26cd723b35 100644 --- a/tests/integration/test_lists/test-db/l0_a10.yml +++ b/tests/integration/test_lists/test-db/l0_a10.yml @@ -146,6 +146,7 @@ l0_a10: # executor - unittest/executor/test_spec_dec_stats_pairing.py - unittest/executor/test_spec_decode_stats_payload.py + - unittest/executor/test_spec_decode_serialization.py - unittest/executor/test_postprocessor_hook.py - unittest/executor/test_proxy_postproc_terminate.py - unittest/executor/test_proxy_fast_death.py diff --git a/tests/unittest/executor/test_spec_decode_serialization.py b/tests/unittest/executor/test_spec_decode_serialization.py new file mode 100644 index 000000000000..e884df8330b4 --- /dev/null +++ b/tests/unittest/executor/test_spec_decode_serialization.py @@ -0,0 +1,129 @@ +# 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. +"""Serialization tests for the per-request speculative_decoding field. + +Per-request spec-decode stats are off by default, so the overwhelmingly common +case is a response that carries none. The field must then be *absent* from the +wire, not present as null: the serving layer dumps responses several different +ways -- plain ``model_dump()`` for non-streaming chat and completions, +``model_dump_json(exclude_unset=False)`` for the completions stream, and +``model_dump_json(exclude_none=True)`` for the chat stream -- and only the last +would have dropped a null on its own. Without a field-scoped serializer every +user would gain a ``"speculative_decoding": null`` key on responses whether or +not they enabled the feature. + +The omission is deliberately scoped to this one field: unrelated optional +fields that clients may rely on being present must still serialize as null. + +GPU-free logic, but importing the serving protocol is unproven on the CPU-only +CI stage, so this file is wired into a GPU list (l0_a10.yml) like its siblings. +""" + +import pytest +from pytest import param + +from tensorrt_llm.serve.openai_protocol import ( + ChatCompletionResponseChoice, ChatCompletionResponseStreamChoice, + ChatMessage, CompletionResponseChoice, CompletionResponseStreamChoice, + DeltaMessage, SpeculativeDecodingStats) + +STATS = SpeculativeDecodingStats( + acceptance_rate=0.5, + total_accepted_draft_tokens=30, + total_draft_tokens=60, + num_spec_steps=20, + acceptance_histogram=[8, 0, 6, 6], + num_spec_tokens=3, +) + + +def _choices(stats): + """One instance of each choice model that can carry the field.""" + return { + "completion": + CompletionResponseChoice(index=0, text="hi", speculative_decoding=stats), + "completion_stream": + CompletionResponseStreamChoice(index=0, + text="hi", + speculative_decoding=stats), + "chat": + ChatCompletionResponseChoice(index=0, + message=ChatMessage(role="assistant", + content="hi"), + speculative_decoding=stats), + "chat_stream": + ChatCompletionResponseStreamChoice(index=0, + delta=DeltaMessage(content="hi"), + speculative_decoding=stats), + } + + +ABSENT = _choices(None) +PRESENT = _choices(STATS) + + +@pytest.mark.parametrize("name", sorted(ABSENT)) # fmt: skip +class TestAbsentStatsAreOmitted: + """Every dump style the serving layer uses must omit the key.""" + + def test_model_dump(self, name): + assert "speculative_decoding" not in ABSENT[name].model_dump() + + def test_model_dump_json(self, name): + assert "speculative_decoding" not in ABSENT[name].model_dump_json() + + def test_model_dump_json_exclude_unset_false(self, name): + # The completions stream serializes this way. + assert "speculative_decoding" not in ABSENT[name].model_dump_json( + exclude_unset=False) + + def test_model_dump_json_exclude_none(self, name): + # The chat stream serializes this way; it would have dropped the null + # anyway, but the field-scoped serializer must not conflict with it. + assert "speculative_decoding" not in ABSENT[name].model_dump_json( + exclude_none=True) + + +@pytest.mark.parametrize("name", sorted(PRESENT)) # fmt: skip +class TestPresentStatsSurvive: + """Omission must not swallow real statistics.""" + + @pytest.mark.parametrize( + "kwargs", + [ + param({}, id="default"), + param({"exclude_unset": False}, id="exclude_unset_false"), + ], + ) # fmt: skip + def test_round_trips(self, name, kwargs): + dumped = PRESENT[name].model_dump_json(**kwargs) + assert '"speculative_decoding"' in dumped + assert '"num_spec_steps":20' in dumped + + def test_values_intact(self, name): + data = PRESENT[name].model_dump()["speculative_decoding"] + assert data["acceptance_histogram"] == [8, 0, 6, 6] + assert data["total_accepted_draft_tokens"] == 30 + + +@pytest.mark.parametrize("name", sorted(ABSENT)) # fmt: skip +def test_unrelated_optional_fields_still_serialize_as_null(name): + """The exclusion is scoped to one field, not blanket exclude_none. + + ``avg_decoded_tokens_per_iter`` is the neighbouring optional field on these + same models; clients parsing it must keep seeing it. + """ + assert ABSENT[name].model_dump()["avg_decoded_tokens_per_iter"] is None + assert '"avg_decoded_tokens_per_iter":null' in ABSENT[name].model_dump_json() From 997c71d5b233567c678cfe4938203944bb9ba2e1 Mon Sep 17 00:00:00 2001 From: Elias Bermudez Date: Thu, 17 Sep 2026 10:19:14 -0700 Subject: [PATCH 4/6] None: test: cover the spec-decode opt-in wiring and pin no-reallocation Two review findings, both about tests that could pass while the behaviour they name was broken. The no-reallocation test asserted only array length, which still holds if the accumulator rebuilds a same-sized list on every step -- exactly the per-step cost the test exists to rule out. Assert object identity instead, which is what pins in-place growth. Nothing covered the server opt-in wiring: the payload tests drove _build_spec_decode_stats with synthetic arguments, so a server that resolved its configuration correctly but never applied it would emit nothing and no test would notice. Covering that needed a refactor first. The bound resolution was inline in OpenAIServer.__init__, reachable only by constructing a server, so extract it as resolve_spec_decode_num_spec_tokens. That value is emitted as num_spec_tokens and sizes the acceptance histogram, so a wrong answer makes every histogram the wrong width -- worth testing directly rather than through a server fixture. Tests cover both halves: resolution (no speculative config, fixed max_draft_len, draft_len_schedule reporting no bound) and application (disabled leaves the args untouched, enabled propagates either bound). These reach the two helpers rather than the endpoints, so they pin the opt-in logic but not the presence of the call sites in openai_chat and openai_completion. Driving a request end to end needs tokenizer, dataset and postproc-worker plumbing that is a much larger harness than this change warrants. Signed-off-by: Elias Bermudez --- tensorrt_llm/serve/openai_server.py | 27 +++++--- .../executor/test_spec_dec_stats_pairing.py | 8 ++- .../test_spec_decode_stats_payload.py | 69 +++++++++++++++++++ 3 files changed, 93 insertions(+), 11 deletions(-) diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index 6be9e1809cd0..74ebabfa24a6 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -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 @@ -791,16 +806,8 @@ def __init__( # numbers silently mounts the Prometheus endpoint too. self._per_request_spec_decode_stats = bool( args and getattr(args, "per_request_spec_decode_stats", False)) - spec_config = getattr(args, "speculative_config", - None) if args else None - if spec_config is None or getattr(spec_config, "draft_len_schedule", - None): - # No spec decode, or draft_len_schedule makes the per-step bound - # vary by batch size; None is the honest answer for both. - self._spec_decode_num_spec_tokens = None - else: - self._spec_decode_num_spec_tokens = getattr(spec_config, - "max_draft_len", None) + 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: diff --git a/tests/unittest/executor/test_spec_dec_stats_pairing.py b/tests/unittest/executor/test_spec_dec_stats_pairing.py index 8d7997c4a323..0dbefd425647 100644 --- a/tests/unittest/executor/test_spec_dec_stats_pairing.py +++ b/tests/unittest/executor/test_spec_dec_stats_pairing.py @@ -131,9 +131,15 @@ def test_positions_beyond_initial_capacity_are_not_truncated(self): def test_capacity_not_grown_when_within_initial_size(self): # The common case (max_draft_len <= MAX_SPEC_DECODE_POSITIONS) must not - # reallocate: growth is a tail path, not per-step overhead. + # reallocate: growth is a tail path, not per-step overhead. Identity is + # what pins that -- a length check alone would still pass against an + # implementation that rebuilt a same-sized list on every step, which is + # exactly the per-step cost this is meant to rule out. request = _fake_request(verified=4, accepted=3, draft_buffer_len=4) + drafted, accepted = request.py_per_pos_drafted, request.py_per_pos_accepted _accumulate([request], max_draft_len=4) + assert request.py_per_pos_drafted is drafted + assert request.py_per_pos_accepted is accepted assert len(request.py_per_pos_drafted) == MAX_SPEC_DECODE_POSITIONS assert len(request.py_per_pos_accepted) == MAX_SPEC_DECODE_POSITIONS diff --git a/tests/unittest/executor/test_spec_decode_stats_payload.py b/tests/unittest/executor/test_spec_decode_stats_payload.py index 824b805486a7..69fba31cf115 100644 --- a/tests/unittest/executor/test_spec_decode_stats_payload.py +++ b/tests/unittest/executor/test_spec_decode_stats_payload.py @@ -39,6 +39,8 @@ from pytest import param from tensorrt_llm._torch.pyexecutor.llm_request import MAX_SPEC_DECODE_POSITIONS +from tensorrt_llm.serve.openai_server import ( + OpenAIServer, resolve_spec_decode_num_spec_tokens) from tensorrt_llm.serve.postprocess_handlers import _build_spec_decode_stats @@ -150,3 +152,70 @@ def test_absent_without_per_position_vectors(self): per_pos_drafted=None, spec_dec_totals=None), _args(num_spec_tokens=3), "stop") is None + + +class TestNumSpecTokensResolution: + """What the server reports as the fixed per-step draft bound. + + This value is emitted as ``num_spec_tokens`` and sizes the acceptance + histogram, so getting it wrong makes every histogram the wrong width. + """ + + def test_no_speculative_config_is_none(self): + assert resolve_spec_decode_num_spec_tokens(SimpleNamespace()) is None + + def test_no_args_is_none(self): + assert resolve_spec_decode_num_spec_tokens(None) is None + + def test_fixed_bound_is_reported(self): + args = SimpleNamespace(speculative_config=SimpleNamespace( + max_draft_len=4, draft_len_schedule=None)) + assert resolve_spec_decode_num_spec_tokens(args) == 4 + + def test_draft_len_schedule_reports_no_bound(self): + # The bound varies by batch size, so None is the honest answer rather + # than whichever max_draft_len happens to be configured alongside it. + args = SimpleNamespace(speculative_config=SimpleNamespace( + max_draft_len=4, draft_len_schedule={1: 4, 8: 2})) + assert resolve_spec_decode_num_spec_tokens(args) is None + + +class TestServerOptIn: + """``_apply_spec_decode_stats_opt_in`` is what reaches the handlers. + + The handlers read ``return_spec_decode_stats`` and + ``spec_decode_num_spec_tokens`` off the postproc args; this is the only + place they are set, so a server that resolved its config correctly but + failed to apply it would emit nothing. + """ + + @staticmethod + def _server(enabled, num_spec_tokens=4): + server = object.__new__(OpenAIServer) + server._per_request_spec_decode_stats = enabled + server._spec_decode_num_spec_tokens = num_spec_tokens + return server + + @staticmethod + def _args(): + return SimpleNamespace(return_spec_decode_stats=False, + spec_decode_num_spec_tokens=None) + + def test_disabled_server_leaves_args_untouched(self): + args = self._args() + OpenAIServer._apply_spec_decode_stats_opt_in(self._server(False), args) + assert args.return_spec_decode_stats is False + assert args.spec_decode_num_spec_tokens is None + + def test_enabled_server_propagates_fixed_bound(self): + args = self._args() + OpenAIServer._apply_spec_decode_stats_opt_in(self._server(True), args) + assert args.return_spec_decode_stats is True + assert args.spec_decode_num_spec_tokens == 4 + + def test_enabled_server_propagates_adaptive_bound(self): + args = self._args() + OpenAIServer._apply_spec_decode_stats_opt_in( + self._server(True, num_spec_tokens=None), args) + assert args.return_spec_decode_stats is True + assert args.spec_decode_num_spec_tokens is None From cf87aca5ba29212ffc0f66766a7123b37005ba45 Mon Sep 17 00:00:00 2001 From: Elias Bermudez Date: Wed, 23 Sep 2026 14:29:38 -0700 Subject: [PATCH 5/6] [None][fix] Attribute spec-decode stats to each n > 1 candidate With n > 1 every candidate runs as its own child request with its own spec-decode counters, but GenerationResultBase kept a single request-level copy that each candidate's response overwrote. The formatters built every choice's speculative_decoding from that copy, so all choices reported whichever candidate responded last. _handle_sequence now captures each sequence's counters on its own CompletionOutput (private, init=False, so the LLM API surface is unchanged), and all four formatters read them from there. A new test drives two candidates with different counters through a real GenerationResultBase and each formatter; it fails on the old code with the last candidate's stats on every choice. Also applies the ruff/ruff-format fixes the pre-commit CI check requested on the new spec-decode test files. Co-Authored-By: Claude Opus 5.5 Signed-off-by: Elias Bermudez --- tensorrt_llm/executor/result.py | 23 ++ tensorrt_llm/serve/postprocess_handlers.py | 24 +- .../executor/test_spec_dec_stats_pairing.py | 2 +- .../test_spec_decode_serialization.py | 42 ++- .../test_spec_decode_stats_payload.py | 272 ++++++++++++++---- 5 files changed, 281 insertions(+), 82 deletions(-) diff --git a/tensorrt_llm/executor/result.py b/tensorrt_llm/executor/result.py index 9ccc7808a99e..2a63665d881e 100644 --- a/tensorrt_llm/executor/result.py +++ b/tensorrt_llm/executor/result.py @@ -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. @@ -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, @@ -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)): diff --git a/tensorrt_llm/serve/postprocess_handlers.py b/tensorrt_llm/serve/postprocess_handlers.py index 39eb163d5779..4f1d05f2bc0b 100644 --- a/tensorrt_llm/serve/postprocess_handlers.py +++ b/tensorrt_llm/serve/postprocess_handlers.py @@ -148,10 +148,15 @@ def from_request(cls, request: ChatCompletionRequest): def _build_spec_decode_stats( - rsp: GenerationResultBase, args: PostprocArgs, + counters: Any, args: PostprocArgs, finish_reason: Optional[str]) -> Optional[SpeculativeDecodingStats]: """Derive per-request speculative-decoding acceptance for one sequence. + ``counters`` is the sequence's own ``CompletionOutput._spec_dec_counters``, + never the request-level copies on the GenerationResult: with n > 1 every + candidate reports its own counters and those copies hold whichever + candidate responded last. + Returns None -- meaning the field is omitted entirely -- when the caller has not opted in, when the sequence has not finished (streaming carries this only on the terminal chunk), when the request never drafted, or when the @@ -163,11 +168,12 @@ def _build_spec_decode_stats( spec_dec_totals, which the executor accumulates exactly, rather than from summing the vectors. """ - if not args.return_spec_decode_stats or finish_reason is None: + if (not args.return_spec_decode_stats or finish_reason is None + or counters is None): return None - totals = getattr(rsp, 'spec_dec_totals', None) - per_pos_accepted = getattr(rsp, 'per_pos_accepted', None) - per_pos_drafted = getattr(rsp, 'per_pos_drafted', None) + totals = counters.spec_dec_totals + per_pos_accepted = counters.per_pos_accepted + per_pos_drafted = counters.per_pos_drafted if not totals or not per_pos_drafted or not per_pos_accepted: return None accepted, drafted = totals @@ -584,7 +590,7 @@ def yield_first_chat(num_tokens: int, 'avg_decoded_tokens_per_iter', None), speculative_decoding=_build_spec_decode_stats( - rsp, args, output.finish_reason), + output._spec_dec_counters, args, output.finish_reason), stop_reason=output.stop_reason, ) if args.return_logprobs: @@ -751,7 +757,7 @@ def chat_response_post_processor( 'avg_decoded_tokens_per_iter', None), speculative_decoding=_build_spec_decode_stats( - rsp, args, output.finish_reason), + output._spec_dec_counters, args, output.finish_reason), ) if output.finish_reason == "stop" and args.has_tool_call.get( output.index, False): @@ -881,7 +887,7 @@ def completion_stream_post_processor(rsp: DetokenizedGenerationResultBase, 'avg_decoded_tokens_per_iter', None), speculative_decoding=_build_spec_decode_stats( - rsp, args, output.finish_reason), + output._spec_dec_counters, args, output.finish_reason), ) if args.return_logprobs: logprobs = output.logprobs_diff @@ -951,7 +957,7 @@ def completion_response_post_processor( 'avg_decoded_tokens_per_iter', None), speculative_decoding=_build_spec_decode_stats( - rsp, args, output.finish_reason), + output._spec_dec_counters, args, output.finish_reason), ) if args.return_logprobs: logprobs = output.logprobs diff --git a/tests/unittest/executor/test_spec_dec_stats_pairing.py b/tests/unittest/executor/test_spec_dec_stats_pairing.py index 0dbefd425647..aaeef373fff8 100644 --- a/tests/unittest/executor/test_spec_dec_stats_pairing.py +++ b/tests/unittest/executor/test_spec_dec_stats_pairing.py @@ -124,7 +124,7 @@ def test_positions_beyond_initial_capacity_are_not_truncated(self): request = _fake_request(verified=deep, accepted=deep - 2, draft_buffer_len=deep) _accumulate([request], max_draft_len=deep) assert request.py_per_pos_drafted[:deep] == [1] * deep - assert request.py_per_pos_accepted[:deep - 2] == [1] * (deep - 2) + assert request.py_per_pos_accepted[: deep - 2] == [1] * (deep - 2) # The survival array must sum to the exact accepted total: that identity # is what a consumer derives the acceptance histogram from. assert sum(request.py_per_pos_accepted) == request.py_total_accepted_draft_tokens diff --git a/tests/unittest/executor/test_spec_decode_serialization.py b/tests/unittest/executor/test_spec_decode_serialization.py index e884df8330b4..f513b0a1ab48 100644 --- a/tests/unittest/executor/test_spec_decode_serialization.py +++ b/tests/unittest/executor/test_spec_decode_serialization.py @@ -35,9 +35,14 @@ from pytest import param from tensorrt_llm.serve.openai_protocol import ( - ChatCompletionResponseChoice, ChatCompletionResponseStreamChoice, - ChatMessage, CompletionResponseChoice, CompletionResponseStreamChoice, - DeltaMessage, SpeculativeDecodingStats) + ChatCompletionResponseChoice, + ChatCompletionResponseStreamChoice, + ChatMessage, + CompletionResponseChoice, + CompletionResponseStreamChoice, + DeltaMessage, + SpeculativeDecodingStats, +) STATS = SpeculativeDecodingStats( acceptance_rate=0.5, @@ -52,21 +57,16 @@ def _choices(stats): """One instance of each choice model that can carry the field.""" return { - "completion": - CompletionResponseChoice(index=0, text="hi", speculative_decoding=stats), - "completion_stream": - CompletionResponseStreamChoice(index=0, - text="hi", - speculative_decoding=stats), - "chat": - ChatCompletionResponseChoice(index=0, - message=ChatMessage(role="assistant", - content="hi"), - speculative_decoding=stats), - "chat_stream": - ChatCompletionResponseStreamChoice(index=0, - delta=DeltaMessage(content="hi"), - speculative_decoding=stats), + "completion": CompletionResponseChoice(index=0, text="hi", speculative_decoding=stats), + "completion_stream": CompletionResponseStreamChoice( + index=0, text="hi", speculative_decoding=stats + ), + "chat": ChatCompletionResponseChoice( + index=0, message=ChatMessage(role="assistant", content="hi"), speculative_decoding=stats + ), + "chat_stream": ChatCompletionResponseStreamChoice( + index=0, delta=DeltaMessage(content="hi"), speculative_decoding=stats + ), } @@ -86,14 +86,12 @@ def test_model_dump_json(self, name): def test_model_dump_json_exclude_unset_false(self, name): # The completions stream serializes this way. - assert "speculative_decoding" not in ABSENT[name].model_dump_json( - exclude_unset=False) + assert "speculative_decoding" not in ABSENT[name].model_dump_json(exclude_unset=False) def test_model_dump_json_exclude_none(self, name): # The chat stream serializes this way; it would have dropped the null # anyway, but the field-scoped serializer must not conflict with it. - assert "speculative_decoding" not in ABSENT[name].model_dump_json( - exclude_none=True) + assert "speculative_decoding" not in ABSENT[name].model_dump_json(exclude_none=True) @pytest.mark.parametrize("name", sorted(PRESENT)) # fmt: skip diff --git a/tests/unittest/executor/test_spec_decode_stats_payload.py b/tests/unittest/executor/test_spec_decode_stats_payload.py index 69fba31cf115..04433fdce238 100644 --- a/tests/unittest/executor/test_spec_decode_stats_payload.py +++ b/tests/unittest/executor/test_spec_decode_stats_payload.py @@ -33,41 +33,51 @@ test_spec_dec_stats_pairing.py. """ +import json from types import SimpleNamespace import pytest from pytest import param +from tensorrt_llm import SamplingParams from tensorrt_llm._torch.pyexecutor.llm_request import MAX_SPEC_DECODE_POSITIONS -from tensorrt_llm.serve.openai_server import ( - OpenAIServer, resolve_spec_decode_num_spec_tokens) -from tensorrt_llm.serve.postprocess_handlers import _build_spec_decode_stats +from tensorrt_llm.bindings import executor as tllm +from tensorrt_llm.executor.result import GenerationResultBase +from tensorrt_llm.serve.openai_server import OpenAIServer, resolve_spec_decode_num_spec_tokens +from tensorrt_llm.serve.postprocess_handlers import ( + ChatPostprocArgs, + CompletionPostprocArgs, + _build_spec_decode_stats, + chat_response_post_processor, + chat_stream_post_processor, + completion_response_post_processor, + completion_stream_post_processor, +) def _rsp(survival, num_spec_steps, totals): """Result stub. `survival` is the leading non-zero part of per_pos_accepted.""" accepted = list(survival) + [0] * (MAX_SPEC_DECODE_POSITIONS - len(survival)) drafted = [num_spec_steps] + [0] * (MAX_SPEC_DECODE_POSITIONS - 1) - return SimpleNamespace(per_pos_accepted=accepted, - per_pos_drafted=drafted, - spec_dec_totals=totals) + return SimpleNamespace( + per_pos_accepted=accepted, per_pos_drafted=drafted, spec_dec_totals=totals + ) def _args(*, enabled=True, num_spec_tokens=None): - return SimpleNamespace(return_spec_decode_stats=enabled, - spec_decode_num_spec_tokens=num_spec_tokens) + return SimpleNamespace( + return_spec_decode_stats=enabled, spec_decode_num_spec_tokens=num_spec_tokens + ) def _assert_identities(stats): histogram = stats.acceptance_histogram assert sum(histogram) == stats.num_spec_steps - assert sum(j * c for j, c in enumerate( - histogram)) == stats.total_accepted_draft_tokens + assert sum(j * c for j, c in enumerate(histogram)) == stats.total_accepted_draft_tokens assert stats.total_accepted_draft_tokens <= stats.total_draft_tokens class TestHistogramDerivation: - @pytest.mark.parametrize( "survival, steps, totals, num_spec_tokens, expected", [ @@ -77,11 +87,10 @@ class TestHistogramDerivation: param([5, 5, 5], 5, (15, 15), 3, [0, 0, 0, 5], id="all_accepted"), ], ) # fmt: skip - def test_histogram(self, survival, steps, totals, num_spec_tokens, - expected): - stats = _build_spec_decode_stats(_rsp(survival, steps, totals), - _args(num_spec_tokens=num_spec_tokens), - "stop") + def test_histogram(self, survival, steps, totals, num_spec_tokens, expected): + stats = _build_spec_decode_stats( + _rsp(survival, steps, totals), _args(num_spec_tokens=num_spec_tokens), "stop" + ) assert stats.acceptance_histogram == expected _assert_identities(stats) @@ -89,24 +98,24 @@ def test_mean_acceptance_length_is_derivable(self): # Not a field: consumers derive it, and it must match the counts. The # field would otherwise duplicate avg_decoded_tokens_per_iter on the # same choice and the two could drift. - stats = _build_spec_decode_stats(_rsp([12, 12, 6], 20, (30, 60)), - _args(num_spec_tokens=3), "stop") - assert 1 + (stats.total_accepted_draft_tokens / - stats.num_spec_steps) == 2.5 + stats = _build_spec_decode_stats( + _rsp([12, 12, 6], 20, (30, 60)), _args(num_spec_tokens=3), "stop" + ) + assert 1 + (stats.total_accepted_draft_tokens / stats.num_spec_steps) == 2.5 def test_histogram_padded_to_configured_budget(self): # Length must describe the draft budget, not the depth this particular # request happened to reach, so it is stable across requests. - stats = _build_spec_decode_stats(_rsp([3], 5, (3, 25)), - _args(num_spec_tokens=5), "stop") + stats = _build_spec_decode_stats(_rsp([3], 5, (3, 25)), _args(num_spec_tokens=5), "stop") assert len(stats.acceptance_histogram) == 6 _assert_identities(stats) def test_adaptive_drafting_reports_no_fixed_bound(self): # Under draft_len_schedule there is no fixed per-step bound, so # num_spec_tokens is None and the histogram sizes to observed depth. - stats = _build_spec_decode_stats(_rsp([12, 12, 6], 20, (30, 60)), - _args(num_spec_tokens=None), "stop") + stats = _build_spec_decode_stats( + _rsp([12, 12, 6], 20, (30, 60)), _args(num_spec_tokens=None), "stop" + ) assert stats.num_spec_tokens is None assert stats.acceptance_histogram == [8, 0, 6, 6] _assert_identities(stats) @@ -117,10 +126,14 @@ def test_deep_drafting_beyond_initial_capacity(self): depth = MAX_SPEC_DECODE_POSITIONS + 4 survival = [1] * depth stats = _build_spec_decode_stats( - SimpleNamespace(per_pos_accepted=survival, - per_pos_drafted=[1] + [0] * (depth - 1), - spec_dec_totals=(depth, depth)), - _args(num_spec_tokens=depth), "stop") + SimpleNamespace( + per_pos_accepted=survival, + per_pos_drafted=[1] + [0] * (depth - 1), + spec_dec_totals=(depth, depth), + ), + _args(num_spec_tokens=depth), + "stop", + ) _assert_identities(stats) @@ -128,30 +141,37 @@ class TestOmission: """The field is absent, not null-filled, whenever it cannot be trusted.""" def test_absent_when_not_opted_in(self): - assert _build_spec_decode_stats(_rsp([12], 20, (30, 60)), - _args(enabled=False, - num_spec_tokens=3), - "stop") is None + assert ( + _build_spec_decode_stats( + _rsp([12], 20, (30, 60)), _args(enabled=False, num_spec_tokens=3), "stop" + ) + is None + ) def test_absent_on_non_terminal_stream_chunk(self): # Streaming carries this only on the chunk bearing finish_reason; # intermediate chunks would report a partial request as if complete. - assert _build_spec_decode_stats(_rsp([12], 20, (30, 60)), - _args(num_spec_tokens=3), None) is None + assert ( + _build_spec_decode_stats(_rsp([12], 20, (30, 60)), _args(num_spec_tokens=3), None) + is None + ) def test_absent_when_nothing_drafted(self): - assert _build_spec_decode_stats(_rsp([], 0, (0, 0)), - _args(num_spec_tokens=3), - "stop") is None + assert ( + _build_spec_decode_stats(_rsp([], 0, (0, 0)), _args(num_spec_tokens=3), "stop") is None + ) def test_absent_without_per_position_vectors(self): # The C++/TRT backend populates spec metrics via # updateNumTokensPerIteration and has no per-position vectors. - assert _build_spec_decode_stats( - SimpleNamespace(per_pos_accepted=None, - per_pos_drafted=None, - spec_dec_totals=None), _args(num_spec_tokens=3), - "stop") is None + assert ( + _build_spec_decode_stats( + SimpleNamespace(per_pos_accepted=None, per_pos_drafted=None, spec_dec_totals=None), + _args(num_spec_tokens=3), + "stop", + ) + is None + ) class TestNumSpecTokensResolution: @@ -168,15 +188,17 @@ def test_no_args_is_none(self): assert resolve_spec_decode_num_spec_tokens(None) is None def test_fixed_bound_is_reported(self): - args = SimpleNamespace(speculative_config=SimpleNamespace( - max_draft_len=4, draft_len_schedule=None)) + args = SimpleNamespace( + speculative_config=SimpleNamespace(max_draft_len=4, draft_len_schedule=None) + ) assert resolve_spec_decode_num_spec_tokens(args) == 4 def test_draft_len_schedule_reports_no_bound(self): # The bound varies by batch size, so None is the honest answer rather # than whichever max_draft_len happens to be configured alongside it. - args = SimpleNamespace(speculative_config=SimpleNamespace( - max_draft_len=4, draft_len_schedule={1: 4, 8: 2})) + args = SimpleNamespace( + speculative_config=SimpleNamespace(max_draft_len=4, draft_len_schedule={1: 4, 8: 2}) + ) assert resolve_spec_decode_num_spec_tokens(args) is None @@ -198,8 +220,7 @@ def _server(enabled, num_spec_tokens=4): @staticmethod def _args(): - return SimpleNamespace(return_spec_decode_stats=False, - spec_decode_num_spec_tokens=None) + return SimpleNamespace(return_spec_decode_stats=False, spec_decode_num_spec_tokens=None) def test_disabled_server_leaves_args_untouched(self): args = self._args() @@ -215,7 +236,158 @@ def test_enabled_server_propagates_fixed_bound(self): def test_enabled_server_propagates_adaptive_bound(self): args = self._args() - OpenAIServer._apply_spec_decode_stats_opt_in( - self._server(True, num_spec_tokens=None), args) + OpenAIServer._apply_spec_decode_stats_opt_in(self._server(True, num_spec_tokens=None), args) assert args.return_spec_decode_stats is True assert args.spec_decode_num_spec_tokens is None + + +def _padded(values): + """Pad a per-position vector to the executor's initial capacity, as sent.""" + return list(values) + [0] * (MAX_SPEC_DECODE_POSITIONS - len(values)) + + +def _sequence_response(*, sequence_index, is_final, per_pos_drafted, per_pos_accepted, totals): + """One child request's final response, as GenerationResultBase receives it. + + Mirrors the executor response shape used in test_disaggregated_params.py, + plus the per-request spec-decode counters the PyTorch executor attaches. + """ + result = SimpleNamespace( + is_final=is_final, + decoding_iter=1, + avg_decoded_tokens_per_iter=None, + context_phase_params=None, + finish_reasons=[tllm.FinishReason.END_ID], + output_token_ids=[[5, 6]], + sequence_index=sequence_index, + cum_log_probs=None, + log_probs=None, + generation_logits=None, + context_logits=None, + request_perf_metrics=None, + additional_context_outputs=None, + additional_generation_outputs=None, + per_pos_drafted=_padded(per_pos_drafted), + per_pos_accepted=_padded(per_pos_accepted), + spec_dec_totals=totals, + ) + return SimpleNamespace(result=result, has_error=lambda: False) + + +def _stats_by_index_from_response(response): + return {c.index: c.speculative_decoding.model_dump() for c in response.choices} + + +def _stats_by_index_from_stream(chunks): + stats = {} + for chunk in chunks: + body = chunk.removeprefix("data: ").strip() + if body == "[DONE]": + continue + for choice in json.loads(body).get("choices", []): + if "speculative_decoding" in choice: + stats[choice["index"]] = choice["speculative_decoding"] + return stats + + +_SPEC_DECODE_OPT_IN = dict( + num_choices=2, num_prompt_tokens=3, return_spec_decode_stats=True, spec_decode_num_spec_tokens=2 +) + + +def _chat_args(): + return ChatPostprocArgs(role="assistant", model="m", **_SPEC_DECODE_OPT_IN) + + +def _completion_args(): + return CompletionPostprocArgs(model="m", detokenize=False, **_SPEC_DECODE_OPT_IN) + + +class TestPerSequenceAttribution: + """Each choice of an n > 1 request reports its own sequence's acceptance. + + Every candidate runs as its own child request with its own counters, but + GenerationResultBase keeps a single request-level copy that each arriving + response overwrites. Building every choice from that copy would stamp + whichever candidate reported last onto all of them. Covers all four + formatters, since each reads the counters at its own call site. + """ + + @staticmethod + def _two_sequence_result(): + # n > 1 needs non-greedy sampling, as every real multi-candidate + # request has. + result = GenerationResultBase( + id=1, sampling_params=SamplingParams(max_tokens=2, n=2, temperature=0.8) + ) + # Sequence 0: 4 steps drafting 2 each; 3 steps accepted >= 1 and 1 + # step accepted both -> histogram [1, 2, 1], 4 of 8 accepted. + result._handle_response( + _sequence_response( + sequence_index=0, + is_final=False, + per_pos_drafted=[4, 4], + per_pos_accepted=[3, 1], + totals=(4, 8), + ) + ) + # Sequence 1 reports last: 3 steps, only 1 accepted anything -> + # histogram [2, 1, 0], 1 of 6 accepted. + result._handle_response( + _sequence_response( + sequence_index=1, + is_final=True, + per_pos_drafted=[3, 3], + per_pos_accepted=[1], + totals=(1, 6), + ) + ) + return result + + @pytest.mark.parametrize( + "post_processor, make_args, stats_by_index", + [ + param( + completion_response_post_processor, + _completion_args, + _stats_by_index_from_response, + id="completion", + ), + param( + completion_stream_post_processor, + _completion_args, + _stats_by_index_from_stream, + id="completion_stream", + ), + param( + chat_response_post_processor, _chat_args, _stats_by_index_from_response, id="chat" + ), + param( + chat_stream_post_processor, + _chat_args, + _stats_by_index_from_stream, + id="chat_stream", + ), + ], + ) + def test_each_choice_reports_its_own_sequence(self, post_processor, make_args, stats_by_index): + output = post_processor(self._two_sequence_result(), make_args()) + + assert stats_by_index(output) == { + 0: { + "acceptance_rate": 0.5, + "total_accepted_draft_tokens": 4, + "total_draft_tokens": 8, + "num_spec_steps": 4, + "acceptance_histogram": [1, 2, 1], + "num_spec_tokens": 2, + }, + 1: { + "acceptance_rate": 1 / 6, + "total_accepted_draft_tokens": 1, + "total_draft_tokens": 6, + "num_spec_steps": 3, + "acceptance_histogram": [2, 1, 0], + "num_spec_tokens": 2, + }, + } From 529628c40b528dd0e87a4929de000ab1e31f5dcf Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Fri, 25 Sep 2026 13:36:24 -0400 Subject: [PATCH 6/6] fix: register speculative decoding fields in serve API reference Signed-off-by: Allison Lim --- .../references/trtllm_serve_api.yaml | 24 +++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/tests/unittest/api_stability/references/trtllm_serve_api.yaml b/tests/unittest/api_stability/references/trtllm_serve_api.yaml index e818d9683a5d..a1209f00a6e5 100644 --- a/tests/unittest/api_stability/references/trtllm_serve_api.yaml +++ b/tests/unittest/api_stability/references/trtllm_serve_api.yaml @@ -364,6 +364,12 @@ models: default: null status: beta required: false + speculative_decoding: + kind: extension + type: Optional[SpeculativeDecodingStats] + default: null + status: prototype + required: false CompletionResponse: fields: @@ -454,6 +460,12 @@ models: default: null status: beta required: false + speculative_decoding: + kind: extension + type: Optional[SpeculativeDecodingStats] + default: null + status: prototype + required: false CompletionStreamResponse: fields: @@ -877,6 +889,12 @@ models: default: null status: beta required: false + speculative_decoding: + kind: extension + type: Optional[SpeculativeDecodingStats] + default: null + status: prototype + required: false ChatCompletionResponse: fields: @@ -967,6 +985,12 @@ models: default: null status: beta required: false + speculative_decoding: + kind: extension + type: Optional[SpeculativeDecodingStats] + default: null + status: prototype + required: false ChatCompletionStreamResponse: fields: