Official PyTorch Implementation of "Structure-Aware Contrastive Learning with Fine-Grained Binding Representations for Drug Discovery", ICASSP 2026.
[2026-01-23] Accepted in ICASSP 2026.
We present SaBAN-DTI, a sequence-based Drug-Target Interaction (DTI) framework that injects structural signals while preserving screening speed and accuracy. Our approach represents proteins using a structure-aware vocabulary that pairs each residue token with a compact descriptor of local geometry. By pretraining a large protein language model on sequences annotated with these descriptors, our encoder learns structural context while operating on plain sequence inputs. Small molecules are encoded with SELFIES to guarantee validity and preserve chemical semantics. Both encoders employ attention-based aggregation that maintains input resolution and produces interpretable importance maps over binding-relevant regions. A contrastive learning objective aligns drug and protein embeddings by drawing together regions corresponding to true binding interfaces.
- Structure-aware protein sequence representation that augments each residue token with a compact local-geometry descriptor, enabling structural context learning from plain sequences
- Attention-based aggregation module that preserves resolution, focuses on binding-relevant regions, and produces interpretable importance maps
- High-performance DTI prediction framework that outperforms existing baselines in both accuracy and speed, supporting large-scale virtual screening for drug discovery
CUDA 11.6+
Ubuntu 18.04 / 20.04 or macOS
# create conda env
conda create -n saban python=3.9
conda activate saban
# install PyTorch
pip install torch==1.13.1+cu116 torchvision==0.14.1+cu116 torchaudio==0.13.1 --extra-index-url https://download.pytorch.org/whl/cu116
# install transformers and core dependencies
pip install transformers==4.30.0
pip install scikit-learn
pip install pandas
pip install numpy
pip install tqdm
pip install swanlab
# install chemistry libraries
pip install rdkit
pip install selfies
# install additional dependencies
pip install setuptools==59.5.0The input CSV should contain the following columns:
Protein: Raw protein sequenceSeq: Structure-aware protein sequence (SaProt vocabulary)SMILES: Drug SMILES notationselfies: Drug SELFIES notationlabel: Binary interaction label (0 or 1)
Example:
Protein,Seq,SMILES,selfies,label
MGLSDGEWQLVLNVWGK...,MpGpLpSpDpGpEpWpQp...,CC(=O)Oc1ccccc1C(=O)O,[C][C][=Branch1][C][=O]...,1usage: train.py [-h] [--prot_encoder_path PROT_ENCODER_PATH]
[--drug_encoder_path DRUG_ENCODER_PATH]
[--token_cache TOKEN_CACHE] [--batch_size BATCH_SIZE]
[--lr LR] [--dropout DROPOUT] [--device DEVICE]
[--save_path SAVE_PATH] [--save_name SAVE_NAME]
[--data_path DATA_PATH] [--num_epochs NUM_EPOCHS]
[--patience PATIENCE] [--k_folds K_FOLDS]
[--stratified] [--seed SEED]python train.py \
--data_path dataset/BindingDB/BindingDB.csv \
--batch_size 64 \
--lr 5e-5 \
--dropout 0.05 \
--num_epochs 200 \
--patience 20 \
--k_folds 5 \
--stratified \
--prot_encoder_path westlake-repl/SaProt_650M_AF2 \
--drug_encoder_path HUBioDataLab/SELFormer \
--token_cache dataset/processed_token \
--save_path checkpoint/Or use the provided shell script:
./run.sh dataset/BindingDB/BindingDB.csv 5Evaluate on DUDE and LIT-PCBA benchmarks:
python evaluate.py \
--checkpoint checkpoint/best_model.ckpt \
--dataset both \
--data-path dataset/ \
--output-dir test_results/ \
--batch-size 32Evaluate on DUDE only:
python evaluate.py \
--checkpoint checkpoint/best_model.ckpt \
--dataset dude \
--data-path dataset/ \
--output-dir test_results/Evaluate on LIT-PCBA only:
python evaluate.py \
--checkpoint checkpoint/best_model.ckpt \
--dataset pcba \
--data-path dataset/ \
--output-dir test_results/Note: The evaluation script automatically handles SMILES to SELFIES conversion and uses cached results for faster subsequent runs.
For single prediction:
from model import DTIModel, TokenEncoder
from transformers import EsmTokenizer, EsmForMaskedLM, AutoTokenizer, AutoModel
import torch
# Load model
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = DTIModel(
prot_dim=1280,
drug_dim=768,
latent_dim=1024,
num_heads=8,
dropout=0.05
).to(device)
checkpoint = torch.load('checkpoint/best_model.ckpt', map_location=device)
model.load_state_dict(checkpoint['model_state_dict'] if 'model_state_dict' in checkpoint else checkpoint)
model.eval()
# Load tokenizers and encoders
prot_tokenizer = EsmTokenizer.from_pretrained("westlake-repl/SaProt_650M_AF2")
drug_tokenizer = AutoTokenizer.from_pretrained("HUBioDataLab/SELFormer")
prot_encoder = EsmForMaskedLM.from_pretrained("westlake-repl/SaProt_650M_AF2")
drug_encoder = AutoModel.from_pretrained("HUBioDataLab/SELFormer")
encoder_model = TokenEncoder(prot_encoder, drug_encoder).to(device)
encoder_model.eval()
# Prepare inputs (structure-aware protein sequence and SELFIES)
protein_seq = "MpGpLpSpDpGpEpWpQpLpVpLpNpVpWpGpKp..." # SaProt format
drug_selfies = "[C][C][=C][C][=C][C][=C][Ring1][=Branch1]"
# Tokenize
prot_inputs = prot_tokenizer([protein_seq], padding=True, truncation=True,
return_tensors='pt', max_length=512)
drug_inputs = drug_tokenizer([drug_selfies], padding=True, truncation=True,
return_tensors='pt', max_length=512)
# Predict
with torch.no_grad():
prot_ids = prot_inputs['input_ids'].to(device)
prot_mask = prot_inputs['attention_mask'].to(device).bool()
drug_ids = drug_inputs['input_ids'].to(device)
drug_mask = drug_inputs['attention_mask'].to(device).bool()
# Encode
prot_embed, drug_embed = encoder_model.encoding(prot_ids, prot_mask, drug_ids, drug_mask)
# Predict interaction
score, _ = model(prot_embed, drug_embed, prot_mask, drug_mask)
probability = torch.sigmoid(score).item()
print(f"Binding probability: {probability:.4f}")Convert SMILES to SELFIES for molecular representation:
import selfies as sf
from rdkit import Chem
def smiles_to_selfies(smiles):
try:
# Canonicalize SMILES first
mol = Chem.MolFromSmiles(smiles)
if mol is not None:
canonical_smiles = Chem.MolToSmiles(mol)
return sf.encoder(canonical_smiles)
else:
return sf.encoder(smiles)
except:
return None
# Example
smiles = "CC(=O)Oc1ccccc1C(=O)O"
selfies = smiles_to_selfies(smiles)
print(selfies) # Output: [C][C][=Branch1][C][=O][O][C][=C][C][=C][C][=C][Ring1][=Branch1][C][=Branch1][C][=O][O]-
Protein Encoder: SaProt-650M (structure-aware protein language model)
- Input: Structure-aware protein sequences with 3Di tokens
- Output: 1280-dimensional embeddings per residue
-
Drug Encoder: SELFormer (SELFIES-based molecular transformer)
- Input: SELFIES molecular representations
- Output: 768-dimensional embeddings per token
-
Attention-Based Aggregation: Multi-head attention pooling
- Protein Projector: Aggregates residue-level features → 1024-dim vector
- Drug Projector: Aggregates token-level features → 1024-dim vector
- Produces interpretable attention weights over binding-relevant regions
-
Bilinear Attention Network (BAN): Fine-grained interaction modeling
- Multi-head bilinear attention between drug and protein embeddings
- Captures local binding interactions at residue-token level
- Output: 1024-dimensional co-embedding
-
Contrastive Learning: CLIP-style objective
- Aligns drug and protein representations in shared embedding space
- Encourages similar embeddings for binding pairs
- Improves generalization to unseen drug-target pairs
-
MLP Classifier: Binary interaction prediction
- Input: BAN co-embedding (1024-dim)
- Output: Binding probability (0-1)
| Parameter | Value | Description |
|---|---|---|
| Protein embedding dim | 1280 | SaProt output dimension |
| Drug embedding dim | 768 | SELFormer output dimension |
| Latent dimension | 1024 | Projection and co-embedding dimension |
| Number of attention heads | 8 | For BAN and pooling |
| Dropout rate | 0.05 | Regularization |
| Learning rate | 5e-5 | AdamW optimizer |
| Batch size | 64 | Training batch size |
| Max sequence length | 512 | For both protein and drug |
SaBAN-DTI/
├── model.py # Core DTI model implementation
├── modules.py # Attention modules (BAN, Projectors, CLIP)
├── dataset.py # Data processing and loading utilities
├── train.py # Training script with k-fold cross-validation
├── evaluate.py # Evaluation on DUDE and LIT-PCBA benchmarks
├── run.sh # Shell script for easy training
├── dataset/ # Dataset directory
│ ├── BindingDB/ # BindingDB dataset
│ ├── DAVIS/ # DAVIS dataset
│ ├── Human/ # Human dataset
│ ├── DUDE/ # DUDE benchmark
│ └── LIT-PCBA/ # LIT-PCBA benchmark
└── README.md # This file
If our paper or code is helpful to you, please cite the following:
@INPROCEEDINGS{lan2025structureawarecontrastivelearningfinegrained,
title={Structure-Aware Contrastive Learning with Fine-Grained Binding Representations for Drug Discovery},
booktitle={ICASSP 2026 - 2026 IEEE International Conference on Acoustics, Speech and Signal Processing (ICASSP)},
author={Jing Lan and Hexiao Ding and Hongzhao Chen and Yufeng Jiang and Nga-Chun Ng and Gwing Kei Yip and Gerald W. Y. Cheng and Yunlin Mao and Jing Cai and Liang-ting Lin and Jung Sun Yoo},
year={2026},
url={https://arxiv.org/abs/2509.14788}
}This project is licensed under the Apache License 2.0.
We thank the authors of the following works for making their models and code publicly available:
- SaProt: Structure-aware protein language model (Paper)
- SELFormer: SELFIES-based molecular transformer (HuggingFace)
- SELFIES: Self-referencing embedded strings for molecular representations (Paper)
