//! Cell Transformer for encoding gene expression profiles. //! //! Implements a transformer-based model for learning cell representations //! from sparse gene expression data. /// Configuration for the Cell Transformer. #[derive(Debug, Clone)] pub struct CellTransformerConfig { /// Number of genes in vocabulary pub num_genes: usize, /// Hidden dimension pub hidden_dim: usize, /// Number of attention heads pub num_heads: usize, /// Number of transformer layers pub num_layers: usize, /// Dropout probability pub dropout: f32, /// Maximum sequence length (genes) pub max_seq_len: usize, } impl Default for CellTransformerConfig { fn default() -> Self { Self { num_genes: 20000, hidden_dim: 256, num_heads: 8, num_layers: 6, dropout: 0.1, max_seq_len: 2048, } } } /// Cell Transformer model. /// /// Encodes sparse gene expression into dense cell embeddings using /// self-attention over expressed genes. #[derive(Debug)] pub struct CellTransformer { config: CellTransformerConfig, /// Gene embeddings (`num_genes` x `hidden_dim`) gene_embeddings: Vec>, /// Expression value encoder value_encoder: ValueEncoder, /// Transformer layers layers: Vec, /// Output projection cell_head: CellHead, } impl CellTransformer { /// Create a new Cell Transformer with random initialization. #[must_use] pub fn new(config: CellTransformerConfig) -> Self { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; let mut rng = rand::rngs::StdRng::seed_from_u64(42); let std = (2.0 / config.hidden_dim as f32).sqrt(); let normal = Normal::new(0.0_f32, std).unwrap(); // Initialize gene embeddings let gene_embeddings: Vec> = (0..config.num_genes) .map(|_| { (0..config.hidden_dim) .map(|_| normal.sample(&mut rng)) .collect() }) .collect(); let value_encoder = ValueEncoder::new(config.hidden_dim); let layers: Vec = (0..config.num_layers) .map(|_| TransformerLayer::new(config.hidden_dim, config.num_heads, config.dropout)) .collect(); let cell_head = CellHead::new(config.hidden_dim); Self { config, gene_embeddings, value_encoder, layers, cell_head, } } /// Encode a batch of cells into embeddings. #[must_use] pub fn encode(&self, gene_indices: &[Vec], values: &[Vec]) -> Vec> { let batch_size = gene_indices.len(); let mut embeddings = Vec::with_capacity(batch_size); for i in 0..batch_size { let cell_embedding = self.encode_single(&gene_indices[i], &values[i]); embeddings.push(cell_embedding); } embeddings } /// Encode a single cell. fn encode_single(&self, gene_indices: &[usize], values: &[f32]) -> Vec { let seq_len = gene_indices.len().min(self.config.max_seq_len); // Get gene embeddings and add value encoding let mut hidden: Vec> = (0..seq_len) .map(|i| { let gene_idx = gene_indices[i]; let value = values[i]; let mut emb = self.gene_embeddings[gene_idx.min(self.config.num_genes - 1)].clone(); // Add value encoding let value_enc = self.value_encoder.encode(value); for (j, v) in value_enc.iter().enumerate() { emb[j] += v; } emb }) .collect(); // Apply transformer layers for layer in &self.layers { hidden = layer.forward(&hidden); } // Pool to cell embedding (mean pooling) self.cell_head.forward(&hidden) } /// Get the hidden dimension. #[must_use] pub fn hidden_dim(&self) -> usize { self.config.hidden_dim } } /// Encodes expression values into embeddings. #[derive(Debug)] struct ValueEncoder { dim: usize, // Learnable parameters for value encoding scale: Vec, bias: Vec, } impl ValueEncoder { fn new(dim: usize) -> Self { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; let mut rng = rand::rngs::StdRng::seed_from_u64(123); let normal = Normal::new(0.0_f32, 0.1).unwrap(); Self { dim, scale: (0..dim) .map(|_| 1.0 + normal.sample(&mut rng) * 0.1) .collect(), bias: (0..dim).map(|_| normal.sample(&mut rng)).collect(), } } fn encode(&self, value: f32) -> Vec { // Log-transform the value (common in scRNA-seq) let log_value = (value + 1.0).ln(); // Create sinusoidal encoding similar to positional encoding let mut encoding = Vec::with_capacity(self.dim); for i in 0..self.dim { let freq = (i as f32 / self.dim as f32) * std::f32::consts::PI; let enc = if i % 2 == 0 { (log_value * freq).sin() * self.scale[i] + self.bias[i] } else { (log_value * freq).cos() * self.scale[i] + self.bias[i] }; encoding.push(enc); } encoding } } /// Single transformer layer. #[derive(Debug)] struct TransformerLayer { hidden_dim: usize, num_heads: usize, head_dim: usize, // Attention weights q_weights: Vec>, k_weights: Vec>, v_weights: Vec>, o_weights: Vec>, // FFN weights ffn_w1: Vec>, ffn_w2: Vec>, // Layer norm parameters ln1_scale: Vec, ln1_bias: Vec, ln2_scale: Vec, ln2_bias: Vec, } impl TransformerLayer { fn new(hidden_dim: usize, num_heads: usize, _dropout: f32) -> Self { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; let mut rng = rand::rngs::StdRng::seed_from_u64(456); let std = (2.0 / hidden_dim as f32).sqrt(); let normal = Normal::new(0.0_f32, std).unwrap(); let head_dim = hidden_dim / num_heads; let ffn_dim = hidden_dim * 4; let mut init_weights = |rows: usize, cols: usize| -> Vec> { (0..rows) .map(|_| (0..cols).map(|_| normal.sample(&mut rng)).collect()) .collect() }; Self { hidden_dim, num_heads, head_dim, q_weights: init_weights(hidden_dim, hidden_dim), k_weights: init_weights(hidden_dim, hidden_dim), v_weights: init_weights(hidden_dim, hidden_dim), o_weights: init_weights(hidden_dim, hidden_dim), ffn_w1: init_weights(hidden_dim, ffn_dim), ffn_w2: init_weights(ffn_dim, hidden_dim), ln1_scale: vec![1.0; hidden_dim], ln1_bias: vec![0.0; hidden_dim], ln2_scale: vec![1.0; hidden_dim], ln2_bias: vec![0.0; hidden_dim], } } fn forward(&self, hidden: &[Vec]) -> Vec> { let seq_len = hidden.len(); // Self-attention with residual let normed = self.layer_norm(hidden, &self.ln1_scale, &self.ln1_bias); let attn_out = self.self_attention(&normed); let mut output: Vec> = (0..seq_len) .map(|i| { (0..self.hidden_dim) .map(|j| hidden[i][j] + attn_out[i][j]) .collect() }) .collect(); // FFN with residual let normed = self.layer_norm(&output, &self.ln2_scale, &self.ln2_bias); let ffn_out = self.feed_forward(&normed); for i in 0..seq_len { for j in 0..self.hidden_dim { output[i][j] += ffn_out[i][j]; } } output } fn self_attention(&self, hidden: &[Vec]) -> Vec> { let seq_len = hidden.len(); let scale = (self.head_dim as f32).sqrt(); // Compute Q, K, V let q = self.matmul(hidden, &self.q_weights); let k = self.matmul(hidden, &self.k_weights); let v = self.matmul(hidden, &self.v_weights); // Simplified attention (not split by heads for demo) let mut attn_output = vec![vec![0.0; self.hidden_dim]; seq_len]; for i in 0..seq_len { // Compute attention scores let mut scores = vec![0.0; seq_len]; for j in 0..seq_len { for d in 0..self.hidden_dim { scores[j] += q[i][d] * k[j][d]; } scores[j] /= scale; } // Softmax let max_score = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max); let exp_scores: Vec = scores.iter().map(|s| (s - max_score).exp()).collect(); let sum_exp: f32 = exp_scores.iter().sum(); let attn_weights: Vec = exp_scores.iter().map(|s| s / sum_exp).collect(); // Weighted sum of values for j in 0..seq_len { for d in 0..self.hidden_dim { attn_output[i][d] += attn_weights[j] * v[j][d]; } } } // Output projection self.matmul(&attn_output, &self.o_weights) } fn feed_forward(&self, hidden: &[Vec]) -> Vec> { // FFN: hidden -> 4*hidden -> hidden with GELU activation let intermediate = self.matmul(hidden, &self.ffn_w1); // GELU activation let activated: Vec> = intermediate .iter() .map(|row| { row.iter() .map(|&x| x * 0.5 * (1.0 + (x * 0.797_884_6).tanh())) .collect() }) .collect(); self.matmul(&activated, &self.ffn_w2) } fn layer_norm(&self, hidden: &[Vec], scale: &[f32], bias: &[f32]) -> Vec> { hidden .iter() .map(|row| { let mean: f32 = row.iter().sum::() / row.len() as f32; let var: f32 = row.iter().map(|x| (x - mean).powi(2)).sum::() / row.len() as f32; let std = (var + 1e-5).sqrt(); row.iter() .enumerate() .map(|(i, &x)| (x - mean) / std * scale[i] + bias[i]) .collect() }) .collect() } fn matmul(&self, a: &[Vec], b: &[Vec]) -> Vec> { let m = a.len(); let n = b[0].len(); let k = a[0].len(); let mut result = vec![vec![0.0; n]; m]; for i in 0..m { for j in 0..n { for l in 0..k.min(b.len()) { result[i][j] += a[i][l] * b[l][j]; } } } result } } /// Cell embedding head. #[derive(Debug)] struct CellHead { hidden_dim: usize, // Projection weights proj_weights: Vec>, } impl CellHead { fn new(hidden_dim: usize) -> Self { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; let mut rng = rand::rngs::StdRng::seed_from_u64(789); let std = (2.0 / hidden_dim as f32).sqrt(); let normal = Normal::new(0.0_f32, std).unwrap(); let proj_weights: Vec> = (0..hidden_dim) .map(|_| (0..hidden_dim).map(|_| normal.sample(&mut rng)).collect()) .collect(); Self { hidden_dim, proj_weights, } } fn forward(&self, hidden: &[Vec]) -> Vec { // Mean pooling let mut pooled = vec![0.0; self.hidden_dim]; for h in hidden { for (i, v) in h.iter().enumerate() { pooled[i] += v; } } let n = hidden.len() as f32; for v in &mut pooled { *v /= n; } // Final projection let mut output = vec![0.0; self.hidden_dim]; for i in 0..self.hidden_dim { for j in 0..self.hidden_dim { output[i] += pooled[j] * self.proj_weights[j][i]; } } output } } #[cfg(test)] mod tests { use super::*; #[test] fn test_cell_transformer_creation() { let config = CellTransformerConfig::default(); let transformer = CellTransformer::new(config); assert_eq!(transformer.hidden_dim(), 256); } #[test] fn test_cell_encoding() { let config = CellTransformerConfig { num_genes: 1000, hidden_dim: 64, num_heads: 4, num_layers: 2, dropout: 0.1, max_seq_len: 512, }; let transformer = CellTransformer::new(config); let gene_indices = vec![vec![0, 10, 50, 100]]; let values = vec![vec![1.0, 2.5, 0.5, 3.0]]; let embeddings = transformer.encode(&gene_indices, &values); assert_eq!(embeddings.len(), 1); assert_eq!(embeddings[0].len(), 64); } #[test] fn test_value_encoder() { let encoder = ValueEncoder::new(32); let enc1 = encoder.encode(1.0); let enc2 = encoder.encode(10.0); assert_eq!(enc1.len(), 32); assert_eq!(enc2.len(), 32); // Different values should produce different encodings let diff: f32 = enc1 .iter() .zip(enc2.iter()) .map(|(a, b)| (a - b).abs()) .sum(); assert!(diff > 0.1); } #[test] fn test_batch_encoding() { let config = CellTransformerConfig { num_genes: 500, hidden_dim: 32, num_heads: 2, num_layers: 1, dropout: 0.0, max_seq_len: 256, }; let transformer = CellTransformer::new(config); let gene_indices = vec![vec![0, 10, 20], vec![5, 15, 25, 35], vec![1, 2]]; let values = vec![ vec![1.0, 2.0, 1.5], vec![3.0, 1.0, 2.0, 0.5], vec![5.0, 4.0], ]; let embeddings = transformer.encode(&gene_indices, &values); assert_eq!(embeddings.len(), 3); for emb in &embeddings { assert_eq!(emb.len(), 32); } } }