Parameter-Efficient Drug-Drug Interaction (DDI) link prediction on the Hetionet biomedical knowledge graph using PyTorch Geometric (RGCN/GAT) and LoRA-adapted language embeddings.
This repository provides a PyTorch Geometric (PyG) framework for Heterogeneous Link Prediction (specifically Drug-Drug Interactions) on Hetionet. It extracts text features combining compound names, IUPAC names, and SMILES strings, then leverages microsoft/BiomedNLP-PubMedBERT-base-uncased-abstract-fulltext configured with LoraConfig(task_type=TaskType.FEATURE_EXTRACTION) to generate rich initial node representations for downstream Graph Neural Networks.
In this pipeline, compound textual features (f"Compound {name}. IUPAC: {iupac}. SMILES: {smiles}") are processed using a frozen PubMedBERT backbone integrated with Hugging Face PEFT (LoraConfig(task_type=TaskType.FEATURE_EXTRACTION)). This design choice provides key technical advantages:
- Computational & Memory Efficiency: Generating text embeddings offline bypasses the prohibitive GPU memory footprint required for end-to-end backpropagation through a Transformer model alongside PyG message-passing layers.
- Preservation of Biomedical Semantics: Keeping PubMedBERT's pre-trained weights frozen prevents catastrophic forgetting on target link prediction tasks, fully retaining its rich biomedical domain understanding across chemical compound metadata.
- Standardized PEFT Adapter Pipeline: Wrapping the language model with
LoraConfigstandardizes the feature extraction code structure. This allows the framework to seamlessly switch between static zero-shot embedding extraction and parameter-efficient fine-tuning without changing downstream PyG data-loading modules.
.
├── LICENSE
├── README.md
├── configs/
│ ├── common.yml # Global runtime configurations
│ ├── datasets.yml # Graph dataset parameters
│ ├── models.yml # GNN architecture parameters
│ ├── process.yml # SMILES embedding pipeline config
│ ├── processors.yml # Text feature extractor settings
│ ├── train.yml # Main training configuration
│ └── trainers.yml # Training loop settings
├── datasets/
│ └── hetionet/ # Raw and processed graph datasets
├── hetiopeft/
│ ├── __init__.py
│ ├── __main__.py # Main CLI entrypoint
│ ├── datasets/ # Data loaders and dataset definitions
│ ├── models/ # Heterogeneous GNN architectures & decoders
│ ├── utils/ # Dynamic negative sampling and helper functions
│ ├── process.py # Preprocessing script
│ └── train.py # Training execution pipeline
├── runs/
│ └── mlflow.db # SQLite database for MLflow experiment tracking
└── pyproject.toml # Project dependencies managed via uv
Dependencies and virtual environments are managed using uv.
# Clone the repository
git clone https://github.com/NaughtFound/HetioPEFT.git
cd HetioPEFT
# Install project dependencies
uv syncInitial baseline runs suffered from extreme overfitting (~0.999 Train AUC alongside high Validation Loss). The following key structural adjustments were implemented to resolve message-passing leakage and edge memorization:
-
Balanced Edge Splitting Parameters
-
Problem: Default settings (
test_ratio=0.70) left only 15% of edges for training, forcing the model to overfit on a tiny fraction of graph topological data. - Fix: Rebalanced data splits to 70% Train, 15% Validation, and 15% Test.
-
Problem: Default settings (
-
Undirected DDI Handling & Disjoint Edge Masking
-
Problem: Drug interactions are naturally symmetric (
$A \leftrightarrow B$ ), but settingis_undirected=Falseleaked target edges across splits. Furthermore, message passing included target edges directly inedge_index. -
Fix: Set
is_undirected=Trueand addeddisjoint_train_ratio=0.2inT.RandomLinkSplitto ensure target edges are hidden during message passing.
-
Problem: Drug interactions are naturally symmetric (
-
Dynamic Per-Epoch Negative Resampling
- Problem: Static negative sampling caused the GNN to quickly memorize fixed non-existent edge pairs.
-
Fix: Implemented
resample_train_negatives()to generate brand-new random negative samples dynamically at the start of every training epoch.
-
Optimizer Regularization
-
Problem: Excessive weight decay (
0.04) constrained linear projections, while hard targets drove BCE loss logits toward infinity. -
Fix: Reduced
weight_decayto1e-4inAdamWand added dropout (0.4).
-
Problem: Excessive weight decay (
All experiments are driven via module targets (-m) configured dynamically through kaizo YAML files inside configs/.
Extract textual feature representations and compute PEFT embeddings using PubMedBERT:
uv run -m hetiopeft --config configs/process.ymlRun Baseline GNN (Without PEFT Features):
uv run -m hetiopeft --config configs/train.yml --run_name hetionet_without_peftRun Enhanced GNN (With PEFT SMILES Embeddings):
uv run -m hetiopeft --config configs/train.yml --use_peft --with_embeddings --run_name hetionet_with_peftAll training runs and metric trajectories are logged in MLflow. Launch the local UI server:
uv run mlflow ui --backend-store-uri sqlite:///runs/mlflow.db| Model Variant | Train Loss | Train AUC | Val Loss | Val AUC | Val AP | Test Loss | Test AUC | Test AP |
|---|---|---|---|---|---|---|---|---|
| Without PEFT | 0.0123 | 0.9999 | 2.3595 | 0.8639 | 0.8924 | 2.2503 | 0.8589 | 0.8963 |
| With PEFT | 0.0402 | 0.9975 | 0.6441 | 0.9456 | 0.9536 | 0.6695 | 0.9648 | 0.9605 |
- Validation Generalization: Integrating PubMedBERT SMILES embeddings drops Validation Loss from 2.3595 → 0.6441 while driving Val ROC-AUC to 0.9456 (+8.17% improvement) and Val AP to 0.9536 (+6.12% improvement).
- Controlled Training Loss: The training loss on the PEFT model settled naturally around 0.0402 (compared to 0.0123 in the baseline), proving that dynamic negative sampling successfully prevented the model from trivially memorizing edges.
- Hetionet: Himmelstein, D. S., et al. (2017). Systematic integration of biomedical knowledge prioritizes candidate disease genes. eLife.
- PubMedBERT: Gu, Y., et al. (2021). Domain-Specific Language Model Pretraining for Biomedical Natural Language Processing. ACM Transactions on Computing for Healthcare.
- LoRA: Hu, E. J., et al. (2021). LoRA: Low-Rank Adaptation of Large Language Models. ICLR.
- PyG (PyTorch Geometric): Fey, M., & Lenssen, J. E. (2019). Fast Graph Representation Learning with PyTorch Geometric. ICLR Workshop.
- Kaizo:
NaughtFound/kaizo— Declarative configuration parser.