added model c hange
This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user