"""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