Consistent formatting pass: line wrapping, import sorting, trailing whitespace removal, let-chain indentation, merged derive attributes, and unsafe block reformatting. Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
484 lines
14 KiB
Rust
484 lines
14 KiB
Rust
//! 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<Vec<f32>>,
|
|
/// Expression value encoder
|
|
value_encoder: ValueEncoder,
|
|
/// Transformer layers
|
|
layers: Vec<TransformerLayer>,
|
|
/// 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<Vec<f32>> = (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<TransformerLayer> = (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<usize>], values: &[Vec<f32>]) -> Vec<Vec<f32>> {
|
|
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<f32> {
|
|
let seq_len = gene_indices.len().min(self.config.max_seq_len);
|
|
|
|
// Get gene embeddings and add value encoding
|
|
let mut hidden: Vec<Vec<f32>> = (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<f32>,
|
|
bias: Vec<f32>,
|
|
}
|
|
|
|
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<f32> {
|
|
// 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<Vec<f32>>,
|
|
k_weights: Vec<Vec<f32>>,
|
|
v_weights: Vec<Vec<f32>>,
|
|
o_weights: Vec<Vec<f32>>,
|
|
// FFN weights
|
|
ffn_w1: Vec<Vec<f32>>,
|
|
ffn_w2: Vec<Vec<f32>>,
|
|
// Layer norm parameters
|
|
ln1_scale: Vec<f32>,
|
|
ln1_bias: Vec<f32>,
|
|
ln2_scale: Vec<f32>,
|
|
ln2_bias: Vec<f32>,
|
|
}
|
|
|
|
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<Vec<f32>> {
|
|
(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<f32>]) -> Vec<Vec<f32>> {
|
|
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<Vec<f32>> = (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<f32>]) -> Vec<Vec<f32>> {
|
|
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<f32> = scores.iter().map(|s| (s - max_score).exp()).collect();
|
|
let sum_exp: f32 = exp_scores.iter().sum();
|
|
let attn_weights: Vec<f32> = 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<f32>]) -> Vec<Vec<f32>> {
|
|
// FFN: hidden -> 4*hidden -> hidden with GELU activation
|
|
let intermediate = self.matmul(hidden, &self.ffn_w1);
|
|
|
|
// GELU activation
|
|
let activated: Vec<Vec<f32>> = 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<f32>], scale: &[f32], bias: &[f32]) -> Vec<Vec<f32>> {
|
|
hidden
|
|
.iter()
|
|
.map(|row| {
|
|
let mean: f32 = row.iter().sum::<f32>() / row.len() as f32;
|
|
let var: f32 =
|
|
row.iter().map(|x| (x - mean).powi(2)).sum::<f32>() / 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<f32>], b: &[Vec<f32>]) -> Vec<Vec<f32>> {
|
|
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<Vec<f32>>,
|
|
}
|
|
|
|
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<Vec<f32>> = (0..hidden_dim)
|
|
.map(|_| (0..hidden_dim).map(|_| normal.sample(&mut rng)).collect())
|
|
.collect();
|
|
|
|
Self {
|
|
hidden_dim,
|
|
proj_weights,
|
|
}
|
|
}
|
|
|
|
fn forward(&self, hidden: &[Vec<f32>]) -> Vec<f32> {
|
|
// 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);
|
|
}
|
|
}
|
|
}
|