from __future__ import annotations import re import subprocess from dataclasses import dataclass from pathlib import Path import numpy as np import torch from transformers import AutoModelForCTC, AutoTokenizer SAMPLE_RATE = 16_000 MODEL_ID = "facebook/wav2vec2-base-960h" @dataclass class WordInterval: word: str start_s: float end_s: float quality: float valid: bool def decode_audio(video_path: Path) -> np.ndarray: result = subprocess.run( [ "ffmpeg", "-v", "error", "-i", str(video_path), "-map", "0:a:0", "-ac", "1", "-ar", str(SAMPLE_RATE), "-f", "f32le", "pipe:1", ], check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, ) waveform = np.frombuffer(result.stdout, dtype=" tuple[list[int], list[list[int]]]: vocab = tokenizer.get_vocab() delimiter = int(tokenizer.convert_tokens_to_ids(tokenizer.word_delimiter_token or "|")) unknown = int(tokenizer.unk_token_id) targets: list[int] = [] per_word: list[list[int]] = [[] for _ in words] for word_index, raw_word in enumerate(words): if word_index: targets.append(delimiter) normalized = re.sub(r"[^a-z']", "", raw_word.lower()) for character in normalized: per_word[word_index].append(len(targets)) targets.append(int(vocab.get(character, unknown))) return targets, per_word def _viterbi(log_probs: np.ndarray, targets: list[int], blank: int) -> np.ndarray | None: if not targets or log_probs.ndim != 2: return None states = np.full(2 * len(targets) + 1, blank, dtype=np.int64) states[1::2] = np.asarray(targets, dtype=np.int64) frames, count = log_probs.shape[0], len(states) if frames == 0 or frames < len(targets): return None previous = np.full(count, -np.inf, dtype=np.float64) previous[0] = float(log_probs[0, blank]) previous[1] = float(log_probs[0, states[1]]) back = np.zeros((frames, count), dtype=np.uint8) skip = np.zeros(count, dtype=bool) if count > 2: skip[2:] = (states[2:] != blank) & (states[2:] != states[:-2]) for frame in range(1, frames): stay = previous one = np.full(count, -np.inf, dtype=np.float64) one[1:] = previous[:-1] two = np.full(count, -np.inf, dtype=np.float64) if skip.any(): two[skip] = previous[np.flatnonzero(skip) - 2] candidates = np.stack((stay, one, two), axis=0) choice = candidates.argmax(axis=0).astype(np.uint8) previous = candidates[choice, np.arange(count)] + log_probs[frame, states] back[frame] = choice state = count - 1 if previous[-1] >= previous[-2] else count - 2 path = np.empty(frames, dtype=np.int32) path[-1] = state for frame in range(frames - 1, 0, -1): state -= int(back[frame, state]) path[frame - 1] = state return path def model_time_constants(model) -> tuple[float, float]: config = model.config stride = int(np.prod(config.conv_stride)) receptive = 1 jump = 1 for kernel, local_stride in zip(config.conv_kernel, config.conv_stride): receptive += (int(kernel) - 1) * jump jump *= int(local_stride) return stride / SAMPLE_RATE, receptive / (2 * SAMPLE_RATE) def align_words( waveform: np.ndarray, words: list[str], tokenizer, model, device: torch.device, ) -> list[WordInterval]: targets, word_targets = _ctc_targets(words, tokenizer) blank = int(tokenizer.pad_token_id) if not targets or not len(waveform): return [WordInterval(w, float("nan"), float("nan"), 0.0, False) for w in words] with torch.inference_mode(): values = torch.as_tensor(waveform, dtype=torch.float32, device=device).unsqueeze(0) logits = model(input_values=values).logits[0].float() log_probs = torch.log_softmax(logits, dim=-1).cpu().numpy() path = _viterbi(log_probs, targets, blank) frame_step, center_s = model_time_constants(model) duration = len(waveform) / SAMPLE_RATE intervals: list[WordInterval] = [] if path is None: return [WordInterval(w, float("nan"), float("nan"), 0.0, False) for w in words] for word, target_indices in zip(words, word_targets): states = np.asarray([2 * index + 1 for index in target_indices], dtype=np.int32) frame_indices = np.flatnonzero(np.isin(path, states)) if len(states) else np.empty(0, dtype=np.int64) if not len(frame_indices): intervals.append(WordInterval(word, float("nan"), float("nan"), 0.0, False)) continue first, last = int(frame_indices[0]), int(frame_indices[-1]) start = max(0.0, first * frame_step + center_s - frame_step / 2) end = min(duration, (last + 1) * frame_step + center_s - frame_step / 2) char_scores = [] for target_index in target_indices: selected = np.flatnonzero(path == 2 * target_index + 1) if len(selected): char_scores.extend(log_probs[selected, targets[target_index]].tolist()) quality = float(np.exp(np.mean(char_scores))) if char_scores else 0.0 valid = end > start intervals.append(WordInterval(word, start, end, quality, valid)) return intervals def load_ctc(device: torch.device): tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) model = AutoModelForCTC.from_pretrained(MODEL_ID).to(device).eval() return tokenizer, model