"""Forced alignment — refine word timestamps against an acoustic model. Why this exists -------------- faster-whisper derives word times by cross-attention, which lands every word *start* systematically ~0.3-0.5s early (the word-end is fine). That bias flows straight into the voice timeline and makes zoom/cut land on the wrong frame — measured on real footage in ``Engine/docs/05_EXPERIENCIAS.md`` (#14). Phonetic forced alignment (wav2vec2, via whisperx) re-anchors each word against the audio and brings that error down to ~30ms. Design ------ * The dependency (``whisperx``) is **optional** and imported lazily, exactly like the rest of this stack (librosa, faster-whisper). When it is missing, or any step fails, :meth:`ForcedAligner.align` returns the words unchanged, so transcription never breaks because alignment did. * The aligner is a single responsibility class: it knows how to turn a transcript into the shape whisperx wants, call it, and write the refined times back. ``transcribe.py`` owns the decision of *whether* to align. * Align models are cached per language on the instance so repeated calls (e.g. many short clips) don't reload the wav2vec2 weights each time. """ import logging from typing import List, Optional, Sequence logger = logging.getLogger(__name__) class ForcedAligner: """Refine word-level timestamps with whisperx phonetic forced alignment. Usage:: aligner = ForcedAligner() words = aligner.align(words, raw_segments, media_path, language, models_dir) ``words`` and ``raw_segments`` come straight from :func:`transcribe` — ``raw_segments`` carries the per-segment ``words`` lists (the same dict objects as in ``words``) so the aligner knows which words belong to which audio window. Returns a list of the *same* word dicts, with ``start``/``end`` overwritten in place where alignment produced a usable time. """ def __init__(self, device: Optional[str] = None): self._device = device self._models: dict = {} # -- capability ------------------------------------------------------ @staticmethod def available() -> bool: """Whether whisperx can be imported (the aligner can run at all).""" try: import whisperx # noqa: F401 except Exception: return False return True def _resolve_device(self) -> str: if self._device: return self._device try: import torch if torch.cuda.is_available(): return "cuda" except Exception: pass return "cpu" # -- public API ------------------------------------------------------ def align( self, words: Sequence[dict], raw_segments: Sequence[dict], audio_path: str, language: str, models_dir: Optional[str] = None, ) -> List[dict]: """Return ``words`` with forced-aligned timestamps where possible. Falls back to the unchanged ``words`` on any failure (missing dependency, model load error, audio read error, or a result that doesn't line up with the input). """ if not words or not language: return list(words) try: import whisperx except Exception: logger.info("whisperx not installed; skipping forced alignment") return list(words) try: device = self._resolve_device() align_input = self._build_align_input(words, raw_segments) audio = whisperx.load_audio(audio_path) if language not in self._models: align_model, metadata = whisperx.load_align_model( language_code=language, device=device, model_dir=str(models_dir) if models_dir else None, ) self._models[language] = (align_model, metadata) align_model, metadata = self._models[language] result = whisperx.align( align_input, align_model, metadata, audio, device, return_char_alignments=False, ) return self._merge_result(words, result.get("segments", [])) except Exception: logger.warning( "forced alignment failed for %s; using raw timestamps", audio_path ) return list(words) # -- internals ------------------------------------------------------- @staticmethod def _build_align_input( words: Sequence[dict], raw_segments: Sequence[dict] ) -> List[dict]: """Transcript in whisperx's expected shape: segments -> words. whisperx.align requires each segment to carry ``text``/``start``/``end`` and a ``words`` list whose entries have ``word``/``start``/``end``/``score``. We only read ``words`` from ``raw_segments`` (the flattened ``words`` list is the source of truth for counts), so the two stay consistent. """ align_segments: List[dict] = [] for seg in raw_segments: seg_words = [ { "word": w.get("word", ""), "start": float(w.get("start", 0.0)), "end": float(w.get("end", 0.0)), "score": float(w.get("confidence", 0.0)), } for w in seg.get("words", []) ] align_segments.append( { "text": (seg.get("text") or "").strip(), "start": float(seg.get("start", 0.0)), "end": float(seg.get("end", 0.0)), "words": seg_words, } ) return align_segments @staticmethod def _merge_result(words: Sequence[dict], aligned_segments: Sequence[dict]) -> List[dict]: """Walk the aligned output in order and overwrite word times in place. whisperx preserves word order within and across segments, so a single running index over the output words lines up with ``words``. A word the aligner failed to place gets ``None``/``0`` times — we skip those rather than clobber a good timestamp, and if counts ever diverge we stop and leave the rest untouched. """ out = list(words) wi = 0 for seg in aligned_segments: for aw in seg.get("words", []): if wi >= len(out): return out start = aw.get("start") end = aw.get("end") if start is None or end is None or end < start: wi += 1 continue out[wi]["start"] = float(start) out[wi]["end"] = float(end) wi += 1 return out