niagara-19m-batch.en / tokenization.py
Daniel Rasmussen
Fix device loading logic
36da8c7
Raw
History Blame Contribute Delete
4.65 kB
"""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)