foretoken 0.0.1a1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- benchmarks/__init__.py +2 -0
- benchmarks/arguments.py +444 -0
- benchmarks/client/__init__.py +2 -0
- benchmarks/client/openai_client.py +148 -0
- benchmarks/config.py +358 -0
- benchmarks/deployment/__init__.py +9 -0
- benchmarks/deployment/discovery.py +166 -0
- benchmarks/deployment/lifecycle.py +111 -0
- benchmarks/logger/__init__.py +5 -0
- benchmarks/logger/cli.py +27 -0
- benchmarks/logger/wandb.py +261 -0
- benchmarks/main.py +74 -0
- benchmarks/metrics/__init__.py +2 -0
- benchmarks/metrics/aggregator.py +161 -0
- benchmarks/report/__init__.py +2 -0
- benchmarks/report/pareto.py +145 -0
- benchmarks/report/summary.py +156 -0
- benchmarks/runner/__init__.py +2 -0
- benchmarks/runner/base.py +284 -0
- benchmarks/runner/multi_dataset.py +105 -0
- benchmarks/runner/run_benchmark.py +79 -0
- benchmarks/runner/run_spec.py +22 -0
- benchmarks/runner/select_runner.py +27 -0
- benchmarks/runner/sweep.py +141 -0
- benchmarks/runner/trace_runner.py +225 -0
- benchmarks/storage/__init__.py +2 -0
- benchmarks/storage/result_writer.py +37 -0
- benchmarks/utils/__init__.py +2 -0
- benchmarks/utils/bench_params.py +223 -0
- benchmarks/workload/__init__.py +2 -0
- benchmarks/workload/hf_dataset.py +154 -0
- benchmarks/workload/loader.py +205 -0
- benchmarks/workload/random_dataset.py +273 -0
- benchmarks/workload/trace_loader.py +275 -0
- benchmarks/workload/trace_workload.py +109 -0
- foretoken/__init__.py +26 -0
- foretoken/accelerators/__init__.py +4 -0
- foretoken/accelerators/_exporter.py +102 -0
- foretoken/accelerators/discovery.py +264 -0
- foretoken/accelerators/metax.py +61 -0
- foretoken/accelerators/nvidia.py +184 -0
- foretoken/arguments.py +297 -0
- foretoken/kubernetes.py +664 -0
- foretoken/main.py +188 -0
- foretoken/manifest.py +132 -0
- foretoken/observability.py +349 -0
- foretoken/platform/__init__.py +8 -0
- foretoken/platform/config.py +119 -0
- foretoken/platform/gateway.py +357 -0
- foretoken/platform/gateway_resources.py +192 -0
- foretoken/platform/helm.py +422 -0
- foretoken/platform/helm_client.py +216 -0
- foretoken/platform/lifecycle.py +345 -0
- foretoken/platform/types.py +33 -0
- foretoken/source.py +171 -0
- foretoken-0.0.1a1.dist-info/METADATA +196 -0
- foretoken-0.0.1a1.dist-info/RECORD +61 -0
- foretoken-0.0.1a1.dist-info/WHEEL +5 -0
- foretoken-0.0.1a1.dist-info/entry_points.txt +2 -0
- foretoken-0.0.1a1.dist-info/licenses/LICENSE +201 -0
- foretoken-0.0.1a1.dist-info/top_level.txt +2 -0
benchmarks/__init__.py
ADDED
benchmarks/arguments.py
ADDED
|
@@ -0,0 +1,444 @@
|
|
|
1
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
2
|
+
# SPDX-FileCopyrightText: Copyright contributors to the Foretoken project
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
"""CLI commands and benchmark argument mapping."""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import argparse
|
|
10
|
+
import json
|
|
11
|
+
from collections.abc import Sequence
|
|
12
|
+
from dataclasses import MISSING, dataclass, fields
|
|
13
|
+
from typing import Any
|
|
14
|
+
|
|
15
|
+
from evalscope.perf.arguments import Arguments
|
|
16
|
+
|
|
17
|
+
from benchmarks.config import (
|
|
18
|
+
BenchConfig,
|
|
19
|
+
DatasetConfig,
|
|
20
|
+
GenerationConfig,
|
|
21
|
+
LoadConfig,
|
|
22
|
+
OutputConfig,
|
|
23
|
+
ParamSweepConfig,
|
|
24
|
+
EndpointConfig,
|
|
25
|
+
WandbConfig,
|
|
26
|
+
)
|
|
27
|
+
@dataclass(frozen=True)
|
|
28
|
+
class BenchCommand:
|
|
29
|
+
"""Run a benchmark against a deployment or existing endpoint."""
|
|
30
|
+
|
|
31
|
+
kustomize_path: str
|
|
32
|
+
config: BenchConfig
|
|
33
|
+
wait_timeout: str
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _default(cls: type, name: str) -> Any:
|
|
37
|
+
field_info = next(item for item in fields(cls) if item.name == name)
|
|
38
|
+
if field_info.default_factory is not MISSING:
|
|
39
|
+
return field_info.default_factory()
|
|
40
|
+
if field_info.default is not MISSING:
|
|
41
|
+
return field_info.default
|
|
42
|
+
raise KeyError(name)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def _output_destinations(value: str) -> tuple[str, ...]:
|
|
46
|
+
return tuple(item.strip() for item in value.split(",") if item.strip())
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _json_object(value: str) -> dict[str, Any]:
|
|
50
|
+
"""Parse one CLI JSON object without accepting other JSON value types."""
|
|
51
|
+
parsed = json.loads(value)
|
|
52
|
+
if not isinstance(parsed, dict):
|
|
53
|
+
raise argparse.ArgumentTypeError("must be a JSON object")
|
|
54
|
+
return parsed
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _add_benchmark_arguments(parser: argparse.ArgumentParser) -> None:
|
|
58
|
+
# Service source
|
|
59
|
+
parser.add_argument(
|
|
60
|
+
"kustomize_path",
|
|
61
|
+
nargs="?",
|
|
62
|
+
metavar="PATH",
|
|
63
|
+
help="Kustomize directory to deploy or reuse",
|
|
64
|
+
)
|
|
65
|
+
parser.add_argument(
|
|
66
|
+
"--url",
|
|
67
|
+
default="",
|
|
68
|
+
help="Existing OpenAI-compatible chat-completions URL",
|
|
69
|
+
)
|
|
70
|
+
parser.add_argument(
|
|
71
|
+
"--model",
|
|
72
|
+
default="",
|
|
73
|
+
help="Model name; inferred when the deployment contains one model",
|
|
74
|
+
)
|
|
75
|
+
parser.add_argument(
|
|
76
|
+
"--api-key", default=_default(EndpointConfig, "api_key"), help="API key"
|
|
77
|
+
)
|
|
78
|
+
parser.add_argument(
|
|
79
|
+
"--timeout",
|
|
80
|
+
type=int,
|
|
81
|
+
default=_default(EndpointConfig, "timeout"),
|
|
82
|
+
help="Request timeout seconds",
|
|
83
|
+
)
|
|
84
|
+
parser.add_argument(
|
|
85
|
+
"--max-retries",
|
|
86
|
+
type=int,
|
|
87
|
+
default=_default(EndpointConfig, "max_retries"),
|
|
88
|
+
help="OpenAI client max retries on transient failures",
|
|
89
|
+
)
|
|
90
|
+
parser.add_argument(
|
|
91
|
+
"--wait-timeout",
|
|
92
|
+
default="15m",
|
|
93
|
+
help="Timeout for each deployment readiness stage",
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
# Load
|
|
97
|
+
parser.add_argument(
|
|
98
|
+
"--parallel",
|
|
99
|
+
type=int,
|
|
100
|
+
default=_default(LoadConfig, "parallel"),
|
|
101
|
+
help="Maximum concurrent requests; ignored with --open-loop",
|
|
102
|
+
)
|
|
103
|
+
parser.add_argument(
|
|
104
|
+
"--number",
|
|
105
|
+
type=int,
|
|
106
|
+
default=_default(LoadConfig, "number"),
|
|
107
|
+
help="Requests per run; total across multiple dataset sources",
|
|
108
|
+
)
|
|
109
|
+
parser.add_argument(
|
|
110
|
+
"--rate",
|
|
111
|
+
type=float,
|
|
112
|
+
default=_default(LoadConfig, "rate"),
|
|
113
|
+
help=(
|
|
114
|
+
"Arrival rate (req/s): -1 sends as fast as possible; "
|
|
115
|
+
">0 uses Poisson arrivals"
|
|
116
|
+
),
|
|
117
|
+
)
|
|
118
|
+
parser.add_argument(
|
|
119
|
+
"--open-loop",
|
|
120
|
+
action="store_true",
|
|
121
|
+
default=_default(LoadConfig, "open_loop"),
|
|
122
|
+
help="Remove the concurrency limit; positive --rate still schedules arrivals",
|
|
123
|
+
)
|
|
124
|
+
|
|
125
|
+
# Generation
|
|
126
|
+
parser.add_argument(
|
|
127
|
+
"--max-tokens",
|
|
128
|
+
type=int,
|
|
129
|
+
nargs="+",
|
|
130
|
+
default=_default(GenerationConfig, "max_tokens"),
|
|
131
|
+
help=(
|
|
132
|
+
"Max generation tokens: one value (fixed) or two values "
|
|
133
|
+
"MIN MAX for uniform sampling per request"
|
|
134
|
+
),
|
|
135
|
+
)
|
|
136
|
+
sampling = parser.add_argument_group("sampling parameters")
|
|
137
|
+
sampling.add_argument(
|
|
138
|
+
"--top-p",
|
|
139
|
+
type=float,
|
|
140
|
+
default=_default(GenerationConfig, "top_p"),
|
|
141
|
+
help="Top-p sampling parameter",
|
|
142
|
+
)
|
|
143
|
+
sampling.add_argument(
|
|
144
|
+
"--top-k",
|
|
145
|
+
type=int,
|
|
146
|
+
default=_default(GenerationConfig, "top_k"),
|
|
147
|
+
help="Top-k sampling parameter",
|
|
148
|
+
)
|
|
149
|
+
sampling.add_argument(
|
|
150
|
+
"--min-p",
|
|
151
|
+
type=float,
|
|
152
|
+
default=_default(GenerationConfig, "min_p"),
|
|
153
|
+
help="Min-p sampling parameter",
|
|
154
|
+
)
|
|
155
|
+
sampling.add_argument(
|
|
156
|
+
"--temperature",
|
|
157
|
+
type=float,
|
|
158
|
+
default=_default(GenerationConfig, "temperature"),
|
|
159
|
+
help="Temperature sampling parameter",
|
|
160
|
+
)
|
|
161
|
+
sampling.add_argument(
|
|
162
|
+
"--frequency-penalty",
|
|
163
|
+
type=float,
|
|
164
|
+
default=_default(GenerationConfig, "frequency_penalty"),
|
|
165
|
+
help="Frequency penalty sampling parameter",
|
|
166
|
+
)
|
|
167
|
+
sampling.add_argument(
|
|
168
|
+
"--presence-penalty",
|
|
169
|
+
type=float,
|
|
170
|
+
default=_default(GenerationConfig, "presence_penalty"),
|
|
171
|
+
help="Presence penalty sampling parameter",
|
|
172
|
+
)
|
|
173
|
+
sampling.add_argument(
|
|
174
|
+
"--repetition-penalty",
|
|
175
|
+
type=float,
|
|
176
|
+
default=_default(GenerationConfig, "repetition_penalty"),
|
|
177
|
+
help="Repetition penalty sampling parameter",
|
|
178
|
+
)
|
|
179
|
+
parser.add_argument(
|
|
180
|
+
"--extra-body",
|
|
181
|
+
type=_json_object,
|
|
182
|
+
default=_default(GenerationConfig, "extra_body"),
|
|
183
|
+
help="JSON object of extra body parameters included in each request",
|
|
184
|
+
)
|
|
185
|
+
parser.add_argument(
|
|
186
|
+
"--stream",
|
|
187
|
+
action=argparse.BooleanOptionalAction,
|
|
188
|
+
default=_default(GenerationConfig, "stream"),
|
|
189
|
+
help=(
|
|
190
|
+
"Stream responses (default). --no-stream sends non-streaming "
|
|
191
|
+
"requests and reports latency only (no TTFT/TPOT)"
|
|
192
|
+
),
|
|
193
|
+
)
|
|
194
|
+
|
|
195
|
+
# Dataset
|
|
196
|
+
parser.add_argument(
|
|
197
|
+
"--dataset",
|
|
198
|
+
type=lambda value: [item.strip() for item in value.split(",") if item.strip()],
|
|
199
|
+
default=_default(DatasetConfig, "dataset"),
|
|
200
|
+
help=(
|
|
201
|
+
"Comma-separated sources: random, JSONL path, Hugging Face "
|
|
202
|
+
"org/name:split, or hf://datasets/...; --number is shared"
|
|
203
|
+
),
|
|
204
|
+
)
|
|
205
|
+
parser.add_argument(
|
|
206
|
+
"--trace",
|
|
207
|
+
dest="trace_path",
|
|
208
|
+
default=_default(DatasetConfig, "trace_path"),
|
|
209
|
+
help=(
|
|
210
|
+
"Trace source: JSONL path, supported trace dataset, or "
|
|
211
|
+
"hf://datasets/...; requires --dataset"
|
|
212
|
+
),
|
|
213
|
+
)
|
|
214
|
+
parser.add_argument(
|
|
215
|
+
"--trace-start",
|
|
216
|
+
type=float,
|
|
217
|
+
default=_default(DatasetConfig, "trace_start"),
|
|
218
|
+
help="Start offset from the first trace timestamp, in seconds",
|
|
219
|
+
)
|
|
220
|
+
parser.add_argument(
|
|
221
|
+
"--trace-duration",
|
|
222
|
+
type=float,
|
|
223
|
+
default=_default(DatasetConfig, "trace_duration"),
|
|
224
|
+
help="Trace window duration in seconds; omit to replay to the end",
|
|
225
|
+
)
|
|
226
|
+
parser.add_argument(
|
|
227
|
+
"--trace-max-concurrency",
|
|
228
|
+
type=int,
|
|
229
|
+
default=_default(DatasetConfig, "trace_max_concurrency"),
|
|
230
|
+
help=(
|
|
231
|
+
"Optional cap on active trace requests; timestamps still control "
|
|
232
|
+
"arrival times"
|
|
233
|
+
),
|
|
234
|
+
)
|
|
235
|
+
parser.add_argument(
|
|
236
|
+
"--trace-synthetic-prefix-reuse",
|
|
237
|
+
action="store_true",
|
|
238
|
+
default=_default(DatasetConfig, "trace_synthetic_prefix_reuse"),
|
|
239
|
+
help=(
|
|
240
|
+
"For Mooncake + random, synthesize deterministic 512-token "
|
|
241
|
+
"prefix blocks from trace hash_ids"
|
|
242
|
+
),
|
|
243
|
+
)
|
|
244
|
+
parser.add_argument(
|
|
245
|
+
"--dataset-offset",
|
|
246
|
+
type=int,
|
|
247
|
+
default=_default(DatasetConfig, "dataset_offset"),
|
|
248
|
+
help="Skip first N samples (JSONL/HF) or token-sequence offset (random)",
|
|
249
|
+
)
|
|
250
|
+
parser.add_argument(
|
|
251
|
+
"--tokenizer-path",
|
|
252
|
+
default=_default(DatasetConfig, "tokenizer_path"),
|
|
253
|
+
help="Tokenizer path (required for --dataset random)",
|
|
254
|
+
)
|
|
255
|
+
parser.add_argument(
|
|
256
|
+
"--random-seed",
|
|
257
|
+
type=int,
|
|
258
|
+
default=_default(DatasetConfig, "random_seed"),
|
|
259
|
+
help="Random payload seed (default: 0)",
|
|
260
|
+
)
|
|
261
|
+
parser.add_argument(
|
|
262
|
+
"--min-prompt-length",
|
|
263
|
+
type=int,
|
|
264
|
+
default=_default(DatasetConfig, "min_prompt_length"),
|
|
265
|
+
help="Minimum prompt length in tokens (random: sampled inner length)",
|
|
266
|
+
)
|
|
267
|
+
parser.add_argument(
|
|
268
|
+
"--max-prompt-length",
|
|
269
|
+
type=int,
|
|
270
|
+
default=_default(DatasetConfig, "max_prompt_length"),
|
|
271
|
+
help="Maximum prompt length in tokens (random: sampled inner length)",
|
|
272
|
+
)
|
|
273
|
+
parser.add_argument(
|
|
274
|
+
"--prefix-length",
|
|
275
|
+
type=int,
|
|
276
|
+
default=_default(DatasetConfig, "prefix_length"),
|
|
277
|
+
help="Shared prefix token length (random dataset only)",
|
|
278
|
+
)
|
|
279
|
+
parser.add_argument(
|
|
280
|
+
"--apply-chat-template",
|
|
281
|
+
action=argparse.BooleanOptionalAction,
|
|
282
|
+
default=_default(DatasetConfig, "apply_chat_template"),
|
|
283
|
+
help="Apply the chat template; defaults on for /chat/completions URLs",
|
|
284
|
+
)
|
|
285
|
+
parser.add_argument(
|
|
286
|
+
"--prompt",
|
|
287
|
+
default=_default(DatasetConfig, "prompt"),
|
|
288
|
+
help="Fixed prompt text; overrides dataset",
|
|
289
|
+
)
|
|
290
|
+
parser.add_argument(
|
|
291
|
+
"--max-turns",
|
|
292
|
+
type=int,
|
|
293
|
+
default=_default(DatasetConfig, "max_turns"),
|
|
294
|
+
help="Max user turns for custom_multi_turn",
|
|
295
|
+
)
|
|
296
|
+
|
|
297
|
+
# Output
|
|
298
|
+
parser.add_argument(
|
|
299
|
+
"--sla-auto-tune",
|
|
300
|
+
action=argparse.BooleanOptionalAction,
|
|
301
|
+
default=_default(OutputConfig, "sla_auto_tune"),
|
|
302
|
+
help="Enable SLA auto-tune search",
|
|
303
|
+
)
|
|
304
|
+
parser.add_argument(
|
|
305
|
+
"--output",
|
|
306
|
+
type=_output_destinations,
|
|
307
|
+
default=_default(OutputConfig, "destinations"),
|
|
308
|
+
help="Comma-separated outputs: local, wandb, and quiet",
|
|
309
|
+
)
|
|
310
|
+
parser.add_argument(
|
|
311
|
+
"--output-dir",
|
|
312
|
+
default=_default(OutputConfig, "output_dir"),
|
|
313
|
+
help="Directory for JSON and W&B artifacts",
|
|
314
|
+
)
|
|
315
|
+
|
|
316
|
+
# W&B
|
|
317
|
+
parser.add_argument(
|
|
318
|
+
"--wandb-project",
|
|
319
|
+
default=_default(WandbConfig, "project"),
|
|
320
|
+
help="W&B project",
|
|
321
|
+
)
|
|
322
|
+
parser.add_argument(
|
|
323
|
+
"--wandb-entity",
|
|
324
|
+
default=_default(WandbConfig, "entity"),
|
|
325
|
+
help="W&B entity",
|
|
326
|
+
)
|
|
327
|
+
parser.add_argument(
|
|
328
|
+
"--wandb-run-name",
|
|
329
|
+
default=_default(WandbConfig, "run_name"),
|
|
330
|
+
help=(
|
|
331
|
+
"W&B run-name prefix; child runs append their label. "
|
|
332
|
+
"Default: {model}_{YYYYMMDD_HHMMSS}"
|
|
333
|
+
),
|
|
334
|
+
)
|
|
335
|
+
|
|
336
|
+
# Param sweep
|
|
337
|
+
parser.add_argument(
|
|
338
|
+
"--bench-params",
|
|
339
|
+
default=_default(ParamSweepConfig, "bench_params"),
|
|
340
|
+
help=(
|
|
341
|
+
"JSONL parameter combinations; parallel, number, and rate may be lists"
|
|
342
|
+
),
|
|
343
|
+
)
|
|
344
|
+
parser.add_argument(
|
|
345
|
+
"--num-runs",
|
|
346
|
+
type=int,
|
|
347
|
+
default=_default(ParamSweepConfig, "num_runs"),
|
|
348
|
+
help="Runs per parameter combination",
|
|
349
|
+
)
|
|
350
|
+
parser.add_argument(
|
|
351
|
+
"--experiment-name",
|
|
352
|
+
default=_default(ParamSweepConfig, "experiment_name"),
|
|
353
|
+
help="Sweep directory name under --output-dir",
|
|
354
|
+
)
|
|
355
|
+
|
|
356
|
+
|
|
357
|
+
def _bench_config(namespace: argparse.Namespace) -> BenchConfig:
|
|
358
|
+
return BenchConfig(
|
|
359
|
+
endpoint=EndpointConfig(
|
|
360
|
+
url=namespace.url,
|
|
361
|
+
model=namespace.model,
|
|
362
|
+
api_key=namespace.api_key,
|
|
363
|
+
timeout=namespace.timeout,
|
|
364
|
+
max_retries=namespace.max_retries,
|
|
365
|
+
),
|
|
366
|
+
load=LoadConfig(
|
|
367
|
+
parallel=namespace.parallel,
|
|
368
|
+
number=namespace.number,
|
|
369
|
+
rate=namespace.rate,
|
|
370
|
+
open_loop=namespace.open_loop,
|
|
371
|
+
),
|
|
372
|
+
generation=GenerationConfig(
|
|
373
|
+
max_tokens=Arguments._validate_max_tokens(namespace.max_tokens),
|
|
374
|
+
stream=namespace.stream,
|
|
375
|
+
top_p=namespace.top_p,
|
|
376
|
+
top_k=namespace.top_k,
|
|
377
|
+
min_p=namespace.min_p,
|
|
378
|
+
temperature=namespace.temperature,
|
|
379
|
+
frequency_penalty=namespace.frequency_penalty,
|
|
380
|
+
presence_penalty=namespace.presence_penalty,
|
|
381
|
+
repetition_penalty=namespace.repetition_penalty,
|
|
382
|
+
extra_body=namespace.extra_body,
|
|
383
|
+
),
|
|
384
|
+
dataset=DatasetConfig(
|
|
385
|
+
dataset=namespace.dataset,
|
|
386
|
+
dataset_offset=namespace.dataset_offset,
|
|
387
|
+
tokenizer_path=namespace.tokenizer_path,
|
|
388
|
+
random_seed=namespace.random_seed,
|
|
389
|
+
min_prompt_length=namespace.min_prompt_length,
|
|
390
|
+
max_prompt_length=namespace.max_prompt_length,
|
|
391
|
+
prefix_length=namespace.prefix_length,
|
|
392
|
+
apply_chat_template=namespace.apply_chat_template,
|
|
393
|
+
prompt=namespace.prompt,
|
|
394
|
+
max_turns=namespace.max_turns,
|
|
395
|
+
trace_path=namespace.trace_path,
|
|
396
|
+
trace_start=namespace.trace_start,
|
|
397
|
+
trace_duration=namespace.trace_duration,
|
|
398
|
+
trace_max_concurrency=namespace.trace_max_concurrency,
|
|
399
|
+
trace_synthetic_prefix_reuse=(
|
|
400
|
+
namespace.trace_synthetic_prefix_reuse
|
|
401
|
+
),
|
|
402
|
+
),
|
|
403
|
+
output=OutputConfig(
|
|
404
|
+
destinations=namespace.output,
|
|
405
|
+
output_dir=namespace.output_dir,
|
|
406
|
+
sla_auto_tune=namespace.sla_auto_tune,
|
|
407
|
+
),
|
|
408
|
+
wandb=WandbConfig(
|
|
409
|
+
project=namespace.wandb_project,
|
|
410
|
+
entity=namespace.wandb_entity,
|
|
411
|
+
run_name=namespace.wandb_run_name,
|
|
412
|
+
),
|
|
413
|
+
param_sweep=ParamSweepConfig(
|
|
414
|
+
bench_params=namespace.bench_params,
|
|
415
|
+
num_runs=namespace.num_runs,
|
|
416
|
+
experiment_name=namespace.experiment_name,
|
|
417
|
+
),
|
|
418
|
+
)
|
|
419
|
+
|
|
420
|
+
|
|
421
|
+
def parse_arguments(argv: Sequence[str] | None = None) -> BenchCommand:
|
|
422
|
+
"""Parse the ``foretoken bench`` command."""
|
|
423
|
+
parser = argparse.ArgumentParser(
|
|
424
|
+
prog="foretoken",
|
|
425
|
+
description="Benchmark Foretoken and OpenAI-compatible inference services",
|
|
426
|
+
)
|
|
427
|
+
subparsers = parser.add_subparsers(dest="command", required=True)
|
|
428
|
+
bench = subparsers.add_parser(
|
|
429
|
+
"bench",
|
|
430
|
+
help="Benchmark a deployed Foretoken service or an existing endpoint",
|
|
431
|
+
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
|
432
|
+
)
|
|
433
|
+
_add_benchmark_arguments(bench)
|
|
434
|
+
|
|
435
|
+
parsed_args = parser.parse_args(argv)
|
|
436
|
+
if bool(parsed_args.kustomize_path) == bool(parsed_args.url):
|
|
437
|
+
bench.error("provide either PATH or --url")
|
|
438
|
+
if parsed_args.url and not parsed_args.model:
|
|
439
|
+
bench.error("--model is required with --url")
|
|
440
|
+
return BenchCommand(
|
|
441
|
+
kustomize_path=parsed_args.kustomize_path or "",
|
|
442
|
+
config=_bench_config(parsed_args),
|
|
443
|
+
wait_timeout=parsed_args.wait_timeout,
|
|
444
|
+
)
|
|
@@ -0,0 +1,148 @@
|
|
|
1
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
2
|
+
# SPDX-FileCopyrightText: Copyright contributors to the Foretoken project
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
"""OpenAI-compatible chat client."""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import time
|
|
10
|
+
from typing import Any, Optional
|
|
11
|
+
|
|
12
|
+
import httpx
|
|
13
|
+
from openai import APIError, AsyncOpenAI
|
|
14
|
+
|
|
15
|
+
from benchmarks.metrics.aggregator import compute_tpot
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _base_url(url: str) -> str:
|
|
19
|
+
return url.rstrip("/").removesuffix("/chat/completions")
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def derive_max_connections(
|
|
23
|
+
*, parallel: int, number: int, open_loop: bool
|
|
24
|
+
) -> int:
|
|
25
|
+
"""Size the httpx pool so it does not throttle below the load model.
|
|
26
|
+
|
|
27
|
+
Closed-loop: in-flight ≤ ``parallel``.
|
|
28
|
+
Open-loop: up to ``number`` may be in flight (gather / paced fire).
|
|
29
|
+
"""
|
|
30
|
+
return number if open_loop else parallel
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class OpenAICompatClient:
|
|
34
|
+
"""OpenAI-compatible chat client."""
|
|
35
|
+
|
|
36
|
+
def __init__(
|
|
37
|
+
self,
|
|
38
|
+
url: str,
|
|
39
|
+
model: str,
|
|
40
|
+
timeout: int,
|
|
41
|
+
api_key: str,
|
|
42
|
+
max_connections: Optional[int],
|
|
43
|
+
max_retries: int,
|
|
44
|
+
headers: dict[str, str] | None = None,
|
|
45
|
+
):
|
|
46
|
+
self.model = model
|
|
47
|
+
if max_connections is None:
|
|
48
|
+
limits = httpx.Limits(
|
|
49
|
+
max_connections=None,
|
|
50
|
+
max_keepalive_connections=None,
|
|
51
|
+
)
|
|
52
|
+
else:
|
|
53
|
+
# Keepalive matches max so finished requests can be reused under
|
|
54
|
+
# the same concurrency budget; the client is closed after a run.
|
|
55
|
+
limits = httpx.Limits(
|
|
56
|
+
max_connections=max_connections,
|
|
57
|
+
max_keepalive_connections=max_connections,
|
|
58
|
+
)
|
|
59
|
+
self.client = AsyncOpenAI(
|
|
60
|
+
base_url=_base_url(url),
|
|
61
|
+
api_key=api_key,
|
|
62
|
+
max_retries=max_retries,
|
|
63
|
+
default_headers=headers,
|
|
64
|
+
http_client=httpx.AsyncClient(
|
|
65
|
+
timeout=timeout,
|
|
66
|
+
limits=limits,
|
|
67
|
+
),
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
async def close(self) -> None:
|
|
71
|
+
await self.client.close()
|
|
72
|
+
|
|
73
|
+
async def generate(
|
|
74
|
+
self,
|
|
75
|
+
max_tokens: int,
|
|
76
|
+
*,
|
|
77
|
+
stream: bool,
|
|
78
|
+
extra_body: dict[str, Any],
|
|
79
|
+
prompt: Optional[str] = None,
|
|
80
|
+
messages: Optional[list[dict[str, Any]]] = None,
|
|
81
|
+
tools: Optional[list[dict[str, Any]]] = None,
|
|
82
|
+
) -> dict[str, Any]:
|
|
83
|
+
"""Run one chat completion; ``stream`` controls the request and metrics."""
|
|
84
|
+
if messages is None:
|
|
85
|
+
if prompt is None:
|
|
86
|
+
raise ValueError("Either prompt or messages must be provided")
|
|
87
|
+
messages = [{"role": "user", "content": prompt}]
|
|
88
|
+
if "stream" in extra_body:
|
|
89
|
+
raise ValueError(
|
|
90
|
+
"stream must be set via --stream/--no-stream, not extra_body"
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
kwargs: dict[str, Any] = {
|
|
94
|
+
"model": self.model,
|
|
95
|
+
"messages": messages,
|
|
96
|
+
"max_tokens": max_tokens,
|
|
97
|
+
"stream": stream,
|
|
98
|
+
}
|
|
99
|
+
if extra_body:
|
|
100
|
+
kwargs["extra_body"] = extra_body
|
|
101
|
+
if stream:
|
|
102
|
+
kwargs["stream_options"] = {"include_usage": True}
|
|
103
|
+
if tools:
|
|
104
|
+
kwargs["tools"] = tools
|
|
105
|
+
|
|
106
|
+
start_time = time.perf_counter()
|
|
107
|
+
ttft: Optional[float] = None
|
|
108
|
+
input_tokens = output_tokens = 0
|
|
109
|
+
status: Optional[int] = None
|
|
110
|
+
error_message: Optional[str] = None
|
|
111
|
+
success = True
|
|
112
|
+
try:
|
|
113
|
+
response = await self.client.chat.completions.create(**kwargs)
|
|
114
|
+
status = httpx.codes.OK
|
|
115
|
+
if stream:
|
|
116
|
+
async for chunk in response:
|
|
117
|
+
if chunk.usage is not None:
|
|
118
|
+
input_tokens = int(chunk.usage.prompt_tokens)
|
|
119
|
+
output_tokens = int(chunk.usage.completion_tokens)
|
|
120
|
+
if not chunk.choices:
|
|
121
|
+
continue
|
|
122
|
+
delta = chunk.choices[0].delta
|
|
123
|
+
if (delta.content or delta.tool_calls) and ttft is None:
|
|
124
|
+
ttft = time.perf_counter() - start_time
|
|
125
|
+
else:
|
|
126
|
+
if response.usage is not None:
|
|
127
|
+
input_tokens = int(response.usage.prompt_tokens)
|
|
128
|
+
output_tokens = int(response.usage.completion_tokens)
|
|
129
|
+
except (APIError, httpx.HTTPError) as exc:
|
|
130
|
+
success = False
|
|
131
|
+
status = getattr(exc, "status_code", None)
|
|
132
|
+
error_message = str(exc)
|
|
133
|
+
|
|
134
|
+
latency = time.perf_counter() - start_time
|
|
135
|
+
# TTFT/TPOT are defined only for streaming token arrival.
|
|
136
|
+
if not stream:
|
|
137
|
+
ttft = None
|
|
138
|
+
return {
|
|
139
|
+
"success": success,
|
|
140
|
+
"status_code": status,
|
|
141
|
+
"stream": stream,
|
|
142
|
+
"latency": latency,
|
|
143
|
+
"ttft": ttft,
|
|
144
|
+
"tpot": compute_tpot(latency, ttft, output_tokens),
|
|
145
|
+
"input_tokens": input_tokens,
|
|
146
|
+
"output_tokens": output_tokens,
|
|
147
|
+
"error": error_message,
|
|
148
|
+
}
|