Spaces:
Running on Zero
Running on Zero
Download modeling.py from LocalDoc/azerbaijani-ocr: direct link, hf CLI and curl.
- Browser
- Download file 2.27 kB
-
https://huggingface.co/spaces/LocalDoc/azerbaijani-ocr/resolve/main/modeling.py
- Command line
-
hf download hf://spaces/LocalDoc/azerbaijani-ocr/modeling.py
-
curl -L -o modeling.py https://huggingface.co/spaces/LocalDoc/azerbaijani-ocr/resolve/main/modeling.py
2.27 kB
| """Model definition. Required to load the weights.""" | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| from PIL import Image | |
| HEIGHT = 48 | |
| DOWNSAMPLE = 4 # width is divided by 4; T = width // 4 | |
| class CRNN(nn.Module): | |
| def __init__(self, nclass, hidden=256, layers=2, dropout=0.1): | |
| super().__init__() | |
| def blk(i, o, pool): | |
| m = [nn.Conv2d(i, o, 3, 1, 1, bias=False), nn.BatchNorm2d(o), | |
| nn.ReLU(inplace=True)] | |
| if pool: | |
| m.append(nn.MaxPool2d(pool, pool)) | |
| return m | |
| self.cnn = nn.Sequential( | |
| *blk(1, 64, (2, 2)), | |
| *blk(64, 128, (2, 2)), | |
| *blk(128, 256, None), | |
| *blk(256, 256, (2, 1)), | |
| *blk(256, 512, None), | |
| *blk(512, 512, (2, 1)), | |
| nn.Conv2d(512, 512, (3, 1), 1, 0, bias=False), | |
| nn.BatchNorm2d(512), nn.ReLU(inplace=True), | |
| ) | |
| self.rnn = nn.LSTM(512, hidden, num_layers=layers, bidirectional=True, | |
| batch_first=True, dropout=dropout if layers > 1 else 0) | |
| self.head = nn.Linear(hidden * 2, nclass) | |
| def forward(self, x): | |
| f = self.cnn(x).squeeze(2).permute(0, 2, 1) | |
| f, _ = self.rnn(f) | |
| return self.head(f) | |
| def preprocess(img, max_width=1200): | |
| """PIL Image -> tensor (1,1,48,W). Input is a crop of ONE text line.""" | |
| img = img.convert("L") | |
| if img.height != HEIGHT: | |
| r = HEIGHT / img.height | |
| img = img.resize((max(8, int(img.width * r)), HEIGHT), Image.BILINEAR) | |
| if img.width > max_width: | |
| img = img.crop((0, 0, max_width, HEIGHT)) | |
| x = (np.array(img, dtype=np.float32) / 255.0 - 0.5) / 0.5 | |
| return torch.from_numpy(x)[None, None] | |
| def ctc_greedy(logits, itos, blank=0): | |
| ids = logits.argmax(-1)[0].tolist() | |
| prev, out = -1, [] | |
| for k in ids: | |
| if k != prev and k != blank: | |
| out.append(itos[k]) | |
| prev = k | |
| return "".join(out) | |
| def load(path="model.pt", device="cpu"): | |
| ck = torch.load(path, map_location=device) | |
| cfg = ck.get("cfg", {}) | |
| itos = ck["itos"] | |
| m = CRNN(len(itos), cfg.get("rnn_hidden", 256), cfg.get("rnn_layers", 2), 0.0) | |
| m.load_state_dict(ck["model"]) | |
| return m.eval().to(device), itos | |