use super::embedding::SpeakerEmbedding;
#[cfg(test)]
mod tests;
use crate::error::{WhisperError, WhisperResult};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ClusteringAlgorithm {
#[default]
Spectral,
KMeans,
Agglomerative,
}
#[derive(Debug, Clone)]
pub struct ClusteringConfig {
pub algorithm: ClusteringAlgorithm,
pub distance_threshold: f32,
pub min_cluster_size: usize,
pub max_iterations: usize,
pub convergence_threshold: f32,
pub use_cosine_distance: bool,
}
impl Default for ClusteringConfig {
fn default() -> Self {
Self {
algorithm: ClusteringAlgorithm::default(),
distance_threshold: 0.5,
min_cluster_size: 1,
max_iterations: 100,
convergence_threshold: 1e-4,
use_cosine_distance: true,
}
}
}
impl ClusteringConfig {
#[must_use]
pub fn for_realtime() -> Self {
Self {
algorithm: ClusteringAlgorithm::KMeans,
max_iterations: 50,
..Default::default()
}
}
#[must_use]
pub fn for_accuracy() -> Self {
Self {
algorithm: ClusteringAlgorithm::Spectral,
max_iterations: 200,
distance_threshold: 0.4,
..Default::default()
}
}
#[must_use]
pub fn with_algorithm(mut self, algorithm: ClusteringAlgorithm) -> Self {
self.algorithm = algorithm;
self
}
#[must_use]
pub fn with_distance_threshold(mut self, threshold: f32) -> Self {
self.distance_threshold = threshold;
self
}
}
#[derive(Debug, Clone)]
pub struct SpeakerCluster {
id: usize,
member_indices: Vec<usize>,
centroid: SpeakerEmbedding,
cohesion: f32,
}
impl SpeakerCluster {
#[must_use]
pub fn new(id: usize, member_indices: Vec<usize>, centroid: SpeakerEmbedding) -> Self {
Self {
id,
member_indices,
centroid,
cohesion: 0.0,
}
}
#[must_use]
pub fn with_cohesion(mut self, cohesion: f32) -> Self {
self.cohesion = cohesion;
self
}
#[must_use]
pub fn id(&self) -> usize {
self.id
}
#[must_use]
pub fn member_indices(&self) -> &[usize] {
&self.member_indices
}
#[must_use]
pub fn centroid(&self) -> &SpeakerEmbedding {
&self.centroid
}
#[must_use]
pub fn size(&self) -> usize {
self.member_indices.len()
}
#[must_use]
pub fn cohesion(&self) -> f32 {
self.cohesion
}
}
#[derive(Debug, Clone)]
pub struct ClusteringResult {
labels: Vec<usize>,
clusters: Vec<SpeakerCluster>,
num_clusters: usize,
silhouette_score: f32,
}
impl ClusteringResult {
#[must_use]
pub fn new(labels: Vec<usize>, clusters: Vec<SpeakerCluster>) -> Self {
let num_clusters = clusters.len();
Self {
labels,
clusters,
num_clusters,
silhouette_score: 0.0,
}
}
#[must_use]
pub fn with_silhouette_score(mut self, score: f32) -> Self {
self.silhouette_score = score;
self
}
#[must_use]
pub fn labels(&self) -> &[usize] {
&self.labels
}
#[must_use]
pub fn clusters(&self) -> &[SpeakerCluster] {
&self.clusters
}
#[must_use]
pub fn num_clusters(&self) -> usize {
self.num_clusters
}
#[must_use]
pub fn silhouette_score(&self) -> f32 {
self.silhouette_score
}
#[must_use]
pub fn cluster_centroids(&self) -> Vec<SpeakerEmbedding> {
self.clusters.iter().map(|c| c.centroid().clone()).collect()
}
}
#[derive(Debug)]
pub struct SpectralClustering {
config: ClusteringConfig,
}
impl SpectralClustering {
#[must_use]
pub fn new(config: ClusteringConfig) -> Self {
Self { config }
}
pub fn cluster(
&self,
embeddings: &[SpeakerEmbedding],
max_clusters: Option<usize>,
min_clusters: usize,
) -> WhisperResult<ClusteringResult> {
match embeddings.len() {
0 => return Ok(ClusteringResult::new(Vec::new(), Vec::new())),
1 => {
let cluster = SpeakerCluster::new(0, vec![0], embeddings[0].clone());
return Ok(ClusteringResult::new(vec![0], vec![cluster]));
}
_ => {}
}
let affinity = self.build_affinity_matrix(embeddings);
let num_clusters = self.estimate_num_clusters(&affinity, max_clusters, min_clusters);
let labels = self.spectral_cluster(&affinity, num_clusters)?;
let clusters = self.build_clusters(embeddings, &labels, num_clusters);
let silhouette = self.compute_silhouette(embeddings, &labels);
Ok(ClusteringResult::new(labels, clusters).with_silhouette_score(silhouette))
}
fn build_affinity_matrix(&self, embeddings: &[SpeakerEmbedding]) -> Vec<Vec<f32>> {
let n = embeddings.len();
let mut affinity = vec![vec![0.0f32; n]; n];
for i in 0..n {
for j in i..n {
let sim = if self.config.use_cosine_distance {
embeddings[i].cosine_similarity(&embeddings[j])
} else {
let dist = embeddings[i].euclidean_distance(&embeddings[j]);
(-dist * dist / 2.0).exp()
};
let aff = (sim + 1.0) / 2.0;
affinity[i][j] = aff;
affinity[j][i] = aff;
}
}
affinity
}
fn estimate_num_clusters(
&self,
affinity: &[Vec<f32>],
max_clusters: Option<usize>,
min_clusters: usize,
) -> usize {
let n = affinity.len();
let max_k = max_clusters.unwrap_or_else(|| n.min(10));
if n <= min_clusters {
return min_clusters.min(n);
}
let _degrees: Vec<f32> = affinity.iter().map(|row| row.iter().sum()).collect();
let threshold = self.config.distance_threshold;
let mut groups = 0;
let mut visited = vec![false; n];
for i in 0..n {
if visited[i] {
continue;
}
groups += 1;
let mut stack = vec![i];
while let Some(node) = stack.pop() {
if visited[node] {
continue;
}
visited[node] = true;
for (j, &aff) in affinity[node].iter().enumerate() {
if !visited[j] && aff > threshold {
stack.push(j);
}
}
}
}
groups.clamp(min_clusters, max_k)
}
fn spectral_cluster(
&self,
affinity: &[Vec<f32>],
num_clusters: usize,
) -> WhisperResult<Vec<usize>> {
let n = affinity.len();
if num_clusters >= n {
return Ok((0..n).collect());
}
let labels = self.kmeans_on_affinity(affinity, num_clusters)?;
Ok(labels)
}
fn kmeans_on_affinity(&self, affinity: &[Vec<f32>], k: usize) -> WhisperResult<Vec<usize>> {
let n = affinity.len();
let dim = n;
if k == 0 || k > n {
return Err(WhisperError::Diarization(
"Invalid number of clusters".to_string(),
));
}
let mut centroids: Vec<Vec<f32>> = affinity[..k].to_vec();
let mut labels = vec![0usize; n];
for _iter in 0..self.config.max_iterations {
let old_labels = labels.clone();
for (i, row) in affinity.iter().enumerate() {
let mut min_dist = f32::MAX;
let mut best_cluster = 0;
for (j, centroid) in centroids.iter().enumerate() {
let dist: f32 = row
.iter()
.zip(centroid.iter())
.map(|(&a, &b)| (a - b).powi(2))
.sum::<f32>()
.sqrt();
if dist < min_dist {
min_dist = dist;
best_cluster = j;
}
}
labels[i] = best_cluster;
}
for (j, centroid) in centroids.iter_mut().enumerate() {
let member_count = labels.iter().filter(|&&l| l == j).count();
if member_count == 0 {
continue;
}
for d in 0..dim {
centroid[d] = labels
.iter()
.enumerate()
.filter(|(_, &l)| l == j)
.map(|(i, _)| affinity[i][d])
.sum::<f32>()
/ member_count as f32;
}
}
if labels == old_labels {
break;
}
}
Ok(labels)
}
fn build_clusters(
&self,
embeddings: &[SpeakerEmbedding],
labels: &[usize],
num_clusters: usize,
) -> Vec<SpeakerCluster> {
let mut clusters = Vec::with_capacity(num_clusters);
for cluster_id in 0..num_clusters {
let member_indices: Vec<usize> = labels
.iter()
.enumerate()
.filter(|(_, &l)| l == cluster_id)
.map(|(i, _)| i)
.collect();
if member_indices.is_empty() {
continue;
}
let member_embeddings: Vec<SpeakerEmbedding> = member_indices
.iter()
.map(|&i| embeddings[i].clone())
.collect();
let centroid = SpeakerEmbedding::mean(&member_embeddings)
.unwrap_or_else(|| embeddings[member_indices[0]].clone());
let cohesion = self.compute_cluster_cohesion(&member_embeddings, ¢roid);
clusters.push(
SpeakerCluster::new(cluster_id, member_indices, centroid).with_cohesion(cohesion),
);
}
clusters
}
fn compute_cluster_cohesion(
&self,
members: &[SpeakerEmbedding],
centroid: &SpeakerEmbedding,
) -> f32 {
if members.is_empty() {
return 0.0;
}
let total_dist: f32 = members
.iter()
.map(|m| {
if self.config.use_cosine_distance {
1.0 - m.cosine_similarity(centroid)
} else {
m.euclidean_distance(centroid)
}
})
.sum();
total_dist / members.len() as f32
}
fn compute_silhouette(&self, embeddings: &[SpeakerEmbedding], labels: &[usize]) -> f32 {
let unique_labels: Vec<usize> = {
let mut v: Vec<usize> = labels.to_vec();
v.sort_unstable();
v.dedup();
v
};
if embeddings.len() < 2 || unique_labels.len() < 2 {
return 0.0;
}
let mut total_silhouette = 0.0;
for (i, emb) in embeddings.iter().enumerate() {
let own_cluster = labels[i];
let same_cluster: Vec<f32> = embeddings
.iter()
.enumerate()
.filter(|(j, _)| *j != i && labels[*j] == own_cluster)
.map(|(_, other)| {
if self.config.use_cosine_distance {
1.0 - emb.cosine_similarity(other)
} else {
emb.euclidean_distance(other)
}
})
.collect();
let a = if same_cluster.is_empty() {
0.0
} else {
same_cluster.iter().sum::<f32>() / same_cluster.len() as f32
};
let b = unique_labels
.iter()
.filter(|&&l| l != own_cluster)
.map(|&other_cluster| {
let other_dists: Vec<f32> = embeddings
.iter()
.enumerate()
.filter(|(_, _)| labels.get(i) == Some(&other_cluster))
.map(|(_, other)| {
if self.config.use_cosine_distance {
1.0 - emb.cosine_similarity(other)
} else {
emb.euclidean_distance(other)
}
})
.collect();
if other_dists.is_empty() {
f32::MAX
} else {
other_dists.iter().sum::<f32>() / other_dists.len() as f32
}
})
.fold(f32::MAX, f32::min);
let s = if a.max(b) > 0.0 {
(b - a) / a.max(b)
} else {
0.0
};
total_silhouette += s;
}
total_silhouette / embeddings.len() as f32
}
#[must_use]
pub fn config(&self) -> &ClusteringConfig {
&self.config
}
}