diff --git a/docs/docs/developer_guide/benchmark_and_profiling.mdx b/docs/docs/developer_guide/benchmark_and_profiling.mdx index fa9963d22..94d94d037 100644 --- a/docs/docs/developer_guide/benchmark_and_profiling.mdx +++ b/docs/docs/developer_guide/benchmark_and_profiling.mdx @@ -59,6 +59,7 @@ Use `bench_serving` by default unless there are specific needs. ``` - Pass `--enable-multi-batch` and set `--batch-size` to a multiple of the server's `--max-running-requests` to stabilize throughput measurements. Surplus requests are queued by the scheduler and promoted batch-by-batch, amortizing per-request prefill and first-step transients into steady-state decode. Under this flag, only `overall_throughput` is authoritative; `input_throughput`, `output_throughput`, `last_ttft`, and ITL include cross-batch queueing in their denominators and should be treated as informational. + - Pass `--disable-ignore-eos` when benchmarking with real prompts, where forcing decode past EOS would shift the output distribution toward gibberish. Requests then early-stop, so the whole-run `output_throughput`/ITL include the decaying-batch tail; the report gains steady-state columns measured only over the window where every request is still decoding (from the last request's first token to the first request's finish). Replay recorded traffic with `--dataset-name sharegpt` or `--dataset-name custom --dataset-path ` (recorded completion lens are ignored; `--output-len` is the shared `max_new_tokens` cap). For OpenAI-format traces with per-request parameters, use `bench_serving`. - Pass `--lora-name ` to route every prompt through a pre-loaded LoRA adapter. Requires the server to be launched with `--enable-lora --lora-paths =`. **`bench_offline_throughput`** directly instantiates the `Engine` object in-process (no HTTP server) and submits all requests at once via `engine.generate()`. The engine's scheduler handles batching and execution. This measures maximum achievable throughput without any network overhead. diff --git a/python/sglang/benchmark/one_batch_server.py b/python/sglang/benchmark/one_batch_server.py index 3ea5f80b9..d904992c9 100644 --- a/python/sglang/benchmark/one_batch_server.py +++ b/python/sglang/benchmark/one_batch_server.py @@ -16,6 +16,7 @@ import argparse import dataclasses import itertools import json +import os import random import re import time @@ -31,6 +32,11 @@ from transformers import AutoProcessor, PreTrainedTokenizer from sglang.benchmark.datasets import get_dataset from sglang.benchmark.endpoint import acquire_endpoint +from sglang.benchmark.stream_metrics import ( + BatchStreamRecorder, + SteadyStateWindow, + validate_finish_reason, +) from sglang.benchmark.utils import get_processor, get_tokenizer from sglang.profiler import run_profile from sglang.srt.arg_groups.overrides import resolving_view @@ -43,6 +49,10 @@ from sglang.test.test_utils import is_in_ci, write_github_step_summary DEFAULT_TIMEOUT = 600 +# Replay recorded text prompts (encoded client-side); recorded completion +# lens are ignored -- --output-len caps every request. +REPLAY_TEXT_DATASETS = ("sharegpt", "custom") + def get_cache_tokens_from_metrics(url: str) -> Optional[tuple]: """ @@ -107,6 +117,7 @@ class BenchArgs: input_len: Tuple[int] = (1024,) output_len: Tuple[int] = (16,) temperature: float = 0.0 + disable_ignore_eos: bool = False return_logprob: bool = False client_stream_interval: int = 1 input_len_step_percentage: float = 0.0 @@ -156,6 +167,23 @@ class BenchArgs: "--output-len", type=int, nargs="+", default=BenchArgs.output_len ) parser.add_argument("--temperature", type=float, default=BenchArgs.temperature) + parser.add_argument( + "--disable-ignore-eos", + action="store_true", + help=( + "Let requests finish at EOS/stop tokens instead of sending " + "ignore_eos=true. Use when benchmarking with real prompts " + "(e.g. --dataset-name sharegpt/custom), where forcing decode " + "past EOS shifts the output distribution toward gibberish. " + "Requests then early-stop, so the whole-run output_throughput " + "(actual tokens generated after the last request's first " + "token, divided by the time since then) still includes the " + "decaying-batch tail; read the steady-state metrics instead, " + "which cover only the window where every request is still " + "decoding (from the last request's first token to the first " + "request's finish)." + ), + ) parser.add_argument("--return-logprob", action="store_true") parser.add_argument( "--client-stream-interval", @@ -219,8 +247,19 @@ class BenchArgs: "--dataset-name", type=str, default=BenchArgs.dataset_name, - choices=["mmmu", "random", "random-ids", "generated-shared-prefix"], - help="Name of the dataset to benchmark on.", + choices=[ + "mmmu", + "random", + "random-ids", + "generated-shared-prefix", + "sharegpt", + "custom", + ], + help="Name of the dataset to benchmark on. sharegpt/custom replay " + "recorded text prompts (custom reads --dataset-path JSONL); their " + "recorded output lens are ignored in favor of --output-len. For " + "OpenAI-format traces with per-request parameters, use " + "sglang.benchmark.serving instead.", ) parser.add_argument( "--fixed-prompt-file", @@ -232,8 +271,9 @@ class BenchArgs: parser.add_argument( "--apply-chat-template", action="store_true", - help="Encode the prompt as a single user message through the " - "model's chat template. Requires --fixed-prompt-file.", + help="Encode each prompt as a single user message through the " + "model's chat template. Requires --fixed-prompt-file or a replay " + "text dataset (sharegpt/custom).", ) parser.add_argument( "--gsp-num-groups", @@ -377,6 +417,10 @@ class BenchOneCaseResult(BaseModel): last_ttft: float last_gen_throughput: float acc_length: float + total_output_tokens: Optional[int] = None + num_early_stopped: Optional[int] = None + steady_state_output_throughput: Optional[float] = None + steady_state_window_s: Optional[float] = None cache_hit_rate: Optional[float] = None profile_link: Optional[str] = None @@ -394,6 +438,18 @@ class BenchOneCaseResult(BaseModel): "last_ttft": round(self.last_ttft, 4), "last_gen_throughput": round(self.last_gen_throughput, 2), "acc_length": round(self.acc_length, 2), + "total_output_tokens": self.total_output_tokens, + "num_early_stopped": self.num_early_stopped, + "steady_state_output_throughput": ( + round(self.steady_state_output_throughput, 2) + if self.steady_state_output_throughput is not None + else None + ), + "steady_state_window_s": ( + round(self.steady_state_window_s, 4) + if self.steady_state_window_s is not None + else None + ), "cache_hit_rate": ( round(self.cache_hit_rate, 4) if self.cache_hit_rate is not None @@ -531,6 +587,7 @@ def run_one_case( input_len: int, output_len: int, temperature: float, + ignore_eos: bool, return_logprob: bool, stream_interval: int, input_len_step_percentage: float, @@ -576,7 +633,12 @@ def run_one_case( image_data = None else: # Load input token ids via benchmark.datasets.get_dataset - supported_datasets = ("random", "random-ids", "mmmu", "generated-shared-prefix") + supported_datasets = ( + "random", + "random-ids", + "mmmu", + "generated-shared-prefix", + ) + REPLAY_TEXT_DATASETS if dataset_name not in supported_datasets: raise ValueError( f"Unsupported dataset for batch benchmark: {dataset_name}. " @@ -591,8 +653,12 @@ def run_one_case( random_output_len=output_len, random_range_ratio=1.0, dataset_path=dataset_path, - tokenize_prompt=dataset_name not in ("mmmu", "generated-shared-prefix"), + tokenize_prompt=dataset_name in ("random", "random-ids"), backend=backend, + sharegpt_output_len=None, + sharegpt_context_len=None, + prompt_suffix="", + apply_chat_template=apply_chat_template, seed=BenchArgs.seed, gsp_num_groups=actual_gsp_groups, gsp_prompts_per_group=(batch_size + actual_gsp_groups - 1) @@ -618,6 +684,16 @@ def run_one_case( 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] + elif dataset_name in REPLAY_TEXT_DATASETS: + if len(input_requests) < batch_size: + raise ValueError( + f"{dataset_name} produced only {len(input_requests)} usable " + f"prompts for batch size {batch_size}; use a smaller batch " + f"size or a larger dataset." + ) + input_ids = [tok_inner.encode(req.prompt) for req in input_requests] + input_len = sum(len(ids) for ids in input_ids) // len(input_ids) + image_data = None else: input_ids = [req.prompt for req in input_requests] image_data = None @@ -630,7 +706,7 @@ def run_one_case( "max_tokens": output_len, "temperature": temperature, "stream": True, - "ignore_eos": True, + "ignore_eos": ignore_eos, } if return_logprob: payload["logprobs"] = 1 @@ -654,7 +730,7 @@ def run_one_case( "sampling_params": { "temperature": temperature, "max_new_tokens": output_len, - "ignore_eos": True, + "ignore_eos": ignore_eos, "json_schema": json_schema, "stream_interval": stream_interval, }, @@ -725,6 +801,7 @@ def run_one_case( metrics_before = get_cache_tokens_from_metrics(url) # Run the request + recorder: Optional[BatchStreamRecorder] = None tic = time.perf_counter() with requests.post( gen_url, @@ -755,6 +832,8 @@ def run_one_case( if len(first_token_indices) == batch_size: last_ttft = time.perf_counter() - tic else: + # mmmu can silently return fewer prompts than requested. + recorder = BatchStreamRecorder(batch_size=len(input_ids)) for chunk in response.iter_lines(decode_unicode=False): chunk = chunk.decode("utf-8") if chunk and chunk.startswith("data:"): @@ -764,18 +843,67 @@ def run_one_case( 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" + finish_reason = data["meta_info"]["finish_reason"] + if finish_reason is not None: + validate_finish_reason(finish_reason, ignore_eos=ignore_eos) + recorder.record_chunk( + index=data["index"], + completion_tokens=data["meta_info"]["completion_tokens"], + finish_type=( + finish_reason["type"] if finish_reason is not None else None + ), + now=time.perf_counter(), ) - if data["meta_info"]["completion_tokens"] == 1: - last_ttft = time.perf_counter() - tic + missing = recorder.missing_indices() + if missing: + raise RuntimeError( + f"No stream chunk received for request indices {missing}." + ) + last_ttft = recorder.all_started_time - tic # Compute metrics latency = time.perf_counter() - tic - input_throughput = batch_size * input_len / last_ttft - output_throughput = batch_size * output_len / (latency - last_ttft) - overall_throughput = batch_size * (input_len + output_len) / latency + if recorder is None: + # vllm chunks carry no token counts; assume the requested lengths. + total_output_tokens = batch_size * output_len + decode_window_tokens = float(total_output_tokens) + num_early_stopped = None + steady_state_window: Optional[SteadyStateWindow] = None + else: + total_output_tokens = recorder.total_output_tokens + num_early_stopped = recorder.num_early_stopped + steady_state_window = recorder.steady_state_window() + if ignore_eos and dataset_name not in REPLAY_TEXT_DATASETS: + # Historical numerator; CI throughput floors are calibrated to it. + decode_window_tokens = float(total_output_tokens) + else: + # Count only tokens generated inside the denominator's window. + decode_window_tokens = ( + total_output_tokens - recorder.tokens_before_all_started + ) + steady_state_output_throughput = ( + steady_state_window.output_throughput + if steady_state_window is not None + else None + ) + if ( + ignore_eos + and recorder is not None + and total_output_tokens != batch_size * output_len + ): + print( + f"WARNING: generated {total_output_tokens} tokens but requested " + f"{batch_size * output_len}; the server clamped max_new_tokens or " + f"served fewer prompts, so throughput uses the actual count." + ) + + if dataset_name in REPLAY_TEXT_DATASETS: + total_input_tokens = sum(len(ids) for ids in input_ids) + else: + total_input_tokens = batch_size * input_len + input_throughput = total_input_tokens / last_ttft + output_throughput = decode_window_tokens / (latency - last_ttft) + overall_throughput = (total_input_tokens + total_output_tokens) / latency if backend == "vllm": # vLLM does not expose these metrics via API @@ -814,6 +942,24 @@ def run_one_case( print(f"acc_length: {acc_length:.2f} ") if metrics_cache_hit_rate is not None: print(f"cache hit rate: {metrics_cache_hit_rate:.4f}") + if not ignore_eos and recorder is not None: + print( + f"total output tokens: {total_output_tokens} " + f"(requested {batch_size * output_len})" + ) + print(f"early-stopped requests: {num_early_stopped}/{batch_size}") + if steady_state_window is not None: + print( + f"steady-state output throughput: " + f"{steady_state_output_throughput:.2f} tok/s " + f"over a {steady_state_window.duration:.2f} s full-batch window" + ) + else: + print( + "WARNING: no steady-state full-batch window (a request finished " + "before the last one started decoding); steady-state metrics " + "are n/a." + ) # Dump results result = BenchOneCaseResult( @@ -828,6 +974,12 @@ def run_one_case( last_ttft=last_ttft, last_gen_throughput=last_gen_throughput, acc_length=acc_length, + total_output_tokens=total_output_tokens, + num_early_stopped=num_early_stopped, + steady_state_output_throughput=steady_state_output_throughput, + steady_state_window_s=( + steady_state_window.duration if steady_state_window is not None else None + ), cache_hit_rate=metrics_cache_hit_rate, profile_link=profile_link, ) @@ -896,15 +1048,25 @@ def get_report_summary( "output cost ($/1M)", "cache hit rate", ] + if bench_args.disable_ignore_eos: + headers += [ + "steady output throughput (tok/s)", + "steady ITL (ms)", + "early stops", + ] if bench_args.profile: headers.append("profile") for res in results: hourly_cost = hourly_cost_per_gpu * server_args.tp_size accept_length = f"{res.acc_length:.2f}" if res.acc_length > 0 else "n/a" - itl_ms = 1000 * res.batch_size / res.output_throughput + # 0 when every request's first chunk is also its finish chunk. + if res.output_throughput > 0: + itl_ms = f"{1000 * res.batch_size / res.output_throughput:.2f}" + output_cost = f"{1e6 / res.output_throughput / 3600 * hourly_cost:.2f}" + else: + itl_ms = output_cost = "n/a" input_cost = 1e6 / (res.input_throughput * input_util) / 3600 * hourly_cost - output_cost = 1e6 / res.output_throughput / 3600 * hourly_cost cache_hit_rate = ( f"{res.cache_hit_rate:.4f}" if res.cache_hit_rate is not None else "n/a" ) @@ -916,11 +1078,24 @@ def get_report_summary( f"{res.input_throughput:.2f}", f"{res.output_throughput:.2f}", accept_length, - f"{itl_ms:.2f}", + itl_ms, f"{input_cost:.2f}", - f"{output_cost:.2f}", + output_cost, cache_hit_rate, ] + if bench_args.disable_ignore_eos: + if res.steady_state_output_throughput is not None: + steady_tput = f"{res.steady_state_output_throughput:.2f}" + steady_itl = ( + f"{1000 * res.batch_size / res.steady_state_output_throughput:.2f}" + ) + else: + steady_tput = steady_itl = "n/a" + row += [ + steady_tput, + steady_itl, + f"{res.num_early_stopped}/{res.batch_size}", + ] if bench_args.profile: if res.profile_link: row.append(f"[Profile]({res.profile_link})") @@ -939,6 +1114,58 @@ def run_benchmark_internal( bench_args: BenchArgs, launch_server_func: Callable = launch_server, ): + # Validate flags that depend only on bench_args before launching a server. + if bench_args.disable_ignore_eos: + if bench_args.backend == "vllm": + raise ValueError( + "--disable-ignore-eos requires the sglang backend: vllm stream " + "chunks carry no cumulative token counts, so early-stop-aware " + "accounting is impossible." + ) + if bench_args.fake_prefill: + raise ValueError( + "--disable-ignore-eos is incompatible with --fake-prefill: " + "decode conditions on uninitialized KV, so EOS timing is " + "meaningless." + ) + if bench_args.client_stream_interval != 1: + raise ValueError( + "--disable-ignore-eos requires --client-stream-interval 1: the " + "steady-state window needs per-token cumulative counts, and " + "larger intervals quantize them by up to interval-1 tokens per " + "request." + ) + + if bench_args.dataset_name in REPLAY_TEXT_DATASETS: + if bench_args.backend == "vllm": + raise ValueError( + "Replay datasets require the sglang backend: vllm stream " + "chunks carry no token counts, so the decode-window token " + "accounting replay throughput relies on is impossible." + ) + if bench_args.fixed_prompt_file: + raise ValueError( + "--fixed-prompt-file bypasses --dataset-name; drop one of them." + ) + if bench_args.dataset_name == "custom" and not os.path.isfile( + bench_args.dataset_path + ): + raise ValueError( + "--dataset-name custom requires --dataset-path pointing at a " + f"JSONL conversation file; got {bench_args.dataset_path!r}." + ) + if len(bench_args.input_len) > 1: + raise ValueError( + "Replay datasets take prompt lengths from the recorded data; " + "sweeping --input-len is meaningless. Pass at most one value " + "(used only by the token-capacity skip guard)." + ) + print( + "NOTE: replay prompts have recorded lengths; the token-capacity " + f"skip guard assumes --input-len ({bench_args.input_len[0]}) per " + "prompt, so set it near the trace's mean length." + ) + # set random seed random.seed(bench_args.seed) np.random.seed(bench_args.seed) @@ -1065,11 +1292,14 @@ def run_benchmark_internal( bench_args.lora_zipf_alpha > 1 ), f"--lora-zipf-alpha must be > 1, got {bench_args.lora_zipf_alpha}" - if bench_args.apply_chat_template and not bench_args.fixed_prompt_file: + if bench_args.apply_chat_template and not ( + bench_args.fixed_prompt_file or bench_args.dataset_name in REPLAY_TEXT_DATASETS + ): raise ValueError( - "--apply-chat-template requires --fixed-prompt-file: the other " - "datasets generate token ids directly, so there is no prompt text " - "to run through a chat template." + "--apply-chat-template requires --fixed-prompt-file or a replay " + "text dataset (sharegpt/custom): the other datasets generate " + "token ids directly, so there is no prompt text to run through a " + "chat template." ) gsp_kwargs = dict( @@ -1091,6 +1321,7 @@ def run_benchmark_internal( input_len=1024, output_len=16, temperature=bench_args.temperature, + ignore_eos=True, return_logprob=bench_args.return_logprob, stream_interval=bench_args.client_stream_interval, input_len_step_percentage=bench_args.input_len_step_percentage, @@ -1135,6 +1366,7 @@ def run_benchmark_internal( il, ol, temperature=bench_args.temperature, + ignore_eos=not bench_args.disable_ignore_eos, return_logprob=bench_args.return_logprob, stream_interval=bench_args.client_stream_interval, input_len_step_percentage=bench_args.input_len_step_percentage, @@ -1184,6 +1416,7 @@ def run_benchmark_internal( il, ol, temperature=bench_args.temperature, + ignore_eos=not bench_args.disable_ignore_eos, return_logprob=bench_args.return_logprob, stream_interval=bench_args.client_stream_interval, input_len_step_percentage=bench_args.input_len_step_percentage, @@ -1207,6 +1440,8 @@ def run_benchmark_internal( lora_name=bench_args.lora_name, lora_request_distribution=bench_args.lora_request_distribution, lora_zipf_alpha=bench_args.lora_zipf_alpha, + fixed_prompt_file=bench_args.fixed_prompt_file, + apply_chat_template=bench_args.apply_chat_template, **gsp_kwargs, ) ) diff --git a/python/sglang/benchmark/stream_metrics.py b/python/sglang/benchmark/stream_metrics.py new file mode 100644 index 000000000..9408430be --- /dev/null +++ b/python/sglang/benchmark/stream_metrics.py @@ -0,0 +1,152 @@ +"""Early-stop-aware accounting for one-batch streaming benchmarks. + +Steady-state window = last request's first token -> first request's finish. +Token counts come from the cumulative meta_info["completion_tokens"] (the +server may coalesce chunks); boundary counts are interpolated per request so +same-step chunk delivery order cannot skew the window. +""" + +from typing import List, Optional, Tuple + +import msgspec + +# (arrival time, cumulative completion tokens) of one delivered chunk. +_Obs = Tuple[float, int] + + +class SteadyStateWindow(msgspec.Struct, frozen=True): + """Full-batch decode window: every request started, none finished yet.""" + + start: float + end: float + output_tokens: float + + @property + def duration(self) -> float: + return self.end - self.start + + @property + def output_throughput(self) -> float: + return self.output_tokens / self.duration + + +def validate_finish_reason(finish_reason: dict, *, ignore_eos: bool) -> None: + """Raise on any finish reason the benchmark must not silently accept.""" + finish_type = finish_reason["type"] + if finish_type == "length": + return + if finish_type == "stop": + # The scheduler's NaN detector reports as a "stop" match; never accept it. + if finish_reason["matched"] == "NaN happened": + raise RuntimeError(f"Request hit NaN logits: {finish_reason}.") + if not ignore_eos: + return + raise RuntimeError( + f"Request stopped early despite ignore_eos=True: {finish_reason}." + ) + raise RuntimeError(f"Unexpected finish reason: {finish_reason}.") + + +class _Boundary: + """Per-request observations bracketing one boundary instant.""" + + def __init__(self, time: float, before: List[Optional[_Obs]]): + self.time = time + self.before = before + self.after: List[Optional[_Obs]] = [None] * len(before) + + def observe(self, index: int, obs: _Obs) -> None: + if self.after[index] is None: + self.after[index] = obs + + def tokens_at_boundary(self, index: int) -> float: + """Tokens of `index` at `time`, interpolated between its chunks.""" + before = self.before[index] + if before is None: + return 0.0 + t_before, c_before = before + after = self.after[index] + if after is None or self.time <= t_before: + return float(c_before) + t_after, c_after = after + if self.time >= t_after or t_after <= t_before: + return float(c_after) + return c_before + (c_after - c_before) * (self.time - t_before) / ( + t_after - t_before + ) + + def total_tokens(self) -> float: + return sum(self.tokens_at_boundary(i) for i in range(len(self.before))) + + +class BatchStreamRecorder: + """Tracks per-request progress of one batched streaming /generate call.""" + + def __init__(self, batch_size: int): + self.batch_size = batch_size + self.all_started_time: Optional[float] = None + self._last_obs: List[Optional[_Obs]] = [None] * batch_size + self._first_token_time: List[Optional[float]] = [None] * batch_size + self._completion_tokens: List[int] = [0] * batch_size + self._finish_types: List[Optional[str]] = [None] * batch_size + self._num_started = 0 + self._t0: Optional[_Boundary] = None + self._t1: Optional[_Boundary] = None + + def record_chunk( + self, + *, + index: int, + completion_tokens: int, + finish_type: Optional[str], + now: float, + ) -> None: + obs = (now, completion_tokens) + for boundary in (self._t0, self._t1): + if boundary is not None: + boundary.observe(index, obs) + self._last_obs[index] = obs + self._completion_tokens[index] = completion_tokens + if self._first_token_time[index] is None and completion_tokens > 0: + self._first_token_time[index] = now + self._num_started += 1 + if self._num_started == self.batch_size: + self.all_started_time = now + self._t0 = _Boundary(time=now, before=list(self._last_obs)) + if finish_type is not None and self._finish_types[index] is None: + self._finish_types[index] = finish_type + if self._t1 is None: + self._t1 = _Boundary(time=now, before=list(self._last_obs)) + + def missing_indices(self) -> List[int]: + return [i for i, t in enumerate(self._first_token_time) if t is None] + + @property + def total_output_tokens(self) -> int: + return sum(self._completion_tokens) + + @property + def num_early_stopped(self) -> int: + return sum(1 for t in self._finish_types if t == "stop") + + @property + def tokens_before_all_started(self) -> Optional[float]: + """Batch tokens at the last request's first token; None before then.""" + if self._t0 is None: + return None + return self._t0.total_tokens() + + def steady_state_window(self) -> Optional[SteadyStateWindow]: + """None when the batch never decoded at full size for a nonzero span.""" + if self._t0 is None or self._t1 is None: + return None + if self._t1.time <= self._t0.time: + return None + output_tokens = self._t1.total_tokens() - self._t0.total_tokens() + if output_tokens <= 0: + return None + return SteadyStateWindow( + start=self._t0.time, + end=self._t1.time, + output_tokens=output_tokens, + )