"""Speaker diarization for local transcripts — WHISPERX-inspired, lazy + graceful. pyannote.audio (with an HF token granted access to the gated diarization models) is optional: when it is absent or the token is missing, ``diarize`` returns ``None`` and callers fall back to ``SPEAKER_00`` for every word — never crashing. The speaker-assignment helpers are pure functions over ``(start, end)`` tracks, so they are fully testable without any model. This mirrors the reference WHISPERX app (``app._assign_speakers_to_words``), adapted to our flat ``segments`` + ``words`` transcript shape. """ import logging from typing import List, Optional, Sequence, Tuple logger = logging.getLogger(__name__) DEFAULT_SPEAKER = "SPEAKER_00" # Raw speaker id -> canonical ``SPEAKER_NN`` label. _DEFAULT_LABELS = ("SPEAKER_00", "SPEAKER_01", "SPEAKER_02", "SPEAKER_03") def diarization_capability(token: Optional[str]) -> Tuple[bool, str]: """Whether speaker identification is available. Returns ``(ok, message)``. ``ok`` is only ``True`` when pyannote.audio is installed *and* a token is configured (both are needed for the gated models). """ try: import pyannote.audio # noqa: F401 except Exception: return False, "Diarização indisponível: componente pyannote.audio ausente." if not (token or "").strip(): return False, "Diarização indisponível: nenhum token HuggingFace configurado." return True, "Identificação de participantes disponível." def diarize( path: str, token: Optional[str], num_speakers: str = "", ) -> Optional[List[Tuple[float, float, str]]]: """Run speaker diarization on ``path``. Returns a list of ``(start, end, raw_speaker_id)`` turns, or ``None`` when pyannote is unavailable, the token is missing/invalid, or analysis fails — the graceful-degradation contract shared with ``transcribe``. """ if not diarization_capability(token)[0]: return None try: from pyannote.audio import Pipeline pipe = None try: pipe = Pipeline.from_pretrained("pyannote/speaker-diarization-3.1", token=token) except TypeError: # pyannote.audio < 4.0 used use_auth_token instead of token. pipe = Pipeline.from_pretrained( "pyannote/speaker-diarization-3.1", use_auth_token=token ) if pipe is None: return None kwargs = {} n = str(num_speakers or "").strip() if n.isdigit() and int(n) > 0: kwargs["num_speakers"] = int(n) result = pipe(path, **kwargs) # pyannote.audio >= 4.0 wraps the annotation; normalize to the raw one. if hasattr(result, "exclusive_speaker_diarization"): result = result.exclusive_speaker_diarization elif hasattr(result, "speaker_diarization"): result = result.speaker_diarization tracks: List[Tuple[float, float, str]] = [] for turn, _, speaker in result.itertracks(yield_label=True): tracks.append((float(turn.start), float(turn.end), str(speaker))) return tracks or None except Exception: logger.warning("diarization failed for %s", path) return None def _overlap_speaker( start: float, end: float, tracks: Sequence[Tuple[float, float, str]] ) -> Optional[str]: """Speaker with the largest summed time-overlap with ``[start, end]``.""" overlap: dict[str, float] = {} for t_start, t_end, sid in tracks: o = min(end, t_end) - max(start, t_start) if o > 0: overlap[sid] = overlap.get(sid, 0.0) + o if not overlap: return None return max(overlap.items(), key=lambda kv: kv[1])[0] def assign_speakers( segments: Sequence[dict], words: Sequence[dict], tracks: Optional[Sequence[Tuple[float, float, str]]], default: str = DEFAULT_SPEAKER, ) -> Tuple[List[dict], List[dict]]: """Attach ``speaker_id`` to every segment and word. ``tracks`` is the output of :func:`diarize` (may be ``None``). Raw speaker ids are mapped to stable ``SPEAKER_NN`` labels in first-seen order. Without tracks everything is assigned ``default``. """ norm_tracks: List[Tuple[float, float, str]] = [] if tracks: label_map: dict[str, str] = {} counter = 0 for t_start, t_end, raw in tracks: if raw not in label_map: label = _DEFAULT_LABELS[counter] if counter < len(_DEFAULT_LABELS) else f"SPEAKER_{counter:02d}" label_map[raw] = label counter += 1 norm_tracks.append((t_start, t_end, label_map[raw])) out_segments: List[dict] = [] for seg in segments: s = dict(seg) s["speaker_id"] = ( _overlap_speaker(seg.get("start", 0.0), seg.get("end", 0.0), norm_tracks) or default ) out_segments.append(s) out_words: List[dict] = [] for w in words: ww = dict(w) ww["speaker_id"] = ( _overlap_speaker(w.get("start", 0.0), w.get("end", 0.0), norm_tracks) or default ) out_words.append(ww) return out_segments, out_words def build_speakers(segments: Sequence[dict]) -> List[dict]: """Ordered ``[{"id", "name"}, ...]`` from the speakers present in segments.""" ids: List[str] = [] seen = set() for seg in segments: sid = seg.get("speaker_id", DEFAULT_SPEAKER) if sid not in seen: seen.add(sid) ids.append(sid) return [{"id": sid, "name": f"Speaker {i + 1}"} for i, sid in enumerate(ids)]