Co-authored-by: Omar Sobh <[email protected]> Co-committed-by: Omar Sobh <[email protected]>
343 lines
10 KiB
Rust
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);
|
|
}
|
|
}
|