#!/usr/bin/env python3 """15M char-level seq2seq transformer for Hinglish transliteration. Locked architecture: d_model=384, nhead=6, 4 enc + 4 dec layers, d_ff=1280, dropout 0.1, max_len=64, norm_first, weight-tied embedding/output. Specials: pad=0, bos=1, eos=2, unk=3. Checkpoint schema: {"config": ..., "vocab": ..., "model": ..., "extra": ...} """ import math import torch import torch.nn as nn PAD, BOS, EOS, UNK = 0, 1, 2, 3 SPECIALS = ["", "", "", ""] DEFAULT_CONFIG = { "d_model": 384, "nhead": 6, "num_encoder_layers": 4, "num_decoder_layers": 4, "dim_feedforward": 1280, "dropout": 0.1, "max_len": 64, } class CharTokenizer: """Char-level tokenizer with a frozen vocab stored in the checkpoint.""" def __init__(self, vocab): # vocab: list of tokens, index = id. First 4 must be SPECIALS. assert vocab[:4] == SPECIALS, "vocab must start with pad/bos/eos/unk" self.vocab = list(vocab) self.stoi = {c: i for i, c in enumerate(self.vocab)} @classmethod def build(cls, texts, min_count=5): from collections import Counter cnt = Counter() for t in texts: cnt.update(t) chars = sorted([c for c, n in cnt.items() if n >= min_count]) return cls(SPECIALS + chars) def __len__(self): return len(self.vocab) def encode(self, text, max_len, add_bos=True, add_eos=True): ids = [self.stoi.get(c, UNK) for c in text] budget = max_len - int(add_bos) - int(add_eos) ids = ids[:budget] if add_bos: ids = [BOS] + ids if add_eos: ids = ids + [EOS] return ids def decode(self, ids): out = [] for i in ids: if i == EOS: break if i in (PAD, BOS): continue out.append(self.vocab[i] if i < len(self.vocab) else "") return "".join(out) class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=512): super().__init__() pe = torch.zeros(max_len, d_model) pos = torch.arange(max_len).unsqueeze(1).float() div = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] = torch.sin(pos * div) pe[:, 1::2] = torch.cos(pos * div) self.register_buffer("pe", pe.unsqueeze(0)) # (1, max_len, d_model) def forward(self, x): # x: (B, T, D) return x + self.pe[:, : x.size(1)] class TranslitModel(nn.Module): def __init__(self, vocab_size, config=None): super().__init__() cfg = dict(DEFAULT_CONFIG) if config: cfg.update(config) self.config = cfg d = cfg["d_model"] self.embed = nn.Embedding(vocab_size, d, padding_idx=PAD) nn.init.normal_(self.embed.weight, mean=0.0, std=d ** -0.5) with torch.no_grad(): self.embed.weight[PAD].zero_() self.pos = PositionalEncoding(d, max_len=max(128, cfg["max_len"] * 2)) self.transformer = nn.Transformer( d_model=d, nhead=cfg["nhead"], num_encoder_layers=cfg["num_encoder_layers"], num_decoder_layers=cfg["num_decoder_layers"], dim_feedforward=cfg["dim_feedforward"], dropout=cfg["dropout"], batch_first=True, norm_first=True, ) self.out = nn.Linear(d, vocab_size, bias=False) self.out.weight = self.embed.weight # weight tying self.scale = math.sqrt(d) def forward(self, src, tgt_in): # src: (B, S) roman ids; tgt_in: (B, T) devanagari ids shifted right src_pad = src == PAD tgt_pad = tgt_in == PAD causal = nn.Transformer.generate_square_subsequent_mask( tgt_in.size(1), device=tgt_in.device ) se = self.pos(self.embed(src) * self.scale) te = self.pos(self.embed(tgt_in) * self.scale) h = self.transformer( se, te, tgt_mask=causal, src_key_padding_mask=src_pad, tgt_key_padding_mask=tgt_pad, memory_key_padding_mask=src_pad, ) return self.out(h) # (B, T, V) def encode_src(self, src): src_pad = src == PAD se = self.pos(self.embed(src) * self.scale) mem = self.transformer.encoder(se, src_key_padding_mask=src_pad) return mem, src_pad def decode_step(self, mem, mem_pad, tgt_in): causal = nn.Transformer.generate_square_subsequent_mask( tgt_in.size(1), device=tgt_in.device ) te = self.pos(self.embed(tgt_in) * self.scale) h = self.transformer.decoder( te, mem, tgt_mask=causal, memory_key_padding_mask=mem_pad ) return self.out(h) def save_checkpoint(path, model, tokenizer, extra=None): torch.save( { "config": model.config, "vocab": tokenizer.vocab, "model": model.state_dict(), "extra": extra or {}, }, path, ) def load_checkpoint(path, map_location="cpu"): ckpt = torch.load(path, map_location=map_location, weights_only=False) tok = CharTokenizer(ckpt["vocab"]) model = TranslitModel(len(tok), ckpt["config"]) model.load_state_dict(ckpt["model"]) model.eval() return model, tok, ckpt.get("extra", {}) def count_params(model): return sum(p.numel() for p in model.parameters()) if __name__ == "__main__": tok = CharTokenizer(SPECIALS + [chr(c) for c in range(ord("a"), ord("z") + 1)] + [chr(c) for c in range(0x0900, 0x0980)] + [" ", "'"]) m = TranslitModel(len(tok)) print(f"vocab={len(tok)} params={count_params(m)/1e6:.2f}M")