Initial commit
This commit is contained in:
@@ -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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user