//! Structure module for 3D coordinate prediction. //! //! Implements SE(3)-equivariant attention (Invariant Point Attention) //! for predicting backbone atom coordinates. use crate::encoder::SequenceEmbedding; use alphafold_shared::{AtomCoord, AtomName}; /// Number of IPA iterations. pub const NUM_IPA_ITERATIONS: usize = 8; /// Backbone frame representation. #[derive(Debug, Clone, Copy)] pub struct BackboneFrame { /// Rotation matrix (3x3) pub rotation: [[f32; 3]; 3], /// Translation vector pub translation: [f32; 3], } impl Default for BackboneFrame { fn default() -> Self { Self { rotation: [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]], translation: [0.0, 0.0, 0.0], } } } impl BackboneFrame { /// Apply frame transformation to a point. #[must_use] pub fn apply(&self, point: [f32; 3]) -> [f32; 3] { let rotated = [ self.rotation[0][0] * point[0] + self.rotation[0][1] * point[1] + self.rotation[0][2] * point[2], self.rotation[1][0] * point[0] + self.rotation[1][1] * point[1] + self.rotation[1][2] * point[2], self.rotation[2][0] * point[0] + self.rotation[2][1] * point[1] + self.rotation[2][2] * point[2], ]; [ rotated[0] + self.translation[0], rotated[1] + self.translation[1], rotated[2] + self.translation[2], ] } /// Compose two frames: self * other. #[must_use] pub fn compose(&self, other: &BackboneFrame) -> BackboneFrame { // R_new = R_self * R_other let mut rotation = [[0.0_f32; 3]; 3]; for i in 0..3 { for j in 0..3 { for k in 0..3 { rotation[i][j] += self.rotation[i][k] * other.rotation[k][j]; } } } // t_new = R_self * t_other + t_self let translation = self.apply(other.translation); BackboneFrame { rotation, translation, } } /// Create frame from axis-angle rotation. #[must_use] pub fn from_axis_angle(axis: [f32; 3], angle: f32) -> Self { let c = angle.cos(); let s = angle.sin(); let t = 1.0 - c; let rotation = [ [ t * axis[0] * axis[0] + c, t * axis[0] * axis[1] - s * axis[2], t * axis[0] * axis[2] + s * axis[1], ], [ t * axis[0] * axis[1] + s * axis[2], t * axis[1] * axis[1] + c, t * axis[1] * axis[2] - s * axis[0], ], [ t * axis[0] * axis[2] - s * axis[1], t * axis[1] * axis[2] + s * axis[0], t * axis[2] * axis[2] + c, ], ]; BackboneFrame { rotation, translation: [0.0, 0.0, 0.0], } } } /// Invariant Point Attention module. #[derive(Debug, Clone)] pub struct InvariantPointAttention { /// Number of query points num_query_points: usize, /// Number of attention heads num_heads: usize, /// Head dimension head_dim: usize, } impl Default for InvariantPointAttention { fn default() -> Self { Self::new(4, 8, 16) } } impl InvariantPointAttention { /// Create a new IPA module. #[must_use] pub fn new(num_query_points: usize, num_heads: usize, head_dim: usize) -> Self { Self { num_query_points, num_heads, head_dim, } } /// Apply IPA to update sequence embeddings and frames. pub fn forward( &self, seq_emb: &SequenceEmbedding, frames: &mut [BackboneFrame], ) -> Vec> { let seq_len = seq_emb.sequence_length; let mut outputs = vec![vec![0.0_f32; seq_emb.embedding_dim]; seq_len]; // Simplified IPA: compute attention based on spatial distance for i in 0..seq_len { for j in 0..seq_len { // Compute spatial distance between frames let dist = distance(&frames[i].translation, &frames[j].translation); // Distance-based attention (simplified) let attention = (-dist / 10.0).exp(); // Accumulate weighted embeddings if let Some(emb_j) = seq_emb.get(j) { for (k, &v) in emb_j.iter().enumerate() { outputs[i][k] += attention * v; } } } // Normalize let norm: f32 = outputs[i].iter().map(|x| x * x).sum::().sqrt() + 1e-6; for x in &mut outputs[i] { *x /= norm; } } outputs } } /// Structure module that predicts 3D coordinates. #[derive(Debug, Clone)] pub struct StructureModule { /// IPA layers ipa_layers: Vec, /// Number of recycles num_recycles: usize, } impl StructureModule { /// Create a new structure module. #[must_use] pub fn new(num_layers: usize, num_recycles: usize) -> Self { let ipa_layers = (0..num_layers) .map(|_| InvariantPointAttention::default()) .collect(); Self { ipa_layers, num_recycles, } } /// Predict 3D structure from sequence embeddings. #[must_use] pub fn predict( &self, seq_emb: &SequenceEmbedding, sequence: &str, ) -> (Vec, Vec) { let seq_len = seq_emb.sequence_length; // Initialize frames (identity) let mut frames: Vec = (0..seq_len).map(|_| BackboneFrame::default()).collect(); // Initialize frame positions along a line for (i, frame) in frames.iter_mut().enumerate() { frame.translation = [i as f32 * 3.8, 0.0, 0.0]; // ~3.8Å per residue } // Recycle through IPA layers for _ in 0..self.num_recycles { for ipa in &self.ipa_layers { let updates = ipa.forward(seq_emb, &mut frames); // Update frames based on IPA output for (i, update) in updates.iter().enumerate() { // Convert update to frame adjustment (simplified) let angle_x = update.first().copied().unwrap_or(0.0) * 0.1; let angle_y = update.get(1).copied().unwrap_or(0.0) * 0.1; let delta_frame = BackboneFrame::from_axis_angle([1.0, 0.0, 0.0], angle_x); let delta_frame2 = BackboneFrame::from_axis_angle([0.0, 1.0, 0.0], angle_y); frames[i] = frames[i].compose(&delta_frame).compose(&delta_frame2); } } } // Convert frames to atom coordinates self.frames_to_atoms(&frames, sequence) } /// Convert backbone frames to atom coordinates. fn frames_to_atoms( &self, frames: &[BackboneFrame], sequence: &str, ) -> (Vec, Vec) { // Standard backbone atom positions in local frame (Angstroms) let n_local = [-1.458, 0.0, 0.0]; let ca_local = [0.0, 0.0, 0.0]; let c_local = [1.523, 0.0, 0.0]; let o_local = [2.0, 1.0, 0.0]; let cb_local = [-0.5, 1.5, 0.0]; let mut atoms = Vec::with_capacity(frames.len() * 5); let mut plddt_scores = Vec::with_capacity(frames.len()); for (i, (frame, residue)) in frames.iter().zip(sequence.chars()).enumerate() { // Generate pLDDT based on frame stability (simplified) let stability = frame.rotation[0][0] + frame.rotation[1][1] + frame.rotation[2][2]; let plddt = (50.0 + stability * 15.0).clamp(20.0, 100.0); plddt_scores.push(plddt); // Transform backbone atoms let n_pos = frame.apply(n_local); atoms.push(AtomCoord { residue_idx: i, atom_name: AtomName::N, x: n_pos[0], y: n_pos[1], z: n_pos[2], b_factor: plddt, }); let ca_pos = frame.apply(ca_local); atoms.push(AtomCoord { residue_idx: i, atom_name: AtomName::Ca, x: ca_pos[0], y: ca_pos[1], z: ca_pos[2], b_factor: plddt, }); let c_pos = frame.apply(c_local); atoms.push(AtomCoord { residue_idx: i, atom_name: AtomName::C, x: c_pos[0], y: c_pos[1], z: c_pos[2], b_factor: plddt, }); let o_pos = frame.apply(o_local); atoms.push(AtomCoord { residue_idx: i, atom_name: AtomName::O, x: o_pos[0], y: o_pos[1], z: o_pos[2], b_factor: plddt, }); // CB for non-glycine if residue != 'G' { let cb_pos = frame.apply(cb_local); atoms.push(AtomCoord { residue_idx: i, atom_name: AtomName::Cb, x: cb_pos[0], y: cb_pos[1], z: cb_pos[2], b_factor: plddt, }); } } (atoms, plddt_scores) } } /// Compute Euclidean distance between two points. fn distance(a: &[f32; 3], b: &[f32; 3]) -> f32 { let dx = a[0] - b[0]; let dy = a[1] - b[1]; let dz = a[2] - b[2]; (dx * dx + dy * dy + dz * dz).sqrt() } #[cfg(test)] mod tests { use super::*; #[test] fn test_backbone_frame_default() { let frame = BackboneFrame::default(); let point = [1.0, 2.0, 3.0]; let transformed = frame.apply(point); assert_eq!(transformed, point); // Identity transform } #[test] fn test_frame_translation() { let mut frame = BackboneFrame::default(); frame.translation = [1.0, 2.0, 3.0]; let point = [0.0, 0.0, 0.0]; let transformed = frame.apply(point); assert_eq!(transformed, [1.0, 2.0, 3.0]); } #[test] fn test_frame_composition() { let frame1 = BackboneFrame { rotation: [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]], translation: [1.0, 0.0, 0.0], }; let frame2 = BackboneFrame { rotation: [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]], translation: [0.0, 1.0, 0.0], }; let composed = frame1.compose(&frame2); assert!((composed.translation[0] - 1.0).abs() < 0.001); assert!((composed.translation[1] - 1.0).abs() < 0.001); } #[test] fn test_structure_module() { use crate::encoder::ProteinEncoder; let encoder = ProteinEncoder::new(); let seq_emb = encoder.encode("ACDEF"); let structure_module = StructureModule::new(2, 1); let (atoms, plddt) = structure_module.predict(&seq_emb, "ACDEF"); assert!(!atoms.is_empty()); assert_eq!(plddt.len(), 5); } #[test] fn test_ipa() { use crate::encoder::ProteinEncoder; let encoder = ProteinEncoder::new(); let seq_emb = encoder.encode("ACDEF"); let mut frames: Vec = (0..5).map(|_| BackboneFrame::default()).collect(); let ipa = InvariantPointAttention::default(); let output = ipa.forward(&seq_emb, &mut frames); assert_eq!(output.len(), 5); } }