| from transliteration.model.tokenizer import CharTransliterationTokenizer |
|
|
|
|
| def test_build_from_corpus_includes_language_tags(): |
| tok = CharTransliterationTokenizer.build_from_corpus(["<2hi> namaste", "नमस्ते"]) |
| assert "<2hi>" in tok.get_vocab() |
| assert "<2bn>" in tok.get_vocab() |
|
|
|
|
| def test_tokenize_treats_language_tag_as_single_token(): |
| tok = CharTransliterationTokenizer.build_from_corpus(["<2hi> namaste", "नमस्ते"]) |
| tokens = tok.tokenize("<2hi> namaste") |
| assert tokens[0] == "<2hi>" |
| assert tokens[1] == " " |
|
|
|
|
| def test_roundtrip_encode_decode(): |
| tok = CharTransliterationTokenizer.build_from_corpus(["<2hi> namaste hai", "नमस्ते है"]) |
| text = "<2hi> namaste hai" |
| ids = tok(text)["input_ids"] |
| decoded = tok.decode(ids, skip_special_tokens=True) |
| assert decoded == text |
|
|
|
|
| def test_unknown_char_maps_to_unk(): |
| tok = CharTransliterationTokenizer.build_from_corpus(["abc"]) |
| ids = tok("xyz123")["input_ids"] |
| |
| unk_id = tok.unk_token_id |
| assert any(i == unk_id for i in ids) |
|
|
|
|
| def test_save_and_load_roundtrip(tmp_path): |
| tok = CharTransliterationTokenizer.build_from_corpus(["<2hi> namaste", "नमस्ते"]) |
| tok.save_pretrained(str(tmp_path)) |
| loaded = CharTransliterationTokenizer.from_pretrained(str(tmp_path)) |
| assert loaded.get_vocab() == tok.get_vocab() |
|
|