hindi-translit-v3 / model.py
IB-Emper's picture
Upload model.py with huggingface_hub
cad2518 verified
Raw
History Blame Contribute Delete
5.8 kB
#!/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 = ["<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):
# 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")