151 lines
5.5 KiB
Python
151 lines
5.5 KiB
Python
"""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)]
|