430 lines
13 KiB
Rust
430 lines
13 KiB
Rust
//! 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<u64>,
|
|
/// 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<usize>,
|
|
/// Cluster centroids.
|
|
pub centroids: Vec<DVector<f64>>,
|
|
/// 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<f64>]) -> Result<KMeansResult> {
|
|
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<KMeansResult> = 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<f64>],
|
|
n_features: usize,
|
|
seed: Option<u64>,
|
|
) -> Result<KMeansResult> {
|
|
// 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<f64>],
|
|
_n_features: usize,
|
|
seed: Option<u64>,
|
|
) -> Vec<DVector<f64>> {
|
|
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<DVector<f64>> = Vec::with_capacity(self.config.k);
|
|
|
|
// Choose first centroid randomly
|
|
let first_idx = rng.r#gen::<usize>() % 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<f64> = data
|
|
.iter()
|
|
.map(|point| {
|
|
centroids
|
|
.iter()
|
|
.map(|c: &DVector<f64>| (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::<usize>() % data.len();
|
|
centroids.push(data[idx].clone());
|
|
} else {
|
|
let threshold: f64 = rng.r#gen::<f64>() * 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<f64>],
|
|
centroids: &[DVector<f64>],
|
|
) -> (Vec<usize>, 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<usize> = 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<f64>],
|
|
labels: &[usize],
|
|
n_features: usize,
|
|
) -> Vec<DVector<f64>> {
|
|
let mut sums: Vec<DVector<f64>> = (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<f64>], centroids: &[DVector<f64>]) -> Vec<usize> {
|
|
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<rtx_medical_io::Volume> {
|
|
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<DVector<f64>> = 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<DVector<f64>> = 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<DVector<f64>> = 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<DVector<f64>> = vec![];
|
|
|
|
let kmeans = KMeans::with_k(3);
|
|
assert!(kmeans.fit(&data).is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_predict() {
|
|
let train_data: Vec<DVector<f64>> = 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<DVector<f64>> = 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]);
|
|
}
|
|
}
|