Files

181 lines
6.8 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 _load_waveform(path: str) -> Optional[dict]:
"""Decode ``path`` ourselves into the waveform dict pyannote accepts.
pyannote 4.x decodes audio through torchcodec, which links against a
specific FFmpeg major version and fails outright when the installed one
differs (``libavutil.56.dylib`` not found) — taking diarization down on
an otherwise working machine. Handing it an already-decoded waveform
skips that path entirely and reuses the ffmpeg extraction the acoustic
analysis already relies on, so video containers work too.
Returns ``None`` when decoding is not possible, letting the caller fall
back to passing the path and whatever pyannote can do with it.
"""
try:
import soundfile
import torch
from .voice_features import decodable_audio
with decodable_audio(path) as audio_path:
if audio_path is None:
return None
data, sample_rate = soundfile.read(audio_path, dtype="float32", always_2d=True)
# soundfile gives (samples, channels); pyannote wants (channels, samples)
return {"waveform": torch.from_numpy(data.T), "sample_rate": int(sample_rate)}
except Exception:
logger.info("could not pre-decode %s for diarization", path)
return None
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(_load_waveform(path) or 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)]