| 1 | """Word-level forced alignment (torchaudio MMS_FA, the modern wav2vec2-CTC |
| 2 | approach WhisperX popularized) — precise per-word start/end times, no token. |
| 3 | |
| 4 | Aligned per Whisper segment (small, fast) and offset back to absolute time. |
| 5 | """ |
| 6 | import re |
| 7 | import warnings |
| 8 | |
| 9 | warnings.filterwarnings("ignore") |
| 10 | |
| 11 | import numpy as np |
| 12 | import torch |
| 13 | from torchaudio.pipelines import MMS_FA as B |
| 14 | |
| 15 | _model = _tok = _aligner = None |
| 16 | |
| 17 | |
| 18 | def _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 | |
| 25 | def _norm(word): |
| 26 | return re.sub(r"[^a-z']", "", word.lower()) |
| 27 | |
| 28 | |
| 29 | def 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 |