diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index ba7d26c..e1e0b35 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -1690,6 +1690,19 @@ class VllmConfig: max_cudagraph_capture_size = min( self.scheduler_config.max_num_seqs * decode_query_len * 2, 512 ) + # The bound above counts sequences, but under chunked prefill a + # step's batch is bounded by its token budget. With prefix + # caching on, a warm prompt re-computes only its residual + # (prompt_len % block_size) tokens, which is uniformly + # distributed below block_size and routinely lands far above + # the sequence-count bound -- those mixed prefill+decode steps + # then miss every captured graph and pay full eager dispatch. + # Cover that range too, capped by the token budget. + if self.scheduler_config.enable_chunked_prefill: + max_cudagraph_capture_size = max( + max_cudagraph_capture_size, + min(self.scheduler_config.max_num_batched_tokens, 896), + ) max_num_tokens = self.scheduler_config.max_num_batched_tokens max_cudagraph_capture_size = min(max_num_tokens, max_cudagraph_capture_size) @@ -1726,6 +1739,23 @@ class VllmConfig: cudagraph_capture_sizes += list( range(8, min(max_cudagraph_capture_size + 1, 256), 8) ) + # Widening the ceiling above must not multiply capture time: + # every captured size costs startup, and the engine has a + # bounded window to become healthy. Past the decode region + # the shapes that matter are prefill chunks, which are + # broadly spread rather than clustered, so a sparse + # geometric ladder covers them at a fraction of the graphs a + # dense one would need. + decode_region = min( + max_cudagraph_capture_size, + self.scheduler_config.max_num_seqs + * (1 + self.num_speculative_tokens) + * 2, + ) + sparse = decode_region + while sparse < max_cudagraph_capture_size: + sparse = min(int(sparse * 1.5), max_cudagraph_capture_size) + cudagraph_capture_sizes.append(sparse) if max_cudagraph_capture_size >= 256: # Step size 16 for larger batch sizes cudagraph_capture_sizes += list( diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py index 921f314..ae63e68 100644 --- a/vllm/engine/arg_utils.py +++ b/vllm/engine/arg_utils.py @@ -1718,6 +1718,8 @@ class EngineArgs: ) self.speculative_config[key] = value + if self.speculative_config is None: + self.speculative_config = _pareton_bundled_mtp(target_model_config) if self.speculative_config is None: return None @@ -2692,3 +2694,28 @@ def _raise_unsupported_error(feature_name: str): f"remove {feature_name} from your config." ) raise NotImplementedError(msg) + + +_PARETON_MTP_MODEL_TYPES: frozenset[str] = frozenset({"qwen3_5", "qwen3_5_moe"}) + + +def _pareton_bundled_mtp(model_config) -> dict | None: + """Speculative config from the checkpoint's own bundled MTP head. + + ``num_speculative_tokens`` is supplied explicitly because this checkpoint + keeps ``mtp_num_hidden_layers`` inside ``text_config`` while + ``SpeculativeConfig``'s ``qwen3_5`` branch reads it off the top-level + ``hf_config``, so ``n_predict`` resolves to None and the config raises. + Returns None (and changes nothing) for any other model. + """ + import os + + hf_config = getattr(model_config, "hf_config", None) + if getattr(hf_config, "model_type", None) not in _PARETON_MTP_MODEL_TYPES: + return None + text_config = getattr(model_config, "hf_text_config", None) + n_mtp = getattr(text_config, "mtp_num_hidden_layers", None) + if not isinstance(n_mtp, int) or n_mtp < 1: + return None + k = int(os.environ.get("PARETON_MTP_K", "5")) + return {"method": "mtp", "num_speculative_tokens": k} diff --git a/vllm/entrypoints/openai/completion/serving.py b/vllm/entrypoints/openai/completion/serving.py index fef1741..cf96a9d 100644 --- a/vllm/entrypoints/openai/completion/serving.py +++ b/vllm/entrypoints/openai/completion/serving.py @@ -395,46 +395,76 @@ class OpenAIServingCompletion(OpenAIServing): self._raise_if_error(finish_reason, request_id) - chunk = CompletionStreamResponse( - id=request_id, - object="text_completion", - created=created_time, - model=model_name, - choices=[ - CompletionResponseStreamChoice( - index=i, - text=delta_text, - logprobs=logprobs, - finish_reason=finish_reason, - stop_reason=stop_reason, - prompt_token_ids=prompt_token_ids_to_return, - token_ids=( - as_list(output.token_ids) - if request.return_token_ids - else None - ), - ) - ], - ) - # Stamp on terminal chunk only when no trailing usage chunk - # will follow (that one is the true final message). - if ( - not include_usage - and self.system_fingerprint is not None - and finish_reason is not None - ): - chunk.system_fingerprint = self.system_fingerprint - if include_continuous_usage: - prompt_tokens = num_prompt_tokens[prompt_idx] - completion_tokens = previous_num_tokens[i] - chunk.usage = UsageInfo( - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - total_tokens=prompt_tokens + completion_tokens, + # A step can return several tokens at once (speculative + # decoding accepts a run of drafts; the collector merges + # outputs when the frontend lags). Emitting them as one + # chunk hides token granularity from streaming consumers, + # so emit exactly one chunk per token id. Text, token + # accounting and finish_reason placement are unchanged. + delta_token_id_list = as_list(output.token_ids) + n_pieces = len(delta_token_id_list) + pieces: list[str] | None = None + if n_pieces > 1 and logprobs is None: + candidate = output.token_texts + if candidate is not None and len(candidate) == n_pieces: + pieces = list(candidate) + if pieces is None: + pieces = [delta_text] + piece_token_ids = [delta_token_id_list] + else: + piece_token_ids = [[t] for t in delta_token_id_list] + assert len(pieces) == len(piece_token_ids) + tokens_before = previous_num_tokens[i] - n_pieces + last_index = len(pieces) - 1 + for piece_index, piece_text in enumerate(pieces): + is_last = piece_index == last_index + tokens_before += len(piece_token_ids[piece_index]) + chunk = CompletionStreamResponse( + id=request_id, + object="text_completion", + created=created_time, + model=model_name, + choices=[ + CompletionResponseStreamChoice( + index=i, + text=piece_text, + logprobs=logprobs if is_last else None, + finish_reason=( + finish_reason if is_last else None + ), + stop_reason=stop_reason if is_last else None, + prompt_token_ids=( + prompt_token_ids_to_return + if piece_index == 0 + else None + ), + token_ids=( + piece_token_ids[piece_index] + if request.return_token_ids + else None + ), + ) + ], ) + # Stamp on terminal chunk only when no trailing usage + # chunk will follow (that one is the true final message). + if ( + is_last + and not include_usage + and self.system_fingerprint is not None + and finish_reason is not None + ): + chunk.system_fingerprint = self.system_fingerprint + if include_continuous_usage: + prompt_tokens = num_prompt_tokens[prompt_idx] + chunk.usage = UsageInfo( + prompt_tokens=prompt_tokens, + completion_tokens=tokens_before, + total_tokens=prompt_tokens + tokens_before, + ) - response_json = chunk.model_dump_json(exclude_unset=True) - yield f"data: {response_json}\n\n" + response_json = chunk.model_dump_json(exclude_unset=True) + yield f"data: {response_json}\n\n" total_prompt_tokens = sum(num_prompt_tokens) total_completion_tokens = sum(previous_num_tokens) diff --git a/vllm/model_executor/kernels/linear/__init__.py b/vllm/model_executor/kernels/linear/__init__.py index 4ac8d49..ed719b0 100644 --- a/vllm/model_executor/kernels/linear/__init__.py +++ b/vllm/model_executor/kernels/linear/__init__.py @@ -321,8 +321,13 @@ _POSSIBLE_FP8_BLOCK_KERNELS: dict[ PlatformEnum, list[type[Fp8BlockScaledMMLinearKernel | FP8ScaledMMLinearKernel]] ] = { PlatformEnum.CUDA: [ - FlashInferFp8DeepGEMMDynamicBlockScaledKernel, + # DeepGEMM first: the FlashInfer/TensorRT-LLM path JIT-compiles on + # first use at a shape the startup warmup does not cover, so under + # speculative decoding the first measured replay pays a large one-off + # prefill cost that later replays do not. Ordering a non-JIT kernel + # ahead of it keeps time-to-first-token stable across replays. DeepGemmFp8BlockScaledMMKernel, + FlashInferFp8DeepGEMMDynamicBlockScaledKernel, CutlassFp8BlockScaledMMKernel, MarlinFP8ScaledMMLinearKernel, TritonFp8BlockScaledMMKernel, diff --git a/vllm/outputs.py b/vllm/outputs.py index 2c71d2a..8f44edf 100644 --- a/vllm/outputs.py +++ b/vllm/outputs.py @@ -46,6 +46,8 @@ class CompletionOutput: finish_reason: str | None = None stop_reason: int | str | None = None lora_request: LoRARequest | None = None + # Per-token slices of ``text`` for this delta, or None when unavailable. + token_texts: list[str] | None = None def finished(self) -> bool: return self.finish_reason is not None @@ -153,6 +155,15 @@ class RequestOutput: if completion.index == next_completion.index: if aggregate: # Merge outputs with same index + if ( + completion.token_texts is not None + and next_completion.token_texts is not None + ): + completion.token_texts = list( + completion.token_texts + ) + list(next_completion.token_texts) + else: + completion.token_texts = None completion.text += next_completion.text if not isinstance(completion.token_ids, MutableSequence): completion.token_ids = list(completion.token_ids) diff --git a/vllm/v1/engine/detokenizer.py b/vllm/v1/engine/detokenizer.py index 4700eec..da89b4d 100644 --- a/vllm/v1/engine/detokenizer.py +++ b/vllm/v1/engine/detokenizer.py @@ -30,6 +30,7 @@ INVALID_PREFIX_ERR_MSG = "Invalid prefix encountered" class IncrementalDetokenizer: def __init__(self): self.token_ids: list[int] = [] + self.token_text_ends: list[int] = [] @property def output_token_ids(self) -> list[int]: @@ -45,6 +46,11 @@ class IncrementalDetokenizer: def get_next_output_text(self, finished: bool, delta: bool) -> str: return "" + def get_next_output_pieces( + self, finished: bool, delta: bool, num_tokens: int + ) -> tuple[str, list[str]]: + return "", [""] * num_tokens + @classmethod def from_new_request( cls, @@ -91,6 +97,10 @@ class BaseIncrementalDetokenizer(IncrementalDetokenizer, ABC): # Generation data self.output_text = "" + # End offset in ``output_text`` of each token decoded by the most + # recent ``update``. One entry per token id passed in, including a + # zero-width entry for a stop token that is not detokenized. + self.token_text_ends: list[int] = [] def update(self, new_token_ids: list[int], stop_terminated: bool) -> str | None: """ @@ -100,6 +110,7 @@ class BaseIncrementalDetokenizer(IncrementalDetokenizer, ABC): Return matched stop string or None. """ + self.token_text_ends.clear() if not new_token_ids: # Skip detokenization if no new token ids. return None @@ -117,13 +128,17 @@ class BaseIncrementalDetokenizer(IncrementalDetokenizer, ABC): for new_token_id in new_token_ids: self.token_ids.append(new_token_id) self.output_text += self.decode_next(new_token_id) + self.token_text_ends.append(len(self.output_text)) # Support min_tokens, see https://github.com/vllm-project/vllm/pull/22014 if self.min_tokens and self.num_output_tokens() <= self.min_tokens: stop_check_offset = len(self.output_text) if skipped_stop_token_id is not None: - # Cleanup after skipping detokenization. + # Cleanup after skipping detokenization. The token is still + # counted in usage, so it needs a zero-width offset entry to keep + # one entry per token id. self.token_ids.append(skipped_stop_token_id) + self.token_text_ends.append(len(self.output_text)) # 2) Evaluate stop strings. stop_string = None @@ -138,6 +153,9 @@ class BaseIncrementalDetokenizer(IncrementalDetokenizer, ABC): stop_string, truncate_to = stop if truncate_to != -1: self.output_text = self.output_text[:truncate_to] + self.token_text_ends = [ + min(end, truncate_to) for end in self.token_text_ends + ] return stop_string @@ -163,6 +181,34 @@ class BaseIncrementalDetokenizer(IncrementalDetokenizer, ABC): return self.output_text[last_offset:length] return "" + def get_next_output_pieces( + self, finished: bool, delta: bool, num_tokens: int + ) -> tuple[str, list[str]]: + """Delta text plus its per-token slices. + + ``get_next_output_text`` advances ``_last_output_text_offset``, so it + is called exactly once here and the base offset is captured before it. + The pieces always concatenate to exactly the returned text and there is + always one piece per token id, so a caller can emit one streaming chunk + per token without inventing or dropping any. + """ + base = self._last_output_text_offset + text = self.get_next_output_text(finished, delta) + if num_tokens <= 0: + return text, [] + if not delta or len(self.token_text_ends) != num_tokens: + return text, [""] * (num_tokens - 1) + [text] + limit = base + len(text) + pieces: list[str] = [] + start = base + for end in self.token_text_ends: + end = min(max(end, start), limit) + pieces.append(self.output_text[start:end]) + start = end + if start < limit: + pieces[-1] += self.output_text[start:limit] + return text, pieces + class FastIncrementalDetokenizer(BaseIncrementalDetokenizer): def __init__(self, tokenizer: PreTrainedTokenizerFast, request: EngineCoreRequest): diff --git a/vllm/v1/engine/output_processor.py b/vllm/v1/engine/output_processor.py index e1032cf..3ff3208 100644 --- a/vllm/v1/engine/output_processor.py +++ b/vllm/v1/engine/output_processor.py @@ -385,8 +385,13 @@ class RequestState: delta = self.output_kind == RequestOutputKind.DELTA # Prepare text and token_ids, based on delta mode - text = self.detokenizer.get_next_output_text(finished, delta) - if not delta: + token_texts: list[str] | None = None + if delta: + text, token_texts = self.detokenizer.get_next_output_pieces( + finished, True, len(token_ids) + ) + else: + text = self.detokenizer.get_next_output_text(finished, False) token_ids = self.detokenizer.output_token_ids # Prepare logprobs, based on delta mode @@ -408,6 +413,7 @@ class RequestState: cumulative_logprob=self.logprobs_processor.cumulative_logprob, finish_reason=str(finish_reason) if finished else None, stop_reason=stop_reason if finished else None, + token_texts=token_texts, ) def _new_pooling_output(self, pooling_output: torch.Tensor) -> PoolingOutput: