Spaces:
Sleeping
Sleeping
Download scripts/bench.py from BeardedAmbivert/inference-server: direct link, hf CLI and curl.
- Browser
- Download file 10.1 kB
-
https://huggingface.co/spaces/BeardedAmbivert/inference-server/resolve/main/scripts/bench.py
- Command line
-
hf download hf://spaces/BeardedAmbivert/inference-server/scripts/bench.py
-
curl -L -o bench.py https://huggingface.co/spaces/BeardedAmbivert/inference-server/resolve/main/scripts/bench.py
10.1 kB
| """HTTP benchmark runner for the embedding inference server.""" | |
| from __future__ import annotations | |
| import argparse | |
| import asyncio | |
| import time | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import Any | |
| import httpx | |
| from utils import ( | |
| base_metadata, | |
| build_texts, | |
| format_ms, | |
| length_stats, | |
| load_text_pool, | |
| percentile, | |
| sample_texts, | |
| seconds_to_ms, | |
| write_json, | |
| ) | |
| DEFAULT_URL = "http://localhost:8000/embed" | |
| DEFAULT_REQUESTS = 40 | |
| DEFAULT_CONCURRENCY = 1 | |
| DEFAULT_TEXTS_PER_REQUEST = 2 | |
| DEFAULT_WARMUP = 0 | |
| DEFAULT_TIMEOUT_SECONDS = 30.0 | |
| BACKENDS = ("pytorch", "onnx") | |
| class RequestResult: | |
| latency_seconds: float | |
| status_code: int | None | |
| error: str | None = None | |
| def ok(self) -> bool: | |
| return self.error is None and self.status_code is not None and 200 <= self.status_code < 300 | |
| def build_payloads( | |
| count: int, | |
| texts_per_request: int, | |
| text_source: str, | |
| text_file: Path | None, | |
| seed: int, | |
| ) -> list[dict[str, list[str]]]: | |
| """Pre-build one payload per request. Synthetic payloads are identical; file-sourced | |
| payloads draw fresh texts per request (index-seeded RNG) so the run is varied yet reproducible.""" | |
| if text_source == "synthetic": | |
| texts = build_texts(texts_per_request) | |
| return [{"texts": texts} for _ in range(count)] | |
| pool = load_text_pool(text_file) | |
| return [{"texts": sample_texts(pool, texts_per_request, seed + i)} for i in range(count)] | |
| async def send_one( | |
| client: httpx.AsyncClient, | |
| url: str, | |
| payload: dict[str, list[str]], | |
| ) -> RequestResult: | |
| start = time.perf_counter() | |
| try: | |
| response = await client.post(url, json=payload) | |
| latency = time.perf_counter() - start | |
| return RequestResult( | |
| latency_seconds=latency, | |
| status_code=response.status_code, | |
| error=None if 200 <= response.status_code < 300 else response.text[:200], | |
| ) | |
| except Exception as exc: | |
| latency = time.perf_counter() - start | |
| return RequestResult( | |
| latency_seconds=latency, | |
| status_code=None, | |
| error=f"{type(exc).__name__}: {exc}", | |
| ) | |
| async def run_requests( | |
| url: str, | |
| payloads: list[dict[str, list[str]]], | |
| concurrency: int, | |
| timeout_seconds: float, | |
| ) -> tuple[list[RequestResult], float]: | |
| timeout = httpx.Timeout(timeout_seconds) | |
| limits = httpx.Limits(max_connections=concurrency, max_keepalive_connections=concurrency) | |
| semaphore = asyncio.Semaphore(concurrency) | |
| async with httpx.AsyncClient(timeout=timeout, limits=limits) as client: | |
| start = time.perf_counter() | |
| async def bounded_send(payload: dict[str, list[str]]) -> RequestResult: | |
| async with semaphore: | |
| return await send_one(client, url, payload) | |
| results = await asyncio.gather(*(bounded_send(payload) for payload in payloads)) | |
| wall_time = time.perf_counter() - start | |
| return list(results), wall_time | |
| def summarize( | |
| results: list[RequestResult], | |
| wall_time_seconds: float, | |
| metadata: dict[str, Any], | |
| ) -> dict[str, Any]: | |
| successful = [result.latency_seconds for result in results if result.ok] | |
| failed = [result for result in results if not result.ok] | |
| success_count = len(successful) | |
| successful_sequences = success_count * metadata["texts_per_request"] | |
| summary = { | |
| "metadata": metadata, | |
| "summary": { | |
| "total_requests": len(results), | |
| "successful_requests": success_count, | |
| "failed_requests": len(failed), | |
| "successful_sequences": successful_sequences, | |
| "wall_time_seconds": wall_time_seconds, | |
| "throughput_rps": success_count / wall_time_seconds if wall_time_seconds > 0 else 0.0, | |
| "throughput_sequences_per_sec": successful_sequences / wall_time_seconds if wall_time_seconds > 0 else 0.0, | |
| "avg_latency_ms": seconds_to_ms(sum(successful) / success_count) if success_count else None, | |
| "p50_latency_ms": seconds_to_ms(percentile(successful, 50)), | |
| "p95_latency_ms": seconds_to_ms(percentile(successful, 95)), | |
| "p99_latency_ms": seconds_to_ms(percentile(successful, 99)), | |
| }, | |
| "failures": [ | |
| { | |
| "status_code": result.status_code, | |
| "error": result.error, | |
| "latency_ms": seconds_to_ms(result.latency_seconds), | |
| } | |
| for result in failed[:10] | |
| ], | |
| } | |
| return summary | |
| def print_summary(report: dict[str, Any]) -> None: | |
| summary = report["summary"] | |
| metadata = report["metadata"] | |
| print("\nBenchmark") | |
| print(f"Label: {metadata['label']}") | |
| print(f"Backend: {metadata['backend']}") | |
| print(f"Server batch size: {metadata['server_batch_size']}") | |
| print(f"URL: {metadata['url']}") | |
| print(f"Requests: {metadata['requests']}") | |
| print(f"Concurrency: {metadata['concurrency']}") | |
| print(f"Texts/request: {metadata['texts_per_request']}") | |
| print(f"Warmup requests: {metadata['warmup']}") | |
| print("\nResults") | |
| print(f"Successful: {summary['successful_requests']}") | |
| print(f"Failed: {summary['failed_requests']}") | |
| print(f"Successful sequences: {summary['successful_sequences']}") | |
| print(f"Wall time: {summary['wall_time_seconds']:.3f}s") | |
| print(f"Request throughput: {summary['throughput_rps']:.2f} req/s") | |
| print(f"Sequence throughput: {summary['throughput_sequences_per_sec']:.2f} seq/s") | |
| print(f"Avg latency: {format_ms(summary['avg_latency_ms'])}") | |
| print(f"p50 latency: {format_ms(summary['p50_latency_ms'])}") | |
| print(f"p95 latency: {format_ms(summary['p95_latency_ms'])}") | |
| print(f"p99 latency: {format_ms(summary['p99_latency_ms'])}") | |
| if report["failures"]: | |
| print("\nFirst failures") | |
| for failure in report["failures"]: | |
| print(f"- status={failure['status_code']} error={failure['error']}") | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description="Benchmark the embedding inference HTTP API.") | |
| parser.add_argument("--label", default="http-benchmark", help="Human-readable run label stored in output metadata.") | |
| parser.add_argument("--backend", choices=BACKENDS, default="pytorch", help="Backend label for metadata only.") | |
| parser.add_argument("--device", help="Device the server ran on (cpu/mps). Metadata only.") | |
| parser.add_argument("--server-batch-size", type=int, help="Server MAX_BATCH_SIZE used for this run. Metadata only.") | |
| parser.add_argument("--url", default=DEFAULT_URL, help=f"Endpoint URL. Default: {DEFAULT_URL}") | |
| parser.add_argument("--requests", type=int, default=DEFAULT_REQUESTS, help="Total measured requests.") | |
| parser.add_argument("--concurrency", type=int, default=DEFAULT_CONCURRENCY, help="Concurrent in-flight requests.") | |
| parser.add_argument("--texts-per-request", type=int, default=DEFAULT_TEXTS_PER_REQUEST, help="Texts in each /embed request.") | |
| parser.add_argument("--text-source", choices=("synthetic", "file"), default="synthetic", help="Where request texts come from. Default: synthetic.") | |
| parser.add_argument("--text-file", type=Path, help="JSONL text pool (one {\"text\": ...} per line); used when --text-source=file.") | |
| parser.add_argument("--seed", type=int, default=0, help="Base seed for per-request text sampling.") | |
| parser.add_argument("--warmup", type=int, default=DEFAULT_WARMUP, help="Warmup requests excluded from results.") | |
| parser.add_argument("--timeout", type=float, default=DEFAULT_TIMEOUT_SECONDS, help="Per-request timeout in seconds.") | |
| parser.add_argument("--output", type=Path, help="Optional JSON output path.") | |
| return parser.parse_args() | |
| def validate_args(args: argparse.Namespace) -> None: | |
| if args.requests < 1: | |
| raise SystemExit("--requests must be >= 1") | |
| if args.concurrency < 1: | |
| raise SystemExit("--concurrency must be >= 1") | |
| if args.texts_per_request < 1: | |
| raise SystemExit("--texts-per-request must be >= 1") | |
| if args.warmup < 0: | |
| raise SystemExit("--warmup must be >= 0") | |
| if args.timeout <= 0: | |
| raise SystemExit("--timeout must be > 0") | |
| if args.server_batch_size is not None and args.server_batch_size < 1: | |
| raise SystemExit("--server-batch-size must be >= 1") | |
| if args.text_source == "file" and args.text_file is None: | |
| raise SystemExit("--text-file is required when --text-source=file") | |
| if args.text_file is not None and not args.text_file.exists(): | |
| raise SystemExit(f"--text-file not found: {args.text_file}") | |
| async def main() -> None: | |
| args = parse_args() | |
| validate_args(args) | |
| measured_payloads = build_payloads( | |
| args.requests, args.texts_per_request, args.text_source, args.text_file, args.seed | |
| ) | |
| if args.warmup: | |
| warmup_payloads = build_payloads( | |
| args.warmup, args.texts_per_request, args.text_source, args.text_file, args.seed + 10_000 | |
| ) | |
| await run_requests(args.url, warmup_payloads, args.concurrency, args.timeout) | |
| results, wall_time = await run_requests( | |
| args.url, measured_payloads, args.concurrency, args.timeout | |
| ) | |
| metadata = { | |
| **base_metadata(), | |
| "label": args.label, | |
| "backend": args.backend, | |
| "device": args.device, | |
| "server_batch_size": args.server_batch_size, | |
| "url": args.url, | |
| "requests": args.requests, | |
| "concurrency": args.concurrency, | |
| "texts_per_request": args.texts_per_request, | |
| "text_source": args.text_source, | |
| "text_file": args.text_file.name if args.text_file else None, | |
| "length_stats": length_stats([t for p in measured_payloads for t in p["texts"]]), | |
| "warmup": args.warmup, | |
| "timeout_seconds": args.timeout, | |
| } | |
| report = summarize(results, wall_time, metadata) | |
| print_summary(report) | |
| if args.output: | |
| write_json(report, args.output) | |
| print(f"\nWrote JSON results to {args.output}") | |
| if __name__ == "__main__": | |
| asyncio.run(main()) | |