97 lines
3.3 KiB
Python
97 lines
3.3 KiB
Python
# Note: This feature requires pyannote.audio and a HuggingFace token.
|
|
# If these are not present, this module will likely fail or raise errors.
|
|
# Due to the complexity and weight of pyannote.audio, this is a placeholder
|
|
# for where the logic would sit. Implementing full diarization requires
|
|
# downloading models and handling complex segment merging.
|
|
|
|
import os
|
|
import sys
|
|
import tracker
|
|
import torch
|
|
|
|
def diarize_audio(audio_path, num_speakers=None, hf_token=None):
|
|
"""
|
|
Performs speaker diarization using pyannote.audio.
|
|
"""
|
|
try:
|
|
from pyannote.audio import Pipeline
|
|
except ImportError:
|
|
tracker.logger.error("Error: pyannote.audio not installed. Diarization skipped.")
|
|
return None
|
|
|
|
if not hf_token:
|
|
tracker.logger.error("Error: HuggingFace Token (HF_TOKEN) not found. Diarization skipped.")
|
|
return None
|
|
|
|
tracker.logger.info(f"Loading Diarization Pipeline (pyannote/speaker-diarization-3.1)...")
|
|
try:
|
|
# Fix for PyTorch 2.6+ weights_only issue
|
|
torch.serialization.add_safe_globals([torch.torch_version.TorchVersion])
|
|
|
|
# Note: 'use_auth_token' was deprecated in favor of 'token' in recent versions
|
|
pipeline = Pipeline.from_pretrained(
|
|
"pyannote/speaker-diarization-3.1",
|
|
token=hf_token
|
|
)
|
|
|
|
# Move to GPU if available
|
|
if torch.cuda.is_available():
|
|
pipeline.to(torch.device("cuda"))
|
|
|
|
tracker.logger.info(f"Diarizing {audio_path}...")
|
|
diarization = pipeline(audio_path, num_speakers=num_speakers)
|
|
|
|
results = []
|
|
for turn, _, speaker in diarization.itertracks(yield_label=True):
|
|
results.append({
|
|
"start": turn.start,
|
|
"end": turn.end,
|
|
"speaker": speaker
|
|
})
|
|
|
|
return results
|
|
|
|
except Exception as e:
|
|
tracker.logger.error(f"Error during diarization: {e}")
|
|
return None
|
|
|
|
def merge_diarization_with_transcript(transcript_segments, diarization_segments):
|
|
"""
|
|
Merges Whisper segments with Diarization speaker labels based on time overlap.
|
|
|
|
Args:
|
|
transcript_segments (list): Whisper segments [{'start': 0.0, 'end': 1.0, 'text': 'Hi'}, ...]
|
|
diarization_segments (list): Diarization segments [{'start': 0.1, 'end': 0.9, 'speaker': 'SPEAKER_00'}]
|
|
|
|
Returns:
|
|
list: Enhanced transcript segments with 'speaker' key.
|
|
"""
|
|
if not diarization_segments:
|
|
return transcript_segments
|
|
|
|
# Simple overlap matching logic
|
|
for t_seg in transcript_segments:
|
|
# Find diarization segment with max overlap
|
|
t_start = t_seg['start']
|
|
t_end = t_seg['end']
|
|
|
|
best_speaker = "Unknown"
|
|
max_overlap = 0
|
|
|
|
for d_seg in diarization_segments:
|
|
d_start = d_seg['start']
|
|
d_end = d_seg['end']
|
|
|
|
# Calculate intersection
|
|
overlap_start = max(t_start, d_start)
|
|
overlap_end = min(t_end, d_end)
|
|
overlap_duration = max(0, overlap_end - overlap_start)
|
|
|
|
if overlap_duration > max_overlap:
|
|
max_overlap = overlap_duration
|
|
best_speaker = d_seg['speaker']
|
|
|
|
t_seg['speaker'] = best_speaker
|
|
|
|
return transcript_segments
|