Views
No views yet
best_model.pt: primary demo checkpointbest_val_model.pt: best validation-loss checkpoint (if available)quality_predictor.pt: Task 5 quality predictor (if available)project_config.json: serialized training/inference configsanskrit_src_tokenizer_v1000.jsonsanskrit_tgt_tokenizer_v2000.jsoninference.pymodels/diffusion/1git clone https://huggingface.co/bhsinghgrid/sanskrit-translation
2cd sanskrit-translation
3python inference.py --model best_model.pt --cli1from huggingface_hub import snapshot_download
2repo_dir = snapshot_download("bhsinghgrid/sanskrit-translation")
3print("Downloaded to:", repo_dir)1import os
2import torch
3from huggingface_hub import snapshot_download
4from config import CONFIG
5from inference import load_model, _build_tokenizers
6
7repo_dir = snapshot_download("bhsinghgrid/sanskrit-translation")
8device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
9
10cfg = CONFIG
11model, cfg = load_model(os.path.join(repo_dir, "best_model.pt"), cfg, device)
12src_tok, tgt_tok = _build_tokenizers(cfg)
13
14def translate(text: str):
15 x = torch.tensor([src_tok.encode(text)], dtype=torch.long, device=device)
16 y = model.generate(
17 x,
18 num_steps=cfg["inference"]["num_steps"],
19 temperature=cfg["inference"]["temperature"],
20 top_k=cfg["inference"]["top_k"],
21 repetition_penalty=cfg["inference"]["repetition_penalty"],
22 diversity_penalty=cfg["inference"]["diversity_penalty"],
23 )
24 ids = [i for i in y[0].tolist() if i > 4]
25 return tgt_tok.decode(ids).strip()
26
27print(translate("dharmo rakṣati rakṣitaḥ"))config.json + model.safetensors).AutoModelForCausalLM.from_pretrained(...)pipeline(...)inference.py / API wrapper)./Users/bhsingh/Documents/Final_Paraphrase/Modify/results8/d3pm_cross_attention_neg_False/best_model.pt