Download src/tokenizer.py from shubhexists/asr: direct link, hf CLI and curl.
- Browser
- Download file 2.81 kB
-
https://huggingface.co/shubhexists/asr/resolve/main/src/tokenizer.py
- Command line
-
hf download hf://shubhexists/asr/src/tokenizer.py
-
curl -L -o tokenizer.py https://huggingface.co/shubhexists/asr/resolve/main/src/tokenizer.py
2.81 kB
| """ | |
| SentencePiece BPE tokenizer wrapper. | |
| Fixed special-token id scheme, shared by model, dataset, and training code: | |
| 0 = <pad> (also the CTC blank id -- neither is ever a real token) | |
| 1 = <unk> | |
| 2 = <bos> | |
| 3 = <eos> | |
| 4 = <spk_change> (reserved for phase 2 Serialized Output Training) | |
| 5.. = regular BPE pieces | |
| """ | |
| from typing import Any, cast | |
| import sentencepiece as spm | |
| PAD_ID = 0 | |
| UNK_ID = 1 | |
| BOS_ID = 2 | |
| EOS_ID = 3 | |
| SPK_CHANGE_ID = 4 | |
| SPK_CHANGE_TOKEN = "<spk_change>" | |
| NUM_RESERVED = 5 # ids 0..4 above; regular vocab starts at 5 | |
| class ASRTokenizer: | |
| def __init__(self, model_path: str): | |
| # SentencePieceProcessor is a SWIG extension with no type stubs, hence Any. | |
| self.sp: Any = spm.SentencePieceProcessor() | |
| self.sp.load(model_path) | |
| assert self.sp.pad_id() == PAD_ID | |
| assert self.sp.unk_id() == UNK_ID | |
| assert self.sp.bos_id() == BOS_ID | |
| assert self.sp.eos_id() == EOS_ID | |
| assert self.sp.piece_to_id(SPK_CHANGE_TOKEN) == SPK_CHANGE_ID, ( | |
| f"expected {SPK_CHANGE_TOKEN} at id {SPK_CHANGE_ID}, " | |
| f"got {self.sp.piece_to_id(SPK_CHANGE_TOKEN)} -- tokenizer was not " | |
| f"trained with prepare_tokenizer.py's expected special-token layout" | |
| ) | |
| def vocab_size(self) -> int: | |
| return self.sp.get_piece_size() | |
| def encode(self, text: str) -> list: | |
| """Text -> list of BPE token ids (no bos/eos). Lowercases internally: | |
| the vocab was trained lowercased but LibriSpeech transcripts are ALL | |
| CAPS, and uppercase input would otherwise become <unk>.""" | |
| return self.sp.encode(text.lower(), out_type=int) | |
| def decode(self, ids: list) -> str: | |
| """Token ids -> text. Silently drops pad/bos/eos/spk_change ids.""" | |
| ids = [i for i in ids if i >= NUM_RESERVED] | |
| return self.sp.decode(ids) | |
| def decoder_input_and_target(self, text: str): | |
| """Returns (input_ids, target_ids) for teacher forcing: | |
| input = [BOS, t1, t2, ..., tn] | |
| target = [t1, t2, ..., tn, EOS] | |
| """ | |
| ids = self.encode(text) | |
| input_ids = [BOS_ID] + ids | |
| target_ids = ids + [EOS_ID] | |
| return input_ids, target_ids | |
| def train_tokenizer(corpus_path: str, model_prefix: str, vocab_size: int = 5000): | |
| # Same stub-less SWIG issue as SentencePieceProcessor above. | |
| cast(Any, spm.SentencePieceTrainer).train( | |
| input=corpus_path, | |
| model_prefix=model_prefix, | |
| vocab_size=vocab_size, | |
| model_type="bpe", | |
| pad_id=PAD_ID, | |
| unk_id=UNK_ID, | |
| bos_id=BOS_ID, | |
| eos_id=EOS_ID, | |
| user_defined_symbols=[SPK_CHANGE_TOKEN], | |
| character_coverage=1.0, | |
| input_sentence_size=5_000_000, | |
| shuffle_input_sentence=True, | |
| ) | |