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>
This commit is contained in:
Fu Dai
2026-06-17 10:20:02 +04:00
co-authored by Claude Opus 4.8
commit 1a585693be
147 changed files with 13827 additions and 0 deletions
+54
View File
@@ -0,0 +1,54 @@
"""Post one or more images to an OCR endpoint and save text_lines per image.
Used to freeze the accuracy baseline and to capture candidate outputs for CER.
Usage:
python scripts/capture_ocr_text.py \
--url http://127.0.0.1:5002/v1/api/ai/suya_ocr_vllm/ \
--out-dir baseline_outputs \
temp_image_c3c56948-282e-453b-8fac-6c482243d1e5.jpg
"""
from __future__ import annotations
import argparse
import base64
from pathlib import Path
import requests
def _payload(image_path: Path) -> dict:
suffix = image_path.suffix.lstrip(".").lower() or "png"
return {
"file": base64.b64encode(image_path.read_bytes()).decode("utf-8"),
"type": "jpg" if suffix == "jpeg" else suffix,
"skip_text_detection": False,
"skip_table_detection": False,
"recognize_math": False,
"ocr_with_boxes": True,
}
def capture(url: str, image_path: Path, out_dir: Path, timeout: float) -> str:
resp = requests.post(url, json=_payload(image_path), timeout=timeout)
body = resp.json()
text = (body.get("data") or {}).get("text_lines", "") if isinstance(body.get("data"), dict) else ""
out_dir.mkdir(parents=True, exist_ok=True)
(out_dir / (image_path.stem + ".txt")).write_text(text, encoding="utf-8")
return text
def main() -> None:
parser = argparse.ArgumentParser(description="Capture OCR text_lines per image.")
parser.add_argument("images", nargs="+", type=Path)
parser.add_argument("--url", required=True)
parser.add_argument("--out-dir", type=Path, required=True)
parser.add_argument("--timeout", type=float, default=900)
args = parser.parse_args()
for image in args.images:
capture(args.url, image, args.out_dir, args.timeout)
print(f"captured {image.name} -> {args.out_dir}")
if __name__ == "__main__":
main()
+59
View File
@@ -0,0 +1,59 @@
"""Character error rate (CER) between a reference baseline and a candidate.
CER = levenshtein(reference, hypothesis) / len(reference).
Used to gate work-reduction changes: keep a change only if CER <= 0.005.
Usage:
python scripts/cer_divergence.py baseline_outputs/ candidate_outputs/
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
def _levenshtein(a: str, b: str) -> int:
if a == b:
return 0
if not a:
return len(b)
if not b:
return len(a)
prev = list(range(len(b) + 1))
for i, ca in enumerate(a, 1):
cur = [i]
for j, cb in enumerate(b, 1):
cost = 0 if ca == cb else 1
cur.append(min(prev[j] + 1, cur[j - 1] + 1, prev[j - 1] + cost))
prev = cur
return prev[-1]
def cer(reference: str, hypothesis: str) -> float:
if not reference:
return 0.0 if not hypothesis else 1.0
return _levenshtein(reference, hypothesis) / len(reference)
def cer_over_dirs(baseline_dir: Path, candidate_dir: Path) -> dict:
per_file = {}
for ref_path in sorted(baseline_dir.glob("*.txt")):
cand_path = candidate_dir / ref_path.name
ref = ref_path.read_text(encoding="utf-8")
hyp = cand_path.read_text(encoding="utf-8") if cand_path.exists() else ""
per_file[ref_path.name] = cer(ref, hyp)
mean = sum(per_file.values()) / len(per_file) if per_file else 0.0
return {"mean_cer": mean, "max_cer": max(per_file.values(), default=0.0), "per_file": per_file}
def main() -> None:
parser = argparse.ArgumentParser(description="CER divergence between two output dirs.")
parser.add_argument("baseline_dir", type=Path)
parser.add_argument("candidate_dir", type=Path)
args = parser.parse_args()
print(json.dumps(cer_over_dirs(args.baseline_dir, args.candidate_dir), indent=2))
if __name__ == "__main__":
main()
+18
View File
@@ -0,0 +1,18 @@
#!/usr/bin/env bash
set -euo pipefail
if ! command -v nvidia-ctk >/dev/null 2>&1; then
echo "nvidia-ctk is not installed. Install NVIDIA Container Toolkit first." >&2
exit 1
fi
sudo nvidia-ctk runtime configure --runtime=docker --config=/etc/docker/daemon.json
if command -v systemctl >/dev/null 2>&1; then
sudo systemctl restart docker
else
sudo service docker restart
fi
docker info --format '{{json .Runtimes}}'
docker run --rm --runtime nvidia --gpus all nvidia/cuda:12.4.1-base-ubuntu22.04 nvidia-smi
+23
View File
@@ -0,0 +1,23 @@
#!/usr/bin/env bash
set -euo pipefail
args=()
skip_next=0
for arg in "$@"; do
if [ "$skip_next" -eq 1 ]; then
skip_next=0
if [ "$arg" = "nvidia" ]; then
continue
fi
fi
if [ "$arg" = "--runtime" ]; then
skip_next=1
continue
fi
args+=("$arg")
done
exec /usr/bin/docker "${args[@]}"
+68
View File
@@ -0,0 +1,68 @@
"""Parse surya_timing_summary lines from logs/info.log and aggregate spans.
Usage:
python scripts/parse_timing.py logs/info.log
"""
from __future__ import annotations
import argparse
import ast
import json
import re
from collections import defaultdict
from pathlib import Path
from typing import Any, Dict, List, Optional
_BATCH_RE = re.compile(r"batch_size=(\d+)")
_REQID_RE = re.compile(r"request_id=(\S+)")
def parse_line(line: str) -> Optional[Dict[str, Any]]:
if "surya_timing_summary" not in line or "events=" not in line:
return None
events_str = line.split("events=", 1)[1].strip()
try:
events = ast.literal_eval(events_str)
except (ValueError, SyntaxError):
return None
batch_match = _BATCH_RE.search(line)
reqid_match = _REQID_RE.search(line)
return {
"request_id": reqid_match.group(1) if reqid_match else "-",
"batch_size": int(batch_match.group(1)) if batch_match else 0,
"events": events,
}
def aggregate(records: List[Dict[str, Any]]) -> Dict[str, Dict[str, float]]:
stats: Dict[str, Dict[str, float]] = defaultdict(
lambda: {"count": 0, "total_ms": 0.0, "mean_ms": 0.0, "total_tokens": 0}
)
for rec in records:
for event in rec["events"]:
s = stats[event["name"]]
s["count"] += 1
s["total_ms"] += event.get("duration_ms", 0.0)
tokens = (event.get("metadata") or {}).get("token_count")
if tokens:
s["total_tokens"] += tokens
for s in stats.values():
s["mean_ms"] = round(s["total_ms"] / s["count"], 2) if s["count"] else 0.0
s["total_ms"] = round(s["total_ms"], 2)
return dict(stats)
def main() -> None:
parser = argparse.ArgumentParser(description="Aggregate surya timing spans from a log file.")
parser.add_argument("logfile", type=Path)
args = parser.parse_args()
records = [
rec
for rec in (parse_line(line) for line in args.logfile.read_text(encoding="utf-8").splitlines())
if rec is not None
]
print(json.dumps({"records": len(records), "spans": aggregate(records)}, indent=2))
if __name__ == "__main__":
main()
+54
View File
@@ -0,0 +1,54 @@
"""Sample a vLLM server's Prometheus /metrics and report peak in-flight sequences.
vLLM exposes `vllm:num_requests_running` (currently executing on GPU) and
`vllm:num_requests_waiting` (queued). Peak running vs VLLM_MAX_NUM_SEQS tells us
whether we are saturating or overflowing the GPU's sequence slots.
Run this in the background during a benchmark, then read its printed summary.
Usage:
python scripts/poll_vllm_metrics.py \
--url http://127.0.0.1:8000/metrics --interval 0.25 --duration 120
"""
from __future__ import annotations
import argparse
import re
import time
import requests
_RUNNING_RE = re.compile(r"^vllm:num_requests_running\S*\s+([0-9.]+)", re.MULTILINE)
_WAITING_RE = re.compile(r"^vllm:num_requests_waiting\S*\s+([0-9.]+)", re.MULTILINE)
def _scrape(url: str) -> tuple[float, float]:
text = requests.get(url, timeout=5).text
running = max((float(m) for m in _RUNNING_RE.findall(text)), default=0.0)
waiting = max((float(m) for m in _WAITING_RE.findall(text)), default=0.0)
return running, waiting
def main() -> None:
parser = argparse.ArgumentParser(description="Poll vLLM /metrics for peak in-flight sequences.")
parser.add_argument("--url", default="http://127.0.0.1:8000/metrics")
parser.add_argument("--interval", type=float, default=0.25)
parser.add_argument("--duration", type=float, default=120)
args = parser.parse_args()
peak_running = 0.0
peak_waiting = 0.0
deadline = time.perf_counter() + args.duration
while time.perf_counter() < deadline:
try:
running, waiting = _scrape(args.url)
peak_running = max(peak_running, running)
peak_waiting = max(peak_waiting, waiting)
except Exception:
pass
time.sleep(args.interval)
print(f"peak_running={peak_running} peak_waiting={peak_waiting}")
if __name__ == "__main__":
main()
+1
View File
@@ -0,0 +1 @@
"""Quantization benchmark harness for the Surya-OCR-2 recognition VLM."""
+59
View File
@@ -0,0 +1,59 @@
"""Summary-row contract and rendering to CSV + markdown."""
from __future__ import annotations
import csv
import io
from typing import Any, Dict, List
SUMMARY_FIELDS = [
"method",
"status",
"t4_deployable",
"model_size_mb",
"mean_latency_s",
"p50_latency_s",
"p95_latency_s",
"throughput_rps",
"decode_tok_s",
"mean_cer",
"max_cer",
"mean_bbox_iou",
"mean_missed_lines",
"mean_extra_lines",
"error",
]
def build_row(**kwargs: Any) -> Dict[str, Any]:
unknown = set(kwargs) - set(SUMMARY_FIELDS)
if unknown:
raise KeyError(f"unknown summary field(s): {sorted(unknown)}")
row = {field: None for field in SUMMARY_FIELDS}
row.update(kwargs)
return row
def rows_to_csv(rows: List[Dict[str, Any]]) -> str:
buf = io.StringIO()
writer = csv.DictWriter(buf, fieldnames=SUMMARY_FIELDS)
writer.writeheader()
for row in rows:
writer.writerow(row)
return buf.getvalue()
def _fmt(value: Any) -> str:
if value is None:
return ""
if isinstance(value, float):
return f"{value:.4f}"
return str(value)
def rows_to_markdown(rows: List[Dict[str, Any]]) -> str:
header = "| " + " | ".join(SUMMARY_FIELDS) + " |"
sep = "| " + " | ".join("---" for _ in SUMMARY_FIELDS) + " |"
lines = [header, sep]
for row in rows:
lines.append("| " + " | ".join(_fmt(row[f]) for f in SUMMARY_FIELDS) + " |")
return "\n".join(lines)
+66
View File
@@ -0,0 +1,66 @@
"""Greedy bbox IoU matching between reference and candidate page boxes."""
from __future__ import annotations
import json
from pathlib import Path
from typing import Dict, List, Sequence
Box = Sequence[float]
def iou(a: Box, b: Box) -> float:
ax0, ay0, ax1, ay1 = a
bx0, by0, bx1, by1 = b
ix0, iy0 = max(ax0, bx0), max(ay0, by0)
ix1, iy1 = min(ax1, bx1), min(ay1, by1)
iw, ih = max(0.0, ix1 - ix0), max(0.0, iy1 - iy0)
inter = iw * ih
if inter <= 0:
return 0.0
area_a = max(0.0, ax1 - ax0) * max(0.0, ay1 - ay0)
area_b = max(0.0, bx1 - bx0) * max(0.0, by1 - by0)
union = area_a + area_b - inter
return inter / union if union > 0 else 0.0
def match_boxes(ref: List[Box], cand: List[Box], iou_threshold: float = 0.5) -> Dict[str, float]:
used = [False] * len(cand)
matched_ious: List[float] = []
for r in ref:
best_j, best_iou = -1, 0.0
for j, c in enumerate(cand):
if used[j]:
continue
cur = iou(r, c)
if cur > best_iou:
best_iou, best_j = cur, j
if best_j >= 0 and best_iou >= iou_threshold:
used[best_j] = True
matched_ious.append(best_iou)
matched = len(matched_ious)
return {
"matched": matched,
"missed": len(ref) - matched,
"extra": len(cand) - matched,
"mean_matched_iou": sum(matched_ious) / matched if matched else 0.0,
}
def bbox_iou_over_dirs(ref_dir: Path, cand_dir: Path, iou_threshold: float = 0.5) -> Dict:
per_file: Dict[str, Dict[str, float]] = {}
for ref_path in sorted(ref_dir.glob("*.json")):
ref_boxes = json.loads(ref_path.read_text(encoding="utf-8")).get("boxes", [])
cand_path = cand_dir / ref_path.name
cand_boxes = (
json.loads(cand_path.read_text(encoding="utf-8")).get("boxes", [])
if cand_path.exists()
else []
)
per_file[ref_path.name] = match_boxes(ref_boxes, cand_boxes, iou_threshold)
n = len(per_file) or 1
return {
"mean_bbox_iou": sum(v["mean_matched_iou"] for v in per_file.values()) / n,
"mean_missed_lines": sum(v["missed"] for v in per_file.values()) / n,
"mean_extra_lines": sum(v["extra"] for v in per_file.values()) / n,
"per_file": per_file,
}
+144
View File
@@ -0,0 +1,144 @@
"""Produce vLLM-loadable quantized checkpoints.
compressor methods (int8/awq/gptq) use llm-compressor `oneshot`.
bnb methods (bnb8/bnb4) use transformers + BitsAndBytesConfig + save_pretrained.
baseline/online (bf16/fp8) need no build.
"""
from __future__ import annotations
from pathlib import Path
from typing import List
from scripts.quant.recipes import METHOD_SPECS
def needs_build(method: str) -> bool:
return METHOD_SPECS[method]["kind"] in ("compressor", "bnb")
def _is_built(out_dir: Path) -> bool:
return (out_dir / "config.json").exists()
CALIB_MAX_SEQ_LEN = 8192
# Never quantize the LM head or the vision tower: vision modules are shape-fragile
# (see the FP8 Marlin failure) and contribute little to decode cost.
QUANT_IGNORE = ["re:.*lm_head", "re:.*visual.*", "re:.*vision.*"]
def _calibration_dataset(calib_images: List[Path], processor):
"""HF Dataset of pre-tokenized multimodal samples (batch dim kept), matching
llm-compressor's multimodal-vision examples; re-tensorized by _data_collator.
Activation calibration only needs representative forward passes, not the
exact training prompt."""
from datasets import Dataset
from PIL import Image
samples = []
for path in calib_images:
image = Image.open(path).convert("RGB")
messages = [{
"role": "user",
"content": [
{"type": "image"},
{"type": "text", "text": "OCR this document."},
],
}]
prompt = processor.apply_chat_template(messages, add_generation_prompt=True, tokenize=False)
inputs = processor(
text=[prompt], images=[image],
padding=False, truncation=True, max_length=CALIB_MAX_SEQ_LEN,
)
samples.append({k: (v.tolist() if hasattr(v, "tolist") else v) for k, v in inputs.items()})
return Dataset.from_list(samples)
def _data_collator(batch):
import torch
assert len(batch) == 1
return {key: torch.tensor(value) for key, value in batch[0].items()}
def _build_compressor(method: str, base_model: str, out_dir: Path, calib_images: List[Path]) -> None:
from llmcompressor import oneshot
from llmcompressor.modifiers.quantization import GPTQModifier
from transformers import AutoProcessor
try:
from transformers import AutoModelForImageTextToText as _AutoModel
except ImportError: # older transformers
from transformers import AutoModelForCausalLM as _AutoModel
spec = METHOD_SPECS[method]
model = _AutoModel.from_pretrained(base_model, dtype="auto", device_map="auto")
processor = AutoProcessor.from_pretrained(base_model)
dataset = _calibration_dataset(calib_images, processor)
if spec["modifier"] == "awq":
# llm-compressor main: AWQ is a transform paired with a QuantizationModifier.
from llmcompressor.modifiers.quantization import QuantizationModifier
try:
from llmcompressor.modifiers.transform.awq import AWQModifier
except ImportError: # older layouts keep AWQModifier under modifiers.awq
from llmcompressor.modifiers.awq import AWQModifier
recipe = [
AWQModifier(duo_scaling=False),
QuantizationModifier(scheme=spec["scheme"], ignore=QUANT_IGNORE),
]
else:
recipe = []
if spec.get("smoothquant"):
from llmcompressor.modifiers.smoothquant import SmoothQuantModifier
recipe.append(SmoothQuantModifier(smoothing_strength=0.8))
recipe.append(GPTQModifier(targets="Linear", scheme=spec["scheme"], ignore=QUANT_IGNORE))
oneshot(
model=model,
dataset=dataset,
recipe=recipe,
max_seq_length=CALIB_MAX_SEQ_LEN,
num_calibration_samples=len(dataset),
data_collator=_data_collator,
)
out_dir.mkdir(parents=True, exist_ok=True)
model.save_pretrained(out_dir, save_compressed=True)
processor.save_pretrained(out_dir)
def _build_bnb(method: str, base_model: str, out_dir: Path) -> None:
import torch
from transformers import AutoProcessor, BitsAndBytesConfig
try:
from transformers import AutoModelForImageTextToText as _AutoModel
except ImportError:
from transformers import AutoModelForCausalLM as _AutoModel
bits = METHOD_SPECS[method]["bits"]
if bits == 8:
config = BitsAndBytesConfig(load_in_8bit=True)
else:
config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_use_double_quant=True,
)
model = _AutoModel.from_pretrained(base_model, quantization_config=config, device_map="auto")
model.save_pretrained(out_dir)
AutoProcessor.from_pretrained(base_model).save_pretrained(out_dir)
def build_model(method: str, base_model: str, out_dir: Path, calib_images: List[Path]) -> Path:
if not needs_build(method):
return Path(base_model)
out_dir = Path(out_dir)
if _is_built(out_dir):
return out_dir
kind = METHOD_SPECS[method]["kind"]
if kind == "compressor":
_build_compressor(method, base_model, out_dir, calib_images)
else:
_build_bnb(method, base_model, out_dir)
return out_dir
+77
View File
@@ -0,0 +1,77 @@
"""Capture per-page OCR text, bboxes, and latency from the vLLM OCR endpoint."""
from __future__ import annotations
import argparse
import base64
import json
import sys
from pathlib import Path
from typing import Any, Dict
import requests
def extract_capture(body: Dict[str, Any]) -> Dict[str, Any]:
data = body.get("data") or {}
if not isinstance(data, dict):
data = {}
blocks = ((data.get("ocr_text_json") or {}).get("blocks")) or []
boxes = [list(b["bbox"]) for b in blocks if isinstance(b, dict) and b.get("bbox")]
return {
"text": data.get("text_lines", "") or "",
"boxes": boxes,
"elapsed_seconds": data.get("elapsed_seconds"),
}
def write_capture(cap: Dict[str, Any], out_dir: Path, stem: str) -> None:
out_dir.mkdir(parents=True, exist_ok=True)
(out_dir / f"{stem}.txt").write_text(cap["text"], encoding="utf-8")
(out_dir / f"{stem}.json").write_text(
json.dumps({"boxes": cap["boxes"], "elapsed_seconds": cap["elapsed_seconds"]}),
encoding="utf-8",
)
def _payload(image_path: Path) -> Dict[str, Any]:
suffix = image_path.suffix.lstrip(".").lower() or "png"
return {
"file": base64.b64encode(image_path.read_bytes()).decode("utf-8"),
"type": "jpg" if suffix == "jpeg" else suffix,
"skip_text_detection": False,
"skip_table_detection": False,
"recognize_math": False,
"ocr_with_boxes": True,
}
def capture_page(url: str, image_path: Path, out_dir: Path, timeout: float) -> Dict[str, Any]:
resp = requests.post(url, json=_payload(image_path), timeout=timeout)
cap = extract_capture(resp.json())
if cap["elapsed_seconds"] is None:
# Latency is a headline metric; a successful response with no timing means
# the response contract changed. Surface it loudly instead of silently
# dropping the page from the latency stats.
print(
f"WARNING: no elapsed_seconds in OCR response for {image_path.name}; "
"latency for this page will be missing",
file=sys.stderr,
)
write_capture(cap, out_dir, image_path.stem)
return cap
def main() -> None:
parser = argparse.ArgumentParser(description="Capture OCR text+boxes+latency per image.")
parser.add_argument("images", nargs="+", type=Path)
parser.add_argument("--url", required=True)
parser.add_argument("--out-dir", type=Path, required=True)
parser.add_argument("--timeout", type=float, default=900)
args = parser.parse_args()
for image in args.images:
cap = capture_page(args.url, image, args.out_dir, args.timeout)
print(f"captured {image.name}: {len(cap['boxes'])} boxes, {cap['elapsed_seconds']}s")
if __name__ == "__main__":
main()
+127
View File
@@ -0,0 +1,127 @@
"""Assemble the quantization benchmark summary from captured per-page outputs.
Unlike run_all's inline scoring (which globs the reference and penalizes any page
a method did not capture), this scores **candidate-driven**: for each method it
scores only the pages that method actually produced, against the matching
reference page. That lets fast methods run the full eval set while slow/eager
methods (e.g. bitsandbytes) run a representative subset, without the subset being
unfairly penalized for missing pages.
Per-method status/reason and t4-deployability are passed in via a small spec so
methods that could not be built or served are recorded as explicit failed rows.
Run inside the bench container: python3 -m scripts.quant.finalize
"""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any, Dict, List
from scripts.cer_divergence import cer
from scripts.quant.aggregate import build_row, rows_to_csv, rows_to_markdown
from scripts.quant.bbox_iou import match_boxes
WORK = Path("results/quant")
REF = WORK / "reference"
RESULTS = WORK / "results"
# Per-method outcome metadata for this run (A100, vLLM v0.22.1 container).
# status "ok" methods are scored from their captured pages — if a method has no
# captures on disk it falls back to a failed row with `fallback_error`. "failed"
# methods record why outright. bf16 is the reference (CER 0 / IoU 1 by construction).
#
# int8/awq/gptq were built with llm-compressor main (0.12.x): the PyPI 0.11 line
# pins transformers<5 and cannot load qwen3_5, but main supports transformers 5.x.
# int8 is plain GPTQ W8A8 (SmoothQuant default mappings unresolvable for qwen3_5).
SPECS: List[Dict[str, Any]] = [
{"method": "bf16", "status": "ok", "t4_deployable": True, "is_reference": True},
{"method": "fp8", "status": "failed", "t4_deployable": False,
"error": "vLLM FP8 Marlin kernel: size_n=32 not divisible by tile_n_size=64 "
"(model layer shape incompatible with fp8 dynamic)"},
{"method": "int8", "status": "ok", "t4_deployable": True, "is_reference": False,
"fallback_error": "checkpoint built (llm-compressor main, GPTQ W8A8) but vLLM "
"failed to serve it — see /tmp/serve_int8.log"},
{"method": "awq", "status": "ok", "t4_deployable": True, "is_reference": False,
"fallback_error": "checkpoint built (llm-compressor main, AWQ W4A16) but vLLM "
"failed to serve it — see /tmp/serve_awq.log"},
{"method": "gptq", "status": "ok", "t4_deployable": True, "is_reference": False,
"fallback_error": "checkpoint built (llm-compressor main, GPTQ W4A16) but vLLM "
"failed to serve it — see /tmp/serve_gptq.log"},
{"method": "bnb8", "status": "ok", "t4_deployable": True, "is_reference": False},
{"method": "bnb4", "status": "ok", "t4_deployable": True, "is_reference": False},
]
def _load_page(d: Path, stem: str):
j = json.loads((d / f"{stem}.json").read_text(encoding="utf-8"))
txt = (d / f"{stem}.txt").read_text(encoding="utf-8") if (d / f"{stem}.txt").exists() else ""
return j, txt
def _score(method: str, is_reference: bool) -> Dict[str, Any]:
cand_dir = REF if is_reference else RESULTS / method
stems = sorted(p.stem for p in cand_dir.glob("*.json"))
latencies, cers, ious, missed, extra = [], [], [], [], []
for stem in stems:
cj, ctxt = _load_page(cand_dir, stem)
if cj.get("elapsed_seconds") is not None:
latencies.append(cj["elapsed_seconds"])
rj, rtxt = _load_page(REF, stem)
cers.append(cer(rtxt, ctxt))
m = match_boxes(rj.get("boxes", []), cj.get("boxes", []), iou_threshold=0.5)
ious.append(m["mean_matched_iou"])
missed.append(m["missed"])
extra.append(m["extra"])
latencies.sort()
n = len(latencies)
model_dir = WORK / "models" / method
size = None
if model_dir.exists():
size = round(sum(f.stat().st_size for f in model_dir.rglob("*") if f.is_file()) / (1024 * 1024), 1)
avg = lambda xs: round(sum(xs) / len(xs), 4) if xs else None
return {
"n_pages": len(stems),
"mean_latency_s": round(sum(latencies) / n, 2) if n else None,
"p50_latency_s": round(latencies[n // 2], 2) if n else None,
"p95_latency_s": round(latencies[min(n - 1, int(n * 0.95))], 2) if n else None,
"mean_cer": avg(cers),
"max_cer": round(max(cers), 4) if cers else None,
"mean_bbox_iou": avg(ious),
"mean_missed_lines": avg(missed),
"mean_extra_lines": avg(extra),
"model_size_mb": size,
}
def main() -> None:
rows = []
for spec in SPECS:
if spec["status"] != "ok":
rows.append(build_row(method=spec["method"], status="failed",
t4_deployable=spec["t4_deployable"], error=spec["error"]))
continue
cand_dir = REF if spec.get("is_reference") else RESULTS / spec["method"]
if not any(cand_dir.glob("*.json")):
rows.append(build_row(method=spec["method"], status="failed",
t4_deployable=spec["t4_deployable"],
error=spec.get("fallback_error", "no captures on disk")))
continue
s = _score(spec["method"], spec.get("is_reference", False))
rows.append(build_row(
method=spec["method"], status="ok", t4_deployable=spec["t4_deployable"],
model_size_mb=s["model_size_mb"], mean_latency_s=s["mean_latency_s"],
p50_latency_s=s["p50_latency_s"], p95_latency_s=s["p95_latency_s"],
mean_cer=s["mean_cer"], max_cer=s["max_cer"], mean_bbox_iou=s["mean_bbox_iou"],
mean_missed_lines=s["mean_missed_lines"], mean_extra_lines=s["mean_extra_lines"],
))
print(f"{spec['method']}: {s['n_pages']} pages, "
f"mean_latency={s['mean_latency_s']}s, mean_cer={s['mean_cer']}, iou={s['mean_bbox_iou']}")
(WORK / "summary.csv").write_text(rows_to_csv(rows), encoding="utf-8")
(WORK / "summary.json").write_text(json.dumps(rows, indent=2), encoding="utf-8")
(WORK / "summary.md").write_text(rows_to_markdown(rows), encoding="utf-8")
print(f"wrote summary for {len(rows)} methods to {WORK}")
if __name__ == "__main__":
main()
+18
View File
@@ -0,0 +1,18 @@
"""Load the fixed eval-set manifest (one image filename per line)."""
from __future__ import annotations
from pathlib import Path
from typing import List
def load_manifest(manifest_path: Path, image_root: Path) -> List[Path]:
paths: List[Path] = []
for raw in manifest_path.read_text(encoding="utf-8").splitlines():
name = raw.strip()
if not name or name.startswith("#"):
continue
image = image_root / name
if not image.exists():
raise FileNotFoundError(f"manifest image not found: {name} (looked in {image_root})")
paths.append(image)
return paths
+78
View File
@@ -0,0 +1,78 @@
"""Speed-vs-accuracy Pareto scatter and per-method bar charts."""
from __future__ import annotations
from pathlib import Path
from typing import Any, Dict, List
import matplotlib
matplotlib.use("Agg") # headless
import matplotlib.pyplot as plt # noqa: E402
def pareto_points(rows: List[Dict[str, Any]], x_key: str, y_key: str) -> List[Dict[str, Any]]:
points = []
for row in rows:
if row.get("status") != "ok":
continue
if row.get(x_key) is None or row.get(y_key) is None:
continue
points.append({
"method": row["method"],
"x": row[x_key],
"y": row[y_key],
"t4_deployable": row.get("t4_deployable"),
})
return points
def _scatter(points, x_label, y_label, title, out_path: Path) -> None:
fig, ax = plt.subplots(figsize=(7, 5))
for p in points:
marker = "o" if p["t4_deployable"] else "x"
ax.scatter(p["x"], p["y"], marker=marker, s=80)
ax.annotate(p["method"], (p["x"], p["y"]), textcoords="offset points", xytext=(5, 5))
ax.set_xlabel(x_label)
ax.set_ylabel(y_label)
ax.set_title(title + " (o = T4-deployable, x = A100-only)")
fig.tight_layout()
fig.savefig(out_path, dpi=120)
plt.close(fig)
def _bar(rows, metric_key, y_label, out_path: Path) -> None:
ok = [r for r in rows if r.get("status") == "ok" and r.get(metric_key) is not None]
fig, ax = plt.subplots(figsize=(7, 5))
ax.bar([r["method"] for r in ok], [r[metric_key] for r in ok])
ax.set_ylabel(y_label)
ax.set_title(metric_key)
fig.tight_layout()
fig.savefig(out_path, dpi=120)
plt.close(fig)
def render_all(rows: List[Dict[str, Any]], out_dir: Path) -> List[Path]:
out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
written: List[Path] = []
p1 = out_dir / "pareto_latency_cer.png"
_scatter(pareto_points(rows, "mean_latency_s", "mean_cer"),
"mean latency (s) [lower=faster]", "CER vs BF16 [lower=better]",
"Speed vs recognition accuracy", p1)
written.append(p1)
p2 = out_dir / "pareto_latency_iou.png"
_scatter(pareto_points(rows, "mean_latency_s", "mean_bbox_iou"),
"mean latency (s) [lower=faster]", "bbox IoU vs BF16 [higher=better]",
"Speed vs detection accuracy", p2)
written.append(p2)
for metric, label in [("mean_latency_s", "mean latency (s)"),
("mean_cer", "CER vs BF16"),
("mean_bbox_iou", "bbox IoU vs BF16")]:
bp = out_dir / f"bar_{metric}.png"
_bar(rows, metric, label, bp)
written.append(bp)
return written
+50
View File
@@ -0,0 +1,50 @@
"""Per-method quantization specs and vLLM serve-argument construction.
kind:
baseline - serve the unquantized model as-is (the accuracy reference)
online - vLLM quantizes at load (fp8 dynamic); serve the base model
compressor - llm-compressor produced a checkpoint; quant config travels with it
bnb - transformers+bitsandbytes produced a checkpoint; serve with bnb flags
"""
from __future__ import annotations
from typing import Dict, List
METHOD_SPECS: Dict[str, Dict] = {
"bf16": {"kind": "baseline", "t4_deployable": True},
"fp8": {"kind": "online", "vllm_quant": "fp8", "t4_deployable": False},
# smoothquant=False: SmoothQuant's default mappings fail to resolve for the
# qwen3_5 layer layout (each mapping matches all 24 input_layernorms), so
# int8 is plain GPTQ W8A8.
"int8": {"kind": "compressor", "modifier": "gptq", "scheme": "W8A8", "smoothquant": False, "t4_deployable": True},
"awq": {"kind": "compressor", "modifier": "awq", "scheme": "W4A16", "smoothquant": False, "t4_deployable": True},
"gptq": {"kind": "compressor", "modifier": "gptq", "scheme": "W4A16", "smoothquant": False, "t4_deployable": True},
"bnb8": {"kind": "bnb", "bits": 8, "t4_deployable": True},
"bnb4": {"kind": "bnb", "bits": 4, "t4_deployable": True},
}
def method_names() -> List[str]:
return list(METHOD_SPECS.keys())
def vllm_serve_args(method: str, model_path: str, base_model: str, port: int) -> List[str]:
spec = METHOD_SPECS[method]
kind = spec["kind"]
serve_model = base_model if kind in ("baseline", "online") else model_path
args = [
"--host", "127.0.0.1",
"--port", str(port),
"--model", serve_model,
"--served-model-name", "datalab-to/surya-ocr-2",
"--max-model-len", "18000",
"--max-num-seqs", "16",
"--gpu-memory-utilization", "0.85",
"--enable-prefix-caching",
"--mm-processor-kwargs", '{"min_pixels":3136,"max_pixels":6291456}',
]
if kind == "online":
args += ["--quantization", spec["vllm_quant"]]
elif kind == "bnb":
args += ["--quantization", "bitsandbytes", "--load-format", "bitsandbytes"]
return args
+111
View File
@@ -0,0 +1,111 @@
"""Orchestrate the quantization sweep: baseline -> each method -> aggregate -> plot.
Each method is isolated: any exception becomes a status="failed" row so one bad
method never aborts the sweep.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Any, Dict, List
from scripts.quant.aggregate import build_row, rows_to_csv, rows_to_markdown
from scripts.quant.bbox_iou import bbox_iou_over_dirs
from scripts.quant.build_model import build_model
from scripts.quant.capture import capture_page
from scripts.quant.manifest import load_manifest
from scripts.quant.recipes import METHOD_SPECS, method_names, vllm_serve_args
from scripts.quant.serve import start_server, stop_server, wait_healthy
# Reuse the existing CER metric.
from scripts.cer_divergence import cer_over_dirs
OCR_URL = "http://127.0.0.1:5002/v1/api/ai/suya_ocr_vllm/"
def _dir_size_mb(path: Path) -> float:
total = sum(f.stat().st_size for f in path.rglob("*") if f.is_file())
return round(total / (1024 * 1024), 1)
def _measure_method(method: str, base_model: str, work_dir: Path,
eval_images: List[Path], reference_dir: Path) -> Dict[str, Any]:
"""Build -> serve -> capture -> metrics for one method. Raises on any failure.
Assumes the API process (api.py) is already running and reads SURYA_INFERENCE_URL
from the env to point at the per-method vLLM server (port set by the runbook).
Returns a dict of measured summary fields (no method/status/t4 keys).
"""
model_dir = work_dir / "models" / method
out_capture = work_dir / "results" / method
port = 8000
built = build_model(method, base_model, model_dir, eval_images[: min(8, len(eval_images))])
serve_args = vllm_serve_args(method, str(built), base_model, port)
proc = start_server(serve_args, str(work_dir / f"vllm_{method}.log"))
try:
if not wait_healthy(port):
raise RuntimeError(f"vLLM did not become healthy for {method}")
latencies = []
for image in eval_images:
cap = capture_page(OCR_URL, image, out_capture, timeout=900)
if cap["elapsed_seconds"] is not None:
latencies.append(cap["elapsed_seconds"])
cer = cer_over_dirs(reference_dir, out_capture)
iou = bbox_iou_over_dirs(reference_dir, out_capture)
finally:
stop_server(port, proc)
latencies.sort()
n = len(latencies)
size = _dir_size_mb(built) if METHOD_SPECS[method]["kind"] in ("compressor", "bnb") else None
return {
"mean_latency_s": round(sum(latencies) / n, 3) if n else None,
"p50_latency_s": latencies[n // 2] if n else None,
"p95_latency_s": latencies[min(n - 1, int(n * 0.95))] if n else None,
"mean_cer": cer["mean_cer"],
"max_cer": cer["max_cer"],
"mean_bbox_iou": iou["mean_bbox_iou"],
"mean_missed_lines": iou["mean_missed_lines"],
"mean_extra_lines": iou["mean_extra_lines"],
"model_size_mb": size,
}
def run_method(method: str, base_model: str, work_dir: Path,
eval_images: List[Path], reference_dir: Path) -> Dict[str, Any]:
t4 = METHOD_SPECS[method]["t4_deployable"]
try:
metrics = _measure_method(method, base_model, work_dir, eval_images, reference_dir)
return build_row(method=method, status="ok", t4_deployable=t4, **metrics)
except Exception as exc: # feasibility probe: record and continue
return build_row(method=method, status="failed", t4_deployable=t4, error=str(exc))
def main() -> None:
parser = argparse.ArgumentParser(description="Run the quantization benchmark sweep.")
parser.add_argument("--base-model", default="datalab-to/surya-ocr-2")
parser.add_argument("--manifest", type=Path, default=Path("eval_set/manifest.txt"))
parser.add_argument("--image-root", type=Path, default=Path("."))
parser.add_argument("--work-dir", type=Path, default=Path("results/quant"))
parser.add_argument("--methods", nargs="*", default=method_names())
args = parser.parse_args()
eval_images = load_manifest(args.manifest, args.image_root)
reference_dir = args.work_dir / "reference"
args.work_dir.mkdir(parents=True, exist_ok=True)
rows = []
for method in args.methods:
print(f"=== {method} ===", flush=True)
rows.append(run_method(method, args.base_model, args.work_dir, eval_images, reference_dir))
(args.work_dir / "summary.csv").write_text(rows_to_csv(rows), encoding="utf-8")
(args.work_dir / "summary.json").write_text(json.dumps(rows, indent=2), encoding="utf-8")
(args.work_dir / "summary.md").write_text(rows_to_markdown(rows), encoding="utf-8")
print(f"wrote summary for {len(rows)} methods to {args.work_dir}")
if __name__ == "__main__":
main()
+60
View File
@@ -0,0 +1,60 @@
#!/usr/bin/env bash
# Serve and score compressor-built checkpoints (int8/gptq/awq) over a page subset.
# Each method: hardened server teardown -> serve checkpoint -> health/fail watch ->
# capture N_PAGES through the OCR API. Run INSIDE the bench container:
# bash scripts/quant/score_methods.sh int8 gptq awq
# Requires api.py running on :5002 pointed at :8000.
set -uo pipefail
export LD_LIBRARY_PATH="/usr/local/cuda/compat:${LD_LIBRARY_PATH:-}"
N_PAGES="${N_PAGES:-3}"
OCR_URL="http://127.0.0.1:5002/v1/api/ai/suya_ocr_vllm/"
LOG=/tmp/score_methods.log
mapfile -t ALL < <(sed '/^#/d;/^$/d' eval_set/manifest.txt)
IMAGES=("${ALL[@]:0:$N_PAGES}")
stop_server() {
local pid
pid=$(ss -ltnp 2>/dev/null | grep ":8000 " | grep -oP 'pid=\K[0-9]+' | head -1)
[ -n "$pid" ] && kill -9 "$pid" 2>/dev/null
pkill -9 -f "vllm.entrypoints" 2>/dev/null
# EngineCore outlives the API server and holds the GPU; kill compute procs and
# poll until memory actually frees (kill returns before GPU release).
for p in $(nvidia-smi --query-compute-apps=pid --format=csv,noheader 2>/dev/null); do
kill -9 "$p" 2>/dev/null
done
for _ in $(seq 1 30); do
used=$(nvidia-smi --query-gpu=memory.used --format=csv,noheader,nounits | head -1)
[ "${used:-99999}" -lt 2000 ] && break
sleep 2
done
}
for m in "$@"; do
echo "=== scoring $m ===" | tee -a "$LOG"
stop_server
# VLLM_EXTRA_ARGS lets a caller add e.g. --dtype float16 (Exllama W4A16 kernel
# on Ampere only supports float16 activations).
nohup python3 -m vllm.entrypoints.openai.api_server --host 127.0.0.1 --port 8000 \
--model "results/quant/models/$m" --served-model-name datalab-to/surya-ocr-2 \
--max-model-len 18000 --max-num-seqs 16 --gpu-memory-utilization 0.85 \
--enable-prefix-caching --mm-processor-kwargs '{"min_pixels":3136,"max_pixels":6291456}' \
${VLLM_EXTRA_ARGS:-} \
> "/tmp/serve_${m}.log" 2>&1 &
ok=0
for _ in $(seq 1 84); do
if curl -fs http://127.0.0.1:8000/health >/dev/null 2>&1; then ok=1; break; fi
if grep -qiE "Engine core initialization failed" "/tmp/serve_${m}.log" 2>/dev/null; then break; fi
sleep 5
done
if [ "$ok" -ne 1 ]; then
echo "$m: FAILED to start" | tee -a "$LOG"
continue
fi
python3 -m scripts.quant.capture --url "$OCR_URL" --out-dir "results/quant/results/$m" "${IMAGES[@]}" \
>> "$LOG" 2>&1
echo "$m: captured $(ls "results/quant/results/$m"/*.json 2>/dev/null | wc -l) pages" | tee -a "$LOG"
done
stop_server
echo "SCORE_DONE" | tee -a "$LOG"
+59
View File
@@ -0,0 +1,59 @@
"""Launch a vLLM OpenAI server for one quant variant, health-check it, and tear
it down by port (never by command-line match — that bit us in Approach A)."""
from __future__ import annotations
import os
import re
import subprocess
import time
import urllib.request
from typing import List, Optional
# Default seconds to wait for a server to become healthy. Overridable via env so a
# method that fails to boot is recorded quickly instead of blocking the sweep.
DEFAULT_HEALTH_TIMEOUT = float(os.environ.get("SURYA_QUANT_HEALTH_TIMEOUT", "900"))
def health_url(port: int) -> str:
return f"http://127.0.0.1:{port}/health"
def parse_listening_pid(ss_output: str, port: int) -> Optional[int]:
for line in ss_output.splitlines():
if f":{port} " in line or line.rstrip().endswith(f":{port}"):
m = re.search(r"pid=(\d+)", line)
if m:
return int(m.group(1))
return None
def wait_healthy(port: int, timeout: float = DEFAULT_HEALTH_TIMEOUT) -> bool:
deadline = time.time() + timeout
url = health_url(port)
while time.time() < deadline:
try:
with urllib.request.urlopen(url, timeout=2) as resp:
if resp.status == 200:
return True
except Exception:
time.sleep(3)
return False
def start_server(serve_args: List[str], log_path: str) -> subprocess.Popen:
cmd = ["python", "-m", "vllm.entrypoints.openai.api_server", *serve_args]
log = open(log_path, "w")
return subprocess.Popen(cmd, stdout=log, stderr=subprocess.STDOUT)
def stop_server(port: int, proc: Optional[subprocess.Popen] = None) -> None:
if proc is not None:
proc.terminate()
try:
proc.wait(timeout=30)
except Exception:
proc.kill()
out = subprocess.run(["ss", "-ltnp"], capture_output=True, text=True).stdout
pid = parse_listening_pid(out, port)
if pid:
subprocess.run(["kill", str(pid)])
+76
View File
@@ -0,0 +1,76 @@
#!/usr/bin/env bash
# MTP (multi-token-prediction) speculative-decode sweep for the latency report.
# Serves the BF16 model three ways — no speculation (baseline), MTP with 1
# speculative token, MTP with 2 — and captures the same page subset through the
# OCR API for each. The model ships 1 nextn-predict layer, so MTP=1 is the
# expected-valid setting and MTP=2 is exploratory.
#
# Run INSIDE the bench container: bash scripts/quant/tune_mtp.sh
# Requires: api.py running on :5002 pointed at :8000; compat libs on path.
set -uo pipefail
MODEL="${MODEL:-datalab-to/surya-ocr-2}"
OCR_URL="http://127.0.0.1:5002/v1/api/ai/suya_ocr_vllm/"
N_PAGES="${N_PAGES:-5}"
OUT_ROOT="results/quant/mtp"
LOG=/tmp/tune_mtp.log
export LD_LIBRARY_PATH="/usr/local/cuda/compat:${LD_LIBRARY_PATH:-}"
mapfile -t ALL < <(sed '/^#/d;/^$/d' eval_set/manifest.txt)
IMAGES=("${ALL[@]:0:$N_PAGES}")
# config_name | extra vLLM args. The speculative-config JSON is written compact
# (no spaces) so it survives word-splitting as a single argv token.
CONFIGS=(
"baseline|--enable-prefix-caching"
"mtp1|--enable-prefix-caching --speculative-config {\"method\":\"mtp\",\"num_speculative_tokens\":1}"
"mtp2|--enable-prefix-caching --speculative-config {\"method\":\"mtp\",\"num_speculative_tokens\":2}"
)
base_args() {
echo "--host 127.0.0.1 --port 8000 --model $MODEL --served-model-name $MODEL \
--max-model-len 18000 --max-num-seqs 16 --gpu-memory-utilization 0.85 --no-enforce-eager \
--mm-processor-kwargs {\"min_pixels\":3136,\"max_pixels\":6291456}"
}
stop_server() {
local pid
pid=$(ss -ltnp 2>/dev/null | grep ":8000 " | grep -oP 'pid=\K[0-9]+' | head -1)
[ -n "$pid" ] && kill -9 "$pid" 2>/dev/null
pkill -9 -f "vllm.entrypoints" 2>/dev/null
# EngineCore outlives the API server and holds the GPU; kill every remaining
# CUDA compute process, then poll until the memory is actually freed.
for p in $(nvidia-smi --query-compute-apps=pid --format=csv,noheader 2>/dev/null); do
kill -9 "$p" 2>/dev/null
done
for _ in $(seq 1 30); do
used=$(nvidia-smi --query-gpu=memory.used --format=csv,noheader,nounits | head -1)
[ "${used:-99999}" -lt 2000 ] && break
sleep 2
done
}
for entry in "${CONFIGS[@]}"; do
name="${entry%%|*}"
extra="${entry#*|}"
echo "=== mtp config: $name ($extra) ===" | tee -a "$LOG"
stop_server
# shellcheck disable=SC2046
nohup python3 -m vllm.entrypoints.openai.api_server $(base_args) $extra \
> "/tmp/mtp_${name}.log" 2>&1 &
ok=0
for _ in $(seq 1 96); do
if curl -fs http://127.0.0.1:8000/health >/dev/null 2>&1; then ok=1; break; fi
if grep -qiE "Engine core initialization failed|ValueError|RuntimeError|Traceback" "/tmp/mtp_${name}.log" 2>/dev/null; then break; fi
sleep 5
done
if [ "$ok" -ne 1 ]; then
echo "$name: FAILED to start (see /tmp/mtp_${name}.log)" | tee -a "$LOG"
continue
fi
python3 -m scripts.quant.capture --url "$OCR_URL" --out-dir "$OUT_ROOT/$name" "${IMAGES[@]}" \
>> "$LOG" 2>&1
echo "$name: captured $(ls "$OUT_ROOT/$name"/*.json 2>/dev/null | wc -l) pages" | tee -a "$LOG"
done
stop_server
echo "MTP_DONE" | tee -a "$LOG"
+77
View File
@@ -0,0 +1,77 @@
#!/usr/bin/env bash
# vLLM serving-parameter tuning sweep for the OCR latency report.
# Serves the BF16 model under several vLLM flag combinations (CUDA graph vs eager,
# prefix caching on/off, chunked prefill on/off), captures a fixed page subset
# through the OCR API for each, and records per-config latency.
#
# Run INSIDE the bench container: bash scripts/quant/tune_vllm.sh
# Requires: api.py already running on :5002 pointed at :8000; compat libs on path.
set -uo pipefail
MODEL="${MODEL:-datalab-to/surya-ocr-2}"
OCR_URL="http://127.0.0.1:5002/v1/api/ai/suya_ocr_vllm/"
N_PAGES="${N_PAGES:-3}"
OUT_ROOT="results/quant/tuning"
LOG=/tmp/tune_vllm.log
export LD_LIBRARY_PATH="/usr/local/cuda/compat:${LD_LIBRARY_PATH:-}"
mapfile -t ALL < <(sed '/^#/d;/^$/d' eval_set/manifest.txt)
IMAGES=("${ALL[@]:0:$N_PAGES}")
# config_name | extra vLLM args
CONFIGS=(
"default|--enable-prefix-caching"
"eager|--enable-prefix-caching --enforce-eager"
"no_prefix_cache|--no-enable-prefix-caching"
"no_chunked_prefill|--enable-prefix-caching --no-enable-chunked-prefill"
)
base_args() {
echo "--host 127.0.0.1 --port 8000 --model $MODEL --served-model-name $MODEL \
--max-model-len 18000 --max-num-seqs 16 --gpu-memory-utilization 0.85 \
--mm-processor-kwargs {\"min_pixels\":3136,\"max_pixels\":6291456}"
}
stop_server() {
local pid
pid=$(ss -ltnp 2>/dev/null | grep ":8000 " | grep -oP 'pid=\K[0-9]+' | head -1)
[ -n "$pid" ] && kill -9 "$pid" 2>/dev/null
pkill -9 -f "vllm.entrypoints" 2>/dev/null
# The EngineCore worker holds the GPU and outlives the API server; kill every
# remaining CUDA compute process, then poll until the memory is actually freed
# (kill returns immediately, GPU release lags).
for p in $(nvidia-smi --query-compute-apps=pid --format=csv,noheader 2>/dev/null); do
kill -9 "$p" 2>/dev/null
done
for _ in $(seq 1 30); do
used=$(nvidia-smi --query-gpu=memory.used --format=csv,noheader,nounits | head -1)
[ "${used:-99999}" -lt 2000 ] && break
sleep 2
done
}
for entry in "${CONFIGS[@]}"; do
name="${entry%%|*}"
extra="${entry#*|}"
echo "=== tuning config: $name ($extra) ===" | tee -a "$LOG"
stop_server
# shellcheck disable=SC2046
nohup python3 -m vllm.entrypoints.openai.api_server $(base_args) $extra \
> "/tmp/tune_${name}.log" 2>&1 &
# wait for health or crash (max 7 min)
ok=0
for _ in $(seq 1 84); do
if curl -fs http://127.0.0.1:8000/health >/dev/null 2>&1; then ok=1; break; fi
if grep -qiE "Engine core initialization failed|RuntimeError" "/tmp/tune_${name}.log" 2>/dev/null; then break; fi
sleep 5
done
if [ "$ok" -ne 1 ]; then
echo "$name: FAILED to start" | tee -a "$LOG"
continue
fi
python3 -m scripts.quant.capture --url "$OCR_URL" --out-dir "$OUT_ROOT/$name" "${IMAGES[@]}" \
>> "$LOG" 2>&1
echo "$name: captured $(ls "$OUT_ROOT/$name"/*.json 2>/dev/null | wc -l) pages" | tee -a "$LOG"
done
stop_server
echo "TUNE_DONE" | tee -a "$LOG"
+77
View File
@@ -0,0 +1,77 @@
#!/usr/bin/env bash
set -euo pipefail
export SURYA_MODEL_CHECKPOINT="${SURYA_MODEL_CHECKPOINT:-datalab-to/surya-ocr-2}"
export SURYA_INFERENCE_BACKEND="${SURYA_INFERENCE_BACKEND:-vllm}"
export SURYA_INFERENCE_URL="${SURYA_INFERENCE_URL:-http://127.0.0.1:8000/v1}"
export SURYA_INFERENCE_AUTOSTART="${SURYA_INFERENCE_AUTOSTART:-false}"
export SURYA_INFERENCE_PARALLEL="${SURYA_INFERENCE_PARALLEL:-8}"
export SURYA_INFERENCE_LOGPROBS="${SURYA_INFERENCE_LOGPROBS:-false}"
export SURYA_INFERENCE_MAX_RETRIES="${SURYA_INFERENCE_MAX_RETRIES:-1}"
# Keep in sync with VLLM_MAX_NUM_SEQS so the batcher fills every vLLM sequence
# slot (see docs/diagnosis_baseline.md). Overcommitting past it only queues.
export SURYA_INFERENCE_MAX_INFLIGHT="${SURYA_INFERENCE_MAX_INFLIGHT:-16}"
export SURYA_MAX_TOKENS_FULL_PAGE="${SURYA_MAX_TOKENS_FULL_PAGE:-6144}"
export SURYA_MAX_BLOCKS_PER_PAGE="${SURYA_MAX_BLOCKS_PER_PAGE:-80}"
export SUYA_OCR_MODE="${SUYA_OCR_MODE:-block}"
export SUYA_MAX_BATCH_SIZE="${SUYA_MAX_BATCH_SIZE:-8}"
export SUYA_BATCH_WAIT_MS="${SUYA_BATCH_WAIT_MS:-25}"
export SUYA_VLLM_IMAGE_FORMAT="${SUYA_VLLM_IMAGE_FORMAT:-JPEG}"
export SUYA_VLLM_JPEG_QUALITY="${SUYA_VLLM_JPEG_QUALITY:-92}"
export VLLM_DTYPE="${VLLM_DTYPE:-float16}"
export VLLM_GPU_MEMORY_UTILIZATION="${VLLM_GPU_MEMORY_UTILIZATION:-0.85}"
export VLLM_MAX_MODEL_LEN="${VLLM_MAX_MODEL_LEN:-18000}"
export VLLM_MAX_NUM_SEQS="${VLLM_MAX_NUM_SEQS:-16}"
export VLLM_MAX_BATCHED_TOKENS="${VLLM_MAX_BATCHED_TOKENS:-4096}"
# vLLM serving parameters are tuned for latency — see
# docs/quantization_benchmark_results.md §5 (A100, BF16). Three flags carry the
# win and MUST stay on; do not pass their negations via VLLM_EXTRA_ARGS:
# --no-enforce-eager CUDA-graph capture. The single biggest lever: eager
# mode (--enforce-eager) measured ~9.5x slower
# (48.97s vs 5.15s/page). Set explicitly so a future
# vLLM default flip can't silently disable graphs.
# --enable-prefix-caching ~15% win; OCR prompts share a long fixed prefix.
# chunked prefill (on by default) is REQUIRED — the qwen3_5 mamba/SSM cache
# fails engine init with --no-enable-chunked-prefill.
python3 -m vllm.entrypoints.openai.api_server \
--host 127.0.0.1 \
--port 8000 \
--model "${SURYA_MODEL_CHECKPOINT}" \
--served-model-name "${SURYA_MODEL_CHECKPOINT}" \
--dtype "${VLLM_DTYPE}" \
--max-model-len "${VLLM_MAX_MODEL_LEN}" \
--max-num-seqs "${VLLM_MAX_NUM_SEQS}" \
--max-num-batched-tokens "${VLLM_MAX_BATCHED_TOKENS}" \
--gpu-memory-utilization "${VLLM_GPU_MEMORY_UTILIZATION}" \
--no-enforce-eager \
--enable-prefix-caching \
--mm-processor-kwargs '{"min_pixels":3136,"max_pixels":6291456}' \
${VLLM_EXTRA_ARGS:-} &
VLLM_PID=$!
cleanup() {
kill "${VLLM_PID}" 2>/dev/null || true
}
trap cleanup EXIT INT TERM
python3 - <<'PY'
import os
import time
import urllib.request
base = os.environ.get("SURYA_INFERENCE_URL", "http://127.0.0.1:8000/v1")
health = base[:-3] + "/health" if base.endswith("/v1") else base.rstrip("/") + "/health"
deadline = time.time() + float(os.environ.get("SURYA_INFERENCE_STARTUP_TIMEOUT", "900"))
while time.time() < deadline:
try:
with urllib.request.urlopen(health, timeout=2) as response:
if response.status == 200:
print(f"vLLM health check passed: {health}", flush=True)
raise SystemExit(0)
except Exception:
time.sleep(2)
raise SystemExit(f"vLLM did not become healthy: {health}")
PY
python3 api.py