asr / src /tokenizer.py
shubhexists's picture
Add Zipformer-inspired ASR model: weights, tokenizer, config, and training code
ce3c8df verified
Raw History Blame Contribute Delete
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"
)
@property
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,
)