diff --git a/benchmarks/benchmark_beam_search_server.py b/benchmarks/benchmark_beam_search_server.py new file mode 100644 index 000000000000..051cbbdd6c35 --- /dev/null +++ b/benchmarks/benchmark_beam_search_server.py @@ -0,0 +1,277 @@ +#!/usr/bin/env python3 +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Simple benchmark script to test beam search performance via OpenAI API. + +This script sends beam search requests to a running vLLM server and handles +profiling via the /start_profile and /stop_profile endpoints. + +Usage: + # Start server with profiling enabled: + # VLLM_TORCH_PROFILER_DIR=./vllm_profile vllm serve Qwen/Qwen3-8B --max-logprobs 100 + + # Run this script: + python benchmarks/benchmark_beam_search_server.py \ + --base-url http://localhost:8000 \ + --beam-width 30 \ + --max-tokens 8 \ + --input-len 128 \ + --num-requests 10 --profile +""" + +import argparse +import random +import time + +import requests +from openai import OpenAI +from tqdm import tqdm + + +def start_profile(base_url: str) -> bool: + """Start profiling on the server.""" + try: + response = requests.post(f"{base_url}/start_profile", timeout=30) + response.raise_for_status() + print("✓ Profiling started") + return True + except Exception as e: + print(f"✗ Failed to start profiling: {e}") + return False + + +def stop_profile(base_url: str) -> bool: + """Stop profiling on the server.""" + try: + response = requests.post(f"{base_url}/stop_profile", timeout=30) + response.raise_for_status() + print("✓ Profiling stopped") + return True + except Exception as e: + print(f"✗ Failed to stop profiling: {e}") + return False + + +def send_beam_search_request( + client: OpenAI, + model: str, + prompt: str, + beam_width: int = 30, + max_tokens: int = 8, + temperature: float = 0.0, +) -> dict | None: + """Send a beam search completion request.""" + start_time = time.perf_counter() + completion = client.completions.create( + model=model, + prompt=prompt, + n=beam_width, # Beam width + max_tokens=max_tokens, + temperature=temperature, + extra_body={"use_beam_search": True}, # This triggers beam search + ) + end_time = time.perf_counter() + + latency = end_time - start_time + + # Convert completion object to dict-like structure for compatibility + num_choices = len(completion.choices) if completion.choices else 0 + + return { + "latency": latency, + "result": completion, + "num_choices": num_choices, + } + + +def generate_prompt(input_len: int, seed: int | None = None) -> str: + """Generate a random prompt of approximately the specified length. + + Each call generates a unique prompt using random numbers to ensure + disjoint prompts across requests. + """ + if seed is not None: + random.seed(seed) + + # Approximate numbers per token (numbers are typically 1 token each) + num_numbers = int(input_len * 0.9) # Slightly less to account for spaces + + # Generate random numbers + numbers = [str(random.randint(0, 999999)) for _ in range(num_numbers)] + + prompt = " ".join(numbers) + + # Reset seed if it was set to avoid affecting other random calls + if seed is not None: + random.seed() + + return prompt + + +def main(): + parser = argparse.ArgumentParser( + description="Benchmark beam search performance via OpenAI API" + ) + parser.add_argument( + "--base-url", + type=str, + default="http://localhost:8000", + help="Base URL of the vLLM server (default: http://localhost:8000)", + ) + parser.add_argument( + "--beam-width", + type=int, + default=30, + help="Beam width (n parameter) (default: 30)", + ) + parser.add_argument( + "--max-tokens", + type=int, + default=8, + help="Maximum output tokens (default: 8)", + ) + parser.add_argument( + "--input-len", + type=int, + default=128, + help="Approximate input prompt length in tokens (default: 128)", + ) + parser.add_argument( + "--num-requests", + type=int, + default=10, + help="Number of requests to send (default: 10)", + ) + parser.add_argument( + "--temperature", + type=float, + default=1.0, + help="Temperature for beam search (default: 1.0)", + ) + parser.add_argument( + "--prompt", + type=str, + default=None, + help="Custom prompt (if not provided, generates dummy prompt)", + ) + parser.add_argument( + "--profile", + action="store_true", + help="Add profiling (call start/stop_profile endpoints)", + ) + parser.add_argument( + "--model", + type=str, + default=None, + help="Model name (if not provided, will query server for available models)", + ) + + args = parser.parse_args() + + print("Beam Search Server Benchmark") + print(f"Server URL: {args.base_url}") + print(f"Beam width (n): {args.beam_width}") + print(f"Max tokens: {args.max_tokens}") + print(f"Input length: ~{args.input_len} tokens") + print(f"Number of requests: {args.num_requests}") + print(f"Temperature: {args.temperature}") + + # Initialize OpenAI client + openai_api_base = f"{args.base_url}/v1" + client = OpenAI( + api_key="EMPTY", # vLLM doesn't require a real API key + base_url=openai_api_base, + ) + + # Get model name + if args.model: + model_name = args.model + print(f"Model: {model_name}") + else: + print("\nQuerying server for available models...") + try: + models = client.models.list() + if models.data: + model_name = models.data[0].id + print(f"Model: {model_name}") + else: + print("✗ No models available. Please specify --model") + return + except Exception as e: + print(f"✗ Failed to query models: {e}") + print("Please specify --model explicitly") + return + + # Send test request with n=3 + test_prompt = ( + args.prompt if args.prompt else generate_prompt(args.input_len, seed=None) + ) + print("\nSending test request (n=3)...", end=" ", flush=True) + test_result = send_beam_search_request( + client=client, + model=model_name, + prompt=test_prompt, + beam_width=3, + max_tokens=args.max_tokens, + temperature=args.temperature, + ) + if test_result: + print(f"✓ Test request successful (latency: {test_result['latency']:.4f}s)") + else: + print("✗ Test request failed") + print("Aborting benchmark due to test request failure.") + return + + # Start profiling + if args.profile: + print("\nStarting profiler...") + if not start_profile(args.base_url): + print("Warning: Could not start profiling. Continuing anyway...") + time.sleep(1) # Brief pause to ensure profiler is ready + + print(f"\nSending {args.num_requests} requests...") + latencies = [] + successful_requests = 0 + for _ in tqdm(range(args.num_requests)): + # Generate a unique random prompt for each request + prompt = args.prompt or generate_prompt(args.input_len, seed=None) + result = send_beam_search_request( + client, + model_name, + prompt, + args.beam_width, + args.max_tokens, + args.temperature, + ) + + if result: + latency = result["latency"] + latencies.append(latency) + successful_requests += 1 + else: + print("✗ Failed request") + + # Stop profiling + if args.profile: + print("\nStopping profiler...") + stop_profile(args.base_url) + + # Print statistics + if latencies: + import statistics + + print("\nResults") + print(f"Successful requests: {successful_requests}/{args.num_requests}") + print(f"Average latency: {statistics.mean(latencies) * 1000:.2f}ms") + print(f"Median latency: {statistics.median(latencies) * 1000:.2f}ms") + if len(latencies) > 1: + print(f"Std deviation: {statistics.stdev(latencies) * 1000:.2f}ms") + print(f"Min latency: {min(latencies) * 1000:.2f}ms") + print(f"Max latency: {max(latencies) * 1000:.2f}ms") + print("=" * 80) + else: + print("\n✗ No successful requests. Check server logs for errors.") + + +if __name__ == "__main__": + main() diff --git a/vllm/config/model.py b/vllm/config/model.py index 97cba6ea7295..fb5d1da495f6 100644 --- a/vllm/config/model.py +++ b/vllm/config/model.py @@ -4,6 +4,7 @@ import warnings from collections.abc import Callable from dataclasses import InitVar, field +from functools import cached_property from importlib.util import find_spec from typing import TYPE_CHECKING, Any, Literal, cast, get_args @@ -1592,7 +1593,7 @@ def get_diff_sampling_param(self) -> dict[str, Any]: ) return diff_sampling_param - @property + @cached_property def is_encoder_decoder(self) -> bool: """Extract the HF encoder/decoder model flag.""" return is_encoder_decoder(self.hf_config) diff --git a/vllm/entrypoints/openai/serving_engine.py b/vllm/entrypoints/openai/serving_engine.py index 127b8e6dcb87..b7b1fbedf264 100644 --- a/vllm/entrypoints/openai/serving_engine.py +++ b/vllm/entrypoints/openai/serving_engine.py @@ -395,6 +395,8 @@ async def beam_search( logprobs=logprobs_num, max_tokens=1, temperature=temperature, + detokenize=False, # We detokenize the output after the search + skip_clone=True, # Safe to reuse params without cloning ) all_beams = [ BeamSearchSequence( diff --git a/vllm/inputs/preprocess.py b/vllm/inputs/preprocess.py index 839c13868a16..d9650b715013 100644 --- a/vllm/inputs/preprocess.py +++ b/vllm/inputs/preprocess.py @@ -1,8 +1,9 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -from collections.abc import Mapping -from typing import Any, cast +from collections.abc import Callable, Mapping +from functools import wraps +from typing import Any, TypeVar, cast from typing_extensions import assert_never @@ -41,6 +42,25 @@ logger = init_logger(__name__) +T = TypeVar("T") + + +def cache_if_not_none(func: Callable[..., T]) -> Callable[..., T]: + """Cache function results only if they are not None.""" + cache_attr = f"_cached_{func.__name__}" + + @wraps(func) + def wrapper(*args: Any, **kwargs: Any) -> T: + self = args[0] + if hasattr(self, cache_attr): + return getattr(self, cache_attr) + result = func(*args, **kwargs) + if result is not None: + setattr(self, cache_attr, result) + return result + + return wrapper + class InputPreprocessor: def __init__( @@ -59,6 +79,7 @@ def __init__( self.mm_cache_stats = MultiModalCacheStats() if mm_processor_cache else None + @cache_if_not_none def get_tokenizer(self) -> AnyTokenizer: if self.tokenizer is None: raise ValueError( @@ -67,6 +88,7 @@ def get_tokenizer(self) -> AnyTokenizer: return self.tokenizer + @cache_if_not_none def get_bos_token_id(self) -> int | None: if self.tokenizer is None: logger.warning_once( @@ -76,6 +98,7 @@ def get_bos_token_id(self) -> int | None: return self.tokenizer.bos_token_id + @cache_if_not_none def get_eos_token_id(self) -> int | None: if self.tokenizer is None: logger.warning_once( @@ -85,6 +108,7 @@ def get_eos_token_id(self) -> int | None: return self.tokenizer.eos_token_id + @cache_if_not_none def get_decoder_start_token_id(self) -> int | None: """ Obtain the decoder start token id employed by an encoder/decoder diff --git a/vllm/logprobs.py b/vllm/logprobs.py index 6a820308f523..cb234f5f162d 100644 --- a/vllm/logprobs.py +++ b/vllm/logprobs.py @@ -1,6 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -import itertools from collections.abc import Iterable, Iterator, MutableSequence from dataclasses import dataclass, field from typing import overload @@ -75,7 +74,7 @@ def append_fast( self, token_ids: list[int], logprobs: list[float], - ranks: itertools.chain[int], + ranks: list[int], decoded_tokens: Iterable[str | None], ) -> None: """ @@ -186,21 +185,16 @@ def append_logprobs_for_next_position( # We do not need a special case for the sampled token # being in the topk, since inserting duplicated data # into a dictionary twice is the same as doing it once. - topk_ranks = range(1, num_logprobs + 1) - ranks = itertools.chain((rank,), topk_ranks) + ranks = [rank] + list(range(1, num_logprobs + 1)) if isinstance(request_logprobs, FlatLogprobs): request_logprobs.append_fast(token_ids, logprobs, ranks, decoded_tokens) else: - request_logprobs.append( - { - token_id: Logprob( - logprob=logprob, - rank=rank, - decoded_token=token, - ) - for token_id, logprob, rank, token in zip( - token_ids, logprobs, ranks, decoded_tokens - ) - } - ) + logprobs_dict: dict[int, Logprob] = {} + for token_id, logprob, rank_val, token in zip( + token_ids, logprobs, ranks, decoded_tokens + ): + logprobs_dict[token_id] = Logprob( + logprob=logprob, rank=rank_val, decoded_token=token + ) + request_logprobs.append(logprobs_dict) diff --git a/vllm/sampling_params.py b/vllm/sampling_params.py index 0fb1d67687c8..efc220bd2fa3 100644 --- a/vllm/sampling_params.py +++ b/vllm/sampling_params.py @@ -231,6 +231,12 @@ class SamplingParams( set to an integer k, will use only the last k tokens from the prompt (i.e., left truncation). If set to `None`, truncation is disabled.""" output_kind: RequestOutputKind = RequestOutputKind.CUMULATIVE + skip_clone: bool = False + """Internal flag indicating that this SamplingParams instance is safe to + reuse without cloning. When True, clone() will return self without + performing a deep copy. This should only be set when the params object + is guaranteed to be dedicated to a single request and won't be modified + in ways that would affect other uses.""" # The below fields are not supposed to be used as an input. # They are set in post_init. @@ -618,8 +624,13 @@ def clone(self) -> "SamplingParams": data that is expensive to copy. However, if not copied, the processor needs to support parallel decoding for multiple sequences See https://github.com/vllm-project/vllm/issues/3087 + + If skip_clone is True, returns self without copying (early exit). """ + if self.skip_clone: + return self + logit_processor_refs = ( None if self.logits_processors is None diff --git a/vllm/v1/core/kv_cache_manager.py b/vllm/v1/core/kv_cache_manager.py index 2012c3fef88b..d6f17cc693e7 100644 --- a/vllm/v1/core/kv_cache_manager.py +++ b/vllm/v1/core/kv_cache_manager.py @@ -9,7 +9,7 @@ from vllm.distributed.kv_events import KVCacheEvent from vllm.logger import init_logger from vllm.v1.core.kv_cache_coordinator import get_kv_cache_coordinator -from vllm.v1.core.kv_cache_utils import KVCacheBlock +from vllm.v1.core.kv_cache_utils import BlockHash, KVCacheBlock from vllm.v1.kv_cache_interface import KVCacheConfig from vllm.v1.metrics.stats import PrefixCacheStats from vllm.v1.request import Request @@ -17,6 +17,20 @@ logger = init_logger(__name__) +@dataclass(frozen=True) +class _CacheHitCache: + """ + Stores the result of find_longest_cache_hit to avoid recomputing when + consecutive requests have identical block sequences. + """ + + last_block_hash: BlockHash + num_block_hashes: int + max_cache_hit_length: int + computed_blocks: tuple[list[KVCacheBlock], ...] + num_new_computed_tokens: int + + @dataclass class KVCacheBlocks: """ @@ -151,6 +165,10 @@ def __init__( tuple(() for _ in range(self.num_kv_cache_groups)) ) + # Cache for optimizing get_computed_blocks when consecutive requests + # have the same last block hash and max_cache_hit_length. + self._cache_hit_cache: _CacheHitCache | None = None + @property def usage(self) -> float: """Get the KV cache usage. @@ -198,11 +216,34 @@ def get_computed_blocks(self, request: Request) -> tuple[KVCacheBlocks, int]: # num_computed_tokens to be block-size aligned. Removing this limitation # could slightly improve performance in the future. max_cache_hit_length = request.num_tokens - 1 - computed_blocks, num_new_computed_tokens = ( - self.coordinator.find_longest_cache_hit( - request.block_hashes, max_cache_hit_length + + # Optimization: Reuse previous result when consecutive requests are the same. + if ( + request.block_hashes + and self._cache_hit_cache is not None + and max_cache_hit_length == self._cache_hit_cache.max_cache_hit_length + and len(request.block_hashes) == self._cache_hit_cache.num_block_hashes + and request.block_hashes[-1] == self._cache_hit_cache.last_block_hash + ): + computed_blocks = self._cache_hit_cache.computed_blocks + num_new_computed_tokens = self._cache_hit_cache.num_new_computed_tokens + else: + computed_blocks, num_new_computed_tokens = ( + self.coordinator.find_longest_cache_hit( + request.block_hashes, max_cache_hit_length + ) + ) + self._cache_hit_cache = ( + _CacheHitCache( + request.block_hashes[-1], + len(request.block_hashes), + max_cache_hit_length, + computed_blocks, + num_new_computed_tokens, + ) + if request.block_hashes + else None ) - ) if self.log_stats: assert self.prefix_cache_stats is not None @@ -352,6 +393,7 @@ def reset_prefix_cache(self) -> bool: """ if not self.block_pool.reset_prefix_cache(): return False + self._cache_hit_cache = None if self.log_stats: assert self.prefix_cache_stats is not None self.prefix_cache_stats.reset = True