How to use from the
Use from the
PEFT library
Task type is invalid.

PhageContraMLM: Contrastive Learning for Phage Protein Representations

PhageContraMLM is a fine-tuned version of the Rostlab/prot_t5_xl_uniref50 protein language model, trained with Low-Rank Adaptation (LoRA) using a hybrid objective that combines standard Masked Language Modeling (MLM) with a contrastive loss.

The model is built to improve the embedding space for bacteriophage proteins, clustering them by functional group and PHROG family in a zero-shot setting.

Intended Use

PhageContraMLM is intended for researchers in computational biology and virology who need function-aware embeddings for phage protein sequences.

Primary use cases:

  • Zero-shot functional retrieval: querying unknown phage proteins against a database of known PHROG families using cosine similarity or L2 distance.
  • Embedding extraction: generating dense representations of sequences for downstream tasks such as functional annotation or clustering.

How to Use

The model relies on the Hugging Face transformers and peft libraries. The LoRA adapters live inside this repo under runs/protrans_XL_Full_lora_envhog_ContraMLM_v1_1/lora_adapters, so make sure to pass subfolder when loading.

import torch
from transformers import T5Tokenizer, T5ForConditionalGeneration
from peft import PeftModel

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
base_model_name = "Rostlab/prot_t5_xl_uniref50"
adapter_repo = "Amin-Saeidi/PhageContraMLM"
adapter_subfolder = "runs/protrans_XL_Full_lora_envhog_ContraMLM_v1_1/lora_adapters"

# 1. Load tokenizer and base model
tokenizer = T5Tokenizer.from_pretrained(base_model_name, do_lower_case=False)
model = T5ForConditionalGeneration.from_pretrained(base_model_name, torch_dtype=torch.float16)

# 2. Attach LoRA adapters and merge
model = PeftModel.from_pretrained(model, adapter_repo, subfolder=adapter_subfolder)
model = model.merge_and_unload().to(device).eval()

# 3. Prepare sequence (space-separated, rare amino acids replaced)
seq = "M A K K L K I L L L A A S L V S L S P S V F A"
inputs = tokenizer(seq, return_tensors="pt").to(device)

# 4. Extract mean-pooled embeddings
with torch.no_grad():
    outputs = model.encoder(**inputs)
    hidden = outputs.last_hidden_state
    mask = inputs.attention_mask.unsqueeze(-1).to(hidden.dtype)
    pooled_embedding = (hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1.0)

print(pooled_embedding.shape)

From Source

If you want to retrain or modify the model locally, clone the full repo (adapters, data, and scripts included).

First, create a virtual environment in Python 3.11.5:

conda create -n phagecontramlm_env python=3.11.5
conda activate phagecontramlm_env

Clone the repo. You will need git-lfs: for WSL or Linux use sudo apt-get install git-lfs, for Windows either use git bash or get git-lfs from here. Then:

git lfs install
git clone https://huggingface.co/Amin-Saeidi/PhageContraMLM

Install dependencies:

cd PhageContraMLM
pip install -r requirements.txt

Usage (training):

python src/train.py

Check src/train.py for the available config flags (data paths, LoRA rank/alpha, loss weighting) before launching a run.

Training Data and Process

  • Data: envhog_phrog2 dataset.
  • Architecture: ProtT5-XL (encoder-decoder).
  • Fine-tuning method: LoRA (rank=32, alpha=64, target modules: q, k, v, o).
  • Loss function: convex combination of curriculum-aware MLM loss (80%) and temperature-scaled contrastive loss (20%).
  • Positive sampling: the contrastive adjacency matrix is built on-the-fly from a known VISEQ pair graph, pulling positive pairs that share identical VISEQs or cross-VISEQ structural similarities.

Repository Structure & Scripts

This repository contains everything needed to reproduce training and evaluation, organized around src/, data/, and runs/ directories.

src/train.py Main training loop. Implements a custom PairGraphCollator that samples positive protein pairs on-the-fly and builds a dynamic adjacency matrix for the contrastive loss, combined with a curriculum-aware MLM objective.

src/produce_test_data_embeddings.py High-throughput script for generating mean-pooled encoder embeddings. Loads the base ProtT5 model, attaches the best LoRA adapters from your checkpoints, and processes raw FASTA/CSV sequences in batches, saving results as .pkl and .csv files.

src/eval_EmbeddingSpace.py Generates publication-quality plots analyzing the embedding space:

  • t-SNE grids colored by PhrogCat category
  • Pairwise L2 and cosine scatter plots comparing the base model against the fine-tuned PhageContraMLM model

src/eval_PhrogRetrieval.py Zero-shot functional retrieval benchmarking using hnswlib (Hierarchical Navigable Small World graphs):

  • Precision@k (k = 5, 10, 50)
  • Per-family and per-size-bin (rare, medium, common) statistics
  • Functional group confusion matrices and Seaborn clustermaps

Dependencies

torch==2.6.0
transformers==4.37.2
peft==0.10.0
pandas==3.0.1
numpy==2.3.5
matplotlib==3.10.8
seaborn==0.13.2
scikit-learn==1.8.0
hnswlib==0.8.0
safetensors==0.7.0
sentencepiece==0.2.0

Acknowledgments

This work was conducted by Amin SaeidiKelishami during an internship at the Laboratoire Microorganismes: Génome et Environnement (LMGE), under the supervision of Professor Clovis Galiez and Professor Francois ENAULT.

Downloads last month
-
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support