[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 @@
"""
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
"""Back-compat shim. The implementation now lives in
``sglang.benchmark.one_batch_server``; this module preserves the
``python -m sglang.bench_one_batch_server`` entry point and the
``from sglang.bench_one_batch_server import ...`` imports.
"""
import argparse
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
from sglang.benchmark.one_batch_server import * # noqa: F401,F403
from sglang.benchmark.one_batch_server import main
if __name__ == "__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)
main()
@@ -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 dataclasses
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.server_args import ServerArgs
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
DEFAULT_TIMEOUT = 600
@@ -440,13 +455,16 @@ def _warmup_cache(
def _flush_cache_with_retry(url: str, endpoint: str, max_retries: int = 3):
"""Post to a cache flush endpoint with retries on failure."""
for attempt in range(max_retries):
response = requests.post(url + endpoint, timeout=DEFAULT_TIMEOUT)
if response.status_code == 200:
return
if attempt < max_retries - 1:
time.sleep(2)
else:
response.raise_for_status()
try:
response = requests.post(url + endpoint, timeout=DEFAULT_TIMEOUT)
if response.status_code == 200:
return
if attempt >= max_retries - 1:
response.raise_for_status()
except requests.RequestException:
if attempt >= max_retries - 1:
raise
time.sleep(2)
def run_one_case(
@@ -635,50 +653,50 @@ def run_one_case(
# Run the request
tic = time.perf_counter()
response = requests.post(
with requests.post(
gen_url,
json=payload,
stream=True,
timeout=DEFAULT_TIMEOUT,
)
response.raise_for_status()
) as response:
response.raise_for_status()
# Get the TTFT of the last request in the batch
last_ttft = 0.0
if backend == "vllm":
# Parse OpenAI-compatible streaming format from vLLM
first_token_indices = set()
for chunk in response.iter_lines(decode_unicode=False):
chunk = chunk.decode("utf-8")
if chunk and chunk.startswith("data:"):
data_str = chunk[5:].strip()
if data_str == "[DONE]":
break
data = json.loads(data_str)
if "error" in data:
raise RuntimeError(f"Request has failed. {data}.")
for choice in data.get("choices", []):
idx = choice["index"]
if idx not in first_token_indices:
first_token_indices.add(idx)
if len(first_token_indices) == batch_size:
last_ttft = time.perf_counter() - tic
else:
for chunk in response.iter_lines(decode_unicode=False):
chunk = chunk.decode("utf-8")
if chunk and chunk.startswith("data:"):
if chunk == "data: [DONE]":
break
data = json.loads(chunk[5:].strip("\n"))
if "error" in data:
raise RuntimeError(f"Request has failed. {data}.")
# Get the TTFT of the last request in the batch
last_ttft = 0.0
if backend == "vllm":
# Parse OpenAI-compatible streaming format from vLLM
first_token_indices = set()
for chunk in response.iter_lines(decode_unicode=False):
chunk = chunk.decode("utf-8")
if chunk and chunk.startswith("data:"):
data_str = chunk[5:].strip()
if data_str == "[DONE]":
break
data = json.loads(data_str)
if "error" in data:
raise RuntimeError(f"Request has failed. {data}.")
for choice in data.get("choices", []):
idx = choice["index"]
if idx not in first_token_indices:
first_token_indices.add(idx)
if len(first_token_indices) == batch_size:
last_ttft = time.perf_counter() - tic
else:
for chunk in response.iter_lines(decode_unicode=False):
chunk = chunk.decode("utf-8")
if chunk and chunk.startswith("data:"):
if chunk == "data: [DONE]":
break
data = json.loads(chunk[5:].strip("\n"))
if "error" in data:
raise RuntimeError(f"Request has failed. {data}.")
assert (
data["meta_info"]["finish_reason"] is None
or data["meta_info"]["finish_reason"]["type"] == "length"
)
if data["meta_info"]["completion_tokens"] == 1:
last_ttft = time.perf_counter() - tic
assert (
data["meta_info"]["finish_reason"] is None
or data["meta_info"]["finish_reason"]["type"] == "length"
)
if data["meta_info"]["completion_tokens"] == 1:
last_ttft = time.perf_counter() - tic
# Compute metrics
latency = time.perf_counter() - tic
@@ -694,9 +712,10 @@ def run_one_case(
response = requests.get(url + "/server_info", timeout=DEFAULT_TIMEOUT)
response.raise_for_status()
server_info = response.json()
internal_state = server_info.get("internal_states", [{}])
last_gen_throughput = internal_state[0].get("last_gen_throughput", None) or -1
acc_length = internal_state[0].get("avg_spec_accept_length", None) or -1
internal_states = server_info.get("internal_states", [])
internal_state = internal_states[0] if internal_states else {}
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
metrics_after = get_cache_tokens_from_metrics(url)
@@ -888,22 +907,21 @@ def run_benchmark_internal(
else:
tokenizer = get_tokenizer(tokenizer_path)
internal_state = server_info.get("internal_states", [{}])
dp_size = internal_state[0].get("dp_size", None) or 1
internal_states = server_info.get("internal_states", [])
internal_state = internal_states[0] if internal_states else {}
dp_size = internal_state.get("dp_size", None) or 1
# 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
)
# Get token capacity
skip_token_capacity_threshold = 0
for i in range(dp_size):
skip_token_capacity_threshold += (
internal_state[i]
.get("memory_usage", {})
.get("token_capacity", 1000000000)
for state in internal_states:
skip_token_capacity_threshold += state.get("memory_usage", {}).get(
"token_capacity", 1000000000
)
assert (
@@ -1113,3 +1131,34 @@ def run_benchmark_internal(
write_github_step_summary(summary)
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()