//! Dimensionality reduction and embedding algorithms. //! //! Implements UMAP, t-SNE, and PCA for visualizing single-cell data //! in 2D space. use crate::CellAtlasError; use cellatlas_shared::{AnalysisConfig, Cell, Embedding2D, EmbeddingMethod}; /// Compute 2D embedding for cells. pub fn compute_embedding( cells: &[Cell], config: &AnalysisConfig, ) -> Result, CellAtlasError> { if cells.is_empty() { return Ok(vec![]); } // Extract expression data for embedding let expression_matrix = extract_expression_matrix(cells); // First, compute PCA to reduce dimensionality let pca_result = compute_pca(&expression_matrix, config.dim_reduction.n_pcs)?; // Then compute 2D embedding let embeddings = match config.dim_reduction.embedding_method { EmbeddingMethod::Umap => compute_umap( &pca_result, config.dim_reduction.n_neighbors, config.dim_reduction.min_dist, )?, EmbeddingMethod::Tsne => compute_tsne(&pca_result)?, EmbeddingMethod::Pca => { // Just use first two PCs pca_result .iter() .map(|row| Embedding2D { x: row[0], y: row[1], }) .collect() } EmbeddingMethod::ForceDirected => compute_force_directed(&pca_result)?, }; // Update cells with embeddings let mut embedded_cells = cells.to_vec(); for (i, cell) in embedded_cells.iter_mut().enumerate() { if i < embeddings.len() { cell.embedding = Some(embeddings[i]); } } Ok(embedded_cells) } /// Extract expression matrix from cells. fn extract_expression_matrix(cells: &[Cell]) -> Vec> { let n_genes = cells.first().map_or(1000, |c| c.expression.num_genes); cells .iter() .map(|cell| { let mut row = vec![0.0; n_genes.min(2000)]; // Cap for demo for (i, gene_idx) in cell.expression.gene_indices.iter().enumerate() { let gene_idx = *gene_idx; if gene_idx < row.len() { row[gene_idx] = cell.expression.values.get(i).copied().unwrap_or(0.0); } } row }) .collect() } /// Compute PCA for dimensionality reduction. fn compute_pca(data: &[Vec], n_components: usize) -> Result>, CellAtlasError> { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; if data.is_empty() { return Ok(vec![]); } let n_samples = data.len(); let n_features = data[0].len(); let n_components = n_components.min(n_features).min(n_samples); // Center the data let mut means = vec![0.0; n_features]; for row in data { for (j, val) in row.iter().enumerate() { means[j] += val; } } for m in &mut means { *m /= n_samples as f32; } let centered: Vec> = data .iter() .map(|row| row.iter().enumerate().map(|(j, &v)| v - means[j]).collect()) .collect(); // Simplified PCA using power iteration // For a real implementation, use SVD let mut rng = rand::rngs::StdRng::seed_from_u64(42); let normal = Normal::new(0.0_f32, 1.0).unwrap(); let mut components: Vec> = Vec::with_capacity(n_components); for _component_idx in 0..n_components { // Initialize random vector let mut v: Vec = (0..n_features).map(|_| normal.sample(&mut rng)).collect(); normalize_vector(&mut v); // Power iteration for _ in 0..50 { // Compute X^T X v let mut new_v = vec![0.0; n_features]; for row in ¢ered { let dot: f32 = row.iter().zip(v.iter()).map(|(a, b)| a * b).sum(); for (j, x) in row.iter().enumerate() { new_v[j] += x * dot; } } // Remove previous components (Gram-Schmidt) for prev_component in &components { let proj: f32 = new_v .iter() .zip(prev_component.iter()) .map(|(a, b)| a * b) .sum(); for (j, pc) in prev_component.iter().enumerate() { new_v[j] -= proj * pc; } } normalize_vector(&mut new_v); v = new_v; } components.push(v); } // Project data onto components let projected: Vec> = centered .iter() .map(|row| { components .iter() .map(|comp| row.iter().zip(comp.iter()).map(|(a, b)| a * b).sum()) .collect() }) .collect(); Ok(projected) } /// Compute UMAP embedding. fn compute_umap( data: &[Vec], n_neighbors: usize, min_dist: f32, ) -> Result, CellAtlasError> { use rand::SeedableRng; use rand_distr::{Distribution, Normal, Uniform}; if data.is_empty() { return Ok(vec![]); } let n = data.len(); let mut rng = rand::rngs::StdRng::seed_from_u64(42); // Build k-NN graph let knn = build_knn_graph(data, n_neighbors); // Initialize low-dimensional embedding let normal = Normal::new(0.0_f32, 10.0).unwrap(); let mut embedding: Vec<[f32; 2]> = (0..n) .map(|_| [normal.sample(&mut rng), normal.sample(&mut rng)]) .collect(); // Optimization parameters let n_epochs = 200; let initial_alpha = 1.0_f32; let a = 1.0; let b = 1.0; // Compute high-dimensional probabilities let _sigmas = compute_sigmas(&knn, n_neighbors); // Stochastic gradient descent let uniform = Uniform::new(0, n).unwrap(); for epoch in 0..n_epochs { let alpha = initial_alpha * (1.0 - epoch as f32 / n_epochs as f32); for i in 0..n { // Attractive forces (neighbors) for (j, high_d_prob) in &knn[i] { let j = *j; let dx = embedding[j][0] - embedding[i][0]; let dy = embedding[j][1] - embedding[i][1]; let dist_sq = dx * dx + dy * dy + 0.001; let dist = dist_sq.sqrt(); // Gradient of attractive term let grad_coeff = -2.0 * a * b * dist.powf(b - 1.0) / (a * dist.powf(2.0 * b) + 1.0) * high_d_prob; embedding[i][0] -= alpha * grad_coeff * dx / dist; embedding[i][1] -= alpha * grad_coeff * dy / dist; } // Repulsive forces (random sampling) let n_neg_samples = 5; for _ in 0..n_neg_samples { let j = uniform.sample(&mut rng); if i == j || knn[i].iter().any(|(k, _)| *k == j) { continue; } let dx = embedding[j][0] - embedding[i][0]; let dy = embedding[j][1] - embedding[i][1]; let dist_sq = dx * dx + dy * dy + 0.001; let dist = dist_sq.sqrt(); // Gradient of repulsive term (with min_dist) if dist > min_dist { let grad_coeff = 2.0 * b / ((dist - min_dist).powi(2) + 1.0) / dist; embedding[i][0] += alpha * grad_coeff * dx * 0.1; embedding[i][1] += alpha * grad_coeff * dy * 0.1; } } } } // Center embedding let mean_x: f32 = embedding.iter().map(|e| e[0]).sum::() / n as f32; let mean_y: f32 = embedding.iter().map(|e| e[1]).sum::() / n as f32; Ok(embedding .iter() .map(|e| Embedding2D { x: e[0] - mean_x, y: e[1] - mean_y, }) .collect()) } /// Compute t-SNE embedding. fn compute_tsne(data: &[Vec]) -> Result, CellAtlasError> { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; if data.is_empty() { return Ok(vec![]); } let n = data.len(); let perplexity = 30.0_f32; let mut rng = rand::rngs::StdRng::seed_from_u64(42); // Initialize embedding let normal = Normal::new(0.0_f32, 0.0001).unwrap(); let mut y: Vec<[f32; 2]> = (0..n) .map(|_| [normal.sample(&mut rng), normal.sample(&mut rng)]) .collect(); // Compute pairwise distances in high dimension let p = compute_joint_probabilities(data, perplexity); // Optimization let n_iterations = 500; let initial_momentum = 0.5; let final_momentum = 0.8; let eta = 200.0_f32; let mut y_momentum: Vec<[f32; 2]> = vec![[0.0, 0.0]; n]; for iter in 0..n_iterations { let momentum = if iter < 250 { initial_momentum } else { final_momentum }; // Compute Q (low-dimensional similarities) let mut q = vec![vec![0.0_f32; n]; n]; let mut sum_q = 0.0_f32; for i in 0..n { for j in (i + 1)..n { let dist_sq = (y[i][0] - y[j][0]).powi(2) + (y[i][1] - y[j][1]).powi(2); let q_ij = 1.0 / (1.0 + dist_sq); q[i][j] = q_ij; q[j][i] = q_ij; sum_q += 2.0 * q_ij; } } // Normalize Q for i in 0..n { for j in 0..n { if i != j { q[i][j] /= sum_q + 1e-12; } } } // Compute gradient let mut grad: Vec<[f32; 2]> = vec![[0.0, 0.0]; n]; for i in 0..n { for j in 0..n { if i == j { continue; } let dist_sq = (y[i][0] - y[j][0]).powi(2) + (y[i][1] - y[j][1]).powi(2); let pq_diff = p[i][j] - q[i][j]; let mult = 4.0 * pq_diff / (1.0 + dist_sq); grad[i][0] += mult * (y[i][0] - y[j][0]); grad[i][1] += mult * (y[i][1] - y[j][1]); } } // Update embedding for i in 0..n { y_momentum[i][0] = momentum * y_momentum[i][0] - eta * grad[i][0]; y_momentum[i][1] = momentum * y_momentum[i][1] - eta * grad[i][1]; y[i][0] += y_momentum[i][0]; y[i][1] += y_momentum[i][1]; } // Center embedding let mean_x: f32 = y.iter().map(|p| p[0]).sum::() / n as f32; let mean_y: f32 = y.iter().map(|p| p[1]).sum::() / n as f32; for i in 0..n { y[i][0] -= mean_x; y[i][1] -= mean_y; } } Ok(y.iter().map(|p| Embedding2D { x: p[0], y: p[1] }).collect()) } /// Compute force-directed layout. fn compute_force_directed(data: &[Vec]) -> Result, CellAtlasError> { use rand::SeedableRng; use rand_distr::{Distribution, Uniform}; if data.is_empty() { return Ok(vec![]); } let n = data.len(); let mut rng = rand::rngs::StdRng::seed_from_u64(42); // Build k-NN graph let knn = build_knn_graph(data, 10); // Initialize positions let uniform = Uniform::new(-10.0_f32, 10.0).unwrap(); let mut pos: Vec<[f32; 2]> = (0..n) .map(|_| [uniform.sample(&mut rng), uniform.sample(&mut rng)]) .collect(); // Force-directed iteration let n_iterations = 100; let k = (100.0 / n as f32).sqrt(); // Optimal edge length for _ in 0..n_iterations { let mut displacement: Vec<[f32; 2]> = vec![[0.0, 0.0]; n]; // Repulsive forces between all pairs for i in 0..n { for j in (i + 1)..n { let dx = pos[i][0] - pos[j][0]; let dy = pos[i][1] - pos[j][1]; let dist = (dx * dx + dy * dy + 0.001).sqrt(); let repulsion = k * k / dist; let fx = dx / dist * repulsion; let fy = dy / dist * repulsion; displacement[i][0] += fx; displacement[i][1] += fy; displacement[j][0] -= fx; displacement[j][1] -= fy; } } // Attractive forces for edges for i in 0..n { for (j, _) in &knn[i] { let j = *j; let dx = pos[i][0] - pos[j][0]; let dy = pos[i][1] - pos[j][1]; let dist = (dx * dx + dy * dy + 0.001).sqrt(); let attraction = dist * dist / k; let fx = dx / dist * attraction; let fy = dy / dist * attraction; displacement[i][0] -= fx; displacement[i][1] -= fy; } } // Apply displacement with cooling let temp = 5.0_f32; for i in 0..n { let disp_len = (displacement[i][0].powi(2) + displacement[i][1].powi(2)).sqrt(); if disp_len > 0.001 { let scale = temp.min(disp_len) / disp_len; pos[i][0] += displacement[i][0] * scale; pos[i][1] += displacement[i][1] * scale; } } } Ok(pos .iter() .map(|p| Embedding2D { x: p[0], y: p[1] }) .collect()) } /// Build k-NN graph from data. fn build_knn_graph(data: &[Vec], k: usize) -> Vec> { let n = data.len(); let mut knn = vec![vec![]; n]; for i in 0..n { let mut distances: Vec<(usize, f32)> = (0..n) .filter(|&j| j != i) .map(|j| { let dist: f32 = data[i] .iter() .zip(data[j].iter()) .map(|(a, b)| (a - b).powi(2)) .sum::() .sqrt(); (j, dist) }) .collect(); distances.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap()); knn[i] = distances.into_iter().take(k).collect(); } knn } /// Compute sigmas for UMAP. fn compute_sigmas(knn: &[Vec<(usize, f32)>], _n_neighbors: usize) -> Vec { knn.iter() .map(|neighbors| { if neighbors.is_empty() { 1.0 } else { // Use mean distance as sigma let mean: f32 = neighbors.iter().map(|(_, d)| d).sum::() / neighbors.len() as f32; mean.max(0.001) } }) .collect() } /// Compute joint probabilities for t-SNE. fn compute_joint_probabilities(data: &[Vec], perplexity: f32) -> Vec> { let n = data.len(); // Compute pairwise distances let mut distances = vec![vec![0.0_f32; n]; n]; for i in 0..n { for j in (i + 1)..n { let dist: f32 = data[i] .iter() .zip(data[j].iter()) .map(|(a, b)| (a - b).powi(2)) .sum(); distances[i][j] = dist; distances[j][i] = dist; } } // Compute conditional probabilities with binary search for sigma let mut p = vec![vec![0.0_f32; n]; n]; let target_entropy = perplexity.ln(); for i in 0..n { let mut sigma = 1.0_f32; let mut lo = 0.001_f32; let mut hi = 1000.0_f32; // Binary search for sigma for _ in 0..50 { let mut p_row = vec![0.0_f32; n]; let mut sum = 0.0_f32; for j in 0..n { if i != j { let p_ij = (-distances[i][j] / (2.0 * sigma * sigma)).exp(); p_row[j] = p_ij; sum += p_ij; } } // Normalize for j in 0..n { p_row[j] /= sum + 1e-12; } // Compute entropy let entropy: f32 = -p_row .iter() .filter(|&&x| x > 1e-12) .map(|&x| x * x.ln()) .sum::(); if (entropy - target_entropy).abs() < 0.01 { p[i] = p_row; break; } if entropy > target_entropy { hi = sigma; } else { lo = sigma; } sigma = f32::midpoint(lo, hi); } } // Symmetrize for i in 0..n { for j in (i + 1)..n { let p_ij = (p[i][j] + p[j][i]) / (2.0 * n as f32); p[i][j] = p_ij.max(1e-12); p[j][i] = p_ij.max(1e-12); } } p } /// Normalize a vector in place. fn normalize_vector(v: &mut [f32]) { let norm: f32 = v.iter().map(|x| x * x).sum::().sqrt(); if norm > 1e-10 { for x in v { *x /= norm; } } } #[cfg(test)] mod tests { use super::*; use cellatlas_shared::{ AnnotationConfig, ClusteringConfig, DimReductionConfig, PreprocessingConfig, QualityMetrics, SparseExpression, }; fn create_test_cell(id: &str, expression: Vec<(usize, f32)>) -> Cell { Cell { id: id.to_string(), barcode: None, expression: SparseExpression { gene_indices: expression.iter().map(|(i, _)| *i).collect(), values: expression.iter().map(|(_, v)| *v).collect(), num_genes: 100, }, cell_type: None, state: None, qc_metrics: QualityMetrics { n_genes: expression.len(), total_counts: expression.iter().map(|(_, v)| v).sum(), pct_mito: 3.0, pct_ribo: 10.0, doublet_score: None, }, spatial_coords: None, cluster_id: None, embedding: None, } } #[test] fn test_pca() { let data: Vec> = (0..20) .map(|i| (0..10).map(|j| (i + j) as f32).collect()) .collect(); let result = compute_pca(&data, 3); assert!(result.is_ok()); let projected = result.unwrap(); assert_eq!(projected.len(), 20); assert_eq!(projected[0].len(), 3); } #[test] fn test_umap() { let data: Vec> = (0..30) .map(|i| { let cluster = i / 10; (0..5) .map(|j| (cluster * 10 + j) as f32 + (i % 10) as f32) .collect() }) .collect(); let result = compute_umap(&data, 5, 0.1); assert!(result.is_ok()); let embedding = result.unwrap(); assert_eq!(embedding.len(), 30); } #[test] fn test_tsne() { // Small dataset for t-SNE (it's slow) let data: Vec> = (0..15) .map(|i| (0..5).map(|j| (i + j) as f32).collect()) .collect(); let result = compute_tsne(&data); assert!(result.is_ok()); let embedding = result.unwrap(); assert_eq!(embedding.len(), 15); } #[test] fn test_compute_embedding() { let cells: Vec = (0..20) .map(|i| { create_test_cell( &format!("cell_{}", i), vec![(0, i as f32), (1, (i * 2) as f32), (2, (i % 5) as f32)], ) }) .collect(); let config = AnalysisConfig { preprocessing: PreprocessingConfig::default(), dim_reduction: DimReductionConfig { n_pcs: 5, embedding_method: EmbeddingMethod::Pca, n_neighbors: 5, min_dist: 0.1, }, clustering: ClusteringConfig::default(), annotation: AnnotationConfig::default(), }; let result = compute_embedding(&cells, &config); assert!(result.is_ok()); let embedded = result.unwrap(); assert_eq!(embedded.len(), 20); assert!(embedded.iter().all(|c| c.embedding.is_some())); } #[test] fn test_build_knn_graph() { let data: Vec> = vec![ vec![0.0, 0.0], vec![1.0, 0.0], vec![0.0, 1.0], vec![10.0, 10.0], ]; let knn = build_knn_graph(&data, 2); assert_eq!(knn.len(), 4); // First point should have second and third as neighbors assert!(knn[0].iter().any(|(j, _)| *j == 1 || *j == 2)); } }