Files
rustytorch/demos/rtx-drugbinder-demo/src/lib.rs
T
Omar Sobh d05e28a449 chore (#19)
Co-authored-by: Omar Sobh <[email protected]>
Co-committed-by: Omar Sobh <[email protected]>
2026-05-16 04:46:41 +00:00

343 lines
10 KiB
Rust

//! `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<BindingPrediction, DrugBinderError> {
// 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<ScreeningResult, DrugBinderError> {
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<ScreeningHit> = 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<GenerationResult, DrugBinderError> {
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<Molecule, DrugBinderError> {
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<ProteinTarget, DrugBinderError> {
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);
}
}