96 lines
2.6 KiB
Python
96 lines
2.6 KiB
Python
import logging
|
|
import torch
|
|
from typing import Optional
|
|
from pyannote.audio import Pipeline
|
|
from ..config import settings
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_diarization_pipeline: Optional[Pipeline] = None
|
|
|
|
|
|
def get_diarization_pipeline() -> Pipeline:
|
|
global _diarization_pipeline
|
|
if _diarization_pipeline is None:
|
|
import os
|
|
|
|
hf_token = settings.hf_token or os.environ.get("HUGGINGFACE_TOKEN")
|
|
logger.info("Loading diarization model on %s", settings.diarization_device)
|
|
_diarization_pipeline = Pipeline.from_pretrained(
|
|
settings.diarization_model,
|
|
token=hf_token,
|
|
)
|
|
_diarization_pipeline.to(torch.device(settings.diarization_device))
|
|
return _diarization_pipeline
|
|
|
|
|
|
def diarize(audio_path: str) -> list[dict]:
|
|
import torch
|
|
import soundfile as sf
|
|
|
|
pipeline = get_diarization_pipeline()
|
|
|
|
waveform, sample_rate = sf.read(audio_path, dtype="float32")
|
|
if waveform.ndim > 1:
|
|
waveform = waveform.mean(axis=1)
|
|
audio_input = {
|
|
"waveform": torch.from_numpy(waveform).unsqueeze(0),
|
|
"sample_rate": sample_rate,
|
|
}
|
|
diarization = pipeline(audio_input)
|
|
|
|
segments = []
|
|
for turn, _, speaker in diarization.itertracks(yield_label=True):
|
|
segments.append(
|
|
{
|
|
"speaker_label": speaker,
|
|
"start_time": turn.start,
|
|
"end_time": turn.end,
|
|
}
|
|
)
|
|
|
|
segments.sort(key=lambda s: s["start_time"])
|
|
segments = _merge_adjacent_same_speaker(segments)
|
|
segments = _relabel_speakers(segments)
|
|
|
|
return segments
|
|
|
|
|
|
def _merge_adjacent_same_speaker(
|
|
segments: list[dict], max_gap: float = 0.5
|
|
) -> list[dict]:
|
|
if not segments:
|
|
return []
|
|
|
|
merged = [segments[0]]
|
|
for seg in segments[1:]:
|
|
prev = merged[-1]
|
|
if (
|
|
prev["speaker_label"] == seg["speaker_label"]
|
|
and (seg["start_time"] - prev["end_time"]) <= max_gap
|
|
):
|
|
prev["end_time"] = seg["end_time"]
|
|
else:
|
|
merged.append(seg)
|
|
return merged
|
|
|
|
|
|
def _relabel_speakers(segments: list[dict]) -> list[dict]:
|
|
label_map = {}
|
|
counter = 0
|
|
for seg in segments:
|
|
spk = seg["speaker_label"]
|
|
if spk not in label_map:
|
|
label = chr(65 + counter)
|
|
label_map[spk] = f"Speaker {label}"
|
|
counter += 1
|
|
seg["speaker_label"] = label_map[spk]
|
|
return segments
|
|
|
|
|
|
def unload_model():
|
|
global _diarization_pipeline
|
|
_diarization_pipeline = None
|
|
if torch.cuda.is_available():
|
|
torch.cuda.empty_cache()
|