Download src/recognize.py from lab-ii/sakha-ocr: direct link, hf CLI and curl.
- Browser
- Download file 3.89 kB
-
https://huggingface.co/lab-ii/sakha-ocr/resolve/main/src/recognize.py
- Command line
-
hf download hf://lab-ii/sakha-ocr/src/recognize.py
-
curl -L -o recognize.py https://huggingface.co/lab-ii/sakha-ocr/resolve/main/src/recognize.py
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 | |