//! Protein encoder for binding affinity prediction. //! //! Encodes protein sequences and structures for use in binding prediction. use drugbinder_shared::ProteinTarget; /// Configuration for protein encoder. #[derive(Debug, Clone)] pub struct ProteinEncoderConfig { /// Amino acid embedding dimension pub aa_dim: usize, /// Hidden dimension pub hidden_dim: usize, /// Number of transformer layers pub num_layers: usize, /// Number of attention heads pub num_heads: usize, /// Maximum sequence length pub max_seq_len: usize, } impl Default for ProteinEncoderConfig { fn default() -> Self { Self { aa_dim: 64, hidden_dim: 128, num_layers: 4, num_heads: 4, max_seq_len: 1024, } } } /// Protein encoder using transformer architecture. #[derive(Debug)] pub struct ProteinEncoder { config: ProteinEncoderConfig, aa_embeddings: AminoAcidEmbeddings, positional_encoding: PositionalEncoding, transformer_layers: Vec, pocket_encoder: PocketEncoder, } impl ProteinEncoder { /// Create a new protein encoder. #[must_use] pub fn new(config: ProteinEncoderConfig) -> Self { let aa_embeddings = AminoAcidEmbeddings::new(config.aa_dim); let positional_encoding = PositionalEncoding::new(config.hidden_dim, config.max_seq_len); let transformer_layers: Vec = (0..config.num_layers) .map(|_| TransformerLayer::new(config.hidden_dim, config.num_heads)) .collect(); let pocket_encoder = PocketEncoder::new(config.hidden_dim); Self { config, aa_embeddings, positional_encoding, transformer_layers, pocket_encoder, } } /// Encode a protein target. #[must_use] pub fn encode(&self, protein: &ProteinTarget) -> ProteinEmbedding { // Encode sequence let seq_embedding = self.encode_sequence(&protein.sequence); // Encode pockets if available let pocket_embeddings = if protein.pockets.is_empty() { None } else { Some(self.pocket_encoder.encode(&protein.pockets, &seq_embedding)) }; ProteinEmbedding { sequence_embedding: seq_embedding.clone(), pocket_embeddings, global_embedding: self.pool_sequence(&seq_embedding), } } fn encode_sequence(&self, sequence: &str) -> Vec> { let truncated: String = sequence.chars().take(self.config.max_seq_len).collect(); // Get amino acid embeddings let mut embeddings = self.aa_embeddings.embed(&truncated); // Project to hidden dimension if needed if self.config.aa_dim != self.config.hidden_dim { embeddings = self.project(&embeddings, self.config.aa_dim, self.config.hidden_dim); } // Add positional encoding for (i, emb) in embeddings.iter_mut().enumerate() { let pos_enc = self.positional_encoding.get(i); for (j, &p) in pos_enc.iter().enumerate() { if j < emb.len() { emb[j] += p; } } } // Apply transformer layers let mut hidden = embeddings; for layer in &self.transformer_layers { hidden = layer.forward(&hidden); } hidden } fn project(&self, embeddings: &[Vec], in_dim: usize, out_dim: usize) -> Vec> { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; // Simple projection (in practice would be learned) let mut rng = rand::rngs::StdRng::seed_from_u64(42); let std = (2.0 / in_dim as f32).sqrt(); let normal = Normal::new(0.0_f32, std).unwrap(); let weights: Vec> = (0..in_dim) .map(|_| (0..out_dim).map(|_| normal.sample(&mut rng)).collect()) .collect(); embeddings .iter() .map(|emb| { let mut output = vec![0.0; out_dim]; for (i, &x) in emb.iter().enumerate() { if i < weights.len() { for (j, &w) in weights[i].iter().enumerate() { output[j] += x * w; } } } output }) .collect() } fn pool_sequence(&self, embeddings: &[Vec]) -> Vec { if embeddings.is_empty() { return vec![0.0; self.config.hidden_dim]; } // Mean pooling let mut pooled = vec![0.0; self.config.hidden_dim]; for emb in embeddings { for (i, &v) in emb.iter().enumerate() { if i < pooled.len() { pooled[i] += v; } } } let n = embeddings.len() as f32; for v in &mut pooled { *v /= n; } pooled } } /// Protein embedding result. #[derive(Debug, Clone)] pub struct ProteinEmbedding { /// Residue-level embeddings pub sequence_embedding: Vec>, /// Pocket embeddings (if pockets defined) pub pocket_embeddings: Option>>, /// Global protein embedding pub global_embedding: Vec, } /// Amino acid embeddings. #[derive(Debug)] struct AminoAcidEmbeddings { dim: usize, embeddings: Vec>, } impl AminoAcidEmbeddings { fn new(dim: usize) -> Self { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; let mut rng = rand::rngs::StdRng::seed_from_u64(42); let std = (2.0 / dim as f32).sqrt(); let normal = Normal::new(0.0_f32, std).unwrap(); // 21 amino acids + unknown let embeddings: Vec> = (0..22) .map(|_| (0..dim).map(|_| normal.sample(&mut rng)).collect()) .collect(); Self { dim, embeddings } } fn embed(&self, sequence: &str) -> Vec> { sequence .chars() .map(|c| { let idx = aa_to_index(c); self.embeddings[idx].clone() }) .collect() } } /// Positional encoding. #[derive(Debug)] struct PositionalEncoding { encoding: Vec>, } impl PositionalEncoding { fn new(dim: usize, max_len: usize) -> Self { let mut encoding = Vec::with_capacity(max_len); for pos in 0..max_len { let mut pe = vec![0.0; dim]; for i in 0..dim / 2 { let freq = 1.0 / (10000.0_f32).powf(2.0 * i as f32 / dim as f32); pe[2 * i] = (pos as f32 * freq).sin(); pe[2 * i + 1] = (pos as f32 * freq).cos(); } encoding.push(pe); } Self { encoding } } fn get(&self, position: usize) -> &[f32] { &self.encoding[position.min(self.encoding.len() - 1)] } } /// Transformer layer. #[derive(Debug)] struct TransformerLayer { hidden_dim: usize, num_heads: usize, head_dim: usize, q_weights: Vec>, k_weights: Vec>, v_weights: Vec>, o_weights: Vec>, ffn_w1: Vec>, ffn_w2: Vec>, } impl TransformerLayer { fn new(hidden_dim: usize, num_heads: usize) -> 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), } } fn forward(&self, hidden: &[Vec]) -> Vec> { let seq_len = hidden.len(); if seq_len == 0 { return vec![]; } // Self-attention let q = self.matmul(hidden, &self.q_weights); let k = self.matmul(hidden, &self.k_weights); let v = self.matmul(hidden, &self.v_weights); let scale = (self.head_dim as f32).sqrt(); let mut attn_output = vec![vec![0.0; self.hidden_dim]; seq_len]; for i in 0..seq_len { 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(); 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 + residual let projected = self.matmul(&attn_output, &self.o_weights); let mut output: Vec> = (0..seq_len) .map(|i| { (0..self.hidden_dim) .map(|j| hidden[i].get(j).unwrap_or(&0.0) + projected[i].get(j).unwrap_or(&0.0)) .collect() }) .collect(); // FFN let ffn_hidden = self.matmul(&output, &self.ffn_w1); let ffn_relu: Vec> = ffn_hidden .iter() .map(|row| row.iter().map(|&x| x.max(0.0)).collect()) .collect(); let ffn_out = self.matmul(&ffn_relu, &self.ffn_w2); // Residual for i in 0..seq_len { for j in 0..self.hidden_dim { output[i][j] += ffn_out[i].get(j).unwrap_or(&0.0); } } output } fn matmul(&self, a: &[Vec], b: &[Vec]) -> Vec> { let m = a.len(); if m == 0 || b.is_empty() { return vec![vec![]; m]; } let n = b[0].len(); let k = a[0].len().min(b.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 { result[i][j] += a[i].get(l).unwrap_or(&0.0) * b[l].get(j).unwrap_or(&0.0); } } } result } } /// Pocket encoder. #[derive(Debug)] struct PocketEncoder { hidden_dim: usize, weights: Vec>, } impl PocketEncoder { 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(); // Spatial features (3D coords + volume + druggability) + residue embedding let input_dim = hidden_dim + 5; let weights: Vec> = (0..input_dim) .map(|_| (0..hidden_dim).map(|_| normal.sample(&mut rng)).collect()) .collect(); Self { hidden_dim, weights, } } fn encode( &self, pockets: &[drugbinder_shared::BindingPocket], sequence_embedding: &[Vec], ) -> Vec> { pockets .iter() .map(|pocket| { // Aggregate residue embeddings for pocket let mut pocket_residue_emb = vec![0.0; self.hidden_dim]; let mut count = 0; for &res_idx in &pocket.residues { if res_idx < sequence_embedding.len() { for (i, &v) in sequence_embedding[res_idx].iter().enumerate() { if i < self.hidden_dim { pocket_residue_emb[i] += v; } } count += 1; } } if count > 0 { for v in &mut pocket_residue_emb { *v /= count as f32; } } // Add spatial features let mut features = pocket_residue_emb; features.push(pocket.center.x / 100.0); features.push(pocket.center.y / 100.0); features.push(pocket.center.z / 100.0); features.push(pocket.volume / 1000.0); features.push(pocket.druggability); // Project let mut output = vec![0.0; self.hidden_dim]; for (i, &x) in features.iter().enumerate() { if i < self.weights.len() { for (j, &w) in self.weights[i].iter().enumerate() { output[j] += x * w; } } } output }) .collect() } } /// Map amino acid to embedding index. fn aa_to_index(aa: char) -> usize { match aa.to_ascii_uppercase() { 'A' => 0, // Ala 'R' => 1, // Arg 'N' => 2, // Asn 'D' => 3, // Asp 'C' => 4, // Cys 'Q' => 5, // Gln 'E' => 6, // Glu 'G' => 7, // Gly 'H' => 8, // His 'I' => 9, // Ile 'L' => 10, // Leu 'K' => 11, // Lys 'M' => 12, // Met 'F' => 13, // Phe 'P' => 14, // Pro 'S' => 15, // Ser 'T' => 16, // Thr 'W' => 17, // Trp 'Y' => 18, // Tyr 'V' => 19, // Val 'X' => 20, // Unknown _ => 21, // Other } } #[cfg(test)] mod tests { use super::*; #[test] fn test_protein_encoder_creation() { let config = ProteinEncoderConfig::default(); let encoder = ProteinEncoder::new(config); assert_eq!(encoder.config.hidden_dim, 128); } #[test] fn test_encode_protein() { let config = ProteinEncoderConfig { hidden_dim: 64, num_layers: 2, ..Default::default() }; let encoder = ProteinEncoder::new(config); let targets = drugbinder_shared::get_sample_targets(); let target = &targets[0]; let embedding = encoder.encode(target); assert_eq!(embedding.global_embedding.len(), 64); assert!(!embedding.sequence_embedding.is_empty()); } #[test] fn test_aa_to_index() { assert_eq!(aa_to_index('A'), 0); assert_eq!(aa_to_index('R'), 1); assert_eq!(aa_to_index('K'), 11); } #[test] fn test_positional_encoding() { let pe = PositionalEncoding::new(64, 100); let enc = pe.get(0); assert_eq!(enc.len(), 64); // Different positions should have different encodings let enc0 = pe.get(0); let enc10 = pe.get(10); let diff: f32 = enc0 .iter() .zip(enc10.iter()) .map(|(a, b)| (a - b).abs()) .sum(); assert!(diff > 0.1); } }