dti-fusion: MERGED-pretrained checkpoints

LoRA-adapted ESM Cambrian (ESMC-300M) protein encoder + ECFP4 ligand fingerprints, fused via one of three pooling/fusion heads, trained as binary drug-target-interaction classifiers on SPRINT's MERGED benchmark (PubChem BioAssay + BindingDB + ChEMBL, ~94:1 negative:positive class imbalance). See dn-gh/dti-merged-preprocessed for the training data and github.com/danialgharaie/dti-fusion (train_fusion_merged.py) for the full training script and architecture definitions (CrossAttnFusion, ConcatBaseline, AttnPoolBaseline).

Recommended checkpoint

checkpoint_attn_pool_baseline_lora_negratio3.pt — attn_pool_baseline architecture, supcon_lambda=0.0, neg_ratio=3, seed 44. This is the best-performing config found by a 3-architecture × 2-SupCon sweep, confirmed reproducible across 3 seeds (mean test AUPR 0.7601 ± 0.0034; this exact checkpoint scored 0.7647, the highest of the three). It beats SPRINT's own reported MERGED AUPR (0.526) by +0.234.

All checkpoints

Every checkpoint below is a full sweep result — kept for the ablation record, not just the winner. Each .pt file is a dict with keys lora_state, head_state, and (for SupCon runs) proj_state, saved from the best early-stopped epoch (3-epoch smoothed validation AUPR, patience 7, 10-epoch warm-up floor — see script for details).

Checkpoint Architecture SupCon λ Test AUROC Test AUPR Test F1 (calibrated)
checkpoint_attn_pool_baseline_lora_negratio3.pt attn_pool_baseline 0.0 0.7606 0.7647 0.7126
checkpoint_attn_pool_baseline_lora_negratio3_supcon0.5.pt attn_pool_baseline 0.5 0.7509 0.7549 0.7092
checkpoint_concat_baseline_lora_negratio3_supcon0.5.pt concat_baseline 0.5 0.7411 0.7518 0.6973
checkpoint_concat_baseline_lora_negratio3.pt concat_baseline 0.0 0.7384 0.7498 0.6961
checkpoint_cross_attn_lora_negratio3.pt cross_attn 0.0 0.7366 0.7385 0.7030
checkpoint_cross_attn_lora_negratio3_supcon0.5.pt cross_attn 0.5 0.7066 0.7157 0.6849

The attn_pool_baseline/SupCon=0.0 checkpoint here is specifically the seed-44 run out of a 3-seed reproducibility check (seeds 42, 43, 44 scored 0.7567 / 0.7588 / 0.7647 test AUPR respectively) — the other two seeds' checkpoints were not retained (the training script's output filename does not include the seed, so each reseed run overwrote the previous checkpoint on disk; only the metrics were logged for all three). Full per-run result JSONs (result_merged_*.json) are included in this repo for every row above.

Architecture notes

  • attn_pool_baseline wins on this benchmark, which is a different ranking from what a smaller kinase-only benchmark in the same codebase found for the same architecture set — treat architecture choice as dataset-dependent, not a fixed recommendation.
  • supcon_lambda=0.5 (an auxiliary supervised-contrastive loss weight) is a wash across every architecture here — never more than +0.02 AUPR, sometimes negative. Plausible explanation: this benchmark's label is purely binary, a weaker contrastive signal than a continuous-affinity, multi-bin pseudo-label would give.
  • Learning rate is a flat AdamW schedule (1e-3 head / 1e-4 LoRA, no warmup/decay) and has not been independently tuned for this benchmark.

Loading

import torch
from peft import LoraConfig, get_peft_model
# see train_fusion_merged.py for the exact model class definitions
# (CrossAttnFusion / ConcatBaseline / AttnPoolBaseline) and LoRA config
ckpt = torch.load("checkpoint_attn_pool_baseline_lora_negratio3.pt", map_location="cpu")
# ckpt["lora_state"]: LoRA adapter weights for the ESMC-300M encoder
# ckpt["head_state"]: fusion head weights
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Dataset used to train dn-gh/dti-fusion-merged