[misc] Move bench_one_batch_server into sglang/benchmark/ with a back-compat shim (#28625)

This commit is contained in:
Liangsheng Yin
2026-06-19 14:19:29 -07:00
committed by GitHub
parent c9c2445146
commit d271de64fe
3 changed files with 115 additions and 104 deletions
+7 -45
View File
@@ -1,49 +1,11 @@
""" """Back-compat shim. The implementation now lives in
Benchmark the latency of running a single batch with a server. ``sglang.benchmark.one_batch_server``; this module preserves the
``python -m sglang.bench_one_batch_server`` entry point and the
This script launches a server and uses the HTTP interface. ``from sglang.bench_one_batch_server import ...`` imports.
It accepts server arguments (the same as launch_server.py) and benchmark arguments (e.g., batch size, input lengths).
Usage:
python3 -m sglang.bench_one_batch_server --model meta-llama/Meta-Llama-3.1-8B --batch-size 1 16 64 --input-len 1024 --output-len 8
python3 -m sglang.bench_one_batch_server --model None --base-url http://localhost:30000 --batch-size 16 --input-len 1024 --output-len 8
python3 -m sglang.bench_one_batch_server --model None --base-url http://localhost:30000 --batch-size 16 --input-len 1024 --output-len 8 --show-report --profile --profile-by-stage
python3 -m sglang.bench_one_batch_server --model None --base-url http://localhost:30000 --batch-size 16 --input-len 1024 --output-len 8 --result-filename results.jsonl --profile
""" """
import argparse from sglang.benchmark.one_batch_server import * # noqa: F401,F403
from sglang.benchmark.one_batch_server import main
from sglang.srt.server_args import ServerArgs
from sglang.test.bench_one_batch_server_internal import (
BenchArgs,
run_benchmark_internal,
)
from sglang.test.nightly_bench_utils import save_results_as_pydantic_models
def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
results, server_info = run_benchmark_internal(server_args, bench_args)
# Save results as pydantic models in the JSON format
if bench_args.pydantic_result_filename:
save_results_as_pydantic_models(
results,
pydantic_result_filename=bench_args.pydantic_result_filename,
model_path=server_args.model_path,
server_args=bench_args.server_args_for_metrics,
)
return results, server_info
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser() main()
ServerArgs.add_cli_args(parser)
BenchArgs.add_cli_args(parser)
args = parser.parse_args()
server_args = ServerArgs.from_cli_args(args)
bench_args = BenchArgs.from_cli_args(args)
run_benchmark(server_args, bench_args)
@@ -1,3 +1,17 @@
"""
Benchmark the latency of running a single batch with a server.
This script launches a server and uses the HTTP interface.
It accepts server arguments (the same as launch_server.py) and benchmark arguments (e.g., batch size, input lengths).
Usage:
python3 -m sglang.bench_one_batch_server --model meta-llama/Meta-Llama-3.1-8B --batch-size 1 16 64 --input-len 1024 --output-len 8
python3 -m sglang.bench_one_batch_server --model None --base-url http://localhost:30000 --batch-size 16 --input-len 1024 --output-len 8
python3 -m sglang.bench_one_batch_server --model None --base-url http://localhost:30000 --batch-size 16 --input-len 1024 --output-len 8 --show-report --profile --profile-by-stage
python3 -m sglang.bench_one_batch_server --model None --base-url http://localhost:30000 --batch-size 16 --input-len 1024 --output-len 8 --result-filename results.jsonl --profile
"""
import argparse import argparse
import dataclasses import dataclasses
import itertools import itertools
@@ -22,6 +36,7 @@ from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST
from sglang.srt.entrypoints.http_server import launch_server from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_blackwell from sglang.srt.utils import is_blackwell
from sglang.test.nightly_bench_utils import save_results_as_pydantic_models
from sglang.test.test_utils import is_in_ci, write_github_step_summary from sglang.test.test_utils import is_in_ci, write_github_step_summary
DEFAULT_TIMEOUT = 600 DEFAULT_TIMEOUT = 600
@@ -440,13 +455,16 @@ def _warmup_cache(
def _flush_cache_with_retry(url: str, endpoint: str, max_retries: int = 3): def _flush_cache_with_retry(url: str, endpoint: str, max_retries: int = 3):
"""Post to a cache flush endpoint with retries on failure.""" """Post to a cache flush endpoint with retries on failure."""
for attempt in range(max_retries): for attempt in range(max_retries):
response = requests.post(url + endpoint, timeout=DEFAULT_TIMEOUT) try:
if response.status_code == 200: response = requests.post(url + endpoint, timeout=DEFAULT_TIMEOUT)
return if response.status_code == 200:
if attempt < max_retries - 1: return
time.sleep(2) if attempt >= max_retries - 1:
else: response.raise_for_status()
response.raise_for_status() except requests.RequestException:
if attempt >= max_retries - 1:
raise
time.sleep(2)
def run_one_case( def run_one_case(
@@ -635,50 +653,50 @@ def run_one_case(
# Run the request # Run the request
tic = time.perf_counter() tic = time.perf_counter()
response = requests.post( with requests.post(
gen_url, gen_url,
json=payload, json=payload,
stream=True, stream=True,
timeout=DEFAULT_TIMEOUT, timeout=DEFAULT_TIMEOUT,
) ) as response:
response.raise_for_status() response.raise_for_status()
# Get the TTFT of the last request in the batch # Get the TTFT of the last request in the batch
last_ttft = 0.0 last_ttft = 0.0
if backend == "vllm": if backend == "vllm":
# Parse OpenAI-compatible streaming format from vLLM # Parse OpenAI-compatible streaming format from vLLM
first_token_indices = set() first_token_indices = set()
for chunk in response.iter_lines(decode_unicode=False): for chunk in response.iter_lines(decode_unicode=False):
chunk = chunk.decode("utf-8") chunk = chunk.decode("utf-8")
if chunk and chunk.startswith("data:"): if chunk and chunk.startswith("data:"):
data_str = chunk[5:].strip() data_str = chunk[5:].strip()
if data_str == "[DONE]": if data_str == "[DONE]":
break break
data = json.loads(data_str) data = json.loads(data_str)
if "error" in data: if "error" in data:
raise RuntimeError(f"Request has failed. {data}.") raise RuntimeError(f"Request has failed. {data}.")
for choice in data.get("choices", []): for choice in data.get("choices", []):
idx = choice["index"] idx = choice["index"]
if idx not in first_token_indices: if idx not in first_token_indices:
first_token_indices.add(idx) first_token_indices.add(idx)
if len(first_token_indices) == batch_size: if len(first_token_indices) == batch_size:
last_ttft = time.perf_counter() - tic last_ttft = time.perf_counter() - tic
else: else:
for chunk in response.iter_lines(decode_unicode=False): for chunk in response.iter_lines(decode_unicode=False):
chunk = chunk.decode("utf-8") chunk = chunk.decode("utf-8")
if chunk and chunk.startswith("data:"): if chunk and chunk.startswith("data:"):
if chunk == "data: [DONE]": if chunk == "data: [DONE]":
break break
data = json.loads(chunk[5:].strip("\n")) data = json.loads(chunk[5:].strip("\n"))
if "error" in data: if "error" in data:
raise RuntimeError(f"Request has failed. {data}.") raise RuntimeError(f"Request has failed. {data}.")
assert ( assert (
data["meta_info"]["finish_reason"] is None data["meta_info"]["finish_reason"] is None
or data["meta_info"]["finish_reason"]["type"] == "length" or data["meta_info"]["finish_reason"]["type"] == "length"
) )
if data["meta_info"]["completion_tokens"] == 1: if data["meta_info"]["completion_tokens"] == 1:
last_ttft = time.perf_counter() - tic last_ttft = time.perf_counter() - tic
# Compute metrics # Compute metrics
latency = time.perf_counter() - tic latency = time.perf_counter() - tic
@@ -694,9 +712,10 @@ def run_one_case(
response = requests.get(url + "/server_info", timeout=DEFAULT_TIMEOUT) response = requests.get(url + "/server_info", timeout=DEFAULT_TIMEOUT)
response.raise_for_status() response.raise_for_status()
server_info = response.json() server_info = response.json()
internal_state = server_info.get("internal_states", [{}]) internal_states = server_info.get("internal_states", [])
last_gen_throughput = internal_state[0].get("last_gen_throughput", None) or -1 internal_state = internal_states[0] if internal_states else {}
acc_length = internal_state[0].get("avg_spec_accept_length", None) or -1 last_gen_throughput = internal_state.get("last_gen_throughput", None) or -1
acc_length = internal_state.get("avg_spec_accept_length", None) or -1
# Calculate cache hit rate from before/after metrics delta # Calculate cache hit rate from before/after metrics delta
metrics_after = get_cache_tokens_from_metrics(url) metrics_after = get_cache_tokens_from_metrics(url)
@@ -888,22 +907,21 @@ def run_benchmark_internal(
else: else:
tokenizer = get_tokenizer(tokenizer_path) tokenizer = get_tokenizer(tokenizer_path)
internal_state = server_info.get("internal_states", [{}]) internal_states = server_info.get("internal_states", [])
dp_size = internal_state[0].get("dp_size", None) or 1 internal_state = internal_states[0] if internal_states else {}
dp_size = internal_state.get("dp_size", None) or 1
# Get effective max running requests # Get effective max running requests
max_running_requests_per_dp = internal_state[0].get( max_running_requests_per_dp = internal_state.get(
"effective_max_running_requests_per_dp", -1 "effective_max_running_requests_per_dp", -1
) )
# Get token capacity # Get token capacity
skip_token_capacity_threshold = 0 skip_token_capacity_threshold = 0
for i in range(dp_size): for state in internal_states:
skip_token_capacity_threshold += ( skip_token_capacity_threshold += state.get("memory_usage", {}).get(
internal_state[i] "token_capacity", 1000000000
.get("memory_usage", {})
.get("token_capacity", 1000000000)
) )
assert ( assert (
@@ -1113,3 +1131,34 @@ def run_benchmark_internal(
write_github_step_summary(summary) write_github_step_summary(summary)
return results, server_info return results, server_info
def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
results, server_info = run_benchmark_internal(server_args, bench_args)
# Save results as pydantic models in the JSON format
if bench_args.pydantic_result_filename:
save_results_as_pydantic_models(
results,
pydantic_result_filename=bench_args.pydantic_result_filename,
model_path=server_args.model_path,
server_args=bench_args.server_args_for_metrics,
)
return results, server_info
def main():
parser = argparse.ArgumentParser()
ServerArgs.add_cli_args(parser)
BenchArgs.add_cli_args(parser)
args = parser.parse_args()
server_args = ServerArgs.from_cli_args(args)
bench_args = BenchArgs.from_cli_args(args)
run_benchmark(server_args, bench_args)
if __name__ == "__main__":
main()
@@ -7,13 +7,13 @@ import unittest
from pathlib import Path from pathlib import Path
from typing import ClassVar, Optional from typing import ClassVar, Optional
from sglang.srt.entrypoints.http_server import launch_server from sglang.bench_one_batch_server import (
from sglang.srt.server_args import ServerArgs
from sglang.test.bench_one_batch_server_internal import (
BenchArgs, BenchArgs,
BenchOneCaseResult, BenchOneCaseResult,
run_benchmark_internal, run_benchmark_internal,
) )
from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import DEFAULT_PORT_FOR_SRT_TEST_RUNNER from sglang.test.test_utils import DEFAULT_PORT_FOR_SRT_TEST_RUNNER