File size: 2,805 Bytes
ce3c8df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
"""
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,
    )