diff --git a/python/sglang/test/bench_one_batch_server_internal.py b/python/sglang/test/bench_one_batch_server_internal.py index e2f3ed09b..32c13217a 100644 --- a/python/sglang/test/bench_one_batch_server_internal.py +++ b/python/sglang/test/bench_one_batch_server_internal.py @@ -7,6 +7,7 @@ import os import random import re import time +from types import SimpleNamespace from typing import Callable, List, Optional, Tuple import numpy as np @@ -16,10 +17,9 @@ from tabulate import tabulate from transformers import AutoProcessor, PreTrainedTokenizer from sglang.bench_serving import ( + get_dataset, get_processor, get_tokenizer, - sample_mmmu_requests, - sample_random_requests, ) from sglang.profiler import run_profile from sglang.srt.entrypoints.http_server import launch_server @@ -103,6 +103,10 @@ class BenchArgs: profile_output_dir: Optional[str] = None dataset_path: str = "" dataset_name: str = "random" + gsp_num_groups: int = 1 + gsp_system_prompt_len: int = 2048 + gsp_question_len: int = 128 + gsp_output_len: int = 256 parallel_batch: bool = False result_filename: str = "result.jsonl" pydantic_result_filename: Optional[str] = None @@ -164,9 +168,33 @@ class BenchArgs: "--dataset-name", type=str, default=BenchArgs.dataset_name, - choices=["mmmu", "random"], + choices=["mmmu", "random", "generated-shared-prefix"], help="Name of the dataset to benchmark on.", ) + parser.add_argument( + "--gsp-num-groups", + type=int, + default=BenchArgs.gsp_num_groups, + help="Number of shared prefix groups. batch_size requests are distributed across groups.", + ) + parser.add_argument( + "--gsp-system-prompt-len", + type=int, + default=BenchArgs.gsp_system_prompt_len, + help="Length of the shared system prompt in tokens per group.", + ) + parser.add_argument( + "--gsp-question-len", + type=int, + default=BenchArgs.gsp_question_len, + help="Length of the unique question suffix in tokens per request.", + ) + parser.add_argument( + "--gsp-output-len", + type=int, + default=BenchArgs.gsp_output_len, + help="Output length in tokens for generated-shared-prefix requests.", + ) parser.add_argument("--parallel-batch", action="store_true") parser.add_argument( "--result-filename", @@ -379,6 +407,10 @@ def run_one_case( cache_hit_rate: float = BenchArgs.cache_hit_rate, backend: str = "sglang", model_name: Optional[str] = None, + gsp_num_groups: int = BenchArgs.gsp_num_groups, + gsp_system_prompt_len: int = BenchArgs.gsp_system_prompt_len, + gsp_question_len: int = BenchArgs.gsp_question_len, + gsp_output_len: int = BenchArgs.gsp_output_len, ): if backend == "vllm": # You need to have export VLLM_SERVER_DEV_MODE=1 in your environment to use this endpoint. @@ -388,34 +420,43 @@ def run_one_case( response = requests.post(url + "/flush_cache", timeout=DEFAULT_TIMEOUT) response.raise_for_status() - # Load input token ids - # TODO: reuse bench_serving.get_dataset ? - if dataset_name == "mmmu": - input_requests = sample_mmmu_requests( - num_requests=batch_size, - processor=tokenizer, - fixed_output_len=output_len, - random_sample=False, - ) - elif dataset_name == "random": - input_requests = sample_random_requests( - input_len=input_len, - output_len=output_len, - num_prompts=batch_size, - range_ratio=1.0, - tokenizer=tokenizer, - dataset_path=dataset_path, - random_sample=True, - return_text=False, + # Load input token ids via bench_serving.get_dataset + supported_datasets = ("random", "mmmu", "generated-shared-prefix") + if dataset_name not in supported_datasets: + raise ValueError( + f"Unsupported dataset for batch benchmark: {dataset_name}. " + f"Supported: {supported_datasets}" ) - # Extract input_ids from requests - if dataset_name == "mmmu": - input_ids = [] - # for vlms, tokenizer is an instance of AutoProcessor - tokenizer = tokenizer.tokenizer - for input_req in input_requests: - input_ids += [tokenizer.encode(input_req.prompt)] + actual_gsp_groups = min(gsp_num_groups, batch_size) + dataset_args = SimpleNamespace( + dataset_name=dataset_name, + num_prompts=batch_size, + random_input_len=input_len, + random_output_len=output_len, + random_range_ratio=1.0, + dataset_path=dataset_path, + tokenize_prompt=dataset_name not in ("mmmu", "generated-shared-prefix"), + backend=backend, + seed=BenchArgs.seed, + gsp_num_groups=actual_gsp_groups, + gsp_prompts_per_group=(batch_size + actual_gsp_groups - 1) // actual_gsp_groups, + gsp_system_prompt_len=gsp_system_prompt_len, + gsp_question_len=gsp_question_len, + gsp_output_len=gsp_output_len, + ) + tok_inner = getattr(tokenizer, "tokenizer", tokenizer) + dataset_model_id = model_name or getattr(tok_inner, "name_or_path", None) + input_requests = get_dataset(dataset_args, tokenizer, model_id=dataset_model_id) + + if dataset_name == "generated-shared-prefix": + input_requests = input_requests[:batch_size] + input_ids = [tokenizer.encode(req.prompt) for req in input_requests] + input_len = sum(len(ids) for ids in input_ids) // len(input_ids) + output_len = gsp_output_len + image_data = None + elif dataset_name == "mmmu": + input_ids = [tok_inner.encode(req.prompt) for req in input_requests] image_data = [req.image_data for req in input_requests] else: input_ids = [req.prompt for req in input_requests] @@ -772,6 +813,13 @@ def run_benchmark_internal( ), f"effective_max_running_requests_per_dp is not set, {max_running_requests_per_dp=}" skip_max_running_requests_threshold = max_running_requests_per_dp * dp_size + gsp_kwargs = dict( + gsp_num_groups=bench_args.gsp_num_groups, + gsp_system_prompt_len=bench_args.gsp_system_prompt_len, + gsp_question_len=bench_args.gsp_question_len, + gsp_output_len=bench_args.gsp_output_len, + ) + print(f"{max_running_requests_per_dp=}") print(f"{dp_size=}") print(f"{skip_max_running_requests_threshold=}") @@ -799,6 +847,7 @@ def run_benchmark_internal( parallel_batch=bench_args.parallel_batch, backend=bench_args.backend, model_name=model_name, + **gsp_kwargs, ) print("=" * 8 + " Warmup End " + "=" * 8 + "\n") @@ -834,6 +883,7 @@ def run_benchmark_internal( cache_hit_rate=bench_args.cache_hit_rate, backend=bench_args.backend, model_name=model_name, + **gsp_kwargs, ) ) @@ -876,6 +926,7 @@ def run_benchmark_internal( profile_output_dir=bench_args.profile_output_dir, backend=bench_args.backend, model_name=model_name, + **gsp_kwargs, ) ) except Exception as e: