| |
| |
| |
| |
| |
| |
| |
|
|
| 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 |
|
|