| |
| """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 = ["<pad>", "<bos>", "<eos>", "<unk>"] |
|
|
| 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): |
| |
| 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)) |
|
|
| def forward(self, x): |
| 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 |
| self.scale = math.sqrt(d) |
|
|
| def forward(self, src, tgt_in): |
| |
| 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) |
|
|
| 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") |
|
|