diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index fc4c549f63ef..23f18d5a6afb 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, @@ -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, 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/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/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index d596c1d38a4e..af49b8d79df8 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/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/tensorrt_llm/serve/openai_protocol.py b/tensorrt_llm/serve/openai_protocol.py index 0d85fd554229..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 @@ -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 @@ -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 @@ -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): @@ -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 @@ -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): @@ -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 @@ -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): @@ -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): diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index 689ed7f341e2..b98f128c9f44 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 @@ -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 @@ -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: @@ -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) + async def openai_chat(self, request: ChatCompletionRequest, raw_request: Request) -> Response: @@ -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 @@ -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 @@ -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, diff --git a/tensorrt_llm/serve/postprocess_handlers.py b/tensorrt_llm/serve/postprocess_handlers.py index 3e440543e67d..4f1d05f2bc0b 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,73 @@ def from_request(cls, request: ChatCompletionRequest): ) +def _build_spec_decode_stats( + 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 + 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 + or counters is None): + return 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 + # 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 +589,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( + output._spec_dec_counters, args, output.finish_reason), stop_reason=output.stop_reason, ) if args.return_logprobs: @@ -686,6 +756,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( + output._spec_dec_counters, args, output.finish_reason), ) if output.finish_reason == "stop" and args.has_tool_call.get( output.index, False): @@ -814,6 +886,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( + output._spec_dec_counters, args, output.finish_reason), ) if args.return_logprobs: logprobs = output.logprobs_diff @@ -882,6 +956,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( + output._spec_dec_counters, 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..3d26cd723b35 100644 --- a/tests/integration/test_lists/test-db/l0_a10.yml +++ b/tests/integration/test_lists/test-db/l0_a10.yml @@ -145,6 +145,8 @@ 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_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/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/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: diff --git a/tests/unittest/executor/test_spec_dec_stats_pairing.py b/tests/unittest/executor/test_spec_dec_stats_pairing.py index 10c09afd5317..aaeef373fff8 100644 --- a/tests/unittest/executor/test_spec_dec_stats_pairing.py +++ b/tests/unittest/executor/test_spec_dec_stats_pairing.py @@ -114,6 +114,35 @@ 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. 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 + def _make_llm_request(request_id, seq_slot): return LlmRequest( 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..f513b0a1ab48 --- /dev/null +++ b/tests/unittest/executor/test_spec_decode_serialization.py @@ -0,0 +1,127 @@ +# 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() 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..04433fdce238 --- /dev/null +++ b/tests/unittest/executor/test_spec_decode_stats_payload.py @@ -0,0 +1,393 @@ +# 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. +""" + +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.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 + ) + + +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 + ) + + +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 + + +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, + }, + } 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 = {