From 3c481b9421024de4ecc6840159517f83c6acc84d Mon Sep 17 00:00:00 2001 From: BingjiaWang Date: Tue, 25 Aug 2026 05:49:20 +0800 Subject: [PATCH] [Benchmark] Add optional steady-state window for serving metrics (#30918) --- python/sglang/benchmark/steady_state.py | 268 ++++++++++++++++++ .../sglang/benchmark/steady_state_serving.py | 188 ++++++++++++ .../bench_fn/test_steady_state_benchmark.py | 164 +++++++++++ 3 files changed, 620 insertions(+) create mode 100644 python/sglang/benchmark/steady_state.py create mode 100644 python/sglang/benchmark/steady_state_serving.py create mode 100644 test/registered/bench_fn/test_steady_state_benchmark.py diff --git a/python/sglang/benchmark/steady_state.py b/python/sglang/benchmark/steady_state.py new file mode 100644 index 000000000..a10f286fb --- /dev/null +++ b/python/sglang/benchmark/steady_state.py @@ -0,0 +1,268 @@ +"""Steady-state metrics for online serving benchmarks. + +This module deliberately owns only the interval-selection and interval-throughput +logic. The regular serving benchmark remains responsible for issuing requests +and reporting its full-run metrics. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, List, Optional, Protocol, Sequence, Tuple + +import numpy as np + + +class RequestOutput(Protocol): + """The subset of a serving request result needed by this module.""" + + success: bool + start_time: float + latency: float + ttft: float + itl: List[float] + output_len: int + generated_text: str + + +class InputRequest(Protocol): + prompt_len: int + + +class Tokenizer(Protocol): + def encode(self, text: str, add_special_tokens: bool = False) -> List[Any]: ... + + +@dataclass(frozen=True) +class SteadyStateWindow: + start: float + end: float + duration: float + concurrency_threshold: int + peak_concurrency: int + + +@dataclass(frozen=True) +class SteadyStateMetrics: + concurrency_ratio: float + window_start: float + window_end: float + duration: float + concurrency_threshold: int + completed: int + total_input: Optional[int] + input_throughput: Optional[float] + total_output: float + total_output_retokenized: float + output_throughput: float + output_throughput_retokenized: float + average_concurrency: float + peak_concurrency: int + peak_output_throughput: float + + +def find_steady_state_window( + outputs: Sequence[RequestOutput], concurrency_ratio: float +) -> SteadyStateWindow: + """Select the longest continuous interval above the concurrency threshold.""" + if not 0 < concurrency_ratio <= 1: + raise ValueError("steady-state concurrency ratio must be in (0, 1]") + + successful = [output for output in outputs if output.success and output.latency > 0] + if not successful: + raise ValueError("no successful requests with positive latency") + + events = {} + for output in successful: + events[output.start_time] = events.get(output.start_time, 0) + 1 + end_time = output.start_time + output.latency + events[end_time] = events.get(end_time, 0) - 1 + + concurrency = 0 + intervals = [] + event_times = sorted(events) + for index, event_time in enumerate(event_times[:-1]): + concurrency += events[event_time] + next_time = event_times[index + 1] + if next_time > event_time: + intervals.append((event_time, next_time, concurrency)) + + peak_concurrency = max((item[2] for item in intervals), default=0) + threshold = max(1, int(np.ceil(peak_concurrency * concurrency_ratio))) + + spans: List[Tuple[float, float]] = [] + span_start: Optional[float] = None + span_end: Optional[float] = None + for start, end, active_requests in intervals: + if active_requests >= threshold: + if span_start is None: + span_start = start + span_end = end + elif span_start is not None and span_end is not None: + spans.append((span_start, span_end)) + span_start = span_end = None + if span_start is not None and span_end is not None: + spans.append((span_start, span_end)) + + if not spans: + raise ValueError("no steady-state measurement window could be determined") + + # max() keeps the first span when durations tie, making selection deterministic. + window_start, window_end = max(spans, key=lambda span: span[1] - span[0]) + return SteadyStateWindow( + start=window_start, + end=window_end, + duration=window_end - window_start, + concurrency_threshold=threshold, + peak_concurrency=peak_concurrency, + ) + + +def _token_timestamps(output: RequestOutput) -> List[float]: + timestamps = [output.start_time + output.ttft] + for inter_token_latency in output.itl: + timestamps.append(timestamps[-1] + inter_token_latency) + return timestamps + + +def _output_tokens_in_window( + outputs: Sequence[RequestOutput], + output_lens: Sequence[int], + retokenized_output_lens: Sequence[int], + window: SteadyStateWindow, +) -> Tuple[float, float, List[Tuple[float, float]]]: + output_tokens = 0.0 + retokenized_tokens = 0.0 + token_events: List[Tuple[float, float]] = [] + + for output, output_len, retokenized_len in zip( + outputs, output_lens, retokenized_output_lens + ): + if not output.success or output.latency <= 0: + continue + + timestamps = _token_timestamps(output) + timestamps_in_window = [ + timestamp + for timestamp in timestamps + if window.start <= timestamp <= window.end + ] + fraction_in_window = len(timestamps_in_window) / len(timestamps) + output_tokens += output_len * fraction_in_window + retokenized_tokens += retokenized_len * fraction_in_window + + event_weight = output_len / len(timestamps) + token_events.extend( + (timestamp, event_weight) for timestamp in timestamps_in_window + ) + + return output_tokens, retokenized_tokens, token_events + + +def calculate_steady_state_metrics( + outputs: Sequence[RequestOutput], + tokenizer: Tokenizer, + concurrency_ratio: float, + input_requests: Optional[Sequence[InputRequest]] = None, +) -> SteadyStateMetrics: + """Calculate token throughput metrics for the steady-state window. + + Input tokens do not have per-token timestamps. Their throughput therefore uses + an arrival-based definition: prompt tokens from successful requests that start + inside the half-open measurement window ``[start, end)`` divided by its duration. + ``None`` is reported when aligned input requests are unavailable, as in the + current multi-turn serving benchmark. + """ + window = find_steady_state_window(outputs, concurrency_ratio) + + if input_requests is not None and len(input_requests) != len(outputs): + raise ValueError("input requests and outputs must have the same length") + + if input_requests is None: + total_input = None + input_throughput = None + else: + total_input = sum( + request.prompt_len + for request, output in zip(input_requests, outputs) + if output.success and window.start <= output.start_time < window.end + ) + input_throughput = total_input / window.duration + + output_lens = [output.output_len if output.success else 0 for output in outputs] + retokenized_output_lens = [ + ( + len(tokenizer.encode(output.generated_text, add_special_tokens=False)) + if output.success + else 0 + ) + for output in outputs + ] + total_output, total_output_retokenized, token_events = _output_tokens_in_window( + outputs, + output_lens, + retokenized_output_lens, + window, + ) + + successful = [output for output in outputs if output.success and output.latency > 0] + completed = sum( + window.start <= output.start_time + output.latency <= window.end + for output in successful + ) + overlap_duration = sum( + max( + 0.0, + min(output.start_time + output.latency, window.end) + - max(output.start_time, window.start), + ) + for output in successful + ) + + duration_seconds = max(1, int(np.ceil(window.duration))) + tokens_per_second = np.zeros(duration_seconds) + for token_time, token_weight in token_events: + elapsed = max(0.0, token_time - window.start) + # A token exactly on a second boundary belongs to the second that just + # ended. In particular, a token at window.end must not be folded into + # a bucket that already contains the preceding second's tokens. + bucket = max(0, int(np.ceil(elapsed)) - 1) + bucket = min(bucket, duration_seconds - 1) + tokens_per_second[bucket] += token_weight + + return SteadyStateMetrics( + concurrency_ratio=concurrency_ratio, + window_start=window.start, + window_end=window.end, + duration=window.duration, + concurrency_threshold=window.concurrency_threshold, + completed=completed, + total_input=total_input, + input_throughput=input_throughput, + total_output=total_output, + total_output_retokenized=total_output_retokenized, + output_throughput=total_output / window.duration, + output_throughput_retokenized=total_output_retokenized / window.duration, + average_concurrency=overlap_duration / window.duration, + peak_concurrency=window.peak_concurrency, + peak_output_throughput=float(np.max(tokens_per_second)), + ) + + +def steady_state_output_throughput( + outputs: Sequence[RequestOutput], + output_lens: Sequence[int], + retokenized_output_lens: Sequence[int], + concurrency_ratio: float, +) -> Tuple[float, float, float, int]: + """Compatibility helper for callers interested only in output throughput.""" + window = find_steady_state_window(outputs, concurrency_ratio) + output_tokens, retokenized_tokens, _ = _output_tokens_in_window( + outputs, output_lens, retokenized_output_lens, window + ) + return ( + output_tokens / window.duration, + retokenized_tokens / window.duration, + window.duration, + window.concurrency_threshold, + ) diff --git a/python/sglang/benchmark/steady_state_serving.py b/python/sglang/benchmark/steady_state_serving.py new file mode 100644 index 000000000..91e5e4b8c --- /dev/null +++ b/python/sglang/benchmark/steady_state_serving.py @@ -0,0 +1,188 @@ +"""Dedicated serving benchmark entry point with steady-state reporting. + +Usage example:: + + python -m sglang.benchmark.steady_state_serving \ + --steady-state-concurrency-ratio 0.8 \ + --backend sglang --dataset-name random --num-prompts 1000 + +All arguments other than the two steady-state options are forwarded unchanged to +``sglang.benchmark.serving``. The regular serving module is not modified and its +normal full-run result remains intact. +""" + +from __future__ import annotations + +import argparse +import inspect +import json +import sys +from dataclasses import asdict +from pathlib import Path +from typing import Any, Callable, Dict, Optional, Sequence, Tuple + +from sglang.benchmark import serving +from sglang.benchmark.steady_state import ( + SteadyStateMetrics, + calculate_steady_state_metrics, +) + + +def _print_metrics(metrics: SteadyStateMetrics) -> None: + print("\n{s:{c}^{n}}".format(s=" Steady-State Result ", n=50, c="=")) + print("{:<40} {:<10.2f}".format("Concurrency ratio:", metrics.concurrency_ratio)) + print("{:<40} {:<10.2f}".format("Measurement duration (s):", metrics.duration)) + print("{:<40} {:<10}".format("Minimum concurrency:", metrics.concurrency_threshold)) + print("{:<40} {:<10}".format("Completed requests:", metrics.completed)) + if metrics.input_throughput is None: + print("{:<40} {:<10}".format("Input token throughput (tok/s):", "N/A")) + else: + print("{:<40} {:<10}".format("Input tokens:", metrics.total_input)) + print( + "{:<40} {:<10.2f}".format( + "Input token throughput (tok/s):", metrics.input_throughput + ) + ) + print("{:<40} {:<10.2f}".format("Generated tokens:", metrics.total_output)) + print( + "{:<40} {:<10.2f}".format( + "Output token throughput (tok/s):", metrics.output_throughput + ) + ) + print( + "{:<40} {:<10.2f}".format( + "Output throughput, retokenized (tok/s):", + metrics.output_throughput_retokenized, + ) + ) + print( + "{:<40} {:<10.2f}".format("Average concurrency:", metrics.average_concurrency) + ) + print("{:<40} {:<10}".format("Peak concurrency:", metrics.peak_concurrency)) + print( + "{:<40} {:<10.2f}".format( + "Peak output throughput (tok/s):", metrics.peak_output_throughput + ) + ) + print("=" * 50) + + +def _append_metrics(path: str, metrics: SteadyStateMetrics) -> None: + with Path(path).open("a") as file: + file.write(json.dumps(asdict(metrics)) + "\n") + + +def _run_with_capture( + args: argparse.Namespace, + concurrency_ratio: float, + output_file: Optional[str], + run_serving_benchmark: Callable[[argparse.Namespace], Dict[str, Any]], +) -> Tuple[Dict[str, Any], SteadyStateMetrics]: + captured: Dict[str, Any] = {} + original_calculate_metrics = serving.calculate_metrics + calculate_signature = inspect.signature(original_calculate_metrics) + + def capture_calculate_metrics(*call_args, **call_kwargs): + bound = calculate_signature.bind(*call_args, **call_kwargs) + input_requests = bound.arguments["input_requests"] + captured["input_requests"] = ( + None if input_requests is None else list(input_requests) + ) + captured["outputs"] = list(bound.arguments["outputs"]) + captured["tokenizer"] = bound.arguments["tokenizer"] + return original_calculate_metrics(*call_args, **call_kwargs) + + serving.calculate_metrics = capture_calculate_metrics + try: + benchmark_result = run_serving_benchmark(args) + finally: + serving.calculate_metrics = original_calculate_metrics + + if "outputs" not in captured: + raise RuntimeError( + "serving benchmark finished without producing request results" + ) + + steady_state_metrics = calculate_steady_state_metrics( + outputs=captured["outputs"], + tokenizer=captured["tokenizer"], + concurrency_ratio=concurrency_ratio, + input_requests=captured["input_requests"], + ) + _print_metrics(steady_state_metrics) + if output_file: + _append_metrics(output_file, steady_state_metrics) + + return benchmark_result, steady_state_metrics + + +def run_steady_state_benchmark( + args: argparse.Namespace, + concurrency_ratio: float, + output_file: Optional[str] = None, +) -> Tuple[Dict[str, Any], SteadyStateMetrics]: + """Run the normal serving benchmark and add an isolated steady-state report.""" + return _run_with_capture( + args=args, + concurrency_ratio=concurrency_ratio, + output_file=output_file, + run_serving_benchmark=serving.run_benchmark, + ) + + +def _custom_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description="Add isolated steady-state metrics to a serving benchmark run.", + epilog="All other options are forwarded to sglang.benchmark.serving.", + allow_abbrev=False, + ) + parser.add_argument( + "--steady-state-concurrency-ratio", + type=float, + required=True, + help=( + "Measure the longest continuous interval whose concurrency is at " + "least this fraction of peak concurrency. Must be in (0, 1]." + ), + ) + parser.add_argument( + "--steady-state-output-file", + type=str, + default=None, + help="Append the standalone steady-state result to this JSONL file.", + ) + return parser + + +def cli_main(argv: Optional[Sequence[str]] = None) -> None: + argv = list(sys.argv[1:] if argv is None else argv) + custom_parser = _custom_parser() + custom_args, serving_argv = custom_parser.parse_known_args(argv) + + # Validate before server setup so malformed values fail fast. + if not 0 < custom_args.steady_state_concurrency_ratio <= 1: + custom_parser.error("--steady-state-concurrency-ratio must be in (0, 1]") + + original_run_benchmark = serving.run_benchmark + original_argv = sys.argv + + def run_and_report(args: argparse.Namespace): + result, _ = _run_with_capture( + args=args, + concurrency_ratio=custom_args.steady_state_concurrency_ratio, + output_file=custom_args.steady_state_output_file, + run_serving_benchmark=original_run_benchmark, + ) + return result + + serving.run_benchmark = run_and_report + sys.argv = [original_argv[0], *serving_argv] + try: + serving.cli_main() + finally: + sys.argv = original_argv + serving.run_benchmark = original_run_benchmark + + +if __name__ == "__main__": + cli_main() diff --git a/test/registered/bench_fn/test_steady_state_benchmark.py b/test/registered/bench_fn/test_steady_state_benchmark.py new file mode 100644 index 000000000..534e34221 --- /dev/null +++ b/test/registered/bench_fn/test_steady_state_benchmark.py @@ -0,0 +1,164 @@ +import unittest +from argparse import Namespace +from contextlib import redirect_stdout +from dataclasses import dataclass, field +from io import StringIO +from unittest.mock import patch + +from sglang.benchmark import steady_state_serving +from sglang.benchmark.steady_state import ( + calculate_steady_state_metrics, + find_steady_state_window, + steady_state_output_throughput, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class _StringTokenizer: + def encode(self, text, add_special_tokens=False): + return text.split() + + +@dataclass +class _RequestOutput: + generated_text: str + success: bool + start_time: float + latency: float + ttft: float + itl: list[float] = field(default_factory=list) + output_len: int = 0 + + +@dataclass +class _InputRequest: + prompt_len: int + + +def _request(start_time, latency, ttft, itl, output_len): + return _RequestOutput( + generated_text=" ".join(["token"] * output_len), + success=True, + start_time=start_time, + latency=latency, + ttft=ttft, + itl=itl, + output_len=output_len, + ) + + +class TestSteadyStateMetrics(unittest.TestCase): + def test_trims_ramp_up_and_drain(self): + outputs = [ + _request(0.0, 10.0, 1.0, [2.0, 2.0, 2.0, 2.0], 5), + _request(2.0, 6.0, 1.0, [2.0, 2.0], 3), + _request(2.0, 6.0, 1.0, [2.0, 2.0], 3), + ] + + metrics = calculate_steady_state_metrics( + outputs=outputs, + tokenizer=_StringTokenizer(), + concurrency_ratio=0.8, + input_requests=[ + _InputRequest(prompt_len=100), + _InputRequest(prompt_len=200), + _InputRequest(prompt_len=300), + ], + ) + + self.assertEqual(metrics.window_start, 2.0) + self.assertEqual(metrics.window_end, 8.0) + self.assertEqual(metrics.duration, 6.0) + self.assertEqual(metrics.concurrency_threshold, 3) + self.assertEqual(metrics.completed, 2) + self.assertEqual(metrics.total_input, 500) + self.assertEqual(metrics.input_throughput, 500 / 6) + self.assertEqual(metrics.total_output, 9.0) + self.assertEqual(metrics.output_throughput, 1.5) + self.assertEqual(metrics.output_throughput_retokenized, 1.5) + self.assertEqual(metrics.average_concurrency, 3.0) + self.assertEqual(metrics.peak_concurrency, 3) + self.assertEqual(metrics.peak_output_throughput, 3.0) + + throughput, retokenized_throughput, duration, threshold = ( + steady_state_output_throughput( + outputs, [5, 3, 3], [5, 3, 3], concurrency_ratio=0.8 + ) + ) + self.assertEqual(throughput, 1.5) + self.assertEqual(retokenized_throughput, 1.5) + self.assertEqual(duration, 6.0) + self.assertEqual(threshold, 3) + + def test_selects_longest_continuous_high_concurrency_span(self): + outputs = [ + _request(0.0, 2.0, 0.5, [0.5], 2), + _request(0.0, 2.0, 0.5, [0.5], 2), + _request(3.0, 4.0, 0.5, [0.5], 2), + _request(3.0, 4.0, 0.5, [0.5], 2), + ] + + window = find_steady_state_window(outputs, concurrency_ratio=1.0) + + self.assertEqual((window.start, window.end), (3.0, 7.0)) + self.assertEqual(window.concurrency_threshold, 2) + + def test_rejects_invalid_or_empty_input(self): + with self.assertRaisesRegex(ValueError, "must be in"): + find_steady_state_window([], concurrency_ratio=0.0) + with self.assertRaisesRegex(ValueError, "no successful requests"): + find_steady_state_window([], concurrency_ratio=0.8) + + def test_dedicated_runner_preserves_regular_metrics_calculator(self): + outputs = [ + _request(0.0, 4.0, 1.0, [1.0, 1.0, 1.0], 4), + _request(0.0, 4.0, 1.0, [1.0, 1.0, 1.0], 4), + ] + tokenizer = _StringTokenizer() + input_requests = [_InputRequest(prompt_len=4), _InputRequest(prompt_len=4)] + + def regular_calculate_metrics( + input_requests, outputs, dur_s, tokenizer, backend + ): + return object(), [] + + def regular_run_benchmark(args): + steady_state_serving.serving.calculate_metrics( + input_requests=input_requests, + outputs=outputs, + dur_s=4.0, + tokenizer=tokenizer, + backend="sglang", + ) + return {"normal_result": True} + + with ( + patch.object( + steady_state_serving.serving, + "calculate_metrics", + new=regular_calculate_metrics, + ), + patch.object( + steady_state_serving.serving, + "run_benchmark", + new=regular_run_benchmark, + ), + redirect_stdout(StringIO()), + ): + result, metrics = steady_state_serving.run_steady_state_benchmark( + Namespace(), concurrency_ratio=1.0 + ) + self.assertIs( + steady_state_serving.serving.calculate_metrics, + regular_calculate_metrics, + ) + + self.assertEqual(result, {"normal_result": True}) + self.assertEqual(metrics.input_throughput, 2.0) + self.assertEqual(metrics.output_throughput, 2.0) + + +if __name__ == "__main__": + unittest.main()