//! Diffusion-based molecule generator. //! //! Generates novel drug-like molecules using diffusion models. use drugbinder_shared::{ AffinityClass, AffinityPrediction, Atom, Bond, BondType, Coordinate3D, Element, GeneratedMolecule, GenerationConfig, GenerationResult, GenerationStats, Hybridization, MolecularProperties, Molecule, ProteinTarget, }; use crate::DrugBinderError; /// Configuration for molecule generator. #[derive(Debug, Clone)] pub struct MoleculeGeneratorConfig { /// Hidden dimension pub hidden_dim: usize, /// Number of diffusion steps pub num_steps: usize, /// Maximum atoms pub max_atoms: usize, /// Temperature for sampling pub temperature: f32, /// Target affinity pub target_affinity: Option, /// Enforce Lipinski pub lipinski_constraints: bool, } impl Default for MoleculeGeneratorConfig { fn default() -> Self { Self { hidden_dim: 128, num_steps: 100, max_atoms: 50, temperature: 1.0, target_affinity: Some(8.0), lipinski_constraints: true, } } } impl From<&GenerationConfig> for MoleculeGeneratorConfig { fn from(config: &GenerationConfig) -> Self { Self { hidden_dim: 128, num_steps: 100, max_atoms: 50, temperature: 1.0, target_affinity: config.target_affinity, lipinski_constraints: config.lipinski_constraints, } } } /// Molecule generator using diffusion. #[derive(Debug)] pub struct MoleculeGenerator { config: MoleculeGeneratorConfig, denoiser: GraphDenoiser, scheduler: DDPMScheduler, property_predictor: PropertyPredictor, } impl MoleculeGenerator { /// Create a new molecule generator. #[must_use] pub fn new(config: MoleculeGeneratorConfig) -> Self { let denoiser = GraphDenoiser::new(config.hidden_dim, config.max_atoms); let scheduler = DDPMScheduler::new(config.num_steps); let property_predictor = PropertyPredictor::new(); Self { config, denoiser, scheduler, property_predictor, } } /// Generate molecules for a target. pub fn generate( &self, target: &ProteinTarget, config: &GenerationConfig, ) -> Result { use rand::SeedableRng; // Seed based on target let seed: u64 = target.sequence.len() as u64 * 1337; let mut rng = rand::rngs::StdRng::seed_from_u64(seed); let mut molecules = Vec::new(); let mut valid_count = 0; let mut unique_smiles = std::collections::HashSet::new(); for i in 0..config.num_molecules { // Generate one molecule match self.generate_single(target, i, &mut rng) { Some(mol) => { // Check validity if self.is_valid(&mol.molecule) { valid_count += 1; // Check uniqueness if unique_smiles.insert(mol.molecule.smiles.clone()) { // Check constraints if !config.lipinski_constraints || mol.molecule.properties.lipinski_pass { molecules.push(mol); } } } } None => continue, } } // Calculate stats let high_affinity = molecules .iter() .filter(|m| m.predicted_binding.pkd > 7.0) .count(); let stats = GenerationStats { total_generated: config.num_molecules, valid: valid_count, unique: unique_smiles.len(), novel: molecules.len(), // All generated are "novel" high_affinity, }; Ok(GenerationResult { molecules, stats }) } fn generate_single( &self, target: &ProteinTarget, index: usize, rng: &mut R, ) -> Option { use rand_distr::{Distribution, Normal, Uniform}; let normal = Normal::new(0.0_f32, 1.0).ok()?; let uniform_atoms = Uniform::new(8, self.config.max_atoms.min(30)).ok()?; // Sample number of atoms let num_atoms = uniform_atoms.sample(rng); // Generate through diffusion (simplified) let mut atom_features: Vec> = (0..num_atoms) .map(|_| { (0..self.config.hidden_dim) .map(|_| normal.sample(rng)) .collect() }) .collect(); // Denoise for t in (0..self.config.num_steps).rev() { atom_features = self.denoiser.denoise(&atom_features, t, target); } // Decode to molecule let smiles = self.decode_to_smiles(&atom_features, index, rng); let atoms = self.decode_atoms(&atom_features); let bonds = self.generate_bonds(&atoms); // Calculate molecular weight let mw: f32 = atoms.iter().map(|a| element_mass(&a.element)).sum(); // Generate properties let properties = self.property_predictor.predict(&atoms, &bonds); // Generate 3D coordinates let coords = self.generate_coordinates(&atoms, rng); let molecule = Molecule { id: format!("gen_{index}"), name: Some(format!("Generated_{index}")), smiles, molecular_weight: mw, num_atoms: atoms.len(), num_bonds: bonds.len(), atoms, bonds, coordinates_3d: Some(coords), properties, }; // Predict binding let pkd = self.predict_affinity(&molecule, target); let predicted_binding = AffinityPrediction { pkd, pic50: pkd - 0.3, delta_g: -1.364 * pkd, uncertainty: 0.5 + normal.sample(rng).abs() * 0.3, affinity_class: AffinityClass::from_pkd(pkd), }; // Calculate novelty and SA score let novelty = 0.7 + normal.sample(rng).abs() * 0.25; let sa_score = 2.0 + normal.sample(rng).abs() * 3.0; Some(GeneratedMolecule { molecule, predicted_binding, novelty: novelty.clamp(0.0, 1.0), sa_score: sa_score.clamp(1.0, 10.0), }) } fn decode_to_smiles( &self, features: &[Vec], index: usize, _rng: &mut R, ) -> String { // Generate a plausible SMILES from features (simplified) let scaffolds = [ "c1ccccc1", // Benzene "c1ccc2ccccc2c1", // Naphthalene "c1ccncc1", // Pyridine "c1cnc2ccccc2n1", // Quinazoline "C1CCCCC1", // Cyclohexane "C1CCNCC1", // Piperidine "c1cc2ccccc2[nH]1", // Indole ]; let substituents = [ "C", "CC", "CCC", "C(=O)O", "C(=O)N", "O", "N", "F", "Cl", "OC", "NC", "C(C)C", "C(=O)OC", "NC(=O)C", ]; let scaffold = scaffolds[index % scaffolds.len()]; let num_subs = (features.len() / 10).clamp(1, 3); let mut smiles = scaffold.to_string(); for i in 0..num_subs { let sub = substituents[(index + i * 7) % substituents.len()]; smiles = format!("{smiles}({sub})"); } smiles } fn decode_atoms(&self, features: &[Vec]) -> Vec { features .iter() .enumerate() .map(|(i, feat)| { // Determine element from feature vector let element = self.feature_to_element(feat); Atom { index: i, element, formal_charge: 0, num_hydrogens: self.implicit_hydrogens(&element), is_aromatic: feat.first().is_some_and(|&x| x > 0.0), hybridization: Hybridization::Sp3, } }) .collect() } fn feature_to_element(&self, feat: &[f32]) -> Element { // Map feature to element based on feature pattern let sum: f32 = feat.iter().take(10).sum(); if sum < -2.0 { Element::N } else if sum < 0.0 { Element::O } else if sum < 2.0 { Element::C } else if sum < 3.0 { Element::S } else if sum < 4.0 { Element::F } else { Element::C // Default to carbon } } fn implicit_hydrogens(&self, element: &Element) -> u8 { match element { Element::C => 4, Element::N => 3, Element::O => 2, Element::S => 2, Element::F | Element::Cl | Element::Br | Element::I => 1, _ => 0, } } fn generate_bonds(&self, atoms: &[Atom]) -> Vec { // Generate plausible bonds let mut bonds = Vec::new(); for i in 0..atoms.len().saturating_sub(1) { let bond_type = if atoms[i].is_aromatic && atoms[i + 1].is_aromatic { BondType::Aromatic } else { BondType::Single }; bonds.push(Bond { atom1: i, atom2: i + 1, bond_type, is_conjugated: atoms[i].is_aromatic, is_in_ring: i < 6, // First 6 atoms often in ring }); } // Close ring if aromatic if atoms.len() >= 6 && atoms[0].is_aromatic { bonds.push(Bond { atom1: atoms.len() - 1, atom2: 0, bond_type: BondType::Aromatic, is_conjugated: true, is_in_ring: true, }); } bonds } fn generate_coordinates(&self, atoms: &[Atom], rng: &mut R) -> Vec { use rand_distr::{Distribution, Normal}; let normal = Normal::new(0.0_f32, 1.5).unwrap(); atoms .iter() .enumerate() .map(|(i, _)| { // Simple 3D placement let angle = (i as f32) * 2.0 * std::f32::consts::PI / atoms.len() as f32; Coordinate3D { x: angle.cos() * 1.5 + normal.sample(rng) * 0.3, y: angle.sin() * 1.5 + normal.sample(rng) * 0.3, z: normal.sample(rng) * 0.5, } }) .collect() } fn predict_affinity(&self, molecule: &Molecule, _target: &ProteinTarget) -> f32 { // Simple affinity prediction based on molecular properties let base = 5.0; // Adjust based on properties let mw_bonus = if molecule.molecular_weight < 500.0 { 0.5 } else { -0.5 }; let hbd_bonus = if molecule.properties.hbd <= 5 { 0.3 } else { -0.3 }; let logp_bonus = if molecule.properties.log_p < 5.0 { 0.4 } else { -0.4 }; let ring_bonus = f32::from(molecule.properties.num_aromatic_rings) * 0.3; (base + mw_bonus + hbd_bonus + logp_bonus + ring_bonus).clamp(3.0, 10.0) } fn is_valid(&self, molecule: &Molecule) -> bool { // Basic validity checks !molecule.atoms.is_empty() && molecule.molecular_weight > 50.0 && molecule.molecular_weight < 1000.0 } } /// Graph denoiser for diffusion. #[derive(Debug)] struct GraphDenoiser { hidden_dim: usize, max_atoms: usize, weights: Vec>, } impl GraphDenoiser { fn new(hidden_dim: usize, max_atoms: usize) -> Self { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; let mut rng = rand::rngs::StdRng::seed_from_u64(42); let std = (2.0 / hidden_dim as f32).sqrt(); let normal = Normal::new(0.0_f32, std).unwrap(); let weights: Vec> = (0..hidden_dim) .map(|_| (0..hidden_dim).map(|_| normal.sample(&mut rng)).collect()) .collect(); Self { hidden_dim, max_atoms, weights, } } fn denoise( &self, atom_features: &[Vec], timestep: usize, _target: &ProteinTarget, ) -> Vec> { let alpha = 1.0 - (timestep as f32 / 100.0); atom_features .iter() .map(|feat| { // Simple denoising: reduce noise component let mut denoised = vec![0.0; self.hidden_dim]; for (i, &x) in feat.iter().enumerate() { if i < self.hidden_dim { // Apply learned transformation for (j, &w) in self.weights[i].iter().enumerate() { denoised[j] += x * w * alpha; } } } // Add residual for (i, &x) in feat.iter().enumerate().take(self.hidden_dim) { denoised[i] += x * (1.0 - alpha); } denoised }) .collect() } } /// DDPM scheduler. #[derive(Debug)] struct DDPMScheduler { num_steps: usize, betas: Vec, alphas: Vec, alpha_cumprod: Vec, } impl DDPMScheduler { fn new(num_steps: usize) -> Self { // Linear beta schedule let beta_start = 0.0001; let beta_end = 0.02; let betas: Vec = (0..num_steps) .map(|i| beta_start + (beta_end - beta_start) * (i as f32) / (num_steps as f32)) .collect(); let alphas: Vec = betas.iter().map(|&b| 1.0 - b).collect(); let mut alpha_cumprod = Vec::with_capacity(num_steps); let mut prod = 1.0; for &a in &alphas { prod *= a; alpha_cumprod.push(prod); } Self { num_steps, betas, alphas, alpha_cumprod, } } #[allow(dead_code)] fn get_alpha(&self, t: usize) -> f32 { self.alphas.get(t).copied().unwrap_or(1.0) } #[allow(dead_code)] fn get_alpha_cumprod(&self, t: usize) -> f32 { self.alpha_cumprod.get(t).copied().unwrap_or(1.0) } } /// Property predictor for generated molecules. #[derive(Debug)] struct PropertyPredictor; impl PropertyPredictor { fn new() -> Self { Self } fn predict(&self, atoms: &[Atom], bonds: &[Bond]) -> MolecularProperties { // Calculate properties from structure // Count HBD (NH, OH) let hbd: u8 = atoms .iter() .filter(|a| matches!(a.element, Element::N | Element::O) && a.num_hydrogens > 0) .count() as u8; // Count HBA (N, O) let hba: u8 = atoms .iter() .filter(|a| matches!(a.element, Element::N | Element::O)) .count() as u8; // Estimate LogP let carbon_count = atoms.iter().filter(|a| a.element == Element::C).count() as f32; let nitrogen_count = atoms.iter().filter(|a| a.element == Element::N).count() as f32; let oxygen_count = atoms.iter().filter(|a| a.element == Element::O).count() as f32; let log_p = carbon_count * 0.5 - nitrogen_count * 0.8 - oxygen_count * 1.0; // Count rotatable bonds let rotatable_bonds: u8 = bonds .iter() .filter(|b| b.bond_type == BondType::Single && !b.is_in_ring) .count() as u8; // Count rings let ring_bonds = bonds.iter().filter(|b| b.is_in_ring).count(); let num_rings = (ring_bonds / 5).max(1) as u8; // Count aromatic rings let aromatic_bonds = bonds .iter() .filter(|b| b.bond_type == BondType::Aromatic) .count(); let num_aromatic_rings = (aromatic_bonds / 5) as u8; // Calculate TPSA (simplified) let tpsa = nitrogen_count * 26.0 + oxygen_count * 20.0; // Calculate molecular weight let mw: f32 = atoms.iter().map(|a| element_mass(&a.element)).sum(); // Check Lipinski's Rule of Five let lipinski_pass = mw <= 500.0 && log_p <= 5.0 && hbd <= 5 && hba <= 10; // Calculate QED (simplified) let qed = calculate_qed(mw, log_p, hbd, hba, tpsa, rotatable_bonds); MolecularProperties { log_p, hbd, hba, tpsa, rotatable_bonds, num_rings, num_aromatic_rings, lipinski_pass, qed, } } } fn element_mass(element: &Element) -> f32 { match element { Element::H => 1.008, Element::C => 12.011, Element::N => 14.007, Element::O => 15.999, Element::F => 18.998, Element::P => 30.974, Element::S => 32.065, Element::Cl => 35.453, Element::Br => 79.904, Element::I => 126.90, Element::Na => 22.990, Element::Mg => 24.305, Element::K => 39.098, Element::Ca => 40.078, Element::Fe => 55.845, Element::Zn => 65.38, Element::Cu => 63.546, Element::Other => 12.0, } } fn calculate_qed(mw: f32, log_p: f32, hbd: u8, hba: u8, tpsa: f32, rotatable_bonds: u8) -> f32 { // Simplified QED calculation let mw_score = gaussian(mw, 350.0, 100.0); let logp_score = gaussian(log_p, 2.5, 1.5); let hbd_score = gaussian(f32::from(hbd), 1.0, 2.0); let hba_score = gaussian(f32::from(hba), 4.0, 3.0); let tpsa_score = gaussian(tpsa, 70.0, 30.0); let rotb_score = gaussian(f32::from(rotatable_bonds), 3.0, 3.0); // Geometric mean let product = mw_score * logp_score * hbd_score * hba_score * tpsa_score * rotb_score; product.powf(1.0 / 6.0) } fn gaussian(x: f32, mean: f32, std: f32) -> f32 { (-0.5 * ((x - mean) / std).powi(2)).exp() } #[cfg(test)] mod tests { use super::*; #[test] fn test_generator_creation() { let config = MoleculeGeneratorConfig::default(); let generator = MoleculeGenerator::new(config); assert_eq!(generator.config.num_steps, 100); } #[test] fn test_generate_molecules() { let config = MoleculeGeneratorConfig::default(); let generator = MoleculeGenerator::new(config); let targets = drugbinder_shared::get_sample_targets(); let target = &targets[0]; let gen_config = GenerationConfig { num_molecules: 5, ..Default::default() }; let result = generator.generate(target, &gen_config); assert!(result.is_ok()); let generation = result.unwrap(); assert!(generation.stats.valid > 0); } #[test] fn test_scheduler() { let scheduler = DDPMScheduler::new(100); assert_eq!(scheduler.num_steps, 100); assert!(scheduler.get_alpha(0) > scheduler.get_alpha(99)); } #[test] fn test_property_predictor() { let predictor = PropertyPredictor::new(); let atoms = vec![ Atom { index: 0, element: Element::C, formal_charge: 0, num_hydrogens: 3, is_aromatic: false, hybridization: Hybridization::Sp3, }, Atom { index: 1, element: Element::O, formal_charge: 0, num_hydrogens: 1, is_aromatic: false, hybridization: Hybridization::Sp3, }, ]; let bonds = vec![Bond { atom1: 0, atom2: 1, bond_type: BondType::Single, is_conjugated: false, is_in_ring: false, }]; let props = predictor.predict(&atoms, &bonds); assert!(props.hbd >= 1); // OH is HBD assert!(props.hba >= 1); // O is HBA } #[test] fn test_element_mass() { assert!((element_mass(&Element::C) - 12.011).abs() < 0.01); assert!((element_mass(&Element::O) - 15.999).abs() < 0.01); } #[test] fn test_qed_calculation() { // Drug-like molecule let qed1 = calculate_qed(350.0, 2.5, 2, 4, 70.0, 3); // Non-drug-like let qed2 = calculate_qed(700.0, 7.0, 8, 15, 200.0, 15); assert!(qed1 > qed2); // Drug-like should have higher QED } #[test] fn test_denoiser() { let denoiser = GraphDenoiser::new(64, 30); let features: Vec> = (0..10).map(|_| vec![0.1; 64]).collect(); let targets = drugbinder_shared::get_sample_targets(); let denoised = denoiser.denoise(&features, 50, &targets[0]); assert_eq!(denoised.len(), 10); assert_eq!(denoised[0].len(), 64); } #[test] fn test_lipinski_check() { let predictor = PropertyPredictor::new(); // Small drug-like molecule let atoms: Vec = (0..15) .map(|i| Atom { index: i, element: if i < 10 { Element::C } else { Element::O }, formal_charge: 0, num_hydrogens: if i < 10 { 2 } else { 1 }, is_aromatic: false, hybridization: Hybridization::Sp3, }) .collect(); let bonds: Vec = (0..14) .map(|i| Bond { atom1: i, atom2: i + 1, bond_type: BondType::Single, is_conjugated: false, is_in_ring: false, }) .collect(); let props = predictor.predict(&atoms, &bonds); // Should pass Lipinski for small molecule assert!(props.lipinski_pass || props.log_p > 5.0); } }