# 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