Files
rustytorch/crates/specialized/rtx-segmentation/src/kmeans.rs
T
2026-03-04 00:08:42 +00:00

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, &centroids);
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, &centroids);
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]);
}
}