KHD / Scripts /inference.py
Karez's picture
Upload folder using huggingface_hub
960dd1b verified
Raw
History Blame Contribute Delete
18.9 kB
"""
Kurdish Handwritten Paragraph Recognition - Inference Script
Usage:
# Single image
python inference.py --image sample.tif --model_path model.safetensors --vocab_path vocab.json
# Directory of images
python inference.py --image_dir ./test_images --model_path model.safetensors --vocab_path vocab.json
# With .pth checkpoint
python inference.py --image sample.tif --model_path finetuned_model.pth --vocab_path vocab.json
# KHATT Arabic model (different vocab)
python inference.py --image arabic_sample.tif --model_path khatt_model.safetensors \
--vocab_path khatt_vocab.json
"""
import os
import glob
import json
import math
import time
import argparse
from PIL import Image
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms as transforms
import torchvision.models as models
# ===============================
# Argument Parser
# ===============================
def parse_args():
parser = argparse.ArgumentParser(
description="Kurdish Handwritten Paragraph Recognition - Inference")
# Input
parser.add_argument("--image", type=str, default=None,
help="Path to a single paragraph image")
parser.add_argument("--image_dir", type=str, default=None,
help="Directory of paragraph images to process")
# Model and vocabulary
parser.add_argument("--model_path", type=str, required=True,
help="Path to model weights (.pth or .safetensors)")
parser.add_argument("--vocab_path", type=str, required=True,
help="Path to vocabulary JSON file (vocab.json)")
parser.add_argument("--config_path", type=str, default=None,
help="Path to config.json (auto-loads architecture settings)")
# Image dimensions
parser.add_argument("--img_height", type=int, default=600)
parser.add_argument("--img_width", type=int, default=1235)
# Model architecture (overridden by config.json if provided)
parser.add_argument("--hidden_size", type=int, default=256)
parser.add_argument("--encoder_layers", type=int, default=3)
parser.add_argument("--decoder_layers", type=int, default=6)
parser.add_argument("--num_heads", type=int, default=8)
parser.add_argument("--ff_dim", type=int, default=2048)
parser.add_argument("--max_seq_len", type=int, default=555)
parser.add_argument("--use_upsample", action="store_true", default=True)
parser.add_argument("--no_upsample", action="store_true")
# Output
parser.add_argument("--output_file", type=str, default=None,
help="Save predictions to text file")
parser.add_argument("--show_timing", action="store_true",
help="Show per-image inference time")
# Device
parser.add_argument("--device", type=str, default=None,
help="Device (cuda/cpu, auto-detected if not set)")
return parser.parse_args()
# ===============================
# Vocabulary
# ===============================
PAD_TOKEN = 0
SOS_TOKEN = 1
EOS_TOKEN = 2
def load_vocabulary(vocab_path):
"""Load vocabulary from JSON file."""
with open(vocab_path, "r", encoding="utf-8") as f:
vocab_data = json.load(f)
if "vocab_list" in vocab_data:
char_list = vocab_data["vocab_list"]
elif "char_to_idx" in vocab_data:
mapping = vocab_data["char_to_idx"]
char_list = [None] * len(mapping)
for char, idx in mapping.items():
char_list[idx] = char
else:
raise ValueError("Vocabulary JSON must contain 'vocab_list' or 'char_to_idx'")
idx_to_char = {idx: char for idx, char in enumerate(char_list)}
return char_list, idx_to_char
def decode_output(tensor, idx_to_char):
"""Convert tensor of character indices to text."""
if isinstance(tensor, torch.Tensor):
tensor = tensor.cpu().tolist()
text = ""
for idx in tensor:
if idx == PAD_TOKEN or idx == SOS_TOKEN:
continue
if idx == EOS_TOKEN:
break
if idx in idx_to_char:
text += idx_to_char[idx]
return text
# ===============================
# Positional Encodings
# ===============================
class PositionalEncoding2D(nn.Module):
"""2D sinusoidal positional encoding for visual feature maps."""
def __init__(self, d_model, max_h=100, max_w=300):
super().__init__()
pe = torch.zeros(max_h, max_w, d_model)
d_half = d_model // 2
pos_h = torch.arange(0, max_h, dtype=torch.float).unsqueeze(1)
div_h = torch.exp(torch.arange(0, d_half, 2).float() * (-math.log(10000.0) / d_half))
pe_h = torch.zeros(max_h, d_half)
pe_h[:, 0::2] = torch.sin(pos_h * div_h)
pe_h[:, 1::2] = torch.cos(pos_h * div_h)
pos_w = torch.arange(0, max_w, dtype=torch.float).unsqueeze(1)
div_w = torch.exp(torch.arange(0, d_half, 2).float() * (-math.log(10000.0) / d_half))
pe_w = torch.zeros(max_w, d_half)
pe_w[:, 0::2] = torch.sin(pos_w * div_w)
pe_w[:, 1::2] = torch.cos(pos_w * div_w)
for h in range(max_h):
for w in range(max_w):
pe[h, w, :d_half] = pe_h[h]
pe[h, w, d_half:] = pe_w[w]
self.register_buffer('pe', pe)
def forward(self, x, height, width):
_, seq_len, d_model = x.shape
pe_2d = self.pe[:height, :width, :].reshape(height * width, d_model)
if seq_len <= pe_2d.size(0):
pe_2d = pe_2d[:seq_len]
else:
pad = torch.zeros(seq_len - pe_2d.size(0), d_model, device=x.device)
pe_2d = torch.cat([pe_2d, pad], dim=0)
return x + pe_2d.unsqueeze(0)
class PositionalEncoding1D(nn.Module):
"""1D sinusoidal positional encoding for decoder sequences."""
def __init__(self, d_model, max_len=1000):
super().__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer('pe', pe.unsqueeze(0))
def forward(self, x):
return x + self.pe[:, :x.size(1), :]
# ===============================
# CNN Feature Extractor
# ===============================
class CNNFeatureExtractor(nn.Module):
"""DenseNet-121 backbone with optional horizontal upsampling."""
def __init__(self, output_dim=256, use_upsample=True):
super().__init__()
densenet = models.densenet121(weights=models.DenseNet121_Weights.DEFAULT)
self.features = densenet.features
backbone_channels = 1024
if use_upsample:
self.upsample = nn.Sequential(
nn.ConvTranspose2d(backbone_channels, 512,
kernel_size=(1, 4), stride=(1, 2), padding=(0, 1)),
nn.BatchNorm2d(512),
nn.ReLU(inplace=True))
adapt_in = 512
else:
self.upsample = None
adapt_in = backbone_channels
self.adaptation = nn.Sequential(
nn.Conv2d(adapt_in, output_dim, kernel_size=1),
nn.BatchNorm2d(output_dim),
nn.ReLU(inplace=True))
def forward(self, x):
features = F.relu(self.features(x), inplace=True)
if self.upsample is not None:
features = self.upsample(features)
features = self.adaptation(features)
b, c, h, w = features.shape
return features.view(b, c, h * w).permute(0, 2, 1), h, w
# ===============================
# Transformer OCR Model
# ===============================
class TransformerOCRParagraphModel(nn.Module):
"""DenseNet121-Transformer for end-to-end paragraph recognition."""
def __init__(self, vocab_size, hidden_size=256, nhead=8,
num_encoder_layers=3, num_decoder_layers=6,
dim_feedforward=2048, dropout=0.0,
use_upsample=True, max_seq_len=555):
super().__init__()
self.max_seq_len = max_seq_len
self.vocab_size = vocab_size
self.feature_extractor = CNNFeatureExtractor(
output_dim=hidden_size, use_upsample=use_upsample)
self.pos_encoder_2d = PositionalEncoding2D(hidden_size)
self.pos_decoder_1d = PositionalEncoding1D(hidden_size, max_len=max_seq_len)
encoder_layer = nn.TransformerEncoderLayer(
d_model=hidden_size, nhead=nhead,
dim_feedforward=dim_feedforward, dropout=dropout,
batch_first=True)
self.transformer_encoder = nn.TransformerEncoder(
encoder_layer, num_layers=num_encoder_layers)
decoder_layer = nn.TransformerDecoderLayer(
d_model=hidden_size, nhead=nhead,
dim_feedforward=dim_feedforward, dropout=dropout,
batch_first=True)
self.transformer_decoder = nn.TransformerDecoder(
decoder_layer, num_layers=num_decoder_layers)
self.token_embedding = nn.Embedding(vocab_size, hidden_size)
self.output_projection = nn.Linear(hidden_size, vocab_size)
def _generate_square_subsequent_mask(self, sz):
mask = (torch.triu(torch.ones(sz, sz)) == 1).transpose(0, 1)
return mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, 0.0)
def generate(self, img, max_length=None):
"""Auto-regressive greedy generation for a single image."""
if max_length is None:
max_length = self.max_seq_len
self.eval()
with torch.no_grad():
if img.dim() == 3:
img = img.unsqueeze(0)
memory, feat_h, feat_w = self.feature_extractor(img)
memory = self.pos_encoder_2d(memory, feat_h, feat_w)
memory = self.transformer_encoder(memory)
ys = torch.ones(1, 1).fill_(SOS_TOKEN).long().to(img.device)
for _ in range(max_length - 1):
tgt_embedded = self.pos_decoder_1d(self.token_embedding(ys))
tgt_mask = self._generate_square_subsequent_mask(ys.size(1)).to(img.device)
out = self.transformer_decoder(tgt_embedded, memory, tgt_mask=tgt_mask)
out = self.output_projection(out)
next_word = out[0, -1].argmax().item()
ys = torch.cat([ys, torch.ones(1, 1).long().fill_(next_word).to(img.device)], dim=1)
if next_word == EOS_TOKEN:
break
return ys[0]
def generate_batch(self, imgs, max_length=None):
"""Auto-regressive greedy batch generation."""
if max_length is None:
max_length = self.max_seq_len
self.eval()
batch_size = imgs.size(0)
with torch.no_grad():
memory, feat_h, feat_w = self.feature_extractor(imgs)
memory = self.pos_encoder_2d(memory, feat_h, feat_w)
memory = self.transformer_encoder(memory)
ys = torch.ones(batch_size, 1).fill_(SOS_TOKEN).long().to(imgs.device)
finished = torch.zeros(batch_size, dtype=torch.bool, device=imgs.device)
for _ in range(max_length - 1):
tgt_embedded = self.pos_decoder_1d(self.token_embedding(ys))
tgt_mask = self._generate_square_subsequent_mask(ys.size(1)).to(imgs.device)
out = self.transformer_decoder(tgt_embedded, memory, tgt_mask=tgt_mask)
out = self.output_projection(out)
next_tokens = out[:, -1].argmax(dim=-1)
next_tokens[finished] = PAD_TOKEN
ys = torch.cat([ys, next_tokens.unsqueeze(1)], dim=1)
finished = finished | (next_tokens == EOS_TOKEN)
if finished.all():
break
return ys
# ===============================
# Image Preprocessing
# ===============================
def preprocess_image(image_path, img_height, img_width):
"""Load and preprocess a paragraph image.
Aspect-ratio-preserving resize, right-aligned on white canvas for RTL."""
image = Image.open(image_path).convert("RGB")
orig_w, orig_h = image.size
scale = min(img_width / orig_w, img_height / orig_h)
new_w = int(orig_w * scale)
new_h = int(orig_h * scale)
image = image.resize((new_w, new_h), Image.Resampling.LANCZOS)
canvas = Image.new("RGB", (img_width, img_height), color=(255, 255, 255))
x_offset = img_width - new_w # Right-align for RTL
canvas.paste(image, (x_offset, 0))
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
])
return transform(canvas)
# ===============================
# Config Loader
# ===============================
def load_config(config_path):
"""Load architecture settings from config.json."""
with open(config_path, "r", encoding="utf-8") as f:
return json.load(f)
# ===============================
# Main
# ===============================
def main():
args = parse_args()
if args.image is None and args.image_dir is None:
print("Error: Provide --image or --image_dir")
return
# Device
if args.device:
device = torch.device(args.device)
else:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Device: {device}")
# Load config if provided (overrides CLI args)
if args.config_path and os.path.exists(args.config_path):
config = load_config(args.config_path)
print(f"Loaded config from: {args.config_path}")
args.hidden_size = config.get("hidden_size", args.hidden_size)
args.encoder_layers = config.get("num_encoder_layers", args.encoder_layers)
args.decoder_layers = config.get("num_decoder_layers", args.decoder_layers)
args.num_heads = config.get("num_attention_heads", args.num_heads)
args.ff_dim = config.get("feed_forward_dim", args.ff_dim)
args.max_seq_len = config.get("max_sequence_length", args.max_seq_len)
args.img_height = config.get("image_height", args.img_height)
args.img_width = config.get("image_width", args.img_width)
if "use_upsample" in config:
args.use_upsample = config["use_upsample"]
args.no_upsample = not config["use_upsample"]
use_upsample = args.use_upsample and not args.no_upsample
# Vocabulary
char_list, idx_to_char = load_vocabulary(args.vocab_path)
vocab_size = len(char_list)
print(f"Vocabulary: {vocab_size} tokens")
# Model
model = TransformerOCRParagraphModel(
vocab_size=vocab_size,
hidden_size=args.hidden_size,
nhead=args.num_heads,
num_encoder_layers=args.encoder_layers,
num_decoder_layers=args.decoder_layers,
dim_feedforward=args.ff_dim,
use_upsample=use_upsample,
max_seq_len=args.max_seq_len
).to(device)
# Load weights
print(f"Loading weights: {args.model_path}")
if args.model_path.endswith(".safetensors"):
from safetensors.torch import load_file
state_dict = load_file(args.model_path)
else:
checkpoint = torch.load(args.model_path, map_location=device)
state_dict = checkpoint.get("model_state_dict", checkpoint)
# Handle PE size mismatches
model_state = model.state_dict()
filtered = {}
for key, value in state_dict.items():
if key in model_state:
if value.shape == model_state[key].shape:
filtered[key] = value
model.load_state_dict(filtered, strict=False)
model.eval()
total_params = sum(p.numel() for n, p in model.named_parameters() if '.pe' not in n)
print(f"Model loaded: {total_params:,} parameters")
print(f"Upsample: {'ON' if use_upsample else 'OFF'}")
print(f"Image size: {args.img_height} x {args.img_width}")
print(f"Max sequence length: {args.max_seq_len}")
# Collect images
image_paths = []
if args.image:
image_paths = [args.image]
elif args.image_dir:
for ext in ("*.tif", "*.tiff", "*.png", "*.jpg", "*.jpeg", "*.bmp"):
image_paths.extend(glob.glob(os.path.join(args.image_dir, ext)))
image_paths.extend(glob.glob(os.path.join(args.image_dir, ext.upper())))
image_paths = sorted(list(set(image_paths)))
if not image_paths:
print("No images found.")
return
print(f"\nProcessing {len(image_paths)} image(s)...\n")
# Output file
out_file = None
if args.output_file:
out_file = open(args.output_file, "w", encoding="utf-8")
total_time = 0
for img_path in image_paths:
filename = os.path.basename(img_path)
# Preprocess
tensor = preprocess_image(img_path, args.img_height, args.img_width).to(device)
# Inference with timing
if torch.cuda.is_available():
torch.cuda.synchronize()
start = time.perf_counter()
output = model.generate(tensor)
if torch.cuda.is_available():
torch.cuda.synchronize()
elapsed = time.perf_counter() - start
total_time += elapsed
# Decode
text = decode_output(output, idx_to_char)
lines = text.split('\n')
# Display
print(f"{'='*60}")
print(f"File: {filename}")
if args.show_timing:
print(f"Time: {elapsed*1000:.1f} ms")
print(f"Lines detected: {len(lines)}")
print(f"{'─'*60}")
for i, line in enumerate(lines):
print(f" Line {i+1}: {line}")
print()
# Save to file
if out_file:
out_file.write(f"# {filename}\n")
out_file.write(text + "\n\n")
# Summary
print(f"{'='*60}")
print(f"Done. {len(image_paths)} image(s) processed.")
if args.show_timing:
avg_ms = (total_time / len(image_paths)) * 1000
print(f"Average inference: {avg_ms:.1f} ms/image")
print(f"Total time: {total_time:.2f} s")
if out_file:
out_file.close()
print(f"Predictions saved to: {args.output_file}")
if __name__ == "__main__":
main()