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:
@@ -0,0 +1,404 @@
|
||||
"""RecognitionPredictor: per-block OCR via BLOCK_PROMPT.
|
||||
|
||||
Given page images and corresponding LayoutResult (or any list of LayoutBox),
|
||||
crops each block, runs BLOCK_PROMPT, returns PageOCRResult per page.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from surya.common.blank import is_blank_region
|
||||
from surya.inference import SuryaInferenceManager, get_default_manager
|
||||
from surya.inference.parsers import clean_block_html, parse_full_page_html
|
||||
from surya.inference.prompts import (
|
||||
PROMPT_TYPE_BLOCK,
|
||||
PROMPT_TYPE_HIGH_ACCURACY_BBOX,
|
||||
SKIP_OCR_LABELS,
|
||||
)
|
||||
from surya.inference.schema import BatchInputItem
|
||||
from surya.inference.util import image_token_budget
|
||||
from surya.layout.label import LAYOUT_PRED_RELABEL, TEXT_LABELS
|
||||
from surya.layout.schema import LayoutResult
|
||||
from surya.logging import get_logger
|
||||
from surya.recognition.schema import (
|
||||
BlockOCRResult,
|
||||
PageOCRResult,
|
||||
)
|
||||
from surya.settings import settings
|
||||
from surya.timing import timing_span
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
# Surya's canonical labels we shouldn't OCR (mirrors model-emitted SKIP_OCR_LABELS
|
||||
# after canonicalization).
|
||||
SKIP_CANON_LABELS = {LAYOUT_PRED_RELABEL.get(lbl, lbl) for lbl in SKIP_OCR_LABELS}
|
||||
|
||||
|
||||
def _crop_block(image: Image.Image, polygon, pad: int = 4) -> Image.Image:
|
||||
xs = [p[0] for p in polygon]
|
||||
ys = [p[1] for p in polygon]
|
||||
x0 = max(0, int(min(xs)) - pad)
|
||||
y0 = max(0, int(min(ys)) - pad)
|
||||
x1 = min(image.size[0], int(max(xs)) + pad)
|
||||
y1 = min(image.size[1], int(max(ys)) + pad)
|
||||
if x1 <= x0 or y1 <= y0:
|
||||
return image.crop((0, 0, 1, 1))
|
||||
return image.crop((x0, y0, x1, y1))
|
||||
|
||||
|
||||
def _drop_blank_text_blocks(
|
||||
image: Image.Image,
|
||||
blocks: List[BlockOCRResult],
|
||||
) -> List[BlockOCRResult]:
|
||||
"""Drop text-labeled blocks whose source page region is essentially blank.
|
||||
|
||||
Full-page OCR can emit text divs for regions that are visually empty
|
||||
(margins, gutter space) — the model hallucinates a paragraph where there
|
||||
is none. We crop the region, count near-white pixels, and drop the block
|
||||
when the fraction exceeds ``blank_pixel_fraction``. Only text-like labels
|
||||
(see ``TEXT_LABELS``) are eligible: tables, forms, equations, and visual
|
||||
blocks may legitimately contain large whitespace and are left untouched.
|
||||
"""
|
||||
kept: List[BlockOCRResult] = []
|
||||
dropped = 0
|
||||
for blk in blocks:
|
||||
if blk.label not in TEXT_LABELS or blk.skipped or blk.error:
|
||||
kept.append(blk)
|
||||
continue
|
||||
crop = _crop_block(image, blk.polygon)
|
||||
if not is_blank_region(crop):
|
||||
kept.append(blk)
|
||||
continue
|
||||
dropped += 1
|
||||
if dropped:
|
||||
logger.info(f"dropped {dropped} blank text block(s) from full-page OCR")
|
||||
return kept
|
||||
|
||||
|
||||
def _detect_repeat_loop(
|
||||
text: str,
|
||||
base_max_repeats: int = 4,
|
||||
window_size: int = 500,
|
||||
scaling_factor: float = 3.0,
|
||||
) -> bool:
|
||||
"""True iff the tail of ``text`` ends in a repeating sequence.
|
||||
|
||||
Ported from chandra's detect_repeat_token. For each candidate length
|
||||
1..window_size/2, takes that many trailing chars and counts consecutive
|
||||
identical preceding blocks. Shorter loops need many repeats to count;
|
||||
longer ones only need a few. Catches the typical decoder failure mode
|
||||
where a page output gets stuck emitting the same div / phrase until it
|
||||
hits max_tokens.
|
||||
"""
|
||||
if not text:
|
||||
return False
|
||||
for seq_len in range(1, window_size // 2 + 1):
|
||||
candidate = text[-seq_len:]
|
||||
max_repeats = int(base_max_repeats * (1 + scaling_factor / seq_len))
|
||||
repeats = 0
|
||||
pos = len(text) - seq_len
|
||||
while pos >= 0 and text[pos : pos + seq_len] == candidate:
|
||||
repeats += 1
|
||||
pos -= seq_len
|
||||
if repeats > max_repeats:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class RecognitionPredictor:
|
||||
"""Per-block OCR. Construct with a SuryaInferenceManager (or rely on default)."""
|
||||
|
||||
def __init__(self, manager: Optional[SuryaInferenceManager] = None):
|
||||
self.manager = manager
|
||||
self._disable_tqdm = settings.DISABLE_TQDM
|
||||
|
||||
@property
|
||||
def disable_tqdm(self) -> bool:
|
||||
return self._disable_tqdm
|
||||
|
||||
@disable_tqdm.setter
|
||||
def disable_tqdm(self, value: bool) -> None:
|
||||
self._disable_tqdm = bool(value)
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
return
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
images: List[Image.Image],
|
||||
layout_results: Optional[List[LayoutResult]] = None,
|
||||
*,
|
||||
full_page: Optional[bool] = None,
|
||||
) -> List[PageOCRResult]:
|
||||
"""Run OCR on each page.
|
||||
|
||||
Mode resolution:
|
||||
- ``full_page=None`` (default): block mode if ``layout_results`` is
|
||||
given, else full-page mode. This is the most-do-what-I-mean form.
|
||||
- ``full_page=True``: full-page OCR (single HIGH_ACCURACY_BBOX_PROMPT
|
||||
request per page). ``layout_results`` is ignored — a warning is
|
||||
logged if it was supplied.
|
||||
- ``full_page=False``: block mode (per-layout-block OCR request).
|
||||
``layout_results`` is required.
|
||||
|
||||
Full-page is the more accurate path; block mode is for callers that
|
||||
specifically need per-block crops (e.g. for downstream merging with
|
||||
text-line detection).
|
||||
"""
|
||||
if not images:
|
||||
return []
|
||||
if full_page is None:
|
||||
full_page = layout_results is None
|
||||
if full_page:
|
||||
if layout_results is not None:
|
||||
logger.info(
|
||||
"RecognitionPredictor called with full_page=True and "
|
||||
"layout_results; layout will be used as fallback if the "
|
||||
"full-page output devolves into a repetition loop."
|
||||
)
|
||||
return self._full_page_ocr(images, fallback_layout=layout_results)
|
||||
if layout_results is None:
|
||||
raise ValueError("layout_results required when full_page=False")
|
||||
if len(images) != len(layout_results):
|
||||
raise ValueError(
|
||||
f"images and layout_results must be same length "
|
||||
f"({len(images)} vs {len(layout_results)})"
|
||||
)
|
||||
manager = self.manager or get_default_manager()
|
||||
|
||||
# Build a flat batch across all pages for max concurrency
|
||||
batch: List[BatchInputItem] = []
|
||||
block_index_map: List[tuple[int, int]] = [] # (page_idx, block_idx)
|
||||
skipped_flags: List[bool] = []
|
||||
|
||||
with timing_span("recognition_build_block_batch", page_count=len(images)):
|
||||
for page_idx, (img, layout) in enumerate(zip(images, layout_results)):
|
||||
page_boxes = sorted(layout.bboxes, key=lambda b: (b.position, b.bbox[1], b.bbox[0]))
|
||||
page_boxes = page_boxes[: settings.SURYA_MAX_BLOCKS_PER_PAGE]
|
||||
if len(layout.bboxes) > len(page_boxes):
|
||||
logger.info(
|
||||
f"capped OCR blocks for page {page_idx}: "
|
||||
f"{len(layout.bboxes)} -> {len(page_boxes)}"
|
||||
)
|
||||
for block_idx, box in enumerate(page_boxes):
|
||||
skip = box.label in SKIP_CANON_LABELS
|
||||
skipped_flags.append(skip)
|
||||
if skip:
|
||||
continue
|
||||
crop = _crop_block(img, box.polygon)
|
||||
max_tokens = image_token_budget(
|
||||
box.count, ceiling=settings.SURYA_MAX_TOKENS_BLOCK_CEILING
|
||||
)
|
||||
batch.append(
|
||||
BatchInputItem(
|
||||
image=crop,
|
||||
prompt_type=PROMPT_TYPE_BLOCK,
|
||||
max_tokens=max_tokens,
|
||||
metadata={"page_idx": page_idx, "block_idx": block_idx},
|
||||
)
|
||||
)
|
||||
block_index_map.append((page_idx, block_idx))
|
||||
|
||||
with timing_span(
|
||||
"recognition_manager_generate",
|
||||
item_count=len(batch),
|
||||
skipped_count=sum(1 for flag in skipped_flags if flag),
|
||||
max_tokens_sum=sum(item.max_tokens or 0 for item in batch),
|
||||
):
|
||||
outputs = manager.generate(batch) if batch else []
|
||||
|
||||
# Index outputs by (page_idx, block_idx)
|
||||
with timing_span(
|
||||
"recognition_assemble_pages",
|
||||
output_count=len(outputs),
|
||||
token_count=sum(out.token_count or 0 for out in outputs),
|
||||
):
|
||||
out_by_key = {}
|
||||
for out in outputs:
|
||||
key = (out.metadata["page_idx"], out.metadata["block_idx"])
|
||||
out_by_key[key] = out
|
||||
|
||||
# Assemble PageOCRResult per page
|
||||
results: List[PageOCRResult] = []
|
||||
for page_idx, (img, layout) in enumerate(zip(images, layout_results)):
|
||||
w, h = img.size
|
||||
blocks: List[BlockOCRResult] = []
|
||||
page_boxes = sorted(layout.bboxes, key=lambda b: (b.position, b.bbox[1], b.bbox[0]))
|
||||
page_boxes = page_boxes[: settings.SURYA_MAX_BLOCKS_PER_PAGE]
|
||||
for block_idx, box in enumerate(page_boxes):
|
||||
skip = box.label in SKIP_CANON_LABELS
|
||||
if skip:
|
||||
blocks.append(
|
||||
BlockOCRResult(
|
||||
polygon=box.polygon,
|
||||
label=box.label,
|
||||
raw_label=box.raw_label,
|
||||
reading_order=box.position,
|
||||
html="",
|
||||
skipped=True,
|
||||
confidence=1.0,
|
||||
)
|
||||
)
|
||||
continue
|
||||
out = out_by_key.get((page_idx, block_idx))
|
||||
if out is None or out.error:
|
||||
blocks.append(
|
||||
BlockOCRResult(
|
||||
polygon=box.polygon,
|
||||
label=box.label,
|
||||
raw_label=box.raw_label,
|
||||
reading_order=box.position,
|
||||
html="",
|
||||
skipped=False,
|
||||
error=True,
|
||||
confidence=0.0,
|
||||
)
|
||||
)
|
||||
continue
|
||||
html = clean_block_html(out.raw)
|
||||
conf = out.mean_token_prob if out.mean_token_prob is not None else 1.0
|
||||
blocks.append(
|
||||
BlockOCRResult(
|
||||
polygon=box.polygon,
|
||||
label=box.label,
|
||||
raw_label=box.raw_label,
|
||||
reading_order=box.position,
|
||||
html=html,
|
||||
skipped=False,
|
||||
error=False,
|
||||
confidence=conf,
|
||||
raw_logprobs=out.logprobs,
|
||||
)
|
||||
)
|
||||
results.append(
|
||||
PageOCRResult(blocks=blocks, image_bbox=[0, 0, float(w), float(h)])
|
||||
)
|
||||
return results
|
||||
|
||||
def _full_page_ocr(
|
||||
self,
|
||||
images: List[Image.Image],
|
||||
fallback_layout: Optional[List[LayoutResult]] = None,
|
||||
) -> List[PageOCRResult]:
|
||||
"""One HIGH_ACCURACY_BBOX_PROMPT request per page; parses divs into blocks.
|
||||
|
||||
On per-page failure (parse error, empty output, or a detected
|
||||
repetition loop in the decoder output), falls back to layout +
|
||||
block-mode OCR for that page only. ``fallback_layout``, if given,
|
||||
provides per-page LayoutResults to use on fallback; otherwise the
|
||||
LayoutPredictor is invoked lazily for just the affected pages.
|
||||
"""
|
||||
manager = self.manager or get_default_manager()
|
||||
with timing_span("recognition_build_full_page_batch", page_count=len(images)):
|
||||
batch = [
|
||||
BatchInputItem(
|
||||
image=img,
|
||||
prompt_type=PROMPT_TYPE_HIGH_ACCURACY_BBOX,
|
||||
max_tokens=settings.SURYA_MAX_TOKENS_FULL_PAGE,
|
||||
metadata={"page_idx": i},
|
||||
)
|
||||
for i, img in enumerate(images)
|
||||
]
|
||||
with timing_span(
|
||||
"recognition_full_page_generate",
|
||||
item_count=len(batch),
|
||||
max_tokens=settings.SURYA_MAX_TOKENS_FULL_PAGE,
|
||||
):
|
||||
outputs = manager.generate(batch)
|
||||
out_by_page = {o.metadata["page_idx"]: o for o in outputs}
|
||||
|
||||
results: List[Optional[PageOCRResult]] = [None] * len(images)
|
||||
needs_fallback: List[int] = []
|
||||
for page_idx, img in enumerate(images):
|
||||
w, h = img.size
|
||||
page_bbox = [0, 0, float(w), float(h)]
|
||||
out = out_by_page.get(page_idx)
|
||||
if out is None or out.error:
|
||||
# Hard failure (request lost / server error). Always fallback.
|
||||
needs_fallback.append(page_idx)
|
||||
continue
|
||||
if not out.raw:
|
||||
# Empty model output. If the page is genuinely blank, the
|
||||
# model is correct — return an empty result. Only fall back
|
||||
# when the page has content the model failed to emit.
|
||||
if is_blank_region(img):
|
||||
results[page_idx] = PageOCRResult(blocks=[], image_bbox=page_bbox)
|
||||
else:
|
||||
logger.info(
|
||||
f"empty full-page output for non-blank page {page_idx}; "
|
||||
f"falling back to layout + block OCR"
|
||||
)
|
||||
needs_fallback.append(page_idx)
|
||||
continue
|
||||
if _detect_repeat_loop(out.raw):
|
||||
logger.info(
|
||||
f"full-page output for page {page_idx} appears to loop; "
|
||||
f"falling back to layout + block OCR"
|
||||
)
|
||||
needs_fallback.append(page_idx)
|
||||
continue
|
||||
try:
|
||||
parsed = parse_full_page_html(out.raw)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Full-page parse failed for page {page_idx}: {e}; "
|
||||
f"falling back to layout + block OCR"
|
||||
)
|
||||
needs_fallback.append(page_idx)
|
||||
continue
|
||||
confidence = out.mean_token_prob if out.mean_token_prob is not None else 1.0
|
||||
blocks: List[BlockOCRResult] = []
|
||||
for idx, item in enumerate(parsed):
|
||||
x0 = item.bbox[0] / settings.BBOX_SCALE * w
|
||||
y0 = item.bbox[1] / settings.BBOX_SCALE * h
|
||||
x1 = item.bbox[2] / settings.BBOX_SCALE * w
|
||||
y1 = item.bbox[3] / settings.BBOX_SCALE * h
|
||||
polygon = [[x0, y0], [x1, y0], [x1, y1], [x0, y1]]
|
||||
canon = LAYOUT_PRED_RELABEL.get(item.label, item.label)
|
||||
skipped = canon in SKIP_CANON_LABELS
|
||||
blocks.append(
|
||||
BlockOCRResult(
|
||||
polygon=polygon,
|
||||
label=canon,
|
||||
raw_label=item.label,
|
||||
reading_order=idx,
|
||||
html="" if skipped else item.html,
|
||||
skipped=skipped,
|
||||
error=False,
|
||||
confidence=confidence,
|
||||
)
|
||||
)
|
||||
blocks = _drop_blank_text_blocks(img, blocks)
|
||||
results[page_idx] = PageOCRResult(blocks=blocks, image_bbox=page_bbox)
|
||||
|
||||
# Block-mode fallback for any pages whose full-page output failed or looped.
|
||||
if needs_fallback:
|
||||
fb_images = [images[i] for i in needs_fallback]
|
||||
if fallback_layout is not None:
|
||||
fb_layouts = [fallback_layout[i] for i in needs_fallback]
|
||||
else:
|
||||
# Lazy import to avoid the surya.layout ↔ surya.recognition cycle.
|
||||
from surya.layout import LayoutPredictor
|
||||
|
||||
logger.info(
|
||||
f"running layout for {len(fb_images)} page(s) requiring "
|
||||
f"block-mode fallback"
|
||||
)
|
||||
fb_layouts = LayoutPredictor(self.manager)(fb_images)
|
||||
fb_results = self.__call__(fb_images, fb_layouts, full_page=False)
|
||||
for fb_idx, page_idx in enumerate(needs_fallback):
|
||||
results[page_idx] = fb_results[fb_idx]
|
||||
|
||||
# Backfill any still-None pages with empty results (defensive — shouldn't happen).
|
||||
out_results: List[PageOCRResult] = []
|
||||
for page_idx, img in enumerate(images):
|
||||
r = results[page_idx]
|
||||
if r is None:
|
||||
w, h = img.size
|
||||
r = PageOCRResult(blocks=[], image_bbox=[0, 0, float(w), float(h)])
|
||||
out_results.append(r)
|
||||
return out_results
|
||||
@@ -0,0 +1,19 @@
|
||||
from typing import List
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from surya.common.polygon import PolygonBox
|
||||
|
||||
|
||||
class BlockOCRResult(PolygonBox):
|
||||
label: str # canonicalized layout label (Picture, Text, ...)
|
||||
raw_label: str = "" # original model label
|
||||
reading_order: int # 0-indexed position in layout output
|
||||
html: str = "" # block HTML (BLOCK_PROMPT output, "" if skipped)
|
||||
skipped: bool = False # True if label was in SKIP_OCR_LABELS
|
||||
error: bool = False
|
||||
|
||||
|
||||
class PageOCRResult(BaseModel):
|
||||
blocks: List[BlockOCRResult]
|
||||
image_bbox: List[float]
|
||||
Reference in New Issue
Block a user