Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -0,0 +1,533 @@
//! 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<TransformerLayer>,
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<TransformerLayer> = (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<Vec<f32>> {
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<f32>], in_dim: usize, out_dim: usize) -> Vec<Vec<f32>> {
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<Vec<f32>> = (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<f32>]) -> Vec<f32> {
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<Vec<f32>>,
/// Pocket embeddings (if pockets defined)
pub pocket_embeddings: Option<Vec<Vec<f32>>>,
/// Global protein embedding
pub global_embedding: Vec<f32>,
}
/// Amino acid embeddings.
#[derive(Debug)]
struct AminoAcidEmbeddings {
dim: usize,
embeddings: Vec<Vec<f32>>,
}
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<Vec<f32>> = (0..22)
.map(|_| (0..dim).map(|_| normal.sample(&mut rng)).collect())
.collect();
Self { dim, embeddings }
}
fn embed(&self, sequence: &str) -> Vec<Vec<f32>> {
sequence
.chars()
.map(|c| {
let idx = aa_to_index(c);
self.embeddings[idx].clone()
})
.collect()
}
}
/// Positional encoding.
#[derive(Debug)]
struct PositionalEncoding {
encoding: Vec<Vec<f32>>,
}
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<Vec<f32>>,
k_weights: Vec<Vec<f32>>,
v_weights: Vec<Vec<f32>>,
o_weights: Vec<Vec<f32>>,
ffn_w1: Vec<Vec<f32>>,
ffn_w2: Vec<Vec<f32>>,
}
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<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),
}
}
fn forward(&self, hidden: &[Vec<f32>]) -> Vec<Vec<f32>> {
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<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();
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<Vec<f32>> = (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<Vec<f32>> = 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<f32>], b: &[Vec<f32>]) -> Vec<Vec<f32>> {
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<Vec<f32>>,
}
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<Vec<f32>> = (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<f32>],
) -> Vec<Vec<f32>> {
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);
}
}