70 lines
2.4 KiB
Python
70 lines
2.4 KiB
Python
from pathlib import Path
|
|
from typing import List, Dict
|
|
|
|
def diarize_with_resemblyzer(wav_path: Path, window_s: float = 1.5, hop_s: float = 0.75, distance_threshold: float = 0.6) -> List[Dict]:
|
|
"""Lightweight embedding-based diarization using resemblyzer + sklearn.
|
|
|
|
Returns list of segments: {'start': float, 'end': float, 'speaker': 'spk_N'}
|
|
"""
|
|
try:
|
|
import numpy as np
|
|
import librosa
|
|
from resemblyzer import VoiceEncoder
|
|
from sklearn.cluster import AgglomerativeClustering
|
|
except Exception as e:
|
|
raise RuntimeError(f"Missing dependency for embedding diarization: {e}")
|
|
|
|
sr = 16000
|
|
wav, sr_loaded = librosa.load(str(wav_path), sr=sr)
|
|
n = wav.shape[0]
|
|
win = int(window_s * sr)
|
|
hop = int(hop_s * sr)
|
|
if n <= win:
|
|
# short file: embed whole
|
|
encoder = VoiceEncoder()
|
|
emb = encoder.embed_utterance(wav)
|
|
labels = [0]
|
|
centers = [ (0.0 + n / sr) / 2.0 ]
|
|
windows = [(0.0, n / sr)]
|
|
else:
|
|
encoder = VoiceEncoder()
|
|
embeddings = []
|
|
centers = []
|
|
windows = []
|
|
for start in range(0, n - win + 1, hop):
|
|
chunk = wav[start:start+win]
|
|
try:
|
|
e = encoder.embed_utterance(chunk)
|
|
except Exception:
|
|
# fallback: mean pooling
|
|
e = np.mean(chunk)
|
|
embeddings.append(e)
|
|
t0 = start / sr
|
|
t1 = (start + win) / sr
|
|
centers.append((t0 + t1) / 2.0)
|
|
windows.append((t0, t1))
|
|
|
|
X = np.vstack(embeddings)
|
|
# Agglomerative clustering with distance threshold
|
|
model = AgglomerativeClustering(n_clusters=None, distance_threshold=distance_threshold, affinity='euclidean', linkage='average')
|
|
labels = model.fit_predict(X)
|
|
|
|
# merge consecutive windows with same label into segments
|
|
segs = []
|
|
if len(labels) == 0:
|
|
return segs
|
|
cur_label = labels[0]
|
|
cur_start = windows[0][0]
|
|
cur_end = windows[0][1]
|
|
for i, lab in enumerate(labels[1:], start=1):
|
|
if lab == cur_label:
|
|
cur_end = windows[i][1]
|
|
else:
|
|
segs.append({"start": cur_start, "end": cur_end, "speaker": f"spk_{int(cur_label)+1}"})
|
|
cur_label = lab
|
|
cur_start = windows[i][0]
|
|
cur_end = windows[i][1]
|
|
segs.append({"start": cur_start, "end": cur_end, "speaker": f"spk_{int(cur_label)+1}"})
|
|
|
|
return segs
|