//! K-means clustering implementation. //! //! This module provides K-means clustering for image segmentation, //! used to segment MRE images into tissue regions. use crate::error::{Result, SegmentationError}; use nalgebra::DVector; use rand::prelude::*; use rayon::prelude::*; use serde::{Deserialize, Serialize}; /// K-means clustering configuration. #[derive(Debug, Clone)] pub struct KMeansConfig { /// Number of clusters. pub k: usize, /// Maximum number of iterations. pub max_iterations: usize, /// Convergence tolerance (relative change in centroids). pub tolerance: f64, /// Random seed for initialization. pub seed: Option, /// Number of initializations (runs with different starting points). pub n_init: usize, } impl Default for KMeansConfig { fn default() -> Self { Self { k: 3, max_iterations: 300, tolerance: 1e-4, seed: None, n_init: 10, } } } /// Result of K-means clustering. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct KMeansResult { /// Cluster assignments for each sample. pub labels: Vec, /// Cluster centroids. pub centroids: Vec>, /// Within-cluster sum of squares (inertia). pub inertia: f64, /// Number of iterations performed. pub n_iterations: usize, } /// K-means clustering algorithm. #[derive(Debug, Clone)] pub struct KMeans { config: KMeansConfig, } impl KMeans { /// Create a new K-means instance with the given configuration. pub fn new(config: KMeansConfig) -> Self { Self { config } } /// Create a new K-means instance with k clusters. pub fn with_k(k: usize) -> Self { Self::new(KMeansConfig { k, ..Default::default() }) } /// Fit the model to the data and return cluster assignments. /// /// # Arguments /// * `data` - Samples as row vectors (n_samples x n_features) pub fn fit(&self, data: &[DVector]) -> Result { if data.is_empty() { return Err(SegmentationError::InvalidInput("Empty data".to_string())); } if self.config.k == 0 { return Err(SegmentationError::InvalidInput("k must be > 0".to_string())); } if self.config.k > data.len() { return Err(SegmentationError::InvalidInput(format!( "k ({}) must be <= n_samples ({})", self.config.k, data.len() ))); } let n_features = data[0].len(); // Run multiple initializations and keep the best let mut best_result: Option = None; for init in 0..self.config.n_init { let seed = self.config.seed.map(|s| s + init as u64); let result = self.fit_once(data, n_features, seed)?; if best_result.is_none() || result.inertia < best_result.as_ref().unwrap().inertia { best_result = Some(result); } } Ok(best_result.unwrap()) } fn fit_once( &self, data: &[DVector], n_features: usize, seed: Option, ) -> Result { // Initialize centroids using k-means++ let mut centroids = self.init_centroids_plusplus(data, n_features, seed); let mut labels = vec![0usize; data.len()]; let mut prev_inertia = f64::INFINITY; for iter in 0..self.config.max_iterations { // Assignment step: assign each point to nearest centroid let (new_labels, inertia) = self.assign_clusters(data, ¢roids); labels = new_labels; // Check for convergence if (prev_inertia - inertia).abs() / prev_inertia.max(1.0) < self.config.tolerance { return Ok(KMeansResult { labels, centroids, inertia, n_iterations: iter + 1, }); } prev_inertia = inertia; // Update step: recompute centroids centroids = self.update_centroids(data, &labels, n_features); } let (labels, inertia) = self.assign_clusters(data, ¢roids); Ok(KMeansResult { labels, centroids, inertia, n_iterations: self.config.max_iterations, }) } /// Initialize centroids using k-means++ algorithm. fn init_centroids_plusplus( &self, data: &[DVector], _n_features: usize, seed: Option, ) -> Vec> { use rand::SeedableRng; let mut rng: rand::rngs::StdRng = if let Some(s) = seed { rand::rngs::StdRng::seed_from_u64(s) } else { rand::rngs::StdRng::from_entropy() }; let mut centroids: Vec> = Vec::with_capacity(self.config.k); // Choose first centroid randomly let first_idx = rng.r#gen::() % data.len(); centroids.push(data[first_idx].clone()); // Choose remaining centroids for _ in 1..self.config.k { // Compute distances to nearest centroid let distances: Vec = data .iter() .map(|point| { centroids .iter() .map(|c: &DVector| (point - c).norm_squared()) .fold(f64::INFINITY, f64::min) }) .collect(); // Sample proportional to distance squared let total: f64 = distances.iter().sum(); if total < 1e-10 { // All points are at distance 0, pick randomly let idx = rng.r#gen::() % data.len(); centroids.push(data[idx].clone()); } else { let threshold: f64 = rng.r#gen::() * total; let mut cumsum = 0.0; for (idx, d) in distances.iter().enumerate() { cumsum += d; if cumsum >= threshold { centroids.push(data[idx].clone()); break; } } } } centroids } /// Assign each point to the nearest centroid. fn assign_clusters( &self, data: &[DVector], centroids: &[DVector], ) -> (Vec, f64) { let results: Vec<(usize, f64)> = data .par_iter() .map(|point| { let mut min_dist = f64::INFINITY; let mut min_idx = 0; for (i, centroid) in centroids.iter().enumerate() { let dist = (point - centroid).norm_squared(); if dist < min_dist { min_dist = dist; min_idx = i; } } (min_idx, min_dist) }) .collect(); let labels: Vec = results.iter().map(|(idx, _)| *idx).collect(); let inertia: f64 = results.iter().map(|(_, dist)| dist).sum(); (labels, inertia) } /// Recompute centroids as the mean of assigned points. fn update_centroids( &self, data: &[DVector], labels: &[usize], n_features: usize, ) -> Vec> { let mut sums: Vec> = (0..self.config.k) .map(|_| DVector::zeros(n_features)) .collect(); let mut counts = vec![0usize; self.config.k]; for (point, &label) in data.iter().zip(labels.iter()) { sums[label] += point; counts[label] += 1; } sums.into_iter() .zip(counts.iter()) .map( |(sum, &count)| { if count > 0 { sum / count as f64 } else { sum } }, ) .collect() } /// Predict cluster labels for new data. pub fn predict(&self, data: &[DVector], centroids: &[DVector]) -> Vec { self.assign_clusters(data, centroids).0 } } /// Segment a volume using K-means clustering on voxel intensities. pub fn segment_volume_kmeans( volume: &rtx_medical_io::Volume, k: usize, mask: Option<&rtx_medical_io::Volume>, ) -> Result { let shape = volume.shape(); let mut data = Vec::new(); let mut indices = Vec::new(); // Collect voxel values for z in 0..shape[2] { for y in 0..shape[1] { for x in 0..shape[0] { if let Some(m) = mask { if m.get(x, y, z).unwrap_or(0.0) <= 0.0 { continue; } } let val = volume.get(x, y, z).unwrap_or(0.0); data.push(DVector::from_element(1, val)); indices.push((x, y, z)); } } } if data.is_empty() { return Err(SegmentationError::InvalidInput( "No voxels to segment".to_string(), )); } // Run K-means let kmeans = KMeans::with_k(k); let result = kmeans.fit(&data)?; // Create output volume let mut output = rtx_medical_io::Volume::zeros(shape); for (idx, (x, y, z)) in indices.into_iter().enumerate() { output.set(x, y, z, result.labels[idx] as f64); } Ok(output) } #[cfg(test)] mod tests { use super::*; #[test] fn test_kmeans_basic() { // Create simple 2D data with 3 clear clusters let data: Vec> = vec![ // Cluster 0: around (0, 0) DVector::from_vec(vec![0.0, 0.1]), DVector::from_vec(vec![0.1, 0.0]), DVector::from_vec(vec![-0.1, 0.0]), DVector::from_vec(vec![0.0, -0.1]), // Cluster 1: around (10, 0) DVector::from_vec(vec![10.0, 0.1]), DVector::from_vec(vec![10.1, 0.0]), DVector::from_vec(vec![9.9, 0.0]), DVector::from_vec(vec![10.0, -0.1]), // Cluster 2: around (5, 10) DVector::from_vec(vec![5.0, 10.1]), DVector::from_vec(vec![5.1, 10.0]), DVector::from_vec(vec![4.9, 10.0]), DVector::from_vec(vec![5.0, 9.9]), ]; let config = KMeansConfig { k: 3, max_iterations: 100, tolerance: 1e-6, seed: Some(42), n_init: 3, }; let kmeans = KMeans::new(config); let result = kmeans.fit(&data).unwrap(); assert_eq!(result.labels.len(), 12); assert_eq!(result.centroids.len(), 3); // Check that points in the same cluster have the same label assert_eq!(result.labels[0], result.labels[1]); assert_eq!(result.labels[0], result.labels[2]); assert_eq!(result.labels[0], result.labels[3]); assert_eq!(result.labels[4], result.labels[5]); assert_eq!(result.labels[4], result.labels[6]); assert_eq!(result.labels[4], result.labels[7]); assert_eq!(result.labels[8], result.labels[9]); assert_eq!(result.labels[8], result.labels[10]); assert_eq!(result.labels[8], result.labels[11]); // Check that different clusters have different labels assert_ne!(result.labels[0], result.labels[4]); assert_ne!(result.labels[0], result.labels[8]); assert_ne!(result.labels[4], result.labels[8]); } #[test] fn test_kmeans_single_cluster() { let data: Vec> = vec![ DVector::from_vec(vec![1.0, 2.0]), DVector::from_vec(vec![1.1, 2.1]), DVector::from_vec(vec![0.9, 1.9]), ]; let kmeans = KMeans::with_k(1); let result = kmeans.fit(&data).unwrap(); assert_eq!(result.labels.len(), 3); assert!(result.labels.iter().all(|&l| l == 0)); assert_eq!(result.centroids.len(), 1); } #[test] fn test_kmeans_invalid_k() { let data: Vec> = vec![DVector::from_vec(vec![1.0, 2.0])]; let kmeans = KMeans::with_k(0); assert!(kmeans.fit(&data).is_err()); let kmeans = KMeans::with_k(5); assert!(kmeans.fit(&data).is_err()); } #[test] fn test_kmeans_empty_data() { let data: Vec> = vec![]; let kmeans = KMeans::with_k(3); assert!(kmeans.fit(&data).is_err()); } #[test] fn test_predict() { let train_data: Vec> = vec![ DVector::from_vec(vec![0.0, 0.0]), DVector::from_vec(vec![10.0, 10.0]), ]; let kmeans = KMeans::with_k(2); let result = kmeans.fit(&train_data).unwrap(); let test_data: Vec> = vec![ DVector::from_vec(vec![0.1, 0.1]), DVector::from_vec(vec![9.9, 9.9]), ]; let labels = kmeans.predict(&test_data, &result.centroids); // The two test points should have different labels assert_ne!(labels[0], labels[1]); } }