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]>
234 lines
7.1 KiB
Rust
234 lines
7.1 KiB
Rust
//! 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<Vec<f32>>,
|
|
/// Positional encoding frequencies
|
|
pub pos_frequencies: Vec<f32>,
|
|
}
|
|
|
|
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<Vec<f32>> = (0..NUM_AMINO_ACIDS)
|
|
.map(|_| {
|
|
(0..EMBEDDING_DIM)
|
|
.map(|_| normal.sample(&mut rng))
|
|
.collect()
|
|
})
|
|
.collect();
|
|
|
|
// Precompute positional encoding frequencies
|
|
let pos_frequencies: Vec<f32> = (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<AminoAcid> =
|
|
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<Vec<f32>>,
|
|
/// 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::<f32>() / emb.len() as f32;
|
|
let var: f32 = emb.iter().map(|x| (x - mean).powi(2)).sum::<f32>() / 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<Vec<Vec<f32>>>,
|
|
/// 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::<f32>() / emb.len() as f32;
|
|
assert!(mean.abs() < 0.01);
|
|
}
|
|
}
|
|
}
|