189 lines
5.4 KiB
Python
189 lines
5.4 KiB
Python
from __future__ import annotations
|
|
|
|
import os
|
|
import tempfile
|
|
from collections.abc import Iterable
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from core.models import Bounds, SceneElement
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class OCRBox:
|
|
text: str
|
|
bounds: Bounds
|
|
confidence: float | None = None
|
|
|
|
def to_scene_element(self, element_id: str) -> SceneElement:
|
|
return SceneElement(
|
|
id=element_id,
|
|
type="text",
|
|
text=self.text,
|
|
bounds=self.bounds,
|
|
confidence=self.confidence,
|
|
source="ocr",
|
|
)
|
|
|
|
|
|
class PaddleOCREngine:
|
|
def __init__(self, **kwargs: Any) -> None:
|
|
self.kwargs = kwargs
|
|
self._engine: Any | None = None
|
|
|
|
def extract(self, image: bytes | str | Path) -> list[OCRBox]:
|
|
engine = self._load()
|
|
image_input, temp_path = _image_input(image)
|
|
try:
|
|
if hasattr(engine, "predict"):
|
|
raw = engine.predict(input=image_input)
|
|
else:
|
|
raw = engine.ocr(image_input, cls=True)
|
|
return parse_paddle_result(raw)
|
|
finally:
|
|
if temp_path:
|
|
temp_path.unlink(missing_ok=True)
|
|
|
|
def _load(self) -> Any:
|
|
if self._engine is None:
|
|
from paddleocr import PaddleOCR
|
|
|
|
self._engine = PaddleOCR(**self.kwargs)
|
|
return self._engine
|
|
|
|
|
|
def run_ocr(
|
|
image: bytes | str | Path,
|
|
*,
|
|
engine: PaddleOCREngine | None = None,
|
|
strict: bool = False,
|
|
) -> list[SceneElement]:
|
|
try:
|
|
boxes = (engine or PaddleOCREngine()).extract(image)
|
|
except ImportError:
|
|
if strict:
|
|
raise
|
|
return []
|
|
return [box.to_scene_element(f"ocr-{index:03d}") for index, box in enumerate(boxes)]
|
|
|
|
|
|
def parse_paddle_result(raw: Any) -> list[OCRBox]:
|
|
boxes: list[OCRBox] = []
|
|
for item in _flatten_pages(raw):
|
|
parsed = _parse_line(item)
|
|
if parsed:
|
|
boxes.append(parsed)
|
|
return boxes
|
|
|
|
|
|
def _image_input(image: bytes | str | Path) -> tuple[str, Path | None]:
|
|
if isinstance(image, bytes):
|
|
handle = tempfile.NamedTemporaryFile(delete=False, suffix=".png")
|
|
try:
|
|
handle.write(image)
|
|
finally:
|
|
handle.close()
|
|
return handle.name, Path(handle.name)
|
|
return os.fspath(image), None
|
|
|
|
|
|
def _flatten_pages(raw: Any) -> Iterable[Any]:
|
|
if raw is None:
|
|
return []
|
|
if isinstance(raw, dict):
|
|
return _dict_lines(raw)
|
|
if isinstance(raw, list):
|
|
flattened: list[Any] = []
|
|
for page in raw:
|
|
if isinstance(page, dict):
|
|
flattened.extend(_dict_lines(page))
|
|
elif _looks_like_ocr_line(page):
|
|
flattened.append(page)
|
|
elif isinstance(page, list):
|
|
flattened.extend(page)
|
|
return flattened
|
|
json_attr = getattr(raw, "json", None)
|
|
if callable(json_attr):
|
|
return _flatten_pages(json_attr)
|
|
return []
|
|
|
|
|
|
def _dict_lines(data: dict[str, Any]) -> list[Any]:
|
|
payload = data.get("res") if isinstance(data.get("res"), dict) else data
|
|
texts = payload.get("rec_texts") or payload.get("texts") or []
|
|
scores = payload.get("rec_scores") or payload.get("scores") or []
|
|
boxes = (
|
|
payload.get("rec_boxes")
|
|
or payload.get("rec_polys")
|
|
or payload.get("dt_polys")
|
|
or []
|
|
)
|
|
return [
|
|
{
|
|
"text": text,
|
|
"confidence": scores[index] if index < len(scores) else None,
|
|
"points": boxes[index] if index < len(boxes) else None,
|
|
}
|
|
for index, text in enumerate(texts)
|
|
]
|
|
|
|
|
|
def _looks_like_ocr_line(value: Any) -> bool:
|
|
return isinstance(value, (list, tuple)) and len(value) >= 2
|
|
|
|
|
|
def _parse_line(line: Any) -> OCRBox | None:
|
|
if isinstance(line, dict):
|
|
text = line.get("text")
|
|
points = line.get("points") or line.get("box") or line.get("bounds")
|
|
confidence = line.get("confidence")
|
|
if not text or not points:
|
|
return None
|
|
return OCRBox(str(text), _bounds_from_points(points), _float_or_none(confidence))
|
|
|
|
if not _looks_like_ocr_line(line):
|
|
return None
|
|
|
|
points = line[0]
|
|
text_payload = line[1]
|
|
if isinstance(text_payload, (list, tuple)) and text_payload:
|
|
text = text_payload[0]
|
|
confidence = text_payload[1] if len(text_payload) > 1 else None
|
|
else:
|
|
text = text_payload
|
|
confidence = None
|
|
if not text:
|
|
return None
|
|
return OCRBox(str(text), _bounds_from_points(points), _float_or_none(confidence))
|
|
|
|
|
|
def _bounds_from_points(points: Any) -> Bounds:
|
|
if isinstance(points, dict):
|
|
return Bounds.from_dict(points)
|
|
if (
|
|
isinstance(points, (list, tuple))
|
|
and len(points) == 4
|
|
and all(isinstance(value, (int, float)) for value in points)
|
|
):
|
|
x1, y1, x2, y2 = [float(value) for value in points]
|
|
return Bounds(x1, y1, x2 - x1, y2 - y1)
|
|
|
|
xs: list[float] = []
|
|
ys: list[float] = []
|
|
for point in points:
|
|
if isinstance(point, dict):
|
|
xs.append(float(point["x"]))
|
|
ys.append(float(point["y"]))
|
|
else:
|
|
xs.append(float(point[0]))
|
|
ys.append(float(point[1]))
|
|
return Bounds(min(xs), min(ys), max(xs) - min(xs), max(ys) - min(ys))
|
|
|
|
|
|
def _float_or_none(value: Any) -> float | None:
|
|
if value is None:
|
|
return None
|
|
return float(value)
|
|
|