"""Распознавание строк обученной моделью (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