sakha-ocr / src /recognize.py
loalkota's picture
Детектор строк и распознаватель якутского текста с рецептом обучения
d4d4e7b verified
Raw History Blame Contribute Delete
3.89 kB
"""Распознавание строк обученной моделью (ONNX, CPU)."""
import json, os
import numpy as np
import onnxruntime as ort
from PIL import Image
HERE = os.path.dirname(os.path.abspath(__file__))
class Recognizer:
def __init__(self, model_dir=None):
d = model_dir or os.path.join(HERE, "model_synth")
self.meta = json.load(open(os.path.join(d, "sakha_rec_crnn.json"), encoding="utf-8"))
self.charset = self.meta["charset"]
self.i2c = {i + 1: c for i, c in enumerate(self.charset)}
self.H = self.meta["height"]
so = ort.SessionOptions()
so.intra_op_num_threads = max(1, (os.cpu_count() or 4) // 2)
self.sess = ort.InferenceSession(os.path.join(d, "sakha_rec_crnn.onnx"),
sess_options=so, providers=["CPUExecutionProvider"])
# Сегментатор режет строку вплотную по чернилам, и край кропа модель
# принимала за знак препинания. Модель обучена на случайных полях и к
# ним нечувствительна (0.18% против 0.20% CER на реальной газете), но
# на плотных кропах поле снижает ошибку вдвое — нормализуем вход.
PAD = 0.18
def _prep(self, im):
im = im.convert("L")
p = int(im.height * self.PAD)
if p:
bg = int(np.median(np.asarray(im)[:, -3:]))
canvas = Image.new("L", (im.width + 2*p, im.height + 2*p), bg)
canvas.paste(im, (p, p))
im = canvas
w = max(8, int(im.width * self.H / im.height))
im = im.resize((min(w, 1600), self.H), Image.BILINEAR)
return 1.0 - np.asarray(im, dtype=np.float32) / 255.0
def _decode(self, logits, t_valid):
ids = logits[:t_valid].argmax(-1)
out, prev = [], -1
for k in ids:
if k != prev and k != 0:
out.append(self.i2c.get(int(k), ""))
prev = k
return "".join(out)
def read_batch(self, crops, bs=32, with_logprobs=False):
"""Батчим строки близкой ширины; декодируем каждую по её реальной длине,
иначе CTC читает нулевую добивку и дописывает лишние символы.
with_logprobs дополнительно отдаёт логарифмы вероятностей по кадрам —
они нужны постобработке, чтобы принимать словарные замены только тогда,
когда картинка их поддерживает.
"""
if not crops:
return ([], []) if with_logprobs else []
arrs = [self._prep(c) for c in crops]
order = sorted(range(len(arrs)), key=lambda i: arrs[i].shape[1])
res = [""] * len(arrs)
lps = [None] * len(arrs)
for s in range(0, len(order), bs):
idx = order[s:s + bs]
W = int(np.ceil(max(arrs[i].shape[1] for i in idx) / 8) * 8)
x = np.zeros((len(idx), 1, self.H, W), dtype=np.float32)
for j, i in enumerate(idx):
a = arrs[i]
x[j, 0, :, :a.shape[1]] = a
logits = self.sess.run(None, {"image": x})[0]
T = logits.shape[1]
for j, i in enumerate(idx):
t = max(1, min(T, arrs[i].shape[1] * T // W))
res[i] = self._decode(logits[j], t)
if with_logprobs:
z = logits[j, :t]
z = z - z.max(axis=-1, keepdims=True)
lps[i] = z - np.log(np.exp(z).sum(axis=-1, keepdims=True))
return (res, lps) if with_logprobs else res