1"""Word-level forced alignment (torchaudio MMS_FA, the modern wav2vec2-CTC
2approach WhisperX popularized) — precise per-word start/end times, no token.
3
4Aligned per Whisper segment (small, fast) and offset back to absolute time.
5"""
6import re
7import warnings
8
9warnings.filterwarnings("ignore")
10
11import numpy as np
12import torch
13from torchaudio.pipelines import MMS_FA as B
14
15_model = _tok = _aligner = None
16
17
18def _load():
19 global _model, _tok, _aligner
20 if _model is None:
21 _model, _tok, _aligner = B.get_model(), B.get_tokenizer(), B.get_aligner()
22 return _model, _tok, _aligner
23
24
25def _norm(word):
26 return re.sub(r"[^a-z']", "", word.lower())
27
28
29def align_segment(data, sr, start, end, text):
30 """Return [{word, start, end}] for the words in this segment, in absolute
31 seconds. Words that can't be aligned (pure numbers/symbols) get times
32 interpolated from their neighbours."""
33 a, b = int(start * sr), int(end * sr)
34 seg = data[a:b]
35 raw = text.split()
36 if len(seg) < int(0.2 * sr) or not raw:
37 return [{"word": w, "start": start, "end": end} for w in raw]
38
39 norm = [_norm(w) for w in raw]
40 idx = [i for i, n in enumerate(norm) if n]
41 if not idx:
42 return [{"word": w, "start": start, "end": end} for w in raw]
43
44 model, tok, aligner = _load()
45 wav = torch.from_numpy(np.ascontiguousarray(seg)).unsqueeze(0)
46 with torch.inference_mode():
47 emit, _ = model(wav)
48 try:
49 spans = aligner(emit[0], tok([norm[i] for i in idx]))
50 except Exception:
51 return [{"word": w, "start": start, "end": end} for w in raw]
52
53 ratio = wav.shape[1] / emit.shape[1] / sr
54 times = {}
55 for k, i in enumerate(idx):
56 s = spans[k]
57 times[i] = (round(start + s[0].start * ratio, 3), round(start + s[-1].end * ratio, 3))
58
59 out = [{"word": w, "start": None, "end": None} for w in raw]
60 for i, t in times.items():
61 out[i]["start"], out[i]["end"] = t
62 # Interpolate unaligned words from neighbours.
63 last_end = start
64 for i, o in enumerate(out):
65 if o["start"] is None:
66 o["start"] = last_end
67 nxt = next((out[j]["start"] for j in range(i + 1, len(out)) if out[j]["start"]), end)
68 o["end"] = nxt
69 last_end = o["end"]
70 return out