added model c hange

This commit is contained in:
2026-01-12 16:01:22 -05:00
parent d180d43e18
commit 9befd9c28d
5 changed files with 40 additions and 16 deletions
@@ -6,31 +6,28 @@
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.
Args:
audio_path (str): Path to audio file.
num_speakers (int, optional): Number of speakers if known.
hf_token (str): HuggingFace Auth Token.
Returns:
list: List of segments with speaker labels [(start, end, speaker), ...].
"""
try:
from pyannote.audio import Pipeline
except ImportError:
print("Error: pyannote.audio not installed. Diarization skipped.")
tracker.logger.error("Error: pyannote.audio not installed. Diarization skipped.")
return None
if not hf_token:
print("Error: HuggingFace Token (HF_TOKEN) not found. Diarization skipped.")
tracker.logger.error("Error: HuggingFace Token (HF_TOKEN) not found. Diarization skipped.")
return None
print(f"Loading Diarization Pipeline (pyannote/speaker-diarization-3.1)...")
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",
@@ -38,11 +35,10 @@ def diarize_audio(audio_path, num_speakers=None, hf_token=None):
)
# Move to GPU if available
import torch
if torch.cuda.is_available():
pipeline.to(torch.device("cuda"))
print(f"Diarizing {audio_path}...")
tracker.logger.info(f"Diarizing {audio_path}...")
diarization = pipeline(audio_path, num_speakers=num_speakers)
results = []
@@ -56,7 +52,7 @@ def diarize_audio(audio_path, num_speakers=None, hf_token=None):
return results
except Exception as e:
print(f"Error during diarization: {e}")
tracker.logger.error(f"Error during diarization: {e}")
return None
def merge_diarization_with_transcript(transcript_segments, diarization_segments):