QOCR / qocr_tiny_v1_ready /inference.py
Qarvexium's picture
Upload 4 files
447f27a verified
Raw
History Blame Contribute Delete
5.15 kB
import os
import json
import argparse
import numpy as np
from PIL import Image
import torch
import torch.nn as nn
HERE = os.path.dirname(os.path.abspath(__file__))
MODEL_PATH = os.path.join(HERE, "qocr_tiny_v1.pt")
VOCAB_PATH = os.path.join(HERE, "vocab.json")
CONFIG_PATH = os.path.join(HERE, "config.json")
with open(VOCAB_PATH, "r", encoding="utf-8") as f:
vocab = json.load(f)
with open(CONFIG_PATH, "r", encoding="utf-8") as f:
config = json.load(f)
CHARSET = vocab["charset"]
BLANK_ID = int(vocab["blank"])
VOCAB_SIZE = int(vocab["vocab_size"])
MAX_HEIGHT = int(config["input"]["max_height"])
MAX_WIDTH = int(config["input"]["max_width"])
CNN1, CNN2, CNN3, CNN4, CNN5 = [int(x) for x in config["encoder"]["channels"]]
LATENT_DIM = int(config["encoder"]["latent_dim"])
GRU_HIDDEN = int(config["rnn"]["hidden"])
GRU_LAYERS = int(config["rnn"]["layers"])
ID_TO_CHAR = {i + 1: c for i, c in enumerate(CHARSET)}
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
USE_AMP = (DEVICE.type == "cuda")
class ConvBNAct(nn.Module):
def __init__(self, in_channels, out_channels, stride=(1, 1)):
super().__init__()
self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)
self.bn = nn.BatchNorm2d(out_channels)
self.act = nn.SiLU(inplace=True)
def forward(self, x):
return self.act(self.bn(self.conv(x)))
class DepthwiseSeparable(nn.Module):
def __init__(self, in_channels, out_channels, stride=(1, 1)):
super().__init__()
self.depthwise = nn.Conv2d(in_channels, in_channels, kernel_size=3, stride=stride, padding=1, groups=in_channels, bias=False)
self.depth_bn = nn.BatchNorm2d(in_channels)
self.pointwise = nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)
self.point_bn = nn.BatchNorm2d(out_channels)
self.act = nn.SiLU(inplace=True)
def forward(self, x):
x = self.depthwise(x)
x = self.depth_bn(x)
x = self.act(x)
x = self.pointwise(x)
x = self.point_bn(x)
x = self.act(x)
return x
class QOCRSmall(nn.Module):
def __init__(self, vocab_size):
super().__init__()
self.encoder = nn.Sequential(
ConvBNAct(1, CNN1, stride=(2, 1)),
ConvBNAct(CNN1, CNN2, stride=(2, 1)),
ConvBNAct(CNN2, CNN3, stride=(2, 1)),
DepthwiseSeparable(CNN3, CNN4, stride=(1, 2)),
DepthwiseSeparable(CNN4, CNN5, stride=(1, 1)),
)
self.latent_projection = nn.Sequential(
nn.Conv2d(CNN5, LATENT_DIM, kernel_size=1, bias=False),
nn.BatchNorm2d(LATENT_DIM),
nn.SiLU(inplace=True),
)
self.gru = nn.GRU(
input_size=LATENT_DIM,
hidden_size=GRU_HIDDEN,
num_layers=GRU_LAYERS,
batch_first=True,
bidirectional=True,
dropout=0.15,
)
self.norm = nn.LayerNorm(GRU_HIDDEN * 2)
self.classifier = nn.Linear(GRU_HIDDEN * 2, vocab_size)
def forward(self, x):
x = self.encoder(x)
x = self.latent_projection(x)
x = x.mean(dim=2)
x = x.transpose(1, 2)
x, _ = self.gru(x)
x = self.norm(x)
x = self.classifier(x)
return x.transpose(0, 1)
model = QOCRSmall(VOCAB_SIZE).to(DEVICE)
state_dict = torch.load(MODEL_PATH, map_location=DEVICE)
model.load_state_dict(state_dict)
model.eval()
def preprocess_image(image):
if not isinstance(image, Image.Image):
image = Image.fromarray(np.asarray(image))
image = image.convert("L")
width, height = image.size
if width <= 0 or height <= 0:
raise ValueError("Invalid image dimensions.")
scale = min(1.0, MAX_WIDTH / width, MAX_HEIGHT / height)
if scale < 1.0:
width = max(1, round(width * scale))
height = max(1, round(height * scale))
image = image.resize((width, height), Image.Resampling.LANCZOS)
array = np.asarray(image, dtype=np.float32)
array /= 255.0
tensor = torch.from_numpy(array).unsqueeze(0).unsqueeze(0).to(DEVICE)
return tensor
def decode_logits(logits):
ids = logits.argmax(dim=2)
sequence = ids[:, 0].tolist()
previous = BLANK_ID
output = []
for token in sequence:
if token != BLANK_ID and token != previous:
output.append(ID_TO_CHAR.get(token, ""))
previous = token
return "".join(output)
@torch.inference_mode()
def ocr(image):
tensor = preprocess_image(image)
if USE_AMP:
with torch.autocast(device_type="cuda", dtype=torch.float16):
logits = model(tensor)
else:
logits = model(tensor)
return decode_logits(logits)
def ocr_file(path):
with Image.open(path) as image:
return ocr(image)
def main():
parser = argparse.ArgumentParser(description="QOCR-Tiny v1 OCR")
parser.add_argument("image", help="Path to cropped text image")
args = parser.parse_args()
result = ocr_file(args.image)
print(result)
if __name__ == "__main__":
main()