Files
rustytorch/demos/rtx-alphafold-demo/src/encoder.rs
T
osobhandClaude Opus 4.6 02d382d5f6 style: apply rustfmt across all crates and demos
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]>
2026-04-12 07:01:58 -07:00

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);
}
}
}