Views
No views yet
| Model | Description |
|---|---|
best.pt | Base model trained from scratch on ~7.4M En–Ml sentence pairs |
finetunedcorrected.pt | Base model fine-tuned on BPCC human-annotated data + curated corrections (recommended) |
d_model=512<2ml> (to Malayalam), <2en> (to English)pip install torch sentencepiece1import torch, sentencepiece as spm
2from model import MTModel
3from config import CFG
4
5DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
6sp = spm.SentencePieceProcessor(model_file="checkpoints/spm.model")
7PAD, BOS, EOS = 0, 1, 2
8ML_TAG = sp.piece_to_id("<2ml>"); EN_TAG = sp.piece_to_id("<2en>")
9
10model = MTModel(sp.get_piece_size(), CFG.d_model, CFG.nhead, CFG.layers,
11 CFG.dim_ff, dropout=0.0, max_len=CFG.max_len).to(DEVICE)
12sd = torch.load("checkpoints/finetunedcorrected.pt", map_location=DEVICE)["model"]
13model.load_state_dict({k.replace("_orig_mod.", ""): v for k, v in sd.items()})
14model.eval()
15
16@torch.no_grad()
17def translate(text, to="ml"):
18 tag = ML_TAG if to == "ml" else EN_TAG
19 ids = [BOS, tag] + sp.encode(text, out_type=int)[:CFG.max_len-2] + [EOS]
20 src = torch.tensor([ids], device=DEVICE)
21 ys = torch.tensor([[BOS]], device=DEVICE)
22 for _ in range(128):
23 nxt = model(src, ys)[0, -1].argmax().item()
24 ys = torch.cat([ys, torch.tensor([[nxt]], device=DEVICE)], 1)
25 if nxt == EOS: break
26 out = [i for i in ys[0].tolist() if i not in (PAD, BOS, EOS, ML_TAG, EN_TAG)]
27 return sp.decode(out)
28
29print(translate("The weather is nice today.", "ml"))
30print(translate("എനിക്ക് വിശക്കുന്നു.", "en"))1pip install fastapi uvicorn
2uvicorn translate_server:app --host 0.0.0.0 --port 8081