Initial commit
This commit is contained in:
@@ -0,0 +1,339 @@
|
||||
//! `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::*;
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user