Initial commit
This commit is contained in:
@@ -0,0 +1,600 @@
|
||||
//! 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<Vec<ClusterInfo>, 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<Vec<ClusterInfo>, 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<Option<&Embedding2D>> =
|
||||
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<usize> = (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<usize> = (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<usize, usize> =
|
||||
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<usize> = cluster_assignments.clone();
|
||||
unique_clusters.sort_unstable();
|
||||
unique_clusters.dedup();
|
||||
|
||||
let cluster_map: std::collections::HashMap<usize, usize> = 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<usize> = 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<Vec<ClusterInfo>, 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<Vec<ClusterInfo>, 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<f32> = 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<usize> = 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<usize> = 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<Vec<ClusterInfo>, 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<usize> = (0..n).collect();
|
||||
let mut active_clusters: std::collections::HashSet<usize> = (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<usize> = 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<usize> = cluster_ids
|
||||
.iter()
|
||||
.enumerate()
|
||||
.filter(|(_, c)| **c == c1)
|
||||
.map(|(i, _)| i)
|
||||
.collect();
|
||||
let members2: Vec<usize> = 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<usize> = active_clusters.iter().copied().collect();
|
||||
let cluster_map: std::collections::HashMap<usize, usize> = 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<usize> = 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<Vec<usize>> {
|
||||
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<Embedding2D> {
|
||||
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<String> {
|
||||
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<Cell> = (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<Cell> = (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<Cell> = (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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user