Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
+600
View File
@@ -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);
}
}