document-ocr / ocr_studio /engine.py
alirezaaminzadeh's picture
Fix PaddleOCR-VL image processor pixel limits
819da4b verified
Raw
History Blame Contribute Delete
4.55 kB
from __future__ import annotations
import threading
from dataclasses import dataclass
from typing import Any
from PIL import Image
from ocr_studio.config import (
MODE_PRECISE,
MODEL_ID,
MODEL_REVISION,
OCR_MAX_PIXELS,
SPOTTING_MAX_PIXELS,
SPOTTING_UPSCALE_THRESHOLD,
MAX_NEW_TOKENS,
)
from ocr_studio.spotting import TextSpan, parse_spans, spans_to_text, strip_special_tokens
PROMPTS = {
"ocr": "OCR:",
"spotting": "Spotting:",
}
@dataclass
class InferenceResult:
raw_text: str
display_text: str
spans: list[TextSpan]
task: str
class PaddleOcrVlEngine:
def __init__(self) -> None:
self.model: Any = None
self.processor: Any = None
self.device: Any = None
self._lock = threading.Lock()
def load(self) -> None:
if self.model is not None:
return
import torch
from transformers import AutoConfig, AutoModelForImageTextToText, AutoProcessor
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dtype = torch.bfloat16 if device.type == "cuda" else torch.float32
config = AutoConfig.from_pretrained(MODEL_ID, revision=MODEL_REVISION)
if not hasattr(config, "text_config") and hasattr(config, "get_text_config"):
config.text_config = config.get_text_config()
processor = AutoProcessor.from_pretrained(
MODEL_ID,
revision=MODEL_REVISION,
trust_remote_code=False,
)
model = AutoModelForImageTextToText.from_pretrained(
MODEL_ID,
config=config,
revision=MODEL_REVISION,
torch_dtype=dtype,
trust_remote_code=False,
low_cpu_mem_usage=True,
)
model = model.to(device).eval()
self.model = model
self.processor = processor
self.device = device
def _prepare_image(self, image: Image.Image, task: str) -> Image.Image:
prepared = image.convert("RGB")
if (
task == "spotting"
and prepared.width < SPOTTING_UPSCALE_THRESHOLD
and prepared.height < SPOTTING_UPSCALE_THRESHOLD
):
prepared = prepared.resize(
(prepared.width * 2, prepared.height * 2),
Image.Resampling.LANCZOS,
)
return prepared
def recognize(self, image: Image.Image, mode: str) -> InferenceResult:
self.load()
import torch
task = "spotting" if mode == MODE_PRECISE else "ocr"
work_image = self._prepare_image(image, task)
max_pixels = SPOTTING_MAX_PIXELS if task == "spotting" else OCR_MAX_PIXELS
messages = [
{
"role": "user",
"content": [
{"type": "image", "image": work_image},
{"type": "text", "text": PROMPTS[task]},
],
}
]
image_processor = self.processor.image_processor
min_pixels = getattr(image_processor, "min_pixels", None)
if min_pixels is None:
size_cfg = getattr(image_processor, "size", {}) or {}
min_pixels = size_cfg.get("shortest_edge") or size_cfg.get("min_pixels") or (16 * 28 * 28)
inputs = self.processor.apply_chat_template(
messages,
add_generation_prompt=True,
tokenize=True,
return_dict=True,
return_tensors="pt",
images_kwargs={
"size": {
"shortest_edge": int(min_pixels),
"longest_edge": max_pixels,
}
},
)
inputs = inputs.to(self.model.device)
with self._lock, torch.inference_mode():
generated = self.model.generate(
**inputs,
max_new_tokens=MAX_NEW_TOKENS,
do_sample=False,
)
prompt_len = inputs["input_ids"].shape[-1]
decoded = self.processor.decode(
generated[0][prompt_len:],
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
).strip()
spans = parse_spans(decoded, image.width, image.height)
display = spans_to_text(spans, strip_special_tokens(decoded)).strip()
return InferenceResult(raw_text=decoded, display_text=display, spans=spans, task=task)