[bench] Support real-traffic replay with early-stop-aware steady-state metrics in bench_one_batch_server (#37469)
This commit is contained in:
@@ -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 `--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 <conversations.jsonl>` (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 <name>` to route every prompt through a pre-loaded LoRA adapter. Requires the server to be launched with `--enable-lora --lora-paths <name>=<path>`.
|
- Pass `--lora-name <name>` to route every prompt through a pre-loaded LoRA adapter. Requires the server to be launched with `--enable-lora --lora-paths <name>=<path>`.
|
||||||
|
|
||||||
**`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.
|
**`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.
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import argparse
|
|||||||
import dataclasses
|
import dataclasses
|
||||||
import itertools
|
import itertools
|
||||||
import json
|
import json
|
||||||
|
import os
|
||||||
import random
|
import random
|
||||||
import re
|
import re
|
||||||
import time
|
import time
|
||||||
@@ -31,6 +32,11 @@ from transformers import AutoProcessor, PreTrainedTokenizer
|
|||||||
|
|
||||||
from sglang.benchmark.datasets import get_dataset
|
from sglang.benchmark.datasets import get_dataset
|
||||||
from sglang.benchmark.endpoint import acquire_endpoint
|
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.benchmark.utils import get_processor, get_tokenizer
|
||||||
from sglang.profiler import run_profile
|
from sglang.profiler import run_profile
|
||||||
from sglang.srt.arg_groups.overrides import resolving_view
|
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
|
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]:
|
def get_cache_tokens_from_metrics(url: str) -> Optional[tuple]:
|
||||||
"""
|
"""
|
||||||
@@ -107,6 +117,7 @@ class BenchArgs:
|
|||||||
input_len: Tuple[int] = (1024,)
|
input_len: Tuple[int] = (1024,)
|
||||||
output_len: Tuple[int] = (16,)
|
output_len: Tuple[int] = (16,)
|
||||||
temperature: float = 0.0
|
temperature: float = 0.0
|
||||||
|
disable_ignore_eos: bool = False
|
||||||
return_logprob: bool = False
|
return_logprob: bool = False
|
||||||
client_stream_interval: int = 1
|
client_stream_interval: int = 1
|
||||||
input_len_step_percentage: float = 0.0
|
input_len_step_percentage: float = 0.0
|
||||||
@@ -156,6 +167,23 @@ class BenchArgs:
|
|||||||
"--output-len", type=int, nargs="+", default=BenchArgs.output_len
|
"--output-len", type=int, nargs="+", default=BenchArgs.output_len
|
||||||
)
|
)
|
||||||
parser.add_argument("--temperature", type=float, default=BenchArgs.temperature)
|
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("--return-logprob", action="store_true")
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--client-stream-interval",
|
"--client-stream-interval",
|
||||||
@@ -219,8 +247,19 @@ class BenchArgs:
|
|||||||
"--dataset-name",
|
"--dataset-name",
|
||||||
type=str,
|
type=str,
|
||||||
default=BenchArgs.dataset_name,
|
default=BenchArgs.dataset_name,
|
||||||
choices=["mmmu", "random", "random-ids", "generated-shared-prefix"],
|
choices=[
|
||||||
help="Name of the dataset to benchmark on.",
|
"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(
|
parser.add_argument(
|
||||||
"--fixed-prompt-file",
|
"--fixed-prompt-file",
|
||||||
@@ -232,8 +271,9 @@ class BenchArgs:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--apply-chat-template",
|
"--apply-chat-template",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
help="Encode the prompt as a single user message through the "
|
help="Encode each prompt as a single user message through the "
|
||||||
"model's chat template. Requires --fixed-prompt-file.",
|
"model's chat template. Requires --fixed-prompt-file or a replay "
|
||||||
|
"text dataset (sharegpt/custom).",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--gsp-num-groups",
|
"--gsp-num-groups",
|
||||||
@@ -377,6 +417,10 @@ class BenchOneCaseResult(BaseModel):
|
|||||||
last_ttft: float
|
last_ttft: float
|
||||||
last_gen_throughput: float
|
last_gen_throughput: float
|
||||||
acc_length: 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
|
cache_hit_rate: Optional[float] = None
|
||||||
profile_link: Optional[str] = None
|
profile_link: Optional[str] = None
|
||||||
|
|
||||||
@@ -394,6 +438,18 @@ class BenchOneCaseResult(BaseModel):
|
|||||||
"last_ttft": round(self.last_ttft, 4),
|
"last_ttft": round(self.last_ttft, 4),
|
||||||
"last_gen_throughput": round(self.last_gen_throughput, 2),
|
"last_gen_throughput": round(self.last_gen_throughput, 2),
|
||||||
"acc_length": round(self.acc_length, 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": (
|
"cache_hit_rate": (
|
||||||
round(self.cache_hit_rate, 4)
|
round(self.cache_hit_rate, 4)
|
||||||
if self.cache_hit_rate is not None
|
if self.cache_hit_rate is not None
|
||||||
@@ -531,6 +587,7 @@ def run_one_case(
|
|||||||
input_len: int,
|
input_len: int,
|
||||||
output_len: int,
|
output_len: int,
|
||||||
temperature: float,
|
temperature: float,
|
||||||
|
ignore_eos: bool,
|
||||||
return_logprob: bool,
|
return_logprob: bool,
|
||||||
stream_interval: int,
|
stream_interval: int,
|
||||||
input_len_step_percentage: float,
|
input_len_step_percentage: float,
|
||||||
@@ -576,7 +633,12 @@ def run_one_case(
|
|||||||
image_data = None
|
image_data = None
|
||||||
else:
|
else:
|
||||||
# Load input token ids via benchmark.datasets.get_dataset
|
# 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:
|
if dataset_name not in supported_datasets:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported dataset for batch benchmark: {dataset_name}. "
|
f"Unsupported dataset for batch benchmark: {dataset_name}. "
|
||||||
@@ -591,8 +653,12 @@ def run_one_case(
|
|||||||
random_output_len=output_len,
|
random_output_len=output_len,
|
||||||
random_range_ratio=1.0,
|
random_range_ratio=1.0,
|
||||||
dataset_path=dataset_path,
|
dataset_path=dataset_path,
|
||||||
tokenize_prompt=dataset_name not in ("mmmu", "generated-shared-prefix"),
|
tokenize_prompt=dataset_name in ("random", "random-ids"),
|
||||||
backend=backend,
|
backend=backend,
|
||||||
|
sharegpt_output_len=None,
|
||||||
|
sharegpt_context_len=None,
|
||||||
|
prompt_suffix="",
|
||||||
|
apply_chat_template=apply_chat_template,
|
||||||
seed=BenchArgs.seed,
|
seed=BenchArgs.seed,
|
||||||
gsp_num_groups=actual_gsp_groups,
|
gsp_num_groups=actual_gsp_groups,
|
||||||
gsp_prompts_per_group=(batch_size + actual_gsp_groups - 1)
|
gsp_prompts_per_group=(batch_size + actual_gsp_groups - 1)
|
||||||
@@ -618,6 +684,16 @@ def run_one_case(
|
|||||||
elif dataset_name == "mmmu":
|
elif dataset_name == "mmmu":
|
||||||
input_ids = [tok_inner.encode(req.prompt) for req in input_requests]
|
input_ids = [tok_inner.encode(req.prompt) for req in input_requests]
|
||||||
image_data = [req.image_data 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:
|
else:
|
||||||
input_ids = [req.prompt for req in input_requests]
|
input_ids = [req.prompt for req in input_requests]
|
||||||
image_data = None
|
image_data = None
|
||||||
@@ -630,7 +706,7 @@ def run_one_case(
|
|||||||
"max_tokens": output_len,
|
"max_tokens": output_len,
|
||||||
"temperature": temperature,
|
"temperature": temperature,
|
||||||
"stream": True,
|
"stream": True,
|
||||||
"ignore_eos": True,
|
"ignore_eos": ignore_eos,
|
||||||
}
|
}
|
||||||
if return_logprob:
|
if return_logprob:
|
||||||
payload["logprobs"] = 1
|
payload["logprobs"] = 1
|
||||||
@@ -654,7 +730,7 @@ def run_one_case(
|
|||||||
"sampling_params": {
|
"sampling_params": {
|
||||||
"temperature": temperature,
|
"temperature": temperature,
|
||||||
"max_new_tokens": output_len,
|
"max_new_tokens": output_len,
|
||||||
"ignore_eos": True,
|
"ignore_eos": ignore_eos,
|
||||||
"json_schema": json_schema,
|
"json_schema": json_schema,
|
||||||
"stream_interval": stream_interval,
|
"stream_interval": stream_interval,
|
||||||
},
|
},
|
||||||
@@ -725,6 +801,7 @@ def run_one_case(
|
|||||||
metrics_before = get_cache_tokens_from_metrics(url)
|
metrics_before = get_cache_tokens_from_metrics(url)
|
||||||
|
|
||||||
# Run the request
|
# Run the request
|
||||||
|
recorder: Optional[BatchStreamRecorder] = None
|
||||||
tic = time.perf_counter()
|
tic = time.perf_counter()
|
||||||
with requests.post(
|
with requests.post(
|
||||||
gen_url,
|
gen_url,
|
||||||
@@ -755,6 +832,8 @@ def run_one_case(
|
|||||||
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:
|
||||||
|
# mmmu can silently return fewer prompts than requested.
|
||||||
|
recorder = BatchStreamRecorder(batch_size=len(input_ids))
|
||||||
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:"):
|
||||||
@@ -764,18 +843,67 @@ def run_one_case(
|
|||||||
if "error" in data:
|
if "error" in data:
|
||||||
raise RuntimeError(f"Request has failed. {data}.")
|
raise RuntimeError(f"Request has failed. {data}.")
|
||||||
|
|
||||||
assert (
|
finish_reason = data["meta_info"]["finish_reason"]
|
||||||
data["meta_info"]["finish_reason"] is None
|
if finish_reason is not None:
|
||||||
or data["meta_info"]["finish_reason"]["type"] == "length"
|
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:
|
missing = recorder.missing_indices()
|
||||||
last_ttft = time.perf_counter() - tic
|
if missing:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"No stream chunk received for request indices {missing}."
|
||||||
|
)
|
||||||
|
last_ttft = recorder.all_started_time - tic
|
||||||
|
|
||||||
# Compute metrics
|
# Compute metrics
|
||||||
latency = time.perf_counter() - tic
|
latency = time.perf_counter() - tic
|
||||||
input_throughput = batch_size * input_len / last_ttft
|
if recorder is None:
|
||||||
output_throughput = batch_size * output_len / (latency - last_ttft)
|
# vllm chunks carry no token counts; assume the requested lengths.
|
||||||
overall_throughput = batch_size * (input_len + output_len) / latency
|
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":
|
if backend == "vllm":
|
||||||
# vLLM does not expose these metrics via API
|
# vLLM does not expose these metrics via API
|
||||||
@@ -814,6 +942,24 @@ def run_one_case(
|
|||||||
print(f"acc_length: {acc_length:.2f} ")
|
print(f"acc_length: {acc_length:.2f} ")
|
||||||
if metrics_cache_hit_rate is not None:
|
if metrics_cache_hit_rate is not None:
|
||||||
print(f"cache hit rate: {metrics_cache_hit_rate:.4f}")
|
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
|
# Dump results
|
||||||
result = BenchOneCaseResult(
|
result = BenchOneCaseResult(
|
||||||
@@ -828,6 +974,12 @@ def run_one_case(
|
|||||||
last_ttft=last_ttft,
|
last_ttft=last_ttft,
|
||||||
last_gen_throughput=last_gen_throughput,
|
last_gen_throughput=last_gen_throughput,
|
||||||
acc_length=acc_length,
|
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,
|
cache_hit_rate=metrics_cache_hit_rate,
|
||||||
profile_link=profile_link,
|
profile_link=profile_link,
|
||||||
)
|
)
|
||||||
@@ -896,15 +1048,25 @@ def get_report_summary(
|
|||||||
"output cost ($/1M)",
|
"output cost ($/1M)",
|
||||||
"cache hit rate",
|
"cache hit rate",
|
||||||
]
|
]
|
||||||
|
if bench_args.disable_ignore_eos:
|
||||||
|
headers += [
|
||||||
|
"steady output throughput (tok/s)",
|
||||||
|
"steady ITL (ms)",
|
||||||
|
"early stops",
|
||||||
|
]
|
||||||
if bench_args.profile:
|
if bench_args.profile:
|
||||||
headers.append("profile")
|
headers.append("profile")
|
||||||
|
|
||||||
for res in results:
|
for res in results:
|
||||||
hourly_cost = hourly_cost_per_gpu * server_args.tp_size
|
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"
|
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
|
input_cost = 1e6 / (res.input_throughput * input_util) / 3600 * hourly_cost
|
||||||
output_cost = 1e6 / res.output_throughput / 3600 * hourly_cost
|
|
||||||
cache_hit_rate = (
|
cache_hit_rate = (
|
||||||
f"{res.cache_hit_rate:.4f}" if res.cache_hit_rate is not None else "n/a"
|
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.input_throughput:.2f}",
|
||||||
f"{res.output_throughput:.2f}",
|
f"{res.output_throughput:.2f}",
|
||||||
accept_length,
|
accept_length,
|
||||||
f"{itl_ms:.2f}",
|
itl_ms,
|
||||||
f"{input_cost:.2f}",
|
f"{input_cost:.2f}",
|
||||||
f"{output_cost:.2f}",
|
output_cost,
|
||||||
cache_hit_rate,
|
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 bench_args.profile:
|
||||||
if res.profile_link:
|
if res.profile_link:
|
||||||
row.append(f"[Profile]({res.profile_link})")
|
row.append(f"[Profile]({res.profile_link})")
|
||||||
@@ -939,6 +1114,58 @@ def run_benchmark_internal(
|
|||||||
bench_args: BenchArgs,
|
bench_args: BenchArgs,
|
||||||
launch_server_func: Callable = launch_server,
|
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
|
# set random seed
|
||||||
random.seed(bench_args.seed)
|
random.seed(bench_args.seed)
|
||||||
np.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
|
bench_args.lora_zipf_alpha > 1
|
||||||
), f"--lora-zipf-alpha must be > 1, got {bench_args.lora_zipf_alpha}"
|
), 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(
|
raise ValueError(
|
||||||
"--apply-chat-template requires --fixed-prompt-file: the other "
|
"--apply-chat-template requires --fixed-prompt-file or a replay "
|
||||||
"datasets generate token ids directly, so there is no prompt text "
|
"text dataset (sharegpt/custom): the other datasets generate "
|
||||||
"to run through a chat template."
|
"token ids directly, so there is no prompt text to run through a "
|
||||||
|
"chat template."
|
||||||
)
|
)
|
||||||
|
|
||||||
gsp_kwargs = dict(
|
gsp_kwargs = dict(
|
||||||
@@ -1091,6 +1321,7 @@ def run_benchmark_internal(
|
|||||||
input_len=1024,
|
input_len=1024,
|
||||||
output_len=16,
|
output_len=16,
|
||||||
temperature=bench_args.temperature,
|
temperature=bench_args.temperature,
|
||||||
|
ignore_eos=True,
|
||||||
return_logprob=bench_args.return_logprob,
|
return_logprob=bench_args.return_logprob,
|
||||||
stream_interval=bench_args.client_stream_interval,
|
stream_interval=bench_args.client_stream_interval,
|
||||||
input_len_step_percentage=bench_args.input_len_step_percentage,
|
input_len_step_percentage=bench_args.input_len_step_percentage,
|
||||||
@@ -1135,6 +1366,7 @@ def run_benchmark_internal(
|
|||||||
il,
|
il,
|
||||||
ol,
|
ol,
|
||||||
temperature=bench_args.temperature,
|
temperature=bench_args.temperature,
|
||||||
|
ignore_eos=not bench_args.disable_ignore_eos,
|
||||||
return_logprob=bench_args.return_logprob,
|
return_logprob=bench_args.return_logprob,
|
||||||
stream_interval=bench_args.client_stream_interval,
|
stream_interval=bench_args.client_stream_interval,
|
||||||
input_len_step_percentage=bench_args.input_len_step_percentage,
|
input_len_step_percentage=bench_args.input_len_step_percentage,
|
||||||
@@ -1184,6 +1416,7 @@ def run_benchmark_internal(
|
|||||||
il,
|
il,
|
||||||
ol,
|
ol,
|
||||||
temperature=bench_args.temperature,
|
temperature=bench_args.temperature,
|
||||||
|
ignore_eos=not bench_args.disable_ignore_eos,
|
||||||
return_logprob=bench_args.return_logprob,
|
return_logprob=bench_args.return_logprob,
|
||||||
stream_interval=bench_args.client_stream_interval,
|
stream_interval=bench_args.client_stream_interval,
|
||||||
input_len_step_percentage=bench_args.input_len_step_percentage,
|
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_name=bench_args.lora_name,
|
||||||
lora_request_distribution=bench_args.lora_request_distribution,
|
lora_request_distribution=bench_args.lora_request_distribution,
|
||||||
lora_zipf_alpha=bench_args.lora_zipf_alpha,
|
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,
|
**gsp_kwargs,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user