Files
surya-ocr/vllm_batcher.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

135 lines
4.5 KiB
Python

import logging
import os
import queue
import threading
import time
from concurrent.futures import Future
from dataclasses import dataclass
from typing import Tuple
from PIL import Image
from surya.timing import (
TimingCollector,
log_timing_summary,
reset_current_timing,
set_current_timing,
timing_span,
)
from vllm_tools import ocr_vllm_batch
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class OcrOptions:
skip_text_detection: bool
recognize_math: bool
with_bboxes: bool
@dataclass
class OcrJob:
image: Image.Image
highres_image: Image.Image
options: OcrOptions
future: Future
request_id: str
class VllmOcrBatcher:
def __init__(self) -> None:
self.max_batch_size = int(os.getenv("SUYA_MAX_BATCH_SIZE", "8"))
self.batch_wait_ms = float(os.getenv("SUYA_BATCH_WAIT_MS", "25"))
self.queue_timeout_seconds = float(os.getenv("SUYA_BATCH_QUEUE_TIMEOUT_SECONDS", "900"))
self._queue: queue.Queue[OcrJob] = queue.Queue(
maxsize=int(os.getenv("SUYA_MAX_QUEUE_SIZE", "128"))
)
self._thread = threading.Thread(target=self._run, name="vllm-ocr-batcher", daemon=True)
self._thread.start()
def submit(
self,
image: Image.Image,
highres_image: Image.Image,
*,
skip_text_detection: bool,
recognize_math: bool,
with_bboxes: bool,
request_id: str,
):
future: Future = Future()
job = OcrJob(
image=image,
highres_image=highres_image,
options=OcrOptions(skip_text_detection, recognize_math, with_bboxes),
future=future,
request_id=request_id,
)
with timing_span("batcher_queue_put", request_id=request_id):
self._queue.put(job, timeout=self.queue_timeout_seconds)
with timing_span("batcher_wait_result", request_id=request_id):
return future.result(timeout=self.queue_timeout_seconds)
def _run(self) -> None:
while True:
first = self._queue.get()
jobs = [first]
deadline = time.perf_counter() + (self.batch_wait_ms / 1000.0)
while len(jobs) < self.max_batch_size:
remaining = deadline - time.perf_counter()
if remaining <= 0:
break
try:
jobs.append(self._queue.get(timeout=remaining))
except queue.Empty:
break
self._process_jobs(jobs)
def _process_jobs(self, jobs: list[OcrJob]) -> None:
groups: dict[OcrOptions, list[OcrJob]] = {}
for job in jobs:
groups.setdefault(job.options, []).append(job)
for options, grouped_jobs in groups.items():
start = time.perf_counter()
collector = TimingCollector(
request_id=",".join(job.request_id for job in grouped_jobs),
batch_size=len(grouped_jobs),
)
token = set_current_timing(collector)
try:
with timing_span("batcher_process_jobs", batch_size=len(grouped_jobs)):
results = ocr_vllm_batch(
[job.image for job in grouped_jobs],
[job.highres_image for job in grouped_jobs],
skip_text_detection=options.skip_text_detection,
recognize_math=options.recognize_math,
with_bboxes=options.with_bboxes,
)
for job, result in zip(grouped_jobs, results):
job.future.set_result(result)
logger.info(
"vllm_batch_complete batch_size=%s option=%s duration_ms=%.2f request_ids=%s",
len(grouped_jobs),
options,
(time.perf_counter() - start) * 1000,
",".join(job.request_id for job in grouped_jobs),
)
log_timing_summary(logger, collector)
except Exception as exc:
for job in grouped_jobs:
job.future.set_exception(exc)
logger.exception(
"vllm_batch_failed batch_size=%s option=%s request_ids=%s",
len(grouped_jobs),
options,
",".join(job.request_id for job in grouped_jobs),
exc_info=True,
)
finally:
reset_current_timing(token)
vllm_ocr_batcher = VllmOcrBatcher()