modeling.py
4.8 KB · 109 lines · python Raw
1 """Custom Sanskrit->English Transformer (trained from scratch).
2
3 Not a Hugging Face Transformers architecture, so load it with the helpers here:
4
5 from huggingface_hub import snapshot_download
6 import sys
7 d = snapshot_download("krpraveen/sanskrit-en-custom-transformer")
8 sys.path.insert(0, d)
9 from modeling import load, translate
10 model, sp, cfg = load(d)
11 print(translate(model, sp, cfg, ["बाल: भवत्सु प्रेमं प्रकटयति ।"]))
12 """
13 import math, os, json
14 import torch
15 import torch.nn as nn
16
17
18 class PositionalEncoding(nn.Module):
19 def __init__(self, d_model, dropout=0.0, max_len=1024):
20 super().__init__()
21 self.drop = nn.Dropout(dropout)
22 pe = torch.zeros(max_len, d_model)
23 pos = torch.arange(0, max_len).unsqueeze(1).float()
24 div = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
25 pe[:, 0::2] = torch.sin(pos * div)
26 pe[:, 1::2] = torch.cos(pos * div)
27 self.register_buffer("pe", pe.unsqueeze(0))
28
29 def forward(self, x):
30 return self.drop(x + self.pe[:, :x.size(1)])
31
32
33 class Seq2SeqTransformer(nn.Module):
34 def __init__(self, vocab, d_model, nhead, layers, dim_ff, dropout, pad_id):
35 super().__init__()
36 self.pad_id, self.d_model = pad_id, d_model
37 self.embed = nn.Embedding(vocab, d_model, padding_idx=pad_id)
38 self.pos = PositionalEncoding(d_model, dropout)
39 self.transformer = nn.Transformer(
40 d_model=d_model, nhead=nhead, num_encoder_layers=layers,
41 num_decoder_layers=layers, dim_feedforward=dim_ff,
42 dropout=dropout, batch_first=True, norm_first=True)
43 self.out = nn.Linear(d_model, vocab, bias=False)
44 self.out.weight = self.embed.weight # tied input/output embeddings
45
46 def emb(self, x):
47 return self.pos(self.embed(x) * math.sqrt(self.d_model))
48
49 def forward(self, src, tgt_in):
50 src_kpm = (src == self.pad_id)
51 tgt_kpm = (tgt_in == self.pad_id)
52 tgt_mask = nn.Transformer.generate_square_subsequent_mask(tgt_in.size(1)).to(src.device)
53 mem = self.transformer.encoder(self.emb(src), src_key_padding_mask=src_kpm)
54 dec = self.transformer.decoder(self.emb(tgt_in), mem, tgt_mask=tgt_mask,
55 tgt_key_padding_mask=tgt_kpm, memory_key_padding_mask=src_kpm)
56 return self.out(dec)
57
58
59 def load(path_or_repo, device="cpu"):
60 """Load the model, SentencePiece processor, and config from a local dir or a Hub repo id."""
61 if os.path.isdir(path_or_repo):
62 d = path_or_repo
63 else:
64 from huggingface_hub import snapshot_download
65 d = snapshot_download(path_or_repo)
66 cfg = json.load(open(os.path.join(d, "config.json")))
67 import sentencepiece as spm
68 sp = spm.SentencePieceProcessor(model_file=os.path.join(d, "spm.model"))
69 model = Seq2SeqTransformer(cfg["vocab_size"], cfg["d_model"], cfg["nhead"],
70 cfg["num_layers"], cfg["dim_ff"], cfg["dropout"], cfg["pad"]).to(device)
71 sd = torch.load(os.path.join(d, "pytorch_model.bin"), map_location=device)
72 model.load_state_dict(sd)
73 model.eval()
74 return model, sp, cfg
75
76
77 @torch.no_grad()
78 def translate(model, sp, cfg, sentences, device=None, batch_size=64):
79 """Greedy-decode a list of Sanskrit sentences into English."""
80 device = device or next(model.parameters()).device
81 max_len = cfg["max_len"]
82 PAD, BOS, EOS = cfg["pad"], cfg["bos"], cfg["eos"]
83 sentences = [str(s) for s in sentences]
84 order = sorted(range(len(sentences)), key=lambda i: len(sentences[i]))
85 out = [None] * len(sentences)
86 for k in range(0, len(sentences), batch_size):
87 idx = order[k:k + batch_size]
88 src = [sp.encode(sentences[i], out_type=int)[:max_len - 2] + [EOS] for i in idx]
89 m = max(len(s) for s in src)
90 src_t = torch.tensor([s + [PAD] * (m - len(s)) for s in src], dtype=torch.long, device=device)
91 src_kpm = (src_t == PAD)
92 mem = model.transformer.encoder(model.emb(src_t), src_key_padding_mask=src_kpm)
93 ys = torch.full((len(idx), 1), BOS, dtype=torch.long, device=device)
94 done = torch.zeros(len(idx), dtype=torch.bool, device=device)
95 for _ in range(max_len - 1):
96 tm = nn.Transformer.generate_square_subsequent_mask(ys.size(1)).to(device)
97 dec = model.transformer.decoder(model.emb(ys), mem, tgt_mask=tm, memory_key_padding_mask=src_kpm)
98 nxt = model.out(dec[:, -1]).argmax(-1).masked_fill(done, PAD)
99 ys = torch.cat([ys, nxt.unsqueeze(1)], 1)
100 done = done | (nxt == EOS)
101 if done.all():
102 break
103 for j, i in enumerate(idx):
104 toks = ys[j, 1:].tolist()
105 if EOS in toks:
106 toks = toks[:toks.index(EOS)]
107 out[i] = sp.decode([t for t in toks if t not in (PAD, BOS)])
108 return out
109