//! `DrugBinder` drug-target binding affinity prediction demo. //! //! This crate implements drug-target binding affinity prediction using //! graph neural networks and molecular diffusion models. //! //! # Architecture //! //! The model consists of: //! - **Molecule Encoder**: Graph neural network for molecular graphs //! - **Protein Encoder**: Sequence/structure encoder for proteins //! - **Binding Predictor**: Cross-attention for binding affinity //! - **Molecule Generator**: Diffusion-based molecular generation //! //! # Example //! //! ```ignore //! use rtx_drugbinder_demo::{predict_binding, BindingPredictionRequest}; //! //! let result = predict_binding(request).await?; //! println!("Predicted pKd: {:.2} ± {:.2}", //! result.affinity.pkd, //! result.affinity.uncertainty); //! ``` pub mod binding_predictor; pub mod molecule_encoder; pub mod molecule_generator; pub mod protein_encoder; pub mod sample_data; use drugbinder_shared::{ BindingPrediction, BindingPredictionRequest, GenerationRequest, GenerationResult, Molecule, MoleculeInput, PredictionConfig, ProteinTarget, ScreeningHit, ScreeningRequest, ScreeningResult, TargetInput, }; use thiserror::Error; /// Errors that can occur during drug binding prediction. #[derive(Debug, Error)] pub enum DrugBinderError { /// Invalid molecule #[error("Invalid molecule: {0}")] InvalidMolecule(String), /// Invalid protein #[error("Invalid protein: {0}")] InvalidProtein(String), /// Model inference error #[error("Inference error: {0}")] InferenceError(String), /// Generation error #[error("Generation error: {0}")] GenerationError(String), } /// Predict binding affinity between molecule and target. pub async fn predict_binding( request: BindingPredictionRequest, ) -> Result { // Parse molecule input let molecule = parse_molecule_input(&request.molecule)?; tracing::info!("Predicting binding for molecule: {}", molecule.smiles); // Parse target input let target = parse_target_input(&request.target)?; tracing::info!("Target protein: {}", target.name); // Encode molecule let mol_encoder = molecule_encoder::MoleculeEncoder::new(molecule_encoder::MoleculeEncoderConfig::default()); let mol_embedding = mol_encoder.encode(&molecule); // Encode protein let prot_encoder = protein_encoder::ProteinEncoder::new(protein_encoder::ProteinEncoderConfig::default()); let prot_embedding = prot_encoder.encode(&target); // Predict binding let predictor = binding_predictor::BindingPredictor::new( binding_predictor::BindingPredictorConfig::default(), ); let prediction = predictor.predict(&mol_embedding, &prot_embedding, &request.config)?; Ok(prediction) } /// Screen multiple molecules against a target. pub async fn screen_molecules( request: ScreeningRequest, ) -> Result { use std::time::Instant; let start = Instant::now(); let target = parse_target_input(&request.target)?; tracing::info!( "Screening {} molecules against {}", request.molecules.len(), target.name ); // Encode protein once let prot_encoder = protein_encoder::ProteinEncoder::new(protein_encoder::ProteinEncoderConfig::default()); let prot_embedding = prot_encoder.encode(&target); // Encode molecules and predict let mol_encoder = molecule_encoder::MoleculeEncoder::new(molecule_encoder::MoleculeEncoderConfig::default()); let predictor = binding_predictor::BindingPredictor::new( binding_predictor::BindingPredictorConfig::default(), ); let config = PredictionConfig { predict_pose: false, ..Default::default() }; let mut hits: Vec = Vec::new(); let mut num_passed_filters = 0; for mol_input in &request.molecules { let Ok(molecule) = parse_molecule_input(mol_input) else { continue; }; // Apply Lipinski filter if request.config.lipinski_filter && !molecule.properties.lipinski_pass { continue; } num_passed_filters += 1; let mol_embedding = mol_encoder.encode(&molecule); let prediction = predictor.predict(&mol_embedding, &prot_embedding, &config)?; // Apply affinity threshold if let Some(min_aff) = request.config.min_affinity && prediction.affinity.pkd < min_aff { continue; } hits.push(ScreeningHit { rank: 0, molecule, prediction, }); } // Sort by specified criteria match request.config.sort_by { drugbinder_shared::ScreeningSortBy::Affinity => { hits.sort_by(|a, b| { b.prediction .affinity .pkd .partial_cmp(&a.prediction.affinity.pkd) .unwrap() }); } drugbinder_shared::ScreeningSortBy::Confidence => { hits.sort_by(|a, b| { b.prediction .confidence .partial_cmp(&a.prediction.confidence) .unwrap() }); } drugbinder_shared::ScreeningSortBy::QED => { hits.sort_by(|a, b| { b.molecule .properties .qed .partial_cmp(&a.molecule.properties.qed) .unwrap() }); } drugbinder_shared::ScreeningSortBy::LipinskiScore => { // Sort by number of Lipinski rules passed (simplified) hits.sort_by(|a, b| { let score_a = lipinski_score(&a.molecule); let score_b = lipinski_score(&b.molecule); score_b.cmp(&score_a) }); } } // Take top k hits.truncate(request.config.top_k); // Assign ranks for (i, hit) in hits.iter_mut().enumerate() { hit.rank = i + 1; } let elapsed = start.elapsed(); Ok(ScreeningResult { hits, total_screened: request.molecules.len(), num_passed_filters, processing_time_ms: elapsed.as_millis() as u64, }) } /// Generate novel molecules for a target. pub async fn generate_molecules( request: GenerationRequest, ) -> Result { let target = parse_target_input(&request.target)?; tracing::info!("Generating molecules for target: {}", target.name); let generator = molecule_generator::MoleculeGenerator::new( molecule_generator::MoleculeGeneratorConfig::from(&request.config), ); let result = generator.generate(&target, &request.config)?; Ok(result) } /// Parse molecule input into a full Molecule. fn parse_molecule_input(input: &MoleculeInput) -> Result { match input { MoleculeInput::Smiles { smiles, name } => { sample_data::smiles_to_molecule(smiles, name.clone()) } MoleculeInput::Full(mol) => Ok(mol.clone()), } } /// Parse target input into a `ProteinTarget`. fn parse_target_input(input: &TargetInput) -> Result { match input { TargetInput::Id { protein_id } => sample_data::get_protein_by_id(protein_id), TargetInput::Sequence { sequence, name } => { Ok(sample_data::sequence_to_protein(sequence, name.clone())) } TargetInput::Full(target) => Ok(target.clone()), } } /// Calculate Lipinski score (number of rules passed). fn lipinski_score(mol: &Molecule) -> u8 { let mut score = 0; if mol.molecular_weight <= 500.0 { score += 1; } if mol.properties.log_p <= 5.0 { score += 1; } if mol.properties.hbd <= 5 { score += 1; } if mol.properties.hba <= 10 { score += 1; } score } #[cfg(test)] mod tests { use super::*; use drugbinder_shared::{GenerationConfig, ScreeningConfig}; #[tokio::test] async fn test_predict_binding() { let request = BindingPredictionRequest { molecule: MoleculeInput::Smiles { smiles: "CC(=O)OC1=CC=CC=C1C(=O)O".to_string(), name: Some("Aspirin".to_string()), }, target: TargetInput::Id { protein_id: "P00918".to_string(), }, config: PredictionConfig::default(), }; let result = predict_binding(request).await; assert!(result.is_ok()); let prediction = result.unwrap(); assert!(prediction.affinity.pkd > 0.0); assert!(prediction.confidence > 0.0); } #[tokio::test] async fn test_screen_molecules() { let molecules = drugbinder_shared::get_sample_molecules(); let request = ScreeningRequest { molecules: molecules .iter() .map(|m| MoleculeInput::Full(m.clone())) .collect(), target: TargetInput::Id { protein_id: "P00918".to_string(), }, config: ScreeningConfig { top_k: 10, min_affinity: None, lipinski_filter: false, ..Default::default() }, }; let result = screen_molecules(request).await; assert!(result.is_ok()); let screening = result.unwrap(); assert!(screening.total_screened > 0); } #[tokio::test] async fn test_generate_molecules() { let request = GenerationRequest { target: TargetInput::Id { protein_id: "P00918".to_string(), }, method: drugbinder_shared::GenerationMethod::Diffusion, config: GenerationConfig { num_molecules: 10, ..Default::default() }, }; let result = generate_molecules(request).await; assert!(result.is_ok()); let generation = result.unwrap(); assert!(!generation.molecules.is_empty()); } #[test] fn test_lipinski_score() { let mol = &drugbinder_shared::get_sample_molecules()[0]; let score = lipinski_score(mol); assert!(score > 0); } }