File size: 4,647 Bytes
10e6832
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
36da8c7
10e6832
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
"""ASR tokenizer implementation using SentencePiece."""

import os

import sentencepiece as sp
import torch
from transformers import AutoConfig, PreTrainedTokenizer
from transformers.utils import cached_file


class Tokenizer(PreTrainedTokenizer):
    """Minimal SentencePiece tokenizer wrapper for ASR model."""

    def __init__(self, vocab_file=None, **kwargs):
        self.vocab_file = vocab_file
        self.sp_model = sp.SentencePieceProcessor()
        if vocab_file:
            self.sp_model.Load(vocab_file)
        super().__init__(**kwargs)

    @property
    def vocab_size(self):
        """Return vocabulary size."""
        return len(self.sp_model)

    def get_vocab(self):
        """Return the vocabulary as a dictionary."""
        if len(self.sp_model) == 0:
            return {}
        return {self.sp_model.IdToPiece(i): i for i in range(len(self.sp_model))}

    def decode(self, token_ids):
        """Decode token ids to text.

        Supports batch decoding (list of lists).
        """
        return self.sp_model.Decode(token_ids)

    def decode_from_logits(self, logits, mask=None):
        """Decode CTC logits to text.

        Parameters
        ----------
        logits : torch.Tensor
            Model logits of shape (batch_size, time_steps, vocab_size).
        mask : torch.Tensor, optional
            Attention mask of shape (batch_size, time_steps).
            If None, all logits are assumed to be unmasked.

        Returns
        -------
        list of str
            Decoded text strings.
        """
        batch_size, max_length = logits.shape[:2]
        device = logits.device

        # Compute lengths from mask
        if mask is None:
            # All logits are unmasked - use full length
            lengths = torch.full(
                (batch_size,), max_length, dtype=torch.long, device=device
            )
        else:
            # Ensure mask is on same device as logits
            mask = mask.to(device)
            lengths = mask.sum(dim=1).long()

        # Greedy CTC decode: take argmax over vocab dimension
        predictions = logits.argmax(dim=-1)

        # Create sequence length mask (vectorized)
        seqlen_mask = (
            torch.arange(max_length, device=device)[None, :] >= lengths[:, None]
        )

        # Apply length mask by setting out-of-bounds positions to blank token
        predictions = predictions.masked_fill(seqlen_mask, self.vocab_size)

        # CTC collapse: remove consecutive duplicates (vectorized)
        # Compute where tokens differ from previous token
        repeat_mask = torch.cat(
            [
                torch.zeros((batch_size, 1), dtype=torch.bool, device=device),
                predictions[:, 1:] == predictions[:, :-1],
            ],
            dim=1,
        )

        # Set repeated tokens to blank
        predictions = predictions.masked_fill(repeat_mask, self.vocab_size)

        # Create mask for valid tokens (not blank and > 0)
        valid_mask = (predictions != self.vocab_size) & (predictions > 0)

        # Use argsort trick to pack valid tokens to the left
        # Sort by (not valid, position) to move valid tokens to front
        sort_keys = (~valid_mask).long() * max_length + torch.arange(
            max_length, device=device
        )[None, :]
        sort_indices = torch.argsort(sort_keys, dim=1)
        packed_predictions = torch.gather(predictions, 1, sort_indices)
        packed_valid = torch.gather(valid_mask, 1, sort_indices)

        # Count valid tokens per sequence
        valid_lengths = packed_valid.sum(dim=1)

        # Move to CPU only at the end for conversion to lists
        packed_predictions = packed_predictions.cpu()
        valid_lengths = valid_lengths.cpu()

        # Convert to list of lists (minimal loop, just slicing)
        decoded_seqs = [
            packed_predictions[i, : valid_lengths[i]].tolist()
            for i in range(batch_size)
        ]

        # Decode all sequences to text
        return self.decode(decoded_seqs)

    @classmethod
    def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
        """Load tokenizer from pretrained model."""

        config = AutoConfig.from_pretrained(pretrained_model_name_or_path, **kwargs)
        tokenizer_file = config.tokenizer_file

        if os.path.isdir(pretrained_model_name_or_path):
            vocab_file = os.path.join(pretrained_model_name_or_path, tokenizer_file)
        else:
            vocab_file = cached_file(
                pretrained_model_name_or_path, tokenizer_file, **kwargs
            )

        return cls(vocab_file=vocab_file)