""" Local inference: diacritize Arabic text using a trained checkpoint. Download checkpoint + vocab.json from Google Drive and run: python inference.py --checkpoint checkpoints/best_model.pt --text "النص العربي" """ import json import sys from pathlib import Path import torch import torch.nn as nn # ── Tashkeel constants ──────────────────────────────────────────────────────── TASHKEEL_SET = set('ًٌٍَُِّْٰٕٓٔ') DIACRITIC_LABELS = { '': 0, 'َ': 1, 'ُ': 2, 'ِ': 3, 'ً': 4, 'ٌ': 5, 'ٍ': 6, 'ْ': 7, 'ّ': 8, 'َّ': 9, 'ُّ': 10, 'ِّ': 11, 'ًّ': 12, 'ٌّ': 13, 'ٍّ': 14, 'ٰ': 1, 'ٓ': 1, 'ٔ': 0, 'ٕ': 0, } NUM_LABELS = 15 LABEL_TO_DIAC = { 0: '', 1: 'َ', 2: 'ُ', 3: 'ِ', 4: 'ً', 5: 'ٌ', 6: 'ٍ', 7: 'ْ', 8: 'ّ', 9: 'َّ', 10: 'ُّ', 11: 'ِّ', 12: 'ًّ', 13: 'ٌّ', 14: 'ٍّ', } # ── Model (must match Colab training definition) ────────────────────────────── class TransformerDiacritizer(nn.Module): def __init__(self, vocab_size, embed_dim=768, n_heads=12, n_layers=24, ffn_dim=3072, dropout=0.0, max_len=256, num_labels=NUM_LABELS): super().__init__() self.max_len = max_len self.embed = nn.Embedding(vocab_size, embed_dim, padding_idx=0) self.pos_embed = nn.Embedding(max_len + 2, embed_dim) self.drop = nn.Dropout(dropout) layer = nn.TransformerEncoderLayer( d_model=embed_dim, nhead=n_heads, dim_feedforward=ffn_dim, dropout=dropout, batch_first=True, norm_first=True, ) self.layers = nn.ModuleList([ nn.TransformerEncoderLayer( d_model=embed_dim, nhead=n_heads, dim_feedforward=ffn_dim, dropout=dropout, batch_first=True, norm_first=True, ) for _ in range(n_layers) ]) self.norm = nn.LayerNorm(embed_dim) self.fc = nn.Linear(embed_dim, num_labels) def forward(self, input_ids, lengths): B, T = input_ids.shape pos = torch.arange(T, device=input_ids.device).unsqueeze(0).expand(B, -1) x = self.drop(self.embed(input_ids) + self.pos_embed(pos)) pad_mask = torch.arange(T, device=input_ids.device)[None,:] >= lengths[:,None] for layer in self.layers: x = layer(x, src_key_padding_mask=pad_mask) return self.fc(self.norm(x)) # ── Inference helpers ───────────────────────────────────────────────────────── def strip_tashkeel(text: str) -> str: return ''.join(c for c in text if c not in TASHKEEL_SET) def diacritize(text: str, model: TransformerDiacritizer, vocab: dict, device: torch.device, max_len: int = 256) -> str: model.eval() raw = strip_tashkeel(text) chars = list(raw[:max_len]) if not chars: return text unk = vocab.get('', 1) ids = torch.tensor([[vocab.get(c, unk) for c in chars]], dtype=torch.long, device=device) lengths = torch.tensor([len(chars)]) with torch.no_grad(): logits = model(ids, lengths) pred = logits.argmax(dim=-1)[0].cpu().tolist() return ''.join(c + LABEL_TO_DIAC.get(l, '') for c, l in zip(chars, pred)) # ── CLI ─────────────────────────────────────────────────────────────────────── if __name__ == '__main__': import argparse parser = argparse.ArgumentParser(description='Diacritize Arabic text') parser.add_argument('--checkpoint', required=True, help='Path to best_model.pt') parser.add_argument('--vocab', default=None, help='Path to vocab.json (default: same dir as checkpoint)') parser.add_argument('--text', default=None, help='Arabic text to diacritize') parser.add_argument('--file', default=None, help='Plain text file (one sentence per line)') parser.add_argument('--max_len', type=int, default=256) args = parser.parse_args() ckpt_path = Path(args.checkpoint) vocab_path = Path(args.vocab) if args.vocab else ckpt_path.parent / 'vocab.json' device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') with open(vocab_path, encoding='utf-8') as f: vocab = json.load(f) ckpt = torch.load(ckpt_path, map_location=device) a = ckpt['args'] model = TransformerDiacritizer( vocab_size = len(vocab), embed_dim = a.get('embed_dim', 512), n_heads = a.get('n_heads', 8), n_layers = a.get('n_layers', 12), ffn_dim = a.get('ffn_dim', 2048), dropout = 0.0, max_len = args.max_len, ).to(device) model.load_state_dict(ckpt['model_state']) model.eval() print(f"Loaded: epoch {ckpt['epoch']}, val DER {ckpt['val_der']:.4%}") print(f"Config: {a.get('n_layers')}L × {a.get('embed_dim')}D × FFN{a.get('ffn_dim')}") if args.text: print(diacritize(args.text, model, vocab, device, args.max_len)) elif args.file: with open(args.file, encoding='utf-8') as f: lines = [l.strip() for l in f if l.strip()] for line in lines: print(diacritize(line, model, vocab, device, args.max_len)) else: print('Enter Arabic text (Ctrl+C to quit):') while True: try: line = input('> ').strip() if line: print(diacritize(line, model, vocab, device, args.max_len)) except (KeyboardInterrupt, EOFError): break