PhageContraMLM / README.md
Amin-Saeidi's picture
Upload folder using huggingface_hub
496bac6 verified
|
Raw
History Blame Contribute Delete
5.69 kB
---
pipeline_tag: feature-extraction
tags:
- biology
- genomics
- bioinformatics
- protein-language-model
- pytorch
- peft
- lora
license: apache-2.0
language:
- en
---
# 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.
```python
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](https://git-scm.com/downloads) or get git-lfs from [here](https://github.com/git-lfs/git-lfs/releases). 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.