[bench] Support real-traffic replay with early-stop-aware steady-state metrics in bench_one_batch_server (#37469)
This commit is contained in:
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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