Files
surya-ocr/benchmarks/concurrency_compare.py
T
Fu DaiandClaude Opus 4.8 1a585693be Suya OCR API — vLLM-backed, OpenAI-compatible OCR service
FastAPI service wrapping the Surya-OCR-2 model (datalab-to) served through vLLM:
legacy /v1/api/ai/* endpoints, an OpenAI-compatible /v1/chat/completions endpoint,
a coalescing request batcher, a local OCR CLI, Docker packaging, multilingual
example outputs, and quantization/concurrency benchmarks.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-17 10:20:02 +04:00

268 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Concurrency benchmark + old-vs-new comparison for the OCR service.
Two subcommands:
sweep Drive one running service across a list of concurrency levels and
record latency/throughput/success per level. Reuses the request
harness in concurrency_test.py.
compare Load two sweep result files (old + new) and emit a markdown table,
a CSV, and comparison plots (latency & throughput vs concurrency,
plus a per-level speedup bar).
Typical flow (see scripts/run_concurrency_comparison.sh):
# against the new vLLM service (its optimized endpoint)
python concurrency_compare.py sweep --label new \
--url http://127.0.0.1:5002/v1/api/ai/suya_ocr_vllm/ \
--image temp_image_c3c56948-282e-453b-8fac-6c482243d1e5.jpg \
--levels 1,2,4,8,16 --out results/compare/new.json
# ... swap containers, then against the old service (its OCR endpoint)
python concurrency_compare.py sweep --label old \
--url http://127.0.0.1:5002/v1/api/ai/suya_ocr/ \
--image temp_image_c3c56948-282e-453b-8fac-6c482243d1e5.jpg \
--levels 1,2,4,8,16 --out results/compare/old.json
python concurrency_compare.py compare \
--old results/compare/old.json --new results/compare/new.json \
--out-dir results/compare
"""
from __future__ import annotations
import argparse
import csv
import json
import time
from pathlib import Path
from typing import Any, Dict, List
from concurrency_test import _encode_image, run_endpoint
def _parse_levels(value: str) -> List[int]:
return [int(x.strip()) for x in value.split(",") if x.strip()]
# ---------------------------------------------------------------- sweep -----
def run_sweep(args: argparse.Namespace) -> None:
image_path = Path(args.image)
image_type = image_path.suffix.lstrip(".").lower() or "png"
payload = {
"file": _encode_image(image_path),
"type": "jpg" if image_type == "jpeg" else image_type,
"skip_text_detection": False,
"skip_table_detection": False,
"recognize_math": False,
"ocr_with_boxes": True,
}
levels = _parse_levels(args.levels)
rows: List[Dict[str, Any]] = []
for concurrency in levels:
requests_count = max(concurrency, concurrency * args.reqs_per_level)
print(f"[{args.label}] concurrency={concurrency} requests={requests_count}", flush=True)
row = run_endpoint(
args.label, args.url, payload,
requests_count=requests_count, concurrency=concurrency, timeout=args.timeout,
)
rows.append(row)
lat = row["latency_seconds"]
print(f" -> success={row['success']}/{row['requests']} "
f"rps={row['throughput_rps']:.3f} mean={lat['mean']:.2f}s p95={lat['p95']:.2f}s",
flush=True)
out = {
"label": args.label,
"url": args.url,
"image": str(image_path),
"levels": levels,
"reqs_per_level": args.reqs_per_level,
"started_at": time.strftime("%Y-%m-%dT%H:%M:%S%z"),
"results": rows,
}
out_path = Path(args.out)
out_path.parent.mkdir(parents=True, exist_ok=True)
out_path.write_text(json.dumps(out, indent=2), encoding="utf-8")
print(f"wrote {out_path}")
# -------------------------------------------------------------- compare -----
def _index_by_concurrency(sweep: Dict[str, Any]) -> Dict[int, Dict[str, Any]]:
return {row["concurrency"]: row for row in sweep["results"]}
def _comparison_rows(old: Dict[str, Any], new: Dict[str, Any]) -> List[Dict[str, Any]]:
old_by_c = _index_by_concurrency(old)
new_by_c = _index_by_concurrency(new)
concurrencies = sorted(set(old_by_c) | set(new_by_c))
rows = []
for c in concurrencies:
o = old_by_c.get(c)
n = new_by_c.get(c)
row: Dict[str, Any] = {"concurrency": c}
row["old_mean_s"] = round(o["latency_seconds"]["mean"], 2) if o else None
row["new_mean_s"] = round(n["latency_seconds"]["mean"], 2) if n else None
row["old_p95_s"] = round(o["latency_seconds"]["p95"], 2) if o else None
row["new_p95_s"] = round(n["latency_seconds"]["p95"], 2) if n else None
row["old_rps"] = round(o["throughput_rps"], 3) if o else None
row["new_rps"] = round(n["throughput_rps"], 3) if n else None
row["old_fail"] = o["failed"] if o else None
row["new_fail"] = n["failed"] if n else None
if o and n and o["latency_seconds"]["mean"] and n["latency_seconds"]["mean"]:
row["latency_speedup"] = round(o["latency_seconds"]["mean"] / n["latency_seconds"]["mean"], 2)
else:
row["latency_speedup"] = None
if o and n and o["throughput_rps"]:
row["throughput_gain"] = round(n["throughput_rps"] / o["throughput_rps"], 2)
else:
row["throughput_gain"] = None
rows.append(row)
return rows
def _write_table(rows: List[Dict[str, Any]], old_label: str, new_label: str, path: Path) -> str:
cols = [
("concurrency", "conc"),
("old_mean_s", f"{old_label} mean(s)"),
("new_mean_s", f"{new_label} mean(s)"),
("latency_speedup", "latency ×"),
("old_p95_s", f"{old_label} p95(s)"),
("new_p95_s", f"{new_label} p95(s)"),
("old_rps", f"{old_label} rps"),
("new_rps", f"{new_label} rps"),
("throughput_gain", "rps ×"),
("old_fail", f"{old_label} fail"),
("new_fail", f"{new_label} fail"),
]
header = "| " + " | ".join(label for _, label in cols) + " |"
sep = "| " + " | ".join("---" for _ in cols) + " |"
lines = [header, sep]
for r in rows:
cells = []
for key, _ in cols:
v = r.get(key)
cells.append("" if v is None else str(v))
lines.append("| " + " | ".join(cells) + " |")
md = "\n".join(lines) + "\n"
path.write_text(md, encoding="utf-8")
return md
def _write_csv(rows: List[Dict[str, Any]], path: Path) -> None:
if not rows:
return
with path.open("w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=list(rows[0].keys()))
writer.writeheader()
writer.writerows(rows)
def _write_plots(rows, old_label, new_label, out_dir: Path) -> List[Path]:
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
x = [r["concurrency"] for r in rows]
paths = []
# 1. latency (mean + p95) vs concurrency
fig, ax = plt.subplots(figsize=(9, 6))
ax.plot(x, [r["old_mean_s"] for r in rows], marker="o", color="#c0392b", label=f"{old_label} mean")
ax.plot(x, [r["old_p95_s"] for r in rows], marker="^", color="#c0392b", linestyle="--", label=f"{old_label} p95")
ax.plot(x, [r["new_mean_s"] for r in rows], marker="o", color="#27ae60", label=f"{new_label} mean")
ax.plot(x, [r["new_p95_s"] for r in rows], marker="^", color="#27ae60", linestyle="--", label=f"{new_label} p95")
ax.set_xlabel("Concurrency (simultaneous requests)")
ax.set_ylabel("Latency per request (s)")
ax.set_title("OCR latency vs concurrency — old vs new")
ax.grid(True, alpha=0.3)
ax.legend()
fig.tight_layout()
p = out_dir / "latency_vs_concurrency.png"
fig.savefig(p, dpi=160)
plt.close(fig)
paths.append(p)
# 2. throughput vs concurrency
fig, ax = plt.subplots(figsize=(9, 6))
ax.plot(x, [r["old_rps"] for r in rows], marker="o", color="#c0392b", label=old_label)
ax.plot(x, [r["new_rps"] for r in rows], marker="o", color="#27ae60", label=new_label)
ax.set_xlabel("Concurrency (simultaneous requests)")
ax.set_ylabel("Throughput (requests/s)")
ax.set_title("OCR throughput vs concurrency — old vs new")
ax.grid(True, alpha=0.3)
ax.legend()
fig.tight_layout()
p = out_dir / "throughput_vs_concurrency.png"
fig.savefig(p, dpi=160)
plt.close(fig)
paths.append(p)
# 3. per-level latency speedup bar
fig, ax = plt.subplots(figsize=(9, 6))
speedups = [r["latency_speedup"] or 0 for r in rows]
ax.bar([str(c) for c in x], speedups, color="#2980b9")
ax.axhline(1.0, color="gray", linestyle="--", linewidth=1)
for i, v in enumerate(speedups):
ax.text(i, v, f"{v:.2f}×", ha="center", va="bottom")
ax.set_xlabel("Concurrency")
ax.set_ylabel(f"Latency speedup ({old_label} mean / {new_label} mean)")
ax.set_title("Per-level latency speedup (>1 = new is faster)")
ax.grid(True, axis="y", alpha=0.3)
fig.tight_layout()
p = out_dir / "latency_speedup.png"
fig.savefig(p, dpi=160)
plt.close(fig)
paths.append(p)
return paths
def run_compare(args: argparse.Namespace) -> None:
old = json.loads(Path(args.old).read_text(encoding="utf-8"))
new = json.loads(Path(args.new).read_text(encoding="utf-8"))
old_label = old.get("label", "old")
new_label = new.get("label", "new")
out_dir = Path(args.out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
rows = _comparison_rows(old, new)
md = _write_table(rows, old_label, new_label, out_dir / "comparison_table.md")
_write_csv(rows, out_dir / "comparison_table.csv")
plots = _write_plots(rows, old_label, new_label, out_dir)
print(md)
print("plots:")
for p in plots:
print(f" {p}")
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
sub = parser.add_subparsers(dest="cmd", required=True)
sp = sub.add_parser("sweep", help="benchmark one running service across concurrency levels")
sp.add_argument("--label", required=True, help="short name for this service, e.g. old / new")
sp.add_argument("--url", required=True, help="full endpoint URL to POST to")
sp.add_argument("--image", required=True, help="path to a png/jpg test image")
sp.add_argument("--levels", default="1,2,4,8,16", help="comma-separated concurrency levels")
sp.add_argument("--reqs-per-level", type=int, default=2, help="requests = level * this (>=level)")
sp.add_argument("--timeout", type=float, default=900)
sp.add_argument("--out", required=True, help="output JSON path")
sp.set_defaults(func=run_sweep)
cp = sub.add_parser("compare", help="compare two sweep result files")
cp.add_argument("--old", required=True, help="old-service sweep JSON")
cp.add_argument("--new", required=True, help="new-service sweep JSON")
cp.add_argument("--out-dir", default="results/compare")
cp.set_defaults(func=run_compare)
args = parser.parse_args()
args.func(args)
if __name__ == "__main__":
main()