| """
|
| Kurdish Handwritten Paragraph Recognition - Fine-tuning Script
|
| DenseNet121-Transformer Architecture
|
|
|
| Fine-tunes a pre-trained model on real handwritten paragraph images.
|
| Loads weights from pretrain.py output checkpoint.
|
|
|
| Usage:
|
| python finetune.py --data_dir ./data/UniqueHandwrittenParagraphs \
|
| --vocab_path ./vocab.json \
|
| --pretrained_path ./output/pretrained_model.pth
|
|
|
| python finetune.py --data_dir ./data/DASNUS-Paragraphs \
|
| --vocab_path ./vocab.json \
|
| --pretrained_path ./output/pretrained_model.pth \
|
| --freeze_epochs 10
|
| """
|
|
|
| import os
|
| import glob
|
| import time
|
| import argparse
|
| import json
|
| import math
|
| import random
|
| import re
|
| import numpy as np
|
| from PIL import Image
|
| from datetime import datetime
|
|
|
| import torch
|
| import torch.nn as nn
|
| import torch.optim as optim
|
| import torch.utils.data as data
|
| import torchvision.transforms as transforms
|
| import torchvision.models as models
|
| from torchvision.transforms import InterpolationMode
|
| from torch.nn import functional as F
|
| from torch.amp import autocast, GradScaler
|
| from tqdm import tqdm
|
| import gc
|
|
|
|
|
|
|
|
|
|
|
|
|
| def parse_args():
|
| parser = argparse.ArgumentParser(
|
| description="Kurdish Handwritten Paragraph Recognition - Fine-tuning")
|
|
|
|
|
| parser.add_argument("--data_dir", type=str, required=True,
|
| help="Root directory with Training/, Validation/, Testing/ subfolders")
|
| parser.add_argument("--vocab_path", type=str, required=True,
|
| help="Path to vocabulary JSON file (vocab.json)")
|
| parser.add_argument("--pretrained_path", type=str, required=True,
|
| help="Path to pre-trained model checkpoint (.pth)")
|
|
|
|
|
| parser.add_argument("--img_height", type=int, default=600)
|
| parser.add_argument("--img_width", type=int, default=1235)
|
| parser.add_argument("--max_seq_len", type=int, default=555)
|
|
|
|
|
| parser.add_argument("--batch_size", type=int, default=16)
|
| parser.add_argument("--num_epochs", type=int, default=80)
|
| parser.add_argument("--learning_rate", type=float, default=5e-5)
|
| parser.add_argument("--grad_clip", type=float, default=5.0)
|
| parser.add_argument("--weight_decay", type=float, default=1e-4)
|
| parser.add_argument("--seed", type=int, default=42)
|
|
|
|
|
| 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("--dropout", type=float, default=0.2)
|
| parser.add_argument("--use_upsample", action="store_true", default=True,
|
| help="Enable horizontal upsampling layer (default: True)")
|
| parser.add_argument("--no_upsample", action="store_true",
|
| help="Disable horizontal upsampling layer")
|
|
|
|
|
| parser.add_argument("--tf_noise_rate", type=float, default=0.05,
|
| help="Teacher forcing noise rate (default: 0.05)")
|
|
|
|
|
| parser.add_argument("--freeze_epochs", type=int, default=10,
|
| help="Number of epochs to freeze CNN encoder (default: 10)")
|
| parser.add_argument("--encoder_lr_mult", type=float, default=0.1,
|
| help="Learning rate multiplier for encoder (default: 0.1)")
|
|
|
|
|
| parser.add_argument("--lr_patience", type=int, default=5,
|
| help="ReduceLROnPlateau patience")
|
| parser.add_argument("--lr_factor", type=float, default=0.5,
|
| help="ReduceLROnPlateau factor")
|
|
|
|
|
| parser.add_argument("--patience", type=int, default=15)
|
|
|
|
|
| parser.add_argument("--mixed_precision", action="store_true", default=True)
|
| parser.add_argument("--no_mixed_precision", action="store_true")
|
| parser.add_argument("--no_aug", action="store_true",
|
| help="Disable data augmentation")
|
| parser.add_argument("--clean_text", action="store_true", default=True,
|
| help="Clean invisible Unicode characters from labels")
|
| parser.add_argument("--no_clean_text", action="store_true")
|
|
|
|
|
| parser.add_argument("--cer_every", type=int, default=5,
|
| help="Compute train CER every N epochs (0 to disable)")
|
| parser.add_argument("--cer_max_samples", type=int, default=256,
|
| help="Max samples for train CER computation")
|
|
|
|
|
| parser.add_argument("--output_dir", type=str, default="./output",
|
| help="Directory to save model and logs")
|
| parser.add_argument("--model_name", type=str, default="finetuned_model",
|
| help="Base name for saved model file")
|
|
|
| return parser.parse_args()
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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'")
|
|
|
| char_to_idx = {char: idx for idx, char in enumerate(char_list)}
|
| idx_to_char = {idx: char for idx, char in enumerate(char_list)}
|
|
|
| return char_list, char_to_idx, idx_to_char
|
|
|
|
|
|
|
| PAD_TOKEN = 0
|
| SOS_TOKEN = 1
|
| EOS_TOKEN = 2
|
|
|
|
|
|
|
|
|
|
|
|
|
| INVISIBLE_CHARS = [
|
| '\u200e', '\u200f', '\u200b', '\u200d', '\ufeff', '\u00ad',
|
| '\u2060', '\u2061', '\u2062', '\u2063', '\u2064',
|
| '\u206a', '\u206b', '\u206c', '\u206d', '\u206e', '\u206f',
|
| '\u2028', '\u2029',
|
| ]
|
|
|
|
|
| def clean_text(text):
|
| """Remove invisible Unicode characters that inflate CER.
|
| Preserves ZWNJ (U+200C) which is used in Kurdish."""
|
| for char in INVISIBLE_CHARS:
|
| if char != '\u200c':
|
| text = text.replace(char, '')
|
| text = re.sub(r' +', ' ', text)
|
| lines = text.split('\n')
|
| lines = [line.strip() for line in lines]
|
| return '\n'.join(lines)
|
|
|
|
|
|
|
|
|
|
|
|
|
| def tensor_to_text(tensor, idx_to_char):
|
| """Convert a 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
|
|
|
|
|
|
|
|
|
|
|
|
|
| class KurdishParagraphDataset(data.Dataset):
|
| """Dataset for Kurdish handwritten paragraph images."""
|
|
|
| def __init__(self, root_dir, transform=None, max_seq_len=555,
|
| img_height=600, img_width=1235, char_to_idx=None,
|
| clean_text_enabled=True):
|
| self.transform = transform
|
| self.max_seq_len = max_seq_len
|
| self.img_height = img_height
|
| self.img_width = img_width
|
| self.char_to_idx = char_to_idx
|
| self.clean_text_enabled = clean_text_enabled
|
|
|
| self.data = []
|
| image_files = []
|
| for ext in ["*.tif", "*.tiff", "*.png", "*.jpg", "*.jpeg"]:
|
| image_files.extend(glob.glob(os.path.join(root_dir, ext)))
|
| image_files.extend(glob.glob(os.path.join(root_dir, ext.upper())))
|
| image_files = sorted(list(set(image_files)))
|
|
|
| for img_path in image_files:
|
| label_path = os.path.splitext(img_path)[0] + ".txt"
|
| if not os.path.exists(label_path):
|
| continue
|
| try:
|
| with open(label_path, "r", encoding="utf-8") as f:
|
| text = f.read().strip()
|
| except Exception:
|
| try:
|
| with open(label_path, "r", encoding="utf-8-sig") as f:
|
| text = f.read().strip()
|
| except Exception:
|
| continue
|
|
|
| if self.clean_text_enabled:
|
| text = clean_text(text)
|
| if len(text) > 0:
|
| self.data.append((img_path, text))
|
|
|
| print(f" Loaded {len(self.data)} paragraph images from {root_dir}")
|
|
|
| def __len__(self):
|
| return len(self.data)
|
|
|
| def __getitem__(self, idx):
|
| img_path, text = self.data[idx]
|
|
|
| image = Image.open(img_path).convert("RGB")
|
| orig_width, orig_height = image.size
|
|
|
| scale = min(self.img_width / orig_width, self.img_height / orig_height)
|
| new_width = int(orig_width * scale)
|
| new_height = int(orig_height * scale)
|
| image = image.resize((new_width, new_height), Image.Resampling.LANCZOS)
|
|
|
| canvas = Image.new('RGB', (self.img_width, self.img_height), (255, 255, 255))
|
| x_offset = self.img_width - new_width
|
| canvas.paste(image, (x_offset, 0))
|
|
|
| if self.transform:
|
| canvas = self.transform(canvas)
|
|
|
| indices = ([SOS_TOKEN] +
|
| [self.char_to_idx.get(c, self.char_to_idx.get(" ", 0)) for c in text] +
|
| [EOS_TOKEN])
|
| if len(indices) > self.max_seq_len:
|
| indices = indices[:self.max_seq_len - 1] + [EOS_TOKEN]
|
|
|
| target = torch.LongTensor(indices)
|
| return canvas, target, len(indices), text
|
|
|
|
|
| def collate_fn(batch):
|
| """Collate function with padding for variable-length targets."""
|
| batch.sort(key=lambda x: x[2], reverse=True)
|
| images, targets, lengths, texts = zip(*batch)
|
|
|
| images = torch.stack(images, 0)
|
| max_length = max(lengths)
|
|
|
| padded = torch.ones(len(targets), max_length).long() * PAD_TOKEN
|
| for i, target in enumerate(targets):
|
| padded[i, :lengths[i]] = target[:lengths[i]]
|
|
|
| return images, padded, torch.LongTensor(lengths), texts
|
|
|
|
|
|
|
|
|
|
|
|
|
| def build_train_transform():
|
| """Standard augmentation for fine-tuning."""
|
| class FinetuneTransform:
|
| def __call__(self, img):
|
| if random.random() < 0.5:
|
| img = transforms.ColorJitter(
|
| brightness=0.15, contrast=0.15,
|
| saturation=0.05, hue=0.01)(img)
|
|
|
| if random.random() < 0.4:
|
| img = transforms.RandomAffine(
|
| degrees=2, translate=(0.02, 0.02),
|
| scale=(0.97, 1.03), shear=(-2, 2),
|
| interpolation=InterpolationMode.BILINEAR, fill=255)(img)
|
|
|
| if random.random() < 0.15:
|
| img = transforms.GaussianBlur(
|
| kernel_size=3, sigma=(0.1, 0.5))(img)
|
|
|
| img = transforms.ToTensor()(img)
|
|
|
| if random.random() < 0.2:
|
| noise = torch.randn_like(img) * 0.01
|
| img = torch.clamp(img + noise, 0.0, 1.0)
|
|
|
| img = transforms.Normalize(
|
| (0.485, 0.456, 0.406), (0.229, 0.224, 0.225))(img)
|
| return img
|
|
|
| return FinetuneTransform()
|
|
|
|
|
| def build_eval_transform():
|
| """Evaluation transform (normalisation only)."""
|
| return transforms.Compose([
|
| transforms.ToTensor(),
|
| transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
|
| ])
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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), :]
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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.2,
|
| use_upsample=True, max_seq_len=555,
|
| tf_noise_rate=0.05):
|
| super().__init__()
|
|
|
| self.max_seq_len = max_seq_len
|
| self.vocab_size = vocab_size
|
| self.tf_noise_rate = tf_noise_rate
|
|
|
| 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)
|
| self.hidden_size = hidden_size
|
|
|
| nn.init.xavier_uniform_(self.token_embedding.weight)
|
| nn.init.xavier_uniform_(self.output_projection.weight)
|
|
|
| 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 _add_teacher_forcing_noise(self, tgt_input):
|
| if self.tf_noise_rate <= 0 or not self.training:
|
| return tgt_input
|
| noise_mask = (torch.rand_like(tgt_input.float()) < self.tf_noise_rate)
|
| noise_mask = noise_mask & (tgt_input != PAD_TOKEN) & (tgt_input != SOS_TOKEN)
|
| random_tokens = torch.randint(3, self.vocab_size, tgt_input.shape, device=tgt_input.device)
|
| return torch.where(noise_mask, random_tokens, tgt_input)
|
|
|
| def forward(self, src, tgt, tgt_key_padding_mask=None):
|
| memory, feat_h, feat_w = self.feature_extractor(src)
|
| memory = self.pos_encoder_2d(memory, feat_h, feat_w)
|
| memory = self.transformer_encoder(memory)
|
|
|
| tgt_input = self._add_teacher_forcing_noise(tgt[:, :-1])
|
| tgt_embedded = self.pos_decoder_1d(self.token_embedding(tgt_input))
|
|
|
| tgt_mask = self._generate_square_subsequent_mask(tgt_embedded.size(1)).to(src.device)
|
| tgt_pad_mask = tgt_key_padding_mask[:, :-1] if tgt_key_padding_mask is not None else None
|
|
|
| output = self.transformer_decoder(
|
| tgt_embedded, memory,
|
| tgt_mask=tgt_mask, tgt_key_padding_mask=tgt_pad_mask)
|
|
|
| return self.output_projection(output)
|
|
|
| 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
|
|
|
| def freeze_encoder(self):
|
| """Freeze CNN backbone parameters."""
|
| for param in self.feature_extractor.parameters():
|
| param.requires_grad = False
|
| print(" Encoder (CNN) frozen")
|
|
|
| def unfreeze_encoder(self):
|
| """Unfreeze CNN backbone parameters."""
|
| for param in self.feature_extractor.parameters():
|
| param.requires_grad = True
|
| print(" Encoder (CNN) unfrozen")
|
|
|
|
|
|
|
|
|
|
|
|
|
| def load_pretrained_weights(model, pretrained_path, device):
|
| """Load pre-trained weights, handling PE size mismatches gracefully."""
|
| print(f"\n Loading pre-trained model: {pretrained_path}")
|
|
|
| if not os.path.exists(pretrained_path):
|
| raise FileNotFoundError(f"Checkpoint not found: {pretrained_path}")
|
|
|
| ckpt = torch.load(pretrained_path, map_location=device)
|
|
|
| if 'epoch' in ckpt:
|
| print(f" Pre-trained epoch: {ckpt['epoch']}")
|
| if 'val_cer' in ckpt:
|
| print(f" Pre-trained Val CER: {ckpt['val_cer']:.4f}")
|
|
|
| state_dict = ckpt.get('model_state_dict', ckpt)
|
| model_state = model.state_dict()
|
|
|
| loaded, skipped = {}, []
|
| for key, value in state_dict.items():
|
| if key in model_state:
|
| if value.shape == model_state[key].shape:
|
| loaded[key] = value
|
| else:
|
| skipped.append((key, f"{value.shape} vs {model_state[key].shape}"))
|
| else:
|
| skipped.append((key, "not in model"))
|
|
|
| model.load_state_dict(loaded, strict=False)
|
|
|
| print(f" Loaded: {len(loaded)}/{len(model_state)} parameters")
|
| if skipped:
|
| print(f" Skipped: {len(skipped)} (PE buffers regenerated)")
|
|
|
| return model
|
|
|
|
|
|
|
|
|
|
|
|
|
| def levenshtein_distance(s1, s2):
|
| if len(s1) < len(s2):
|
| return levenshtein_distance(s2, s1)
|
| if len(s2) == 0:
|
| return len(s1)
|
| prev = range(len(s2) + 1)
|
| for c1 in s1:
|
| curr = [prev[0] + 1]
|
| for j, c2 in enumerate(s2):
|
| curr.append(min(prev[j + 1] + 1, curr[j] + 1, prev[j] + (c1 != c2)))
|
| prev = curr
|
| return prev[-1]
|
|
|
|
|
| def calculate_cer(preds, targets):
|
| total_dist = sum(levenshtein_distance(p, t) for p, t in zip(preds, targets))
|
| total_chars = sum(len(t) for t in targets)
|
| return total_dist / max(1, total_chars)
|
|
|
|
|
| def calculate_wer(preds, targets):
|
| total_dist = sum(levenshtein_distance(p.split(), t.split()) for p, t in zip(preds, targets))
|
| total_words = sum(len(t.split()) for t in targets)
|
| return total_dist / max(1, total_words)
|
|
|
|
|
| def calculate_line_accuracy(preds, targets):
|
| total, correct = 0, 0
|
| for pred, true in zip(preds, targets):
|
| pred_lines = pred.split('\n')
|
| true_lines = true.split('\n')
|
| total += len(true_lines)
|
| for pl, tl in zip(pred_lines, true_lines):
|
| if pl.strip() == tl.strip():
|
| correct += 1
|
| return correct / max(1, total)
|
|
|
|
|
| def evaluate_cer_batch(model, dataloader, device, idx_to_char, max_samples=None):
|
| """Compute CER using batch generation."""
|
| model.eval()
|
| all_preds, all_targets = [], []
|
| count = 0
|
|
|
| with torch.no_grad():
|
| for images, _, _, texts in dataloader:
|
| images = images.to(device)
|
| if max_samples and count + images.size(0) > max_samples:
|
| images = images[:max_samples - count]
|
| texts = texts[:max_samples - count]
|
|
|
| batch_output = model.generate_batch(images)
|
| preds = [tensor_to_text(seq, idx_to_char) for seq in batch_output]
|
| all_preds.extend(preds)
|
| all_targets.extend(texts)
|
| count += len(preds)
|
|
|
| if max_samples and count >= max_samples:
|
| break
|
|
|
| return calculate_cer(all_preds, all_targets)
|
|
|
|
|
|
|
|
|
|
|
|
|
| def comprehensive_evaluation(model, dataloader, device, idx_to_char):
|
| """Full evaluation with CER, WER, line accuracy, and timing."""
|
| model.eval()
|
| all_preds, all_targets = [], []
|
| inference_times = []
|
|
|
|
|
| with torch.no_grad():
|
| for images, _, _, _ in dataloader:
|
| images = images.to(device)
|
| _ = model.generate_batch(images[:min(3, images.size(0))])
|
| break
|
|
|
| if torch.cuda.is_available():
|
| torch.cuda.synchronize()
|
|
|
| with torch.no_grad():
|
| for images, _, _, texts in tqdm(dataloader, desc="Evaluating"):
|
| images = images.to(device)
|
| batch_size = images.size(0)
|
|
|
| if torch.cuda.is_available():
|
| torch.cuda.synchronize()
|
| start = time.perf_counter()
|
|
|
| batch_output = model.generate_batch(images)
|
|
|
| if torch.cuda.is_available():
|
| torch.cuda.synchronize()
|
| elapsed = time.perf_counter() - start
|
|
|
| per_sample = elapsed / batch_size
|
| inference_times.extend([per_sample] * batch_size)
|
|
|
| preds = [tensor_to_text(seq, idx_to_char) for seq in batch_output]
|
| all_preds.extend(preds)
|
| all_targets.extend(texts)
|
|
|
| cer = calculate_cer(all_preds, all_targets)
|
| wer = calculate_wer(all_preds, all_targets)
|
| line_acc = calculate_line_accuracy(all_preds, all_targets)
|
|
|
| total_params = sum(p.numel() for p in model.parameters())
|
|
|
| return {
|
| 'cer': cer, 'wer': wer, 'line_accuracy': line_acc,
|
| 'avg_inference_ms': np.mean(inference_times) * 1000,
|
| 'std_inference_ms': np.std(inference_times) * 1000,
|
| 'fps': len(inference_times) / sum(inference_times),
|
| 'total_params': total_params,
|
| 'predictions': all_preds, 'targets': all_targets,
|
| }
|
|
|
|
|
|
|
|
|
|
|
|
|
| class EarlyStopping:
|
| def __init__(self, patience=15):
|
| self.patience = patience
|
| self.counter = 0
|
| self.best_cer = float('inf')
|
| self.early_stop = False
|
|
|
| def __call__(self, val_cer, model, epoch, path):
|
| if val_cer < self.best_cer:
|
| self.best_cer = val_cer
|
| self.counter = 0
|
| torch.save({
|
| 'epoch': epoch,
|
| 'model_state_dict': model.state_dict(),
|
| 'val_cer': val_cer
|
| }, path)
|
| print(f" Model saved (Val CER: {val_cer:.4f})")
|
| else:
|
| self.counter += 1
|
| print(f" Early stopping: {self.counter}/{self.patience}")
|
| if self.counter >= self.patience:
|
| self.early_stop = True
|
| print(" Early stopping triggered.")
|
|
|
|
|
|
|
|
|
|
|
|
|
| def train_epoch(model, dataloader, optimizer, criterion, device, scaler,
|
| use_mixed_precision=True, grad_clip=5.0):
|
| """Train for one epoch."""
|
| model.train()
|
| epoch_loss = 0
|
|
|
| for images, targets, _, _ in tqdm(dataloader, desc="Training"):
|
| images, targets = images.to(device), targets.to(device)
|
| tgt_pad_mask = (targets == PAD_TOKEN).to(device)
|
|
|
| optimizer.zero_grad()
|
|
|
| if use_mixed_precision:
|
| with autocast(device_type='cuda'):
|
| outputs = model(images, targets, tgt_key_padding_mask=tgt_pad_mask)
|
| loss = criterion(outputs.reshape(-1, outputs.shape[-1]),
|
| targets[:, 1:].reshape(-1))
|
| scaler.scale(loss).backward()
|
| scaler.unscale_(optimizer)
|
| torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)
|
| scaler.step(optimizer)
|
| scaler.update()
|
| else:
|
| outputs = model(images, targets, tgt_key_padding_mask=tgt_pad_mask)
|
| loss = criterion(outputs.reshape(-1, outputs.shape[-1]),
|
| targets[:, 1:].reshape(-1))
|
| loss.backward()
|
| torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)
|
| optimizer.step()
|
|
|
| epoch_loss += loss.item()
|
|
|
| return epoch_loss / len(dataloader)
|
|
|
|
|
| def evaluate_loss(model, dataloader, criterion, device, use_mixed_precision=True):
|
| """Evaluate model loss."""
|
| model.eval()
|
| epoch_loss = 0
|
|
|
| with torch.no_grad():
|
| for images, targets, _, _ in dataloader:
|
| images, targets = images.to(device), targets.to(device)
|
| tgt_pad_mask = (targets == PAD_TOKEN).to(device)
|
|
|
| if use_mixed_precision:
|
| with autocast(device_type='cuda'):
|
| outputs = model(images, targets, tgt_key_padding_mask=tgt_pad_mask)
|
| loss = criterion(outputs.reshape(-1, outputs.shape[-1]),
|
| targets[:, 1:].reshape(-1))
|
| else:
|
| outputs = model(images, targets, tgt_key_padding_mask=tgt_pad_mask)
|
| loss = criterion(outputs.reshape(-1, outputs.shape[-1]),
|
| targets[:, 1:].reshape(-1))
|
| epoch_loss += loss.item()
|
|
|
| return epoch_loss / len(dataloader)
|
|
|
|
|
|
|
|
|
|
|
|
|
| def main():
|
| args = parse_args()
|
|
|
|
|
| use_upsample = args.use_upsample and not args.no_upsample
|
| use_mixed_precision = args.mixed_precision and not args.no_mixed_precision
|
| use_clean_text = args.clean_text and not args.no_clean_text
|
|
|
|
|
| torch.manual_seed(args.seed)
|
| random.seed(args.seed)
|
| np.random.seed(args.seed)
|
|
|
|
|
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| print(f"Device: {device}")
|
| if torch.cuda.is_available():
|
| print(f"GPU: {torch.cuda.get_device_name(0)}")
|
|
|
|
|
| os.makedirs(args.output_dir, exist_ok=True)
|
|
|
|
|
| char_list, char_to_idx, idx_to_char = load_vocabulary(args.vocab_path)
|
| vocab_size = len(char_list)
|
| print(f"Vocabulary size: {vocab_size}")
|
|
|
|
|
| train_transform = build_eval_transform() if args.no_aug else build_train_transform()
|
| eval_transform = build_eval_transform()
|
|
|
|
|
| ds_kwargs = dict(
|
| max_seq_len=args.max_seq_len,
|
| img_height=args.img_height,
|
| img_width=args.img_width,
|
| char_to_idx=char_to_idx,
|
| clean_text_enabled=use_clean_text)
|
|
|
|
|
| train_dir = os.path.join(args.data_dir, "Training")
|
| val_dir = os.path.join(args.data_dir, "Validation")
|
| test_dir = os.path.join(args.data_dir, "Testing")
|
|
|
| train_dataset = KurdishParagraphDataset(train_dir, transform=train_transform, **ds_kwargs)
|
| val_dataset = KurdishParagraphDataset(val_dir, transform=eval_transform, **ds_kwargs)
|
| test_dataset = KurdishParagraphDataset(test_dir, transform=eval_transform, **ds_kwargs)
|
|
|
| loader_kwargs = dict(num_workers=0, pin_memory=True, collate_fn=collate_fn)
|
| train_loader = data.DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True, **loader_kwargs)
|
| val_loader = data.DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False, **loader_kwargs)
|
| test_loader = data.DataLoader(test_dataset, batch_size=args.batch_size, shuffle=False, **loader_kwargs)
|
|
|
| print(f"\n Training: {len(train_dataset)} | Validation: {len(val_dataset)} | Testing: {len(test_dataset)}")
|
|
|
|
|
| print("\nInitializing 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,
|
| dropout=args.dropout,
|
| use_upsample=use_upsample,
|
| max_seq_len=args.max_seq_len,
|
| tf_noise_rate=args.tf_noise_rate
|
| ).to(device)
|
|
|
|
|
| model = load_pretrained_weights(model, args.pretrained_path, device)
|
|
|
|
|
| if args.freeze_epochs > 0:
|
| model.freeze_encoder()
|
|
|
| total_params = sum(p.numel() for p in model.parameters())
|
| trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
| print(f" Total parameters: {total_params:,}")
|
| print(f" Trainable parameters: {trainable_params:,}")
|
|
|
|
|
| encoder_params = list(model.feature_extractor.parameters())
|
| other_params = [p for n, p in model.named_parameters() if 'feature_extractor' not in n]
|
|
|
| optimizer = optim.AdamW([
|
| {'params': encoder_params, 'lr': args.learning_rate * args.encoder_lr_mult},
|
| {'params': other_params, 'lr': args.learning_rate}
|
| ], weight_decay=args.weight_decay)
|
|
|
| scheduler = optim.lr_scheduler.ReduceLROnPlateau(
|
| optimizer, mode='min', factor=args.lr_factor,
|
| patience=args.lr_patience, min_lr=1e-7)
|
|
|
| criterion = nn.CrossEntropyLoss(ignore_index=PAD_TOKEN)
|
| scaler = GradScaler('cuda') if use_mixed_precision else None
|
| early_stopping = EarlyStopping(patience=args.patience)
|
|
|
| best_model_path = os.path.join(args.output_dir, f"{args.model_name}.pth")
|
|
|
|
|
| log_path = os.path.join(args.output_dir,
|
| f"{args.model_name}_LOG_{datetime.now():%Y%m%d_%H%M%S}.txt")
|
| log_file = open(log_path, 'w', encoding='utf-8')
|
|
|
| def log(msg):
|
| print(msg)
|
| log_file.write(msg + '\n')
|
| log_file.flush()
|
|
|
| log(f"\nFine-tuning started: {datetime.now():%Y-%m-%d %H:%M:%S}")
|
| log(f"Pre-trained model: {args.pretrained_path}")
|
| log(f"Config: {vars(args)}")
|
|
|
|
|
| initial_cer = evaluate_cer_batch(model, val_loader, device, idx_to_char)
|
| log(f"\n Initial Val CER (pre-trained): {initial_cer:.4f}")
|
|
|
|
|
| best_val_cer = float('inf')
|
|
|
| for epoch in range(1, args.num_epochs + 1):
|
| start_time = time.time()
|
|
|
|
|
| if epoch == args.freeze_epochs + 1 and args.freeze_epochs > 0:
|
| model.unfreeze_encoder()
|
| trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
| log(f"\n Epoch {epoch}: Encoder unfrozen ({trainable:,} trainable params)")
|
|
|
|
|
| train_loss = train_epoch(model, train_loader, optimizer, criterion,
|
| device, scaler, use_mixed_precision, args.grad_clip)
|
|
|
|
|
| train_cer = None
|
| if args.cer_every > 0 and epoch % args.cer_every == 0:
|
| train_cer = evaluate_cer_batch(model, train_loader, device,
|
| idx_to_char, args.cer_max_samples)
|
|
|
|
|
| val_loss = evaluate_loss(model, val_loader, criterion, device, use_mixed_precision)
|
| val_cer = evaluate_cer_batch(model, val_loader, device, idx_to_char)
|
|
|
| scheduler.step(val_cer)
|
| elapsed = time.time() - start_time
|
| mins, secs = divmod(elapsed, 60)
|
|
|
| lr_enc = optimizer.param_groups[0]['lr']
|
| lr_dec = optimizer.param_groups[1]['lr']
|
|
|
| cer_str = f", Train CER: {train_cer:.4f}" if train_cer is not None else ""
|
| log(f"Epoch {epoch}/{args.num_epochs} ({mins:.0f}m {secs:.0f}s) | "
|
| f"Train Loss: {train_loss:.4f}{cer_str} | "
|
| f"Val Loss: {val_loss:.4f} | Val CER: {val_cer:.4f} | "
|
| f"LR: Enc={lr_enc:.2e}, Dec={lr_dec:.2e}")
|
|
|
| if val_cer < best_val_cer:
|
| best_val_cer = val_cer
|
|
|
| early_stopping(val_cer, model, epoch, best_model_path)
|
| if early_stopping.early_stop:
|
| break
|
|
|
| gc.collect()
|
| if torch.cuda.is_available():
|
| torch.cuda.empty_cache()
|
|
|
|
|
| log(f"\nLoading best model for final evaluation...")
|
| ckpt = torch.load(best_model_path, map_location=device)
|
| model.load_state_dict(ckpt['model_state_dict'])
|
| log(f" Best epoch: {ckpt['epoch']}, Best Val CER: {ckpt['val_cer']:.4f}")
|
|
|
|
|
| log(f"\n--- Validation Set ---")
|
| val_results = comprehensive_evaluation(model, val_loader, device, idx_to_char)
|
| log(f" CER: {val_results['cer']:.4f} | WER: {val_results['wer']:.4f} | "
|
| f"Line Acc: {val_results['line_accuracy']:.4f}")
|
| log(f" Inference: {val_results['avg_inference_ms']:.2f} ms | FPS: {val_results['fps']:.2f}")
|
|
|
|
|
| log(f"\n--- Test Set ---")
|
| test_results = comprehensive_evaluation(model, test_loader, device, idx_to_char)
|
| log(f" CER: {test_results['cer']:.4f} ({(1-test_results['cer'])*100:.2f}% accuracy)")
|
| log(f" WER: {test_results['wer']:.4f} ({(1-test_results['wer'])*100:.2f}% accuracy)")
|
| log(f" Line Accuracy: {test_results['line_accuracy']:.4f}")
|
| log(f" Inference: {test_results['avg_inference_ms']:.2f} ± {test_results['std_inference_ms']:.2f} ms")
|
| log(f" FPS: {test_results['fps']:.2f}")
|
| log(f" Parameters: {test_results['total_params']:,}")
|
|
|
|
|
| log(f"\n--- Sample Predictions ---")
|
| for i in range(min(5, len(test_results['predictions']))):
|
| log(f"\nSample {i + 1}:")
|
| pred_preview = test_results['predictions'][i][:200]
|
| true_preview = test_results['targets'][i][:200]
|
| log(f" Predicted: {pred_preview}")
|
| log(f" Actual: {true_preview}")
|
|
|
| log(f"\nFine-tuning complete: {datetime.now():%Y-%m-%d %H:%M:%S}")
|
| log(f"Best model: {best_model_path}")
|
| log(f"Improvement: {initial_cer:.4f} -> {best_val_cer:.4f} "
|
| f"({(initial_cer - best_val_cer)*100:.2f}% absolute)")
|
|
|
| log_file.close()
|
| print(f"Log saved to: {log_path}")
|
|
|
|
|
| if __name__ == "__main__":
|
| main() |