//! Clustering algorithms for single-cell analysis. //! //! Implements Leiden, Louvain, and k-means clustering for //! grouping cells by expression similarity. use crate::CellAtlasError; use cellatlas_shared::{AnalysisConfig, Cell, ClusterInfo, ClusteringAlgorithm, Embedding2D}; /// Cluster cells based on their embeddings. pub fn cluster_cells( cells: &[Cell], config: &AnalysisConfig, ) -> Result, CellAtlasError> { match config.clustering.algorithm { ClusteringAlgorithm::Leiden => leiden_clustering(cells, config.clustering.resolution), ClusteringAlgorithm::Louvain => louvain_clustering(cells, config.clustering.resolution), ClusteringAlgorithm::Kmeans => { let k = config.clustering.n_clusters.unwrap_or(10); kmeans_clustering(cells, k) } ClusteringAlgorithm::Hierarchical => hierarchical_clustering(cells), } } /// Leiden community detection algorithm. /// /// A refinement of Louvain that guarantees well-connected communities. fn leiden_clustering(cells: &[Cell], resolution: f32) -> Result, CellAtlasError> { use rand::SeedableRng; use rand_distr::{Distribution, Uniform}; if cells.is_empty() { return Ok(vec![]); } let n = cells.len(); let mut rng = rand::rngs::StdRng::seed_from_u64(42); // Build k-NN graph based on embeddings let embeddings: Vec> = cells.iter().map(|c| c.embedding.as_ref()).collect(); let adjacency = build_embedding_knn(&embeddings, 15); // Initialize each cell in its own cluster let mut cluster_assignments: Vec = (0..n).collect(); // Modularity-based optimization (simplified Leiden) let max_iterations = 10; for _ in 0..max_iterations { let mut changed = false; // Random order traversal let mut order: Vec = (0..n).collect(); for i in (1..n).rev() { let j = Uniform::new(0, i + 1).unwrap().sample(&mut rng); order.swap(i, j); } for i in &order { let i = *i; // Find best cluster for this cell let current_cluster = cluster_assignments[i]; let neighbors = &adjacency[i]; if neighbors.is_empty() { continue; } // Count neighbor clusters let mut cluster_counts: std::collections::HashMap = std::collections::HashMap::new(); for j in neighbors { *cluster_counts.entry(cluster_assignments[*j]).or_insert(0) += 1; } // Find cluster with most neighbors (weighted by resolution) let mut best_cluster = current_cluster; let mut best_score = 0.0_f32; for (cluster, count) in &cluster_counts { let cluster = *cluster; let count = *count; let score = count as f32 * resolution; if score > best_score { best_score = score; best_cluster = cluster; } } if best_cluster != current_cluster { cluster_assignments[i] = best_cluster; changed = true; } } if !changed { break; } } // Renumber clusters to be contiguous let mut unique_clusters: Vec = cluster_assignments.clone(); unique_clusters.sort_unstable(); unique_clusters.dedup(); let cluster_map: std::collections::HashMap = unique_clusters .iter() .enumerate() .map(|(new_id, &old_id)| (old_id, new_id)) .collect(); for assignment in &mut cluster_assignments { *assignment = cluster_map[assignment]; } // Build cluster info let num_clusters = unique_clusters.len(); let mut cluster_infos = Vec::with_capacity(num_clusters); for cluster_id in 0..num_clusters { let member_indices: Vec = cluster_assignments .iter() .enumerate() .filter(|(_, c)| **c == cluster_id) .map(|(i, _)| i) .collect(); let n_cells = member_indices.len(); // Calculate centroid let centroid = calculate_centroid(cells, &member_indices); // Get top marker genes (simulated for demo) let markers = get_cluster_markers(cluster_id); cluster_infos.push(ClusterInfo { id: cluster_id, n_cells, markers, centroid, }); } Ok(cluster_infos) } /// Louvain community detection algorithm. fn louvain_clustering(cells: &[Cell], resolution: f32) -> Result, CellAtlasError> { // Louvain is similar to Leiden but without the refinement step // For this demo, we use the same implementation leiden_clustering(cells, resolution) } /// K-means clustering. fn kmeans_clustering(cells: &[Cell], k: usize) -> Result, CellAtlasError> { use rand::SeedableRng; use rand_distr::{Distribution, Uniform}; if cells.is_empty() || k == 0 { return Ok(vec![]); } let n = cells.len(); let k = k.min(n); let mut rng = rand::rngs::StdRng::seed_from_u64(42); // Get 2D embeddings let points: Vec<(f32, f32)> = cells .iter() .map(|c| c.embedding.map_or((0.0, 0.0), |e| (e.x, e.y))) .collect(); // Initialize centroids (k-means++) let mut centroids: Vec<(f32, f32)> = Vec::with_capacity(k); // First centroid: random point let first_idx = Uniform::new(0, n).unwrap().sample(&mut rng); centroids.push(points[first_idx]); // Remaining centroids: weighted by distance for _ in 1..k { let distances: Vec = points .iter() .map(|p| { centroids .iter() .map(|c| (p.0 - c.0).powi(2) + (p.1 - c.1).powi(2)) .fold(f32::INFINITY, f32::min) }) .collect(); let total_dist: f32 = distances.iter().sum(); let threshold = Uniform::new(0.0, total_dist).unwrap().sample(&mut rng); let mut cumsum = 0.0; let mut selected_idx = 0; for (i, d) in distances.iter().enumerate() { cumsum += d; if cumsum >= threshold { selected_idx = i; break; } } centroids.push(points[selected_idx]); } // Run k-means iterations let mut assignments: Vec = vec![0; n]; let max_iterations = 50; for _ in 0..max_iterations { // Assign points to nearest centroid let mut changed = false; for (i, p) in points.iter().enumerate() { let nearest = centroids .iter() .enumerate() .min_by(|(_, a), (_, b)| { let da = (p.0 - a.0).powi(2) + (p.1 - a.1).powi(2); let db = (p.0 - b.0).powi(2) + (p.1 - b.1).powi(2); da.partial_cmp(&db).unwrap() }) .map_or(0, |(idx, _)| idx); if assignments[i] != nearest { assignments[i] = nearest; changed = true; } } if !changed { break; } // Update centroids for (c_idx, centroid) in centroids.iter_mut().enumerate() { let members: Vec<&(f32, f32)> = points .iter() .enumerate() .filter(|(i, _)| assignments[*i] == c_idx) .map(|(_, p)| p) .collect(); if !members.is_empty() { let sum_x: f32 = members.iter().map(|p| p.0).sum(); let sum_y: f32 = members.iter().map(|p| p.1).sum(); let count = members.len() as f32; *centroid = (sum_x / count, sum_y / count); } } } // Build cluster info let mut cluster_infos = Vec::with_capacity(k); for cluster_id in 0..k { let member_indices: Vec = assignments .iter() .enumerate() .filter(|(_, c)| **c == cluster_id) .map(|(i, _)| i) .collect(); if member_indices.is_empty() { continue; } let n_cells = member_indices.len(); let centroid = Some(Embedding2D { x: centroids[cluster_id].0, y: centroids[cluster_id].1, }); let markers = get_cluster_markers(cluster_id); cluster_infos.push(ClusterInfo { id: cluster_id, n_cells, markers, centroid, }); } Ok(cluster_infos) } /// Hierarchical clustering. fn hierarchical_clustering(cells: &[Cell]) -> Result, CellAtlasError> { // For demo, use a simplified approach based on embedding distance // Real implementation would use Ward's linkage or similar if cells.is_empty() { return Ok(vec![]); } let n = cells.len(); let target_clusters = (n as f32).sqrt().ceil() as usize; // Get embeddings let points: Vec<(f32, f32)> = cells .iter() .map(|c| c.embedding.map_or((0.0, 0.0), |e| (e.x, e.y))) .collect(); // Initialize each point as its own cluster let mut cluster_ids: Vec = (0..n).collect(); let mut active_clusters: std::collections::HashSet = (0..n).collect(); // Merge clusters until we reach target while active_clusters.len() > target_clusters { // Find two closest clusters let mut min_dist = f32::INFINITY; let mut merge_pair = (0, 0); let active_vec: Vec = active_clusters.iter().copied().collect(); for i in 0..active_vec.len() { for j in (i + 1)..active_vec.len() { let c1 = active_vec[i]; let c2 = active_vec[j]; // Get members of each cluster let members1: Vec = cluster_ids .iter() .enumerate() .filter(|(_, c)| **c == c1) .map(|(i, _)| i) .collect(); let members2: Vec = cluster_ids .iter() .enumerate() .filter(|(_, c)| **c == c2) .map(|(i, _)| i) .collect(); // Calculate average linkage distance let mut total_dist = 0.0; let mut count = 0; for m1 in &members1 { let m1 = *m1; for m2 in &members2 { let m2 = *m2; let dx = points[m1].0 - points[m2].0; let dy = points[m1].1 - points[m2].1; total_dist += (dx * dx + dy * dy).sqrt(); count += 1; } } let avg_dist = if count > 0 { total_dist / count as f32 } else { f32::INFINITY }; if avg_dist < min_dist { min_dist = avg_dist; merge_pair = (c1, c2); } } } // Merge: assign all members of c2 to c1 let (c1, c2) = merge_pair; for id in &mut cluster_ids { if *id == c2 { *id = c1; } } active_clusters.remove(&c2); } // Renumber clusters let unique: Vec = active_clusters.iter().copied().collect(); let cluster_map: std::collections::HashMap = unique .iter() .enumerate() .map(|(new, old)| (*old, new)) .collect(); for id in &mut cluster_ids { *id = cluster_map[id]; } // Build cluster info let num_clusters = unique.len(); let mut cluster_infos = Vec::with_capacity(num_clusters); for cluster_id in 0..num_clusters { let member_indices: Vec = cluster_ids .iter() .enumerate() .filter(|(_, c)| **c == cluster_id) .map(|(i, _)| i) .collect(); let n_cells = member_indices.len(); let centroid = calculate_centroid(cells, &member_indices); let markers = get_cluster_markers(cluster_id); cluster_infos.push(ClusterInfo { id: cluster_id, n_cells, markers, centroid, }); } Ok(cluster_infos) } /// Build k-NN graph from embeddings. fn build_embedding_knn(embeddings: &[Option<&Embedding2D>], k: usize) -> Vec> { let n = embeddings.len(); let mut adjacency = vec![vec![]; n]; for i in 0..n { let Some(emb_i) = embeddings[i] else { continue; }; // Calculate distances to all other points let mut distances: Vec<(usize, f32)> = embeddings .iter() .enumerate() .filter_map(|(j, emb_j)| { if i == j { return None; } emb_j.map(|e| { let dx = emb_i.x - e.x; let dy = emb_i.y - e.y; (j, (dx * dx + dy * dy).sqrt()) }) }) .collect(); // Sort and keep top k distances.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap()); distances.truncate(k); adjacency[i] = distances.iter().map(|(j, _)| *j).collect(); } adjacency } /// Calculate centroid of a cell cluster. fn calculate_centroid(cells: &[Cell], member_indices: &[usize]) -> Option { if member_indices.is_empty() { return None; } let mut sum_x = 0.0; let mut sum_y = 0.0; let mut count = 0; for i in member_indices { let i = *i; if let Some(emb) = cells.get(i).and_then(|c| c.embedding) { sum_x += emb.x; sum_y += emb.y; count += 1; } } if count == 0 { None } else { Some(Embedding2D { x: sum_x / count as f32, y: sum_y / count as f32, }) } } /// Get marker genes for a cluster (simulated for demo). fn get_cluster_markers(cluster_id: usize) -> Vec { let all_markers = vec![ vec!["CD3D", "CD3E", "IL7R", "CCR7"], // T cells vec!["CD19", "MS4A1", "CD79A", "CD79B"], // B cells vec!["CD14", "LYZ", "CST3", "FCGR3A"], // Monocytes vec!["GNLY", "NKG7", "KLRD1", "PRF1"], // NK cells vec!["FCER1A", "CD1C", "CLEC10A"], // Dendritic cells vec!["CD8A", "CD8B", "GZMK", "GZMB"], // CD8 T cells vec!["PPBP", "PF4", "GP9"], // Platelets vec!["HBA1", "HBA2", "HBB"], // Erythrocytes ]; let idx = cluster_id % all_markers.len(); all_markers[idx] .iter() .map(std::string::ToString::to_string) .collect() } #[cfg(test)] mod tests { use super::*; use cellatlas_shared::{ClusteringConfig, QualityMetrics, SparseExpression}; fn create_test_cell(id: &str, x: f32, y: f32) -> Cell { Cell { id: id.to_string(), barcode: None, expression: SparseExpression { gene_indices: vec![0, 1, 2], values: vec![1.0, 2.0, 1.5], num_genes: 100, }, cell_type: None, state: None, qc_metrics: QualityMetrics { n_genes: 1000, total_counts: 5000.0, pct_mito: 3.0, pct_ribo: 10.0, doublet_score: None, }, spatial_coords: None, cluster_id: None, embedding: Some(Embedding2D { x, y }), } } #[test] fn test_leiden_clustering() { let cells: Vec = (0..20) .map(|i| { let cluster = i / 5; let x = (cluster as f32) * 10.0 + (i % 5) as f32; let y = (cluster as f32) * 10.0 + (i % 5) as f32; create_test_cell(&format!("cell_{}", i), x, y) }) .collect(); let result = leiden_clustering(&cells, 1.0); assert!(result.is_ok()); let clusters = result.unwrap(); assert!(!clusters.is_empty()); // Should find roughly 4 clusters assert!(clusters.len() >= 2); } #[test] fn test_kmeans_clustering() { let cells: Vec = (0..30) .map(|i| { let cluster = i / 10; let x = (cluster as f32) * 20.0 + (i % 10) as f32; let y = (cluster as f32) * 20.0 + (i % 10) as f32 * 0.5; create_test_cell(&format!("cell_{}", i), x, y) }) .collect(); let result = kmeans_clustering(&cells, 3); assert!(result.is_ok()); let clusters = result.unwrap(); assert_eq!(clusters.len(), 3); let total_cells: usize = clusters.iter().map(|c| c.n_cells).sum(); assert_eq!(total_cells, 30); } #[test] fn test_hierarchical_clustering() { let cells: Vec = (0..16) .map(|i| { let x = (i % 4) as f32 * 5.0; let y = (i / 4) as f32 * 5.0; create_test_cell(&format!("cell_{}", i), x, y) }) .collect(); let result = hierarchical_clustering(&cells); assert!(result.is_ok()); let clusters = result.unwrap(); assert!(!clusters.is_empty()); } #[test] fn test_cluster_markers() { let markers = get_cluster_markers(0); assert!(!markers.is_empty()); assert!(markers.contains(&"CD3D".to_string())); } #[test] fn test_calculate_centroid() { let cells = vec![ create_test_cell("cell_0", 0.0, 0.0), create_test_cell("cell_1", 10.0, 0.0), create_test_cell("cell_2", 5.0, 10.0), ]; let centroid = calculate_centroid(&cells, &[0, 1, 2]); assert!(centroid.is_some()); let c = centroid.unwrap(); assert!((c.x - 5.0).abs() < 0.001); assert!((c.y - 10.0 / 3.0).abs() < 0.1); } }