import html import os import re import time from typing import Any, Dict, List, Sequence, Tuple from PIL import Image, ImageDraw # These must be set before importing surya.settings, because settings are # instantiated at import time. os.environ.setdefault("SURYA_INFERENCE_BACKEND", "vllm") os.environ.setdefault("SURYA_INFERENCE_URL", "http://127.0.0.1:8000/v1") os.environ.setdefault("SURYA_INFERENCE_AUTOSTART", "false") os.environ.setdefault("SURYA_INFERENCE_KEEP_ALIVE", "false") os.environ.setdefault("SURYA_INFERENCE_LOGPROBS", "false") os.environ.setdefault("SURYA_INFERENCE_MAX_RETRIES", "1") os.environ.setdefault("SURYA_STREAM_MODE", "off") os.environ.setdefault("SURYA_INFERENCE_PARALLEL", "8") os.environ.setdefault("SURYA_MAX_TOKENS_FULL_PAGE", "6144") os.environ.setdefault("SURYA_MAX_BLOCKS_PER_PAGE", "80") os.environ.setdefault("SUYA_VLLM_IMAGE_FORMAT", "JPEG") os.environ.setdefault("SUYA_VLLM_JPEG_QUALITY", "92") from surya.debug.draw import draw_polys_on_image, draw_bboxes_on_image from surya.inference import SuryaInferenceManager from surya.layout import LayoutPredictor from surya.recognition import RecognitionPredictor from surya.table_rec import TableRecPredictor from surya.timing import timing_span _TAG_RE = re.compile(r"<[^>]+>") OCR_MODE = os.getenv("SUYA_OCR_MODE", "block").strip().lower() def _html_to_text(value: str) -> str: return html.unescape(_TAG_RE.sub(" ", value or "")).strip() def _load_predictors_vllm() -> Dict[str, Any]: manager = SuryaInferenceManager(method="vllm") return { "manager": manager, "layout": LayoutPredictor(manager), "recognition": RecognitionPredictor(manager), "table_rec": TableRecPredictor(manager), } predictors_vllm = _load_predictors_vllm() def backend_info() -> Dict[str, Any]: manager = predictors_vllm["manager"] backend = manager.backend handle = getattr(backend, "handle", None) return { "backend": getattr(manager, "method", "vllm"), "started": handle is not None, "base_url": getattr(handle, "base_url", None), "model_name": getattr(handle, "model_name", None), "inference_url": os.getenv("SURYA_INFERENCE_URL"), "ocr_mode": OCR_MODE, "parallel": os.getenv("SURYA_INFERENCE_PARALLEL"), "max_blocks_per_page": os.getenv("SURYA_MAX_BLOCKS_PER_PAGE"), } def vllm_backend_info() -> Dict[str, Any]: return backend_info() def _annotate_page(highres_img: Image.Image, page: Any, with_bboxes: bool) -> Image.Image: if not with_bboxes: return highres_img annotated = highres_img.copy() draw = ImageDraw.Draw(annotated) for block in page.blocks: x0, y0, x1, y1 = block.bbox color = "orange" if block.skipped else ("red" if block.error else "green") draw.rectangle((x0, y0, x1, y1), outline=color, width=3) draw.text((x0 + 4, y0 + 4), str(block.reading_order), fill=color) return annotated def ocr_vllm_batch( images: Sequence[Image.Image], highres_images: Sequence[Image.Image], skip_text_detection: bool = False, recognize_math: bool = True, with_bboxes: bool = True, ) -> List[Tuple[Image.Image, Any, Image.Image]]: start = time.perf_counter() if OCR_MODE not in {"full_page", "block"}: raise ValueError("SUYA_OCR_MODE must be 'full_page' or 'block'") if len(images) != len(highres_images): raise ValueError("images and highres_images must have the same length") if not images: return [] with timing_span( "ocr_vllm_batch_total", image_count=len(images), ocr_mode=OCR_MODE, skip_text_detection=skip_text_detection, ): if OCR_MODE == "full_page" or skip_text_detection: with timing_span("recognition_full_page_call", image_count=len(highres_images)): pages = predictors_vllm["recognition"](list(highres_images), full_page=True) else: target_sizes = [img.size for img in highres_images] with timing_span("layout_predictor_call", image_count=len(images)): layouts = predictors_vllm["layout"](list(images), target_image_sizes=target_sizes) with timing_span("recognition_block_call", image_count=len(highres_images)): pages = predictors_vllm["recognition"](list(highres_images), layouts, full_page=False) elapsed = time.perf_counter() - start results = [] with timing_span("response_assembly", image_count=len(highres_images), with_bboxes=with_bboxes): for highres_img, page in zip(highres_images, pages): object.__setattr__(page, "_elapsed_seconds", elapsed) annotated = _annotate_page(highres_img, page, with_bboxes) results.append((annotated, page, annotated)) return results def ocr_vllm( img: Image.Image, highres_img: Image.Image, skip_text_detection: bool = False, recognize_math: bool = True, with_bboxes: bool = True, ) -> Tuple[Image.Image, Any, Image.Image]: return ocr_vllm_batch( [img], [highres_img], skip_text_detection=skip_text_detection, recognize_math=recognize_math, with_bboxes=with_bboxes, )[0] def page_ocr_to_response(page: Any) -> Dict[str, Any]: page_json = page.model_dump() text_lines = [ text for text in (_html_to_text(block.get("html", "")) for block in page_json.get("blocks", [])) if text ] return { "ocr_text_json": page_json, "text_lines": "\n".join(text_lines), "backend": backend_info(), "elapsed_seconds": getattr(page, "_elapsed_seconds", None), } def layout_detection_vllm(img: Image.Image) -> Tuple[Image.Image, Any]: pred = predictors_vllm["layout"]([img])[0] polygons = [p.polygon for p in pred.bboxes] labels = [ f"{p.label}-{p.position}-{round(getattr(p, 'confidence', 0) or 0, 2)}" for p in pred.bboxes ] layout_img = draw_polys_on_image( polygons, img.copy(), labels=labels, label_font_size=18 ) return layout_img, pred def table_recognition_vllm( img: Image.Image, highres_img: Image.Image, skip_table_detection: bool ) -> Tuple[Image.Image, Any]: if skip_table_detection: layout_tables = [(0, 0, highres_img.size[0], highres_img.size[1])] table_imgs = [highres_img] else: _, layout_pred = layout_detection_vllm(img) layout_tables = [ tuple(map(int, line.bbox)) for line in layout_pred.bboxes if line.label in ["Table", "TableOfContents"] ] table_imgs = [highres_img.crop(tb) for tb in layout_tables] table_preds = predictors_vllm["table_rec"](table_imgs) table_img = highres_img.copy() for result, table_bbox in zip(table_preds, layout_tables): adjusted_bboxes = [] labels = [] colors = [] for item in [*getattr(result, "rows", []), *getattr(result, "cols", []), *getattr(result, "cells", [])]: adjusted_bboxes.append( [ item.bbox[0] + table_bbox[0], item.bbox[1] + table_bbox[1], item.bbox[2] + table_bbox[0], item.bbox[3] + table_bbox[1], ] ) labels.append(item.label) colors.append("blue" if "Row" in item.label else "red") if adjusted_bboxes: table_img = draw_bboxes_on_image( adjusted_bboxes, highres_img, labels=labels, label_font_size=18, color=colors, ) return table_img, table_preds