indic-transliterate / tests /test_model.py
MeghanaKap's picture
Push transliteration pipeline code + prototype checkpoint + model card
f4386ce verified
Raw
History Blame Contribute Delete
2.59 kB
import torch
from transliteration.model.model import TransliterationConfig, TransliterationModel
from transliteration.model.tokenizer import CharTransliterationTokenizer
def _make_model_and_tokenizer():
tok = CharTransliterationTokenizer.build_from_corpus(
["<2hi> namaste hai", "नमस्ते है", "<2bn> ki আছে"]
)
config = TransliterationConfig(
vocab_size=tok.vocab_size,
d_model=32,
nhead=2,
num_encoder_layers=1,
num_decoder_layers=1,
dim_feedforward=64,
max_position_embeddings=64,
pad_token_id=tok.pad_token_id,
bos_token_id=tok.bos_token_id,
eos_token_id=tok.eos_token_id,
)
model = TransliterationModel(config)
return model, tok
def test_forward_returns_loss_when_labels_given():
model, tok = _make_model_and_tokenizer()
enc = tok(["<2hi> namaste"], return_tensors="pt", padding=True)
labels = tok(["नमस्ते"], return_tensors="pt", padding=True)["input_ids"]
out = model(input_ids=enc["input_ids"], attention_mask=enc["attention_mask"], labels=labels)
assert out.loss is not None
assert out.loss.item() > 0
assert out.logits.shape[0] == 1
def test_generate_produces_valid_token_ids():
model, tok = _make_model_and_tokenizer()
enc = tok(["<2hi> namaste"], return_tensors="pt", padding=True)
out_ids = model.generate(enc["input_ids"], enc["attention_mask"], max_new_tokens=10)
assert out_ids.shape[0] == 1
assert out_ids.shape[1] <= 11 # bos + up to 10 generated
assert (out_ids >= 0).all()
assert (out_ids < tok.vocab_size).all()
def test_generate_batched_stops_on_eos_for_all():
model, tok = _make_model_and_tokenizer()
enc = tok(["<2hi> namaste", "<2bn> ki"], return_tensors="pt", padding=True)
out_ids = model.generate(enc["input_ids"], enc["attention_mask"], max_new_tokens=20)
assert out_ids.shape[0] == 2
def test_save_and_load_model_roundtrip(tmp_path):
model, tok = _make_model_and_tokenizer()
model.save_pretrained(str(tmp_path))
loaded_config = TransliterationConfig.from_pretrained(str(tmp_path))
loaded_model = TransliterationModel.from_pretrained(str(tmp_path), config=loaded_config)
assert loaded_config.vocab_size == model.config.vocab_size
enc = tok(["<2hi> namaste"], return_tensors="pt", padding=True)
out1 = model.generate(enc["input_ids"], enc["attention_mask"], max_new_tokens=5)
out2 = loaded_model.generate(enc["input_ids"], enc["attention_mask"], max_new_tokens=5)
assert torch.equal(out1, out2)