diff --git a/python/sglang/bench_one_batch_server.py b/python/sglang/bench_one_batch_server.py index c6e6cf1c1..39da285e9 100644 --- a/python/sglang/bench_one_batch_server.py +++ b/python/sglang/bench_one_batch_server.py @@ -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() diff --git a/python/sglang/test/bench_one_batch_server_internal.py b/python/sglang/benchmark/one_batch_server.py similarity index 89% rename from python/sglang/test/bench_one_batch_server_internal.py rename to python/sglang/benchmark/one_batch_server.py index ed3b958ad..de95aeb7e 100644 --- a/python/sglang/test/bench_one_batch_server_internal.py +++ b/python/sglang/benchmark/one_batch_server.py @@ -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() diff --git a/test/registered/kv_canary/test_self_e2e_bench_speed.py b/test/registered/kv_canary/test_self_e2e_bench_speed.py index 9ebfa13f7..b7be94b33 100644 --- a/test/registered/kv_canary/test_self_e2e_bench_speed.py +++ b/test/registered/kv_canary/test_self_e2e_bench_speed.py @@ -7,13 +7,13 @@ import unittest from pathlib import Path from typing import ClassVar, Optional -from sglang.srt.entrypoints.http_server import launch_server -from sglang.srt.server_args import ServerArgs -from sglang.test.bench_one_batch_server_internal import ( +from sglang.bench_one_batch_server import ( BenchArgs, BenchOneCaseResult, 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.test_utils import DEFAULT_PORT_FOR_SRT_TEST_RUNNER