text_detect/core/paddle_ocr_backend.py
2026-03-08 13:42:59 +03:00

171 lines
5.6 KiB
Python

# -*- coding: utf-8 -*-
"""Optional PaddleOCR backend wrapper."""
import os
import sys
try:
from paddleocr import PaddleOCR
except ImportError: # pragma: no cover
PaddleOCR = None
class PaddleOCRBackend(object):
def __init__(self,
enable=False,
lang='ru',
use_angle_cls=False,
use_gpu=False,
det_model_dir='',
rec_model_dir='',
cls_model_dir='',
cpu_threads=0):
self.enable = bool(enable)
self.lang = lang or 'ru'
self.use_angle_cls = bool(use_angle_cls)
self.use_gpu = bool(use_gpu)
self.det_model_dir = det_model_dir or ''
self.rec_model_dir = rec_model_dir or ''
self.cls_model_dir = cls_model_dir or ''
self.cpu_threads = int(cpu_threads) if cpu_threads else 0
self.available = False
self.engine = None
self.last_error = ''
self._load()
def _load(self):
if not self.enable:
return
if PaddleOCR is None:
self.last_error = 'paddleocr_not_installed'
print('[WARN] PaddleOCR backend requested but package is not installed', file=sys.stderr)
return
kwargs = {
'lang': self.lang,
'use_angle_cls': self.use_angle_cls,
'use_gpu': self.use_gpu,
'show_log': False,
}
det_dir = self._resolve_dir(self.det_model_dir, 'det')
rec_dir = self._resolve_dir(self.rec_model_dir, 'rec')
cls_dir = self._resolve_dir(self.cls_model_dir, 'cls')
if det_dir:
kwargs['det_model_dir'] = det_dir
if rec_dir:
kwargs['rec_model_dir'] = rec_dir
if cls_dir:
kwargs['cls_model_dir'] = cls_dir
if self.cpu_threads > 0:
kwargs['cpu_threads'] = self.cpu_threads
try:
self.engine = PaddleOCR(**kwargs)
self.available = True
print('[INFO] PaddleOCR backend ready (lang={})'.format(self.lang))
except Exception as exc:
self.last_error = str(exc)
self.available = False
print('[WARN] Failed to initialize PaddleOCR backend: {}'.format(exc), file=sys.stderr)
def run_ocr(self, img_rgb, return_boxes=True):
if not self.available or img_rgb is None:
return '', []
try:
if img_rgb.ndim == 3 and img_rgb.shape[2] == 3:
img_bgr = img_rgb[:, :, ::-1].copy()
else:
img_bgr = img_rgb.copy()
raw = self.engine.ocr(img_bgr, cls=self.use_angle_cls)
except Exception as exc:
print('[WARN] PaddleOCR inference failed: {}'.format(exc), file=sys.stderr)
return '', []
return self._parse_result(raw, return_boxes=return_boxes)
def _parse_result(self, raw, return_boxes=True):
lines = []
self._walk_lines(raw, lines)
blocks = []
texts = []
for line in lines:
poly = line[0]
rec = line[1]
text = self._safe_text(rec[0] if len(rec) > 0 else '')
if not text:
continue
conf = self._safe_float(rec[1] if len(rec) > 1 else -1.0, default=-1.0)
bbox = self._poly_to_bbox(poly)
texts.append(text)
if return_boxes:
blocks.append({
'bbox': [bbox[0], bbox[1], bbox[2], bbox[3]],
'text': text,
'conf': conf,
})
return '\n'.join(texts), blocks if return_boxes else []
def _walk_lines(self, node, out):
if self._is_line_entry(node):
out.append(node)
return
if isinstance(node, (list, tuple)):
for item in node:
self._walk_lines(item, out)
def _is_line_entry(self, node):
if not isinstance(node, (list, tuple)) or len(node) < 2:
return False
poly = node[0]
rec = node[1]
return self._is_poly(poly) and self._is_rec(rec)
def _is_poly(self, poly):
if not isinstance(poly, (list, tuple)) or len(poly) < 3:
return False
first = poly[0]
return isinstance(first, (list, tuple)) and len(first) >= 2
def _is_rec(self, rec):
if not isinstance(rec, (list, tuple)) or len(rec) < 1:
return False
return isinstance(rec[0], (str, bytes))
def _poly_to_bbox(self, poly):
xs = []
ys = []
for pt in poly:
if not isinstance(pt, (list, tuple)) or len(pt) < 2:
continue
xs.append(self._safe_int(pt[0], 0))
ys.append(self._safe_int(pt[1], 0))
if not xs or not ys:
return (0, 0, 0, 0)
return (max(0, min(xs)), max(0, min(ys)), max(xs), max(ys))
def _resolve_dir(self, path, name):
if not path:
return ''
if os.path.isdir(path):
return path
print('[WARN] PaddleOCR {} model dir not found: {}'.format(name, path), file=sys.stderr)
return ''
def _safe_text(self, value):
if value is None:
return ''
if isinstance(value, bytes):
try:
value = value.decode('utf-8', 'ignore')
except Exception:
return ''
return str(value).strip()
def _safe_float(self, value, default=0.0):
try:
return float(value)
except Exception:
return float(default)
def _safe_int(self, value, default=0):
try:
return int(round(float(value)))
except Exception:
return int(default)