pageparse.ai / src /pageparse /ocr /handwriting.py
Varun2007's picture
initial clean deployment commit with compilers
8c3e275
Raw
History Blame Contribute Delete
1.75 kB
from __future__ import annotations
import cv2
import numpy as np
import onnxruntime as ort
from tokenizers import Tokenizer
from pageparse.config import settings
class HandwritingOCR:
def __init__(self) -> None:
model_path = settings.model_path(settings.ocr_model)
tokenizer_path = settings.model_path("tr_ocr_tokenizer.json")
self.session = ort.InferenceSession(
str(model_path),
providers=["CPUExecutionProvider"],
)
self.input_name = self.session.get_inputs()[0].name
self.tokenizer = Tokenizer.from_file(str(tokenizer_path))
def recognize(self, image: np.ndarray) -> str:
input_tensor = self._prepare_input(image)
logits = self.session.run(None, {self.input_name: input_tensor})[0]
return self._decode(logits)
def recognize_batch(self, images: list[np.ndarray]) -> list[str]:
batch = np.concatenate([self._prepare_input(img) for img in images], axis=0)
logits = self.session.run(None, {self.input_name: batch})[0]
results = []
for i in range(logits.shape[0]):
result = self._decode(logits[i:i+1])
results.append(result)
return results
def _prepare_input(self, image: np.ndarray) -> np.ndarray:
resized = cv2.resize(image, (384, 384))
normalized = resized.astype(np.float32) / 255.0
return np.expand_dims(np.expand_dims(normalized, 0), 0)
def _decode(self, logits: np.ndarray) -> str:
token_ids = logits.argmax(axis=-1).squeeze()
ids = token_ids.tolist()
if isinstance(ids, int):
ids = [ids]
decoded = self.tokenizer.decode(ids, skip_special_tokens=True)
return decoded.strip()