Add generated-shared-prefix dataset in bench_one_batch (#18986)

This commit is contained in:
Qiaolin Yu
2026-02-20 13:33:10 -08:00
committed by GitHub
parent ab18734375
commit 96bae2355e

View File

@@ -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: