| 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 |
| 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) |
|
|