Molecular Property Prediction
Overview
Molecular property prediction involves predicting chemical, physical, or biological properties of molecules from their structure. TorchDrug provides comprehensive support for both classification and regression tasks on molecular graphs.
Available Datasets
Drug Discovery Datasets
Classification Tasks:
- BACE (1,513 molecules): Binary classification for β-secretase inhibition
- BBBP (2,039 molecules): Blood-brain barrier penetration prediction
- HIV (41,127 molecules): Ability to inhibit HIV replication
- Tox21 (7,831 molecules): Toxicity prediction across 12 targets
- ToxCast (8,576 molecules): Toxicology screening
- ClinTox (1,478 molecules): Clinical trial toxicity
- SIDER (1,427 molecules): Drug side effects (27 system organ classes)
- MUV (93,087 molecules): Maximum unbiased validation for virtual screening
Regression Tasks:
- ESOL (1,128 molecules): Water solubility prediction
- FreeSolv (642 molecules): Hydration free energy
- Lipophilicity (4,200 molecules): Octanol/water distribution coefficient
- SAMPL (643 molecules): Solvation free energies
Large-Scale Datasets
- QM7 (7,165 molecules): Quantum mechanical properties
- QM8 (21,786 molecules): Electronic spectra and excited state properties
- QM9 (133,885 molecules): Geometric, energetic, electronic, and thermodynamic properties
- PCQM4M (3,803,453 molecules): Large-scale quantum chemistry dataset
- ZINC250k/2M (250k/2M molecules): Drug-like compounds for generative modeling
Task Types
PropertyPrediction
Standard task for graph-level property prediction supporting both classification and regression.
Key Parameters:
model: Graph representation model (GNN)
task: "node", "edge", or "graph" level prediction
criterion: Loss function ("mse", "bce", "ce")
metric: Evaluation metrics ("mae", "rmse", "auroc", "auprc")
num_mlp_layer: Number of MLP layers for readout
Example Workflow:
import torch
from torchdrug import core, models, tasks, datasets
# Load dataset
dataset = datasets.BBBP("~/molecule-datasets/")
# Define model
model = models.GIN(input_dim=dataset.node_feature_dim,
hidden_dims=[256, 256, 256, 256],
edge_input_dim=dataset.edge_feature_dim,
batch_norm=True, readout="mean")
# Define task
task = tasks.PropertyPrediction(model, task=dataset.tasks,
criterion="bce",
metric=("auprc", "auroc"))
MultipleBinaryClassification
Specialized task for multi-label scenarios where each molecule can have multiple binary labels (e.g., Tox21, SIDER).
Key Features:
- Handles missing labels gracefully
- Computes metrics per label and averaged
- Supports weighted loss for imbalanced datasets
Model Selection
Recommended Models by Task
Small Molecules (< 1000 molecules):
- GIN (Graph Isomorphism Network)
- SchNet (for 3D structures)
Medium Datasets (1k-100k molecules):
- GCN, GAT, or GIN
- NFP (Neural Fingerprint)
- MPNN (Message Passing Neural Network)
Large Datasets (> 100k molecules):
- Pre-trained models with fine-tuning
- InfoGraph or MultiviewContrast for self-supervised pre-training
- GIN with deeper architectures
3D Structure Available:
- SchNet (continuous-filter convolutions)
- GearNet (geometry-aware relational graph)
Feature Engineering
Node Features
TorchDrug automatically extracts atom features:
- Atom type
- Formal charge
- Explicit/implicit hydrogens
- Hybridization
- Aromaticity
- Chirality
Edge Features
Bond features include:
- Bond type (single, double, triple, aromatic)
- Stereochemistry
- Conjugation
- Ring membership
Custom Features
Add custom node/edge features using transforms:
from torchdrug import data, transforms
# Add custom features
transform = transforms.VirtualNode() # Add virtual node
dataset = datasets.BBBP("~/molecule-datasets/",
transform=transform)
Training Workflow
Basic Pipeline
- Load Dataset: Choose appropriate dataset
- Split Data: Use scaffold split for drug discovery
- Define Model: Select GNN architecture
- Create Task: Configure loss and metrics
- Setup Optimizer: Adam typically works well
- Train: Use PyTorch Lightning or custom loop
Data Splitting Strategies
Random Split: Standard train/val/test split
Scaffold Split: Group molecules by Bemis-Murcko scaffolds (recommended for drug discovery)
Stratified Split: Maintain label distribution across splits
Best Practices
- Use scaffold splitting for realistic drug discovery evaluation
- Apply data augmentation (virtual nodes, edges) for small datasets
- Monitor multiple metrics (AUROC, AUPRC for classification; MAE, RMSE for regression)
- Use early stopping based on validation performance
- Consider ensemble methods for critical applications
- Pre-train on large datasets before fine-tuning on small datasets
Common Issues and Solutions
Issue: Poor performance on imbalanced datasets
- Solution: Use weighted loss, focal loss, or over/under-sampling
Issue: Overfitting on small datasets
- Solution: Increase regularization, use simpler models, apply data augmentation, or pre-train on larger datasets
Issue: Large memory consumption
- Solution: Reduce batch size, use gradient accumulation, or implement graph sampling
Issue: Slow training
- Solution: Use GPU acceleration, optimize data loading with multiple workers, or use mixed precision training
1---2name: molecular-property-prediction3description: Molecular property prediction involves predicting chemical, physical, or biological properties of molecules from their structure.4---5# Molecular Property Prediction67## Overview89Molecular property prediction involves predicting chemical, physical, or biological properties of molecules from their structure. TorchDrug provides comprehensive support for both classification and regression tasks on molecular graphs.1011## Available Datasets1213### Drug Discovery Datasets1415**Classification Tasks:**16- **BACE** (1,513 molecules): Binary classification for β-secretase inhibition17- **BBBP** (2,039 molecules): Blood-brain barrier penetration prediction18- **HIV** (41,127 molecules): Ability to inhibit HIV replication19- **Tox21** (7,831 molecules): Toxicity prediction across 12 targets20- **ToxCast** (8,576 molecules): Toxicology screening21- **ClinTox** (1,478 molecules): Clinical trial toxicity22- **SIDER** (1,427 molecules): Drug side effects (27 system organ classes)23- **MUV** (93,087 molecules): Maximum unbiased validation for virtual screening2425**Regression Tasks:**26- **ESOL** (1,128 molecules): Water solubility prediction27- **FreeSolv** (642 molecules): Hydration free energy28- **Lipophilicity** (4,200 molecules): Octanol/water distribution coefficient29- **SAMPL** (643 molecules): Solvation free energies3031### Large-Scale Datasets3233- **QM7** (7,165 molecules): Quantum mechanical properties34- **QM8** (21,786 molecules): Electronic spectra and excited state properties35- **QM9** (133,885 molecules): Geometric, energetic, electronic, and thermodynamic properties36- **PCQM4M** (3,803,453 molecules): Large-scale quantum chemistry dataset37- **ZINC250k/2M** (250k/2M molecules): Drug-like compounds for generative modeling3839## Task Types4041### PropertyPrediction4243Standard task for graph-level property prediction supporting both classification and regression.4445**Key Parameters:**46- `model`: Graph representation model (GNN)47- `task`: "node", "edge", or "graph" level prediction48- `criterion`: Loss function ("mse", "bce", "ce")49- `metric`: Evaluation metrics ("mae", "rmse", "auroc", "auprc")50- `num_mlp_layer`: Number of MLP layers for readout5152**Example Workflow:**53```python54import torch55from torchdrug import core, models, tasks, datasets5657# Load dataset58dataset = datasets.BBBP("~/molecule-datasets/")5960# Define model61model = models.GIN(input_dim=dataset.node_feature_dim,62 hidden_dims=[256, 256, 256, 256],63 edge_input_dim=dataset.edge_feature_dim,64 batch_norm=True, readout="mean")6566# Define task67task = tasks.PropertyPrediction(model, task=dataset.tasks,68 criterion="bce",69 metric=("auprc", "auroc"))70```7172### MultipleBinaryClassification7374Specialized task for multi-label scenarios where each molecule can have multiple binary labels (e.g., Tox21, SIDER).7576**Key Features:**77- Handles missing labels gracefully78- Computes metrics per label and averaged79- Supports weighted loss for imbalanced datasets8081## Model Selection8283### Recommended Models by Task8485**Small Molecules (< 1000 molecules):**86- GIN (Graph Isomorphism Network)87- SchNet (for 3D structures)8889**Medium Datasets (1k-100k molecules):**90- GCN, GAT, or GIN91- NFP (Neural Fingerprint)92- MPNN (Message Passing Neural Network)9394**Large Datasets (> 100k molecules):**95- Pre-trained models with fine-tuning96- InfoGraph or MultiviewContrast for self-supervised pre-training97- GIN with deeper architectures9899**3D Structure Available:**100- SchNet (continuous-filter convolutions)101- GearNet (geometry-aware relational graph)102103## Feature Engineering104105### Node Features106107TorchDrug automatically extracts atom features:108- Atom type109- Formal charge110- Explicit/implicit hydrogens111- Hybridization112- Aromaticity113- Chirality114115### Edge Features116117Bond features include:118- Bond type (single, double, triple, aromatic)119- Stereochemistry120- Conjugation121- Ring membership122123### Custom Features124125Add custom node/edge features using transforms:126```python127from torchdrug import data, transforms128129# Add custom features130transform = transforms.VirtualNode() # Add virtual node131dataset = datasets.BBBP("~/molecule-datasets/",132 transform=transform)133```134135## Training Workflow136137### Basic Pipeline1381391. **Load Dataset**: Choose appropriate dataset1402. **Split Data**: Use scaffold split for drug discovery1413. **Define Model**: Select GNN architecture1424. **Create Task**: Configure loss and metrics1435. **Setup Optimizer**: Adam typically works well1446. **Train**: Use PyTorch Lightning or custom loop145146### Data Splitting Strategies147148**Random Split**: Standard train/val/test split149**Scaffold Split**: Group molecules by Bemis-Murcko scaffolds (recommended for drug discovery)150**Stratified Split**: Maintain label distribution across splits151152### Best Practices153154- Use scaffold splitting for realistic drug discovery evaluation155- Apply data augmentation (virtual nodes, edges) for small datasets156- Monitor multiple metrics (AUROC, AUPRC for classification; MAE, RMSE for regression)157- Use early stopping based on validation performance158- Consider ensemble methods for critical applications159- Pre-train on large datasets before fine-tuning on small datasets160161## Common Issues and Solutions162163**Issue: Poor performance on imbalanced datasets**164- Solution: Use weighted loss, focal loss, or over/under-sampling165166**Issue: Overfitting on small datasets**167- Solution: Increase regularization, use simpler models, apply data augmentation, or pre-train on larger datasets168169**Issue: Large memory consumption**170- Solution: Reduce batch size, use gradient accumulation, or implement graph sampling171172**Issue: Slow training**173- Solution: Use GPU acceleration, optimize data loading with multiple workers, or use mixed precision training