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
@@ -7,6 +7,7 @@ import os
import random import random
import re import re
import time import time
from types import SimpleNamespace
from typing import Callable, List, Optional, Tuple from typing import Callable, List, Optional, Tuple
import numpy as np import numpy as np
@@ -16,10 +17,9 @@ from tabulate import tabulate
from transformers import AutoProcessor, PreTrainedTokenizer from transformers import AutoProcessor, PreTrainedTokenizer
from sglang.bench_serving import ( from sglang.bench_serving import (
get_dataset,
get_processor, get_processor,
get_tokenizer, get_tokenizer,
sample_mmmu_requests,
sample_random_requests,
) )
from sglang.profiler import run_profile from sglang.profiler import run_profile
from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.entrypoints.http_server import launch_server
@@ -103,6 +103,10 @@ class BenchArgs:
profile_output_dir: Optional[str] = None profile_output_dir: Optional[str] = None
dataset_path: str = "" dataset_path: str = ""
dataset_name: str = "random" 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 parallel_batch: bool = False
result_filename: str = "result.jsonl" result_filename: str = "result.jsonl"
pydantic_result_filename: Optional[str] = None pydantic_result_filename: Optional[str] = None
@@ -164,9 +168,33 @@ class BenchArgs:
"--dataset-name", "--dataset-name",
type=str, type=str,
default=BenchArgs.dataset_name, default=BenchArgs.dataset_name,
choices=["mmmu", "random"], choices=["mmmu", "random", "generated-shared-prefix"],
help="Name of the dataset to benchmark on.", 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("--parallel-batch", action="store_true")
parser.add_argument( parser.add_argument(
"--result-filename", "--result-filename",
@@ -379,6 +407,10 @@ def run_one_case(
cache_hit_rate: float = BenchArgs.cache_hit_rate, cache_hit_rate: float = BenchArgs.cache_hit_rate,
backend: str = "sglang", backend: str = "sglang",
model_name: Optional[str] = None, 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": if backend == "vllm":
# You need to have export VLLM_SERVER_DEV_MODE=1 in your environment to use this endpoint. # 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 = requests.post(url + "/flush_cache", timeout=DEFAULT_TIMEOUT)
response.raise_for_status() response.raise_for_status()
# Load input token ids # Load input token ids via bench_serving.get_dataset
# TODO: reuse bench_serving.get_dataset ? supported_datasets = ("random", "mmmu", "generated-shared-prefix")
if dataset_name == "mmmu": if dataset_name not in supported_datasets:
input_requests = sample_mmmu_requests( raise ValueError(
num_requests=batch_size, f"Unsupported dataset for batch benchmark: {dataset_name}. "
processor=tokenizer, f"Supported: {supported_datasets}"
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,
) )
# Extract input_ids from requests actual_gsp_groups = min(gsp_num_groups, batch_size)
if dataset_name == "mmmu": dataset_args = SimpleNamespace(
input_ids = [] dataset_name=dataset_name,
# for vlms, tokenizer is an instance of AutoProcessor num_prompts=batch_size,
tokenizer = tokenizer.tokenizer random_input_len=input_len,
for input_req in input_requests: random_output_len=output_len,
input_ids += [tokenizer.encode(input_req.prompt)] 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] image_data = [req.image_data for req in input_requests]
else: else:
input_ids = [req.prompt for req in input_requests] 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=}" ), 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 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"{max_running_requests_per_dp=}")
print(f"{dp_size=}") print(f"{dp_size=}")
print(f"{skip_max_running_requests_threshold=}") print(f"{skip_max_running_requests_threshold=}")
@@ -799,6 +847,7 @@ def run_benchmark_internal(
parallel_batch=bench_args.parallel_batch, parallel_batch=bench_args.parallel_batch,
backend=bench_args.backend, backend=bench_args.backend,
model_name=model_name, model_name=model_name,
**gsp_kwargs,
) )
print("=" * 8 + " Warmup End " + "=" * 8 + "\n") print("=" * 8 + " Warmup End " + "=" * 8 + "\n")
@@ -834,6 +883,7 @@ def run_benchmark_internal(
cache_hit_rate=bench_args.cache_hit_rate, cache_hit_rate=bench_args.cache_hit_rate,
backend=bench_args.backend, backend=bench_args.backend,
model_name=model_name, model_name=model_name,
**gsp_kwargs,
) )
) )
@@ -876,6 +926,7 @@ def run_benchmark_internal(
profile_output_dir=bench_args.profile_output_dir, profile_output_dir=bench_args.profile_output_dir,
backend=bench_args.backend, backend=bench_args.backend,
model_name=model_name, model_name=model_name,
**gsp_kwargs,
) )
) )
except Exception as e: except Exception as e: