HY-2012's picture
Upload folder using huggingface_hub
906ada2 verified
Raw
History Blame Contribute Delete
9.39 kB
# Copyright 2026 AXERA-TECH (authors: Magnetar)
#
# Speaker clustering over CAMPPlus embeddings.
#
# Mirrors python/utils/cluster_utils.py + do_clustering()/compressed_seg() of
# python/utils/ax_cam_bin.py (3D-Speaker-MT.axera). CPU post-processing only;
# requires: scipy scikit-learn fastcluster umap-learn hdbscan
import numpy as np
try:
import scipy
from sklearn.cluster._kmeans import k_means
from sklearn.metrics.pairwise import cosine_similarity
import fastcluster
from scipy.cluster.hierarchy import fcluster
from scipy.spatial.distance import squareform
except ImportError as e:
raise ImportError(
"clustering requires: pip install scipy scikit-learn fastcluster"
) from e
try:
import umap
import hdbscan
except ImportError:
raise ImportError(
'Package "umap" or "hdbscan" not found. '
'Please install them first by "pip install umap-learn hdbscan".')
class SpectralCluster:
"""Spectral clustering with unnormalized Laplacian (speechbrain-style)."""
def __init__(self, min_num_spks=1, max_num_spks=10, pval=0.02,
min_pnum=6, oracle_num=None):
self.min_num_spks = min_num_spks
self.max_num_spks = max_num_spks
self.min_pnum = min_pnum
self.pval = pval
self.k = oracle_num
def __call__(self, X, **kwargs):
pval = kwargs.get('pval', None)
oracle_num = kwargs.get('speaker_num', None)
sim_mat = self.get_sim_mat(X)
prunned_sim_mat = self.p_pruning(sim_mat, pval)
sym_prund_sim_mat = 0.5 * (prunned_sim_mat + prunned_sim_mat.T)
laplacian = self.get_laplacian(sym_prund_sim_mat)
emb, num_of_spk = self.get_spec_embs(laplacian, oracle_num)
labels = self.cluster_embs(emb, num_of_spk)
return labels
def get_sim_mat(self, X):
return cosine_similarity(X, X)
def p_pruning(self, A, pval=None):
if pval is None:
pval = self.pval
n_elems = int((1 - pval) * A.shape[0])
n_elems = min(n_elems, A.shape[0] - self.min_pnum)
for i in range(A.shape[0]):
low_indexes = np.argsort(A[i, :])
low_indexes = low_indexes[0:n_elems]
A[i, low_indexes] = 0
return A
def get_laplacian(self, M):
M[np.diag_indices(M.shape[0])] = 0
D = np.sum(np.abs(M), axis=1)
D = np.diag(D)
return D - M
def get_spec_embs(self, L, k_oracle=None):
if k_oracle is None:
k_oracle = self.k
lambdas, eig_vecs = scipy.sparse.linalg.eigsh(
L, k=min(self.max_num_spks + 1, L.shape[0]), which='SM')
if k_oracle is not None:
num_of_spk = k_oracle
else:
lambda_gap_list = self.getEigenGaps(
lambdas[self.min_num_spks - 1:self.max_num_spks + 1])
num_of_spk = np.argmax(lambda_gap_list) + self.min_num_spks
emb = eig_vecs[:, :num_of_spk]
return emb, num_of_spk
def cluster_embs(self, emb, k):
_, labels, _ = k_means(emb, k)
return labels
def getEigenGaps(self, eig_vals):
eig_vals_gap_list = []
for i in range(len(eig_vals) - 1):
gap = float(eig_vals[i + 1]) - float(eig_vals[i])
eig_vals_gap_list.append(gap)
return eig_vals_gap_list
class UmapHdbscan:
def __init__(self, n_neighbors=20, n_components=60, min_samples=20,
min_cluster_size=10, metric='euclidean'):
self.n_neighbors = n_neighbors
self.n_components = n_components
self.min_samples = min_samples
self.min_cluster_size = min_cluster_size
self.metric = metric
def __call__(self, X, **kwargs):
umap_X = umap.UMAP(
n_neighbors=self.n_neighbors, min_dist=0.0,
n_components=min(self.n_components, X.shape[0] - 2),
metric=self.metric,
).fit_transform(X)
labels = hdbscan.HDBSCAN(
min_samples=self.min_samples,
min_cluster_size=self.min_cluster_size).fit_predict(umap_X)
return labels
class AHCluster:
"""Agglomerative hierarchical clustering (VBx-style)."""
def __init__(self, fix_cos_thr=0.4):
self.fix_cos_thr = fix_cos_thr
def __call__(self, X, **kwargs):
scr_mx = cosine_similarity(X)
scr_mx = squareform(-scr_mx, checks=False)
lin_mat = fastcluster.linkage(scr_mx, method='average',
preserve_input='False')
adjust = abs(lin_mat[:, 2].min())
lin_mat[:, 2] += adjust
labels = fcluster(lin_mat, -self.fix_cos_thr + adjust,
criterion='distance') - 1
return labels
class CommonClustering:
"""Performs clustering over embeddings and returns labels.
Mirrors ax_cam_bin.py do_clustering() defaults:
cluster_type='spectral', mer_cos=0.8, min_cluster_size=4, pval=0.012
"""
def __init__(self, cluster_type, cluster_line=40, mer_cos=None,
min_cluster_size=4, **kwargs):
self.cluster_type = cluster_type
self.cluster_line = cluster_line
self.min_cluster_size = min_cluster_size
self.mer_cos = mer_cos
if self.cluster_type == 'spectral':
self.cluster = SpectralCluster(**kwargs)
elif self.cluster_type == 'umap_hdbscan':
kwargs['min_cluster_size'] = min_cluster_size
self.cluster = UmapHdbscan(**kwargs)
elif self.cluster_type == 'AHC':
self.cluster = AHCluster(**kwargs)
else:
raise ValueError('%s is not currently supported.' % cluster_type)
if self.cluster_type != 'AHC':
self.cluster_for_short = AHCluster()
else:
self.cluster_for_short = self.cluster
def __call__(self, X, **kwargs):
assert len(X.shape) == 2, 'Shape of input should be [N, C]'
if X.shape[0] <= 1:
return np.zeros(X.shape[0], dtype=int)
if X.shape[0] < self.cluster_line:
labels = self.cluster_for_short(X)
else:
labels = self.cluster(X, **kwargs)
labels = self.filter_minor_cluster(labels, X, self.min_cluster_size)
if self.mer_cos is not None:
labels = self.merge_by_cos(labels, X, self.mer_cos)
return labels
def filter_minor_cluster(self, labels, x, min_cluster_size):
cset = np.unique(labels)
csize = np.array([(labels == i).sum() for i in cset])
minor_idx = np.where(csize <= self.min_cluster_size)[0]
if len(minor_idx) == 0:
return labels
minor_cset = cset[minor_idx]
major_idx = np.where(csize > self.min_cluster_size)[0]
if len(major_idx) == 0:
return np.zeros_like(labels)
major_cset = cset[major_idx]
major_center = np.stack([x[labels == i].mean(0) for i in major_cset])
for i in range(len(labels)):
if labels[i] in minor_cset:
cos_sim = cosine_similarity(x[i][np.newaxis], major_center)
labels[i] = major_cset[cos_sim.argmax()]
return labels
def merge_by_cos(self, labels, x, cos_thr):
assert cos_thr > 0 and cos_thr <= 1
while True:
cset = np.unique(labels)
if len(cset) == 1:
break
centers = np.stack([x[labels == i].mean(0) for i in cset])
affinity = cosine_similarity(centers, centers)
affinity = np.triu(affinity, 1)
idx = np.unravel_index(np.argmax(affinity), affinity.shape)
if affinity[idx] < cos_thr:
break
c1, c2 = cset[np.array(idx)]
labels[labels == c2] = c1
return labels
def compressed_seg(seg_list):
"""Mirrors ax_cam_bin.py compressed_seg(): merge adjacent same-speaker
segments [[st, ed, cluster_id], ...]."""
new_seg_list = []
for i, seg in enumerate(seg_list):
seg_st, seg_ed, cluster_id = seg
if i == 0:
new_seg_list.append([seg_st, seg_ed, cluster_id])
elif cluster_id == new_seg_list[-1][2]:
if seg_st > new_seg_list[-1][1]:
new_seg_list.append([seg_st, seg_ed, cluster_id])
else:
new_seg_list[-1][1] = seg_ed
else:
if seg_st < new_seg_list[-1][1]:
p = (new_seg_list[-1][1] + seg_st) / 2
new_seg_list[-1][1] = p
seg_st = p
new_seg_list.append([seg_st, seg_ed, cluster_id])
return new_seg_list
def do_clustering(chunks, embeddings, speaker_num=None):
"""Mirrors ax_cam_bin.py do_clustering().
Returns (speaker_num, output_field_labels):
output_field_labels = [[st, ed, speaker_id], ...]
"""
cluster = CommonClustering(
cluster_type='spectral', mer_cos=0.8, min_num_spks=1,
max_num_spks=15, min_cluster_size=4, oracle_num=None, pval=0.012)
cluster_labels = cluster(
embeddings,
speaker_num=speaker_num if speaker_num is not None else speaker_num)
speaker_num = cluster_labels.max() + 1
output_field_labels = [[i[0], i[1], int(j)]
for i, j in zip(chunks, cluster_labels)]
output_field_labels = compressed_seg(output_field_labels)
return speaker_num, output_field_labels