#!/usr/bin/env python3 from __future__ import annotations import argparse import asyncio import json import statistics import time from pathlib import Path from typing import Any import httpx DEFAULT_PROMPTS = [ "Write a concise Python function that checks whether an integer is prime.", "Explain how you would debug a flaky pytest test in a large codebase.", "Given a list of file paths, write Python code to group them by extension.", ] def percentile(values: list[float], pct: float) -> float | None: if not values: return None values = sorted(values) idx = min(len(values) - 1, max(0, round((pct / 100) * (len(values) - 1)))) return values[idx] def load_prompts(path: Path | None) -> list[str]: if path is None: return DEFAULT_PROMPTS prompts: list[str] = [] for line in path.read_text().splitlines(): line = line.strip() if not line: continue if line.startswith("{"): obj = json.loads(line) prompts.append(obj.get("prompt") or obj.get("content") or obj["text"]) else: prompts.append(line) return prompts async def one_request( client: httpx.AsyncClient, base_url: str, model: str, prompt: str, max_tokens: int, temperature: float, top_p: float, stream: bool, ) -> dict[str, Any]: url = base_url.rstrip("/") + "/chat/completions" payload: dict[str, Any] = { "model": model, "messages": [{"role": "user", "content": prompt}], "temperature": temperature, "top_p": top_p, "max_tokens": max_tokens, } start = time.perf_counter() ttft = None output_text = "" usage: dict[str, Any] | None = None if stream: payload["stream"] = True payload["stream_options"] = {"include_usage": True} async with client.stream("POST", url, json=payload) as response: response.raise_for_status() async for line in response.aiter_lines(): if not line.startswith("data:"): continue data = line.removeprefix("data:").strip() if data == "[DONE]": break chunk = json.loads(data) if chunk.get("usage"): usage = chunk["usage"] choices = chunk.get("choices") or [] if not choices: continue delta = choices[0].get("delta") or {} piece = ( delta.get("content") or delta.get("reasoning_content") or delta.get("reasoning") or "" ) if piece and ttft is None: ttft = time.perf_counter() - start output_text += piece else: response = await client.post(url, json=payload) response.raise_for_status() body = response.json() usage = body.get("usage") message = body["choices"][0]["message"] output_text = ( (message.get("reasoning_content") or "") + (message.get("reasoning") or "") + (message.get("content") or "") ) elapsed = time.perf_counter() - start completion_tokens = usage.get("completion_tokens") if usage else None prompt_tokens = usage.get("prompt_tokens") if usage else None return { "ok": True, "elapsed_s": elapsed, "ttft_s": ttft, "prompt_chars": len(prompt), "output_chars": len(output_text), "prompt_tokens": prompt_tokens, "completion_tokens": completion_tokens, "output_tok_s": (completion_tokens / elapsed) if completion_tokens and elapsed > 0 else None, } async def worker( name: int, queue: asyncio.Queue[str], results: list[dict[str, Any]], args: argparse.Namespace, ) -> None: timeout = httpx.Timeout(args.timeout_s) async with httpx.AsyncClient(timeout=timeout) as client: while True: try: prompt = queue.get_nowait() except asyncio.QueueEmpty: return try: result = await one_request( client=client, base_url=args.base_url, model=args.model, prompt=prompt, max_tokens=args.max_tokens, temperature=args.temperature, top_p=args.top_p, stream=args.stream, ) result["worker"] = name except Exception as exc: # noqa: BLE001 result = { "ok": False, "error": repr(exc), "worker": name, "prompt_chars": len(prompt), } results.append(result) queue.task_done() async def run(args: argparse.Namespace) -> list[dict[str, Any]]: prompts = load_prompts(args.prompts) expanded = [prompts[i % len(prompts)] for i in range(args.requests)] queue: asyncio.Queue[str] = asyncio.Queue() for prompt in expanded: queue.put_nowait(prompt) results: list[dict[str, Any]] = [] workers = [ asyncio.create_task(worker(i, queue, results, args)) for i in range(args.concurrency) ] await queue.join() for task in workers: task.cancel() return results def summarize(results: list[dict[str, Any]], wall_s: float | None = None) -> dict[str, Any]: oks = [row for row in results if row.get("ok")] elapsed = [row["elapsed_s"] for row in oks] ttft = [row["ttft_s"] for row in oks if row.get("ttft_s") is not None] tok_s = [row["output_tok_s"] for row in oks if row.get("output_tok_s") is not None] completion_tokens = [ row["completion_tokens"] for row in oks if row.get("completion_tokens") is not None ] total_completion_tokens = sum(completion_tokens) return { "requests": len(results), "ok": len(oks), "failed": len(results) - len(oks), "wall_s": wall_s, "completion_tokens_total": total_completion_tokens, "aggregate_output_tok_s": ( total_completion_tokens / wall_s if wall_s and total_completion_tokens else None ), "latency_s_mean": statistics.mean(elapsed) if elapsed else None, "latency_s_p50": percentile(elapsed, 50), "latency_s_p95": percentile(elapsed, 95), "ttft_s_p50": percentile(ttft, 50), "ttft_s_p95": percentile(ttft, 95), "output_tok_s_mean": statistics.mean(tok_s) if tok_s else None, } def main() -> None: parser = argparse.ArgumentParser(description="Small OpenAI-compatible serving benchmark.") parser.add_argument("--base-url", default="http://localhost:8000/v1") parser.add_argument("--model", default="Ornith-1.0-35B") parser.add_argument("--prompts", type=Path) parser.add_argument("--requests", type=int, default=3) parser.add_argument("--concurrency", type=int, default=1) parser.add_argument("--max-tokens", type=int, default=256) parser.add_argument("--temperature", type=float, default=0.6) parser.add_argument("--top-p", type=float, default=0.95) parser.add_argument("--timeout-s", type=float, default=600.0) parser.add_argument("--stream", action="store_true") parser.add_argument("--output", type=Path) args = parser.parse_args() start = time.perf_counter() results = asyncio.run(run(args)) wall_s = time.perf_counter() - start summary = summarize(results, wall_s=wall_s) payload = {"summary": summary, "results": results} print(json.dumps(payload, indent=2, sort_keys=True)) if args.output: args.output.parent.mkdir(parents=True, exist_ok=True) with args.output.open("w") as handle: for row in results: handle.write(json.dumps(row, sort_keys=True) + "\n") handle.write(json.dumps({"summary": summary}, sort_keys=True) + "\n") if __name__ == "__main__": main()