Files
rustytorch/demos/rtx-cellatlas-demo/src/cell_transformer.rs
T
osobhandClaude Opus 4.6 02d382d5f6 style: apply rustfmt across all crates and demos
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]>
2026-04-12 07:01:58 -07:00

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);
}
}
}