Skip to content
1anjPublic

About

Structure-Aware Contrastive Learning with Fine-Grained Binding Representations for Drug Discovery (ICASSP 2026)

Resources

Stars

0 stars

Watchers

0 watching

Forks

Repository files navigation

SaBAN-DTI

Official PyTorch Implementation of "Structure-Aware Contrastive Learning with Fine-Grained Binding Representations for Drug Discovery", ICASSP 2026.

arXiv

Framework

News

[2026-01-23] Accepted in ICASSP 2026.

Abstract

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.

Key Contributions

  1. Structure-aware protein sequence representation that augments each residue token with a compact local-geometry descriptor, enabling structural context learning from plain sequences
  2. Attention-based aggregation module that preserves resolution, focuses on binding-relevant regions, and produces interpretable importance maps
  3. High-performance DTI prediction framework that outperforms existing baselines in both accuracy and speed, supporting large-scale virtual screening for drug discovery

Environments

1. GPU environment

CUDA 11.6+

Ubuntu 18.04 / 20.04 or macOS

2. Create conda environment

# 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.0

Training on Downstream Tasks

Dataset Format

The input CSV should contain the following columns:

  • Protein: Raw protein sequence
  • Seq: Structure-aware protein sequence (SaProt vocabulary)
  • SMILES: Drug SMILES notation
  • selfies: Drug SELFIES notation
  • label: Binary interaction label (0 or 1)

Example:

Protein,Seq,SMILES,selfies,label
MGLSDGEWQLVLNVWGK...,MpGpLpSpDpGpEpWpQp...,CC(=O)Oc1ccccc1C(=O)O,[C][C][=Branch1][C][=O]...,1

Usage

usage: 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]

Run command for training with k-fold cross-validation:

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 5

Evaluation on Benchmark Datasets

DUDE and LIT-PCBA Benchmarks

Evaluate 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 32

Evaluate 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.

Inference

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}")

SMILES to SELFIES Conversion

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]

Model Architecture

Components

  1. Protein Encoder: SaProt-650M (structure-aware protein language model)

    • Input: Structure-aware protein sequences with 3Di tokens
    • Output: 1280-dimensional embeddings per residue
  2. Drug Encoder: SELFormer (SELFIES-based molecular transformer)

    • Input: SELFIES molecular representations
    • Output: 768-dimensional embeddings per token
  3. 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
  4. 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
  5. 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
  6. MLP Classifier: Binary interaction prediction

    • Input: BAN co-embedding (1024-dim)
    • Output: Binding probability (0-1)

Key Hyperparameters

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

Repository Structure

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

Citation

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}
}

License

This project is licensed under the Apache License 2.0.

Acknowledgments

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)

About

Structure-Aware Contrastive Learning with Fine-Grained Binding Representations for Drug Discovery (ICASSP 2026)

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages