//! Protein sequence encoder with amino acid embeddings and positional encoding. use alphafold_shared::AminoAcid; use serde::{Deserialize, Serialize}; /// Embedding dimension for amino acids. pub const EMBEDDING_DIM: usize = 256; /// Maximum sequence length supported. pub const MAX_SEQ_LEN: usize = 2500; /// Number of amino acid types (20 standard + 1 unknown). pub const NUM_AMINO_ACIDS: usize = 21; /// Protein sequence encoder. #[derive(Debug, Clone)] pub struct ProteinEncoder { /// Amino acid embedding weights [21 x `EMBEDDING_DIM`] pub aa_embeddings: Vec>, /// Positional encoding frequencies pub pos_frequencies: Vec, } impl Default for ProteinEncoder { fn default() -> Self { Self::new() } } impl ProteinEncoder { /// Create a new protein encoder with initialized embeddings. #[must_use] pub fn new() -> Self { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; let mut rng = rand::rngs::StdRng::seed_from_u64(42); let normal = Normal::new(0.0_f32, 0.02).unwrap(); // Initialize amino acid embeddings let aa_embeddings: Vec> = (0..NUM_AMINO_ACIDS) .map(|_| { (0..EMBEDDING_DIM) .map(|_| normal.sample(&mut rng)) .collect() }) .collect(); // Precompute positional encoding frequencies let pos_frequencies: Vec = (0..EMBEDDING_DIM / 2) .map(|i| 1.0 / 10000_f32.powf(2.0 * i as f32 / EMBEDDING_DIM as f32)) .collect(); Self { aa_embeddings, pos_frequencies, } } /// Encode a sequence of amino acids to embeddings. pub fn encode(&self, sequence: &str) -> SequenceEmbedding { let amino_acids: Vec = sequence.chars().filter_map(AminoAcid::from_code).collect(); let seq_len = amino_acids.len(); let mut embeddings = Vec::with_capacity(seq_len); for (pos, aa) in amino_acids.iter().enumerate() { let aa_idx = aa.embedding_idx(); let aa_emb = &self.aa_embeddings[aa_idx]; // Add positional encoding let mut emb = vec![0.0_f32; EMBEDDING_DIM]; for (i, &freq) in self.pos_frequencies.iter().enumerate() { let pos_f = pos as f32; emb[2 * i] = aa_emb[2 * i] + (pos_f * freq).sin(); emb[2 * i + 1] = aa_emb[2 * i + 1] + (pos_f * freq).cos(); } embeddings.push(emb); } SequenceEmbedding { embeddings, sequence_length: seq_len, embedding_dim: EMBEDDING_DIM, } } /// Encode pairwise residue features. #[must_use] pub fn encode_pairs(&self, seq_len: usize) -> PairEmbedding { let pair_dim = EMBEDDING_DIM / 2; let mut pair_embeddings = vec![vec![vec![0.0_f32; pair_dim]; seq_len]; seq_len]; for i in 0..seq_len { for j in 0..seq_len { // Relative position encoding let rel_pos = (i as i32 - j as i32).clamp(-32, 32); let rel_pos_idx = (rel_pos + 32) as usize; // Simple relative position feature for k in 0..pair_dim { let freq = self.pos_frequencies[k.min(self.pos_frequencies.len() - 1)]; pair_embeddings[i][j][k] = ((rel_pos_idx as f32) * freq).sin(); } } } PairEmbedding { embeddings: pair_embeddings, sequence_length: seq_len, pair_dim, } } } /// Encoded sequence representation. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SequenceEmbedding { /// Embeddings for each residue [`seq_len` x `embedding_dim`] pub embeddings: Vec>, /// Sequence length pub sequence_length: usize, /// Embedding dimension pub embedding_dim: usize, } impl SequenceEmbedding { /// Get embedding for a specific residue. #[must_use] pub fn get(&self, idx: usize) -> Option<&[f32]> { self.embeddings.get(idx).map(std::vec::Vec::as_slice) } /// Apply layer normalization to embeddings. pub fn layer_norm(&mut self, eps: f32) { for emb in &mut self.embeddings { let mean: f32 = emb.iter().sum::() / emb.len() as f32; let var: f32 = emb.iter().map(|x| (x - mean).powi(2)).sum::() / emb.len() as f32; let std = (var + eps).sqrt(); for x in emb.iter_mut() { *x = (*x - mean) / std; } } } } /// Pairwise residue embedding. #[derive(Debug, Clone)] pub struct PairEmbedding { /// Pairwise embeddings [`seq_len` x `seq_len` x `pair_dim`] pub embeddings: Vec>>, /// Sequence length pub sequence_length: usize, /// Pair embedding dimension pub pair_dim: usize, } impl PairEmbedding { /// Get pairwise embedding for residues i and j. #[must_use] pub fn get(&self, i: usize, j: usize) -> Option<&[f32]> { self.embeddings .get(i) .and_then(|row| row.get(j)) .map(std::vec::Vec::as_slice) } /// Update pair embedding with outer product of sequence embeddings. pub fn outer_product_update(&mut self, seq_emb: &SequenceEmbedding) { let seq_len = self.sequence_length; for i in 0..seq_len { for j in 0..seq_len { if let (Some(emb_i), Some(emb_j)) = (seq_emb.get(i), seq_emb.get(j)) { // Simple outer product projection for k in 0..self.pair_dim { let idx_i = k * 2 % seq_emb.embedding_dim; let idx_j = (k * 2 + 1) % seq_emb.embedding_dim; self.embeddings[i][j][k] += emb_i[idx_i] * emb_j[idx_j] * 0.1; } } } } } } #[cfg(test)] mod tests { use super::*; #[test] fn test_encoder_creation() { let encoder = ProteinEncoder::new(); assert_eq!(encoder.aa_embeddings.len(), NUM_AMINO_ACIDS); assert_eq!(encoder.aa_embeddings[0].len(), EMBEDDING_DIM); } #[test] fn test_sequence_encoding() { let encoder = ProteinEncoder::new(); let embedding = encoder.encode("ACDEF"); assert_eq!(embedding.sequence_length, 5); assert_eq!(embedding.embedding_dim, EMBEDDING_DIM); assert_eq!(embedding.embeddings.len(), 5); } #[test] fn test_pair_encoding() { let encoder = ProteinEncoder::new(); let pairs = encoder.encode_pairs(10); assert_eq!(pairs.sequence_length, 10); assert_eq!(pairs.embeddings.len(), 10); assert_eq!(pairs.embeddings[0].len(), 10); } #[test] fn test_layer_norm() { let encoder = ProteinEncoder::new(); let mut embedding = encoder.encode("AAA"); embedding.layer_norm(1e-5); // After layer norm, mean should be ~0 and std ~1 for emb in &embedding.embeddings { let mean: f32 = emb.iter().sum::() / emb.len() as f32; assert!(mean.abs() < 0.01); } } }