Metadata-Version: 2.4
Name: esm_sumo
Version: 0.1.0
Summary: ESM-2 fine-tuning, evaluation, and interpretability package for SUMOylation site prediction.
Author: Şevki Aybars Türel
Requires-Python: >=3.9
Description-Content-Type: text/markdown
Requires-Dist: torch>=2.0.0
Requires-Dist: transformers>=4.30.0
Requires-Dist: scikit-learn>=1.0.0
Requires-Dist: numpy>=1.21.0
Requires-Dist: matplotlib>=3.5.0

This package was developed for the thesis titled **ACCURATE AND INTERPRETABLE PREDICTION OF SUMOYLATION SITES FROM PROTEIN SEQUENCES**.

# ESM-SUMO

`esm_sumo` is a Python package built on top of the ESM-2 protein language model (650M parameters) designed for SUMOylation site prediction, cross-validation benchmarking, and batched *in-silico* deep mutational scanning.

---

## Installation

Install the package directly from source in editable mode:

```bash
git clone https://github.com/AybarsTurel/esm_sumo.git
cd esm_sumo
pip install -e .
```

---

## Quickstart & Core Functions

### 1. Model Loading (`esm_sumo.models`)

**`load_esm_sumo_model(finetuned=True, model_name_or_path=None, device=None)`**

Loads the ESM-2 sequence classification model and its corresponding tokenizer.

- Set `finetuned=True` to load fine-tuned SUMOylation weights (`aybarsturel/esm2-650m-sumoylation`).
- Set `finetuned=False` to load the base/raw ESM-2 architecture (`facebook/esm2_t33_650M_UR50D`).

```python
from esm_sumo import load_esm_sumo_model

# Load default fine-tuned SUMOylation model
model, tokenizer = load_esm_sumo_model(finetuned=True)

# Load raw base ESM-2 model
base_model, tokenizer = load_esm_sumo_model(finetuned=False)
```

Warning: Loading models may take time because of the size of the models.

---

### 2. Dataset Management & Utilities (`esm_sumo.data`)

**`load_default_dataset(name="128mer", pos_filename="pos_train.fasta", neg_filename="neg_train.fasta")`**

Loads one of the pre-packaged benchmark datasets by key name. Available keys:

- `"128mer"` — Standard 128-mer window length
- `"21mer_matched"` — 21-mer window length matched
- `"128mer_full_homology"` — Full homology sequence set
- `"128mer_representative_homology"` — Representative homology-reduced set

```python
from esm_sumo import load_default_dataset


# Load packaged 128-mer dataset
X, y = load_default_dataset("128mer")
```

**`load_fasta_dataset(pos_fasta_path, neg_fasta_path)`**

Parses custom positive and negative FASTA files from disk into NumPy arrays (`X` sequences and `y` binary labels).

```python
from esm_sumo import load_fasta_dataset

X_custom, y_custom = load_fasta_dataset("path/to/pos.fasta", "path/to/neg.fasta")
```

**`ProteinSequenceDataset(sequences, labels=None, tokenizer=None, max_length=130)`**

PyTorch `Dataset` wrapper that tokenizes protein sequences and prepares input tensors (`input_ids`, `attention_mask`, `labels`) for training or inference.

```python
from torch.utils.data import DataLoader
from esm_sumo import ProteinSequenceDataset

dataset = ProteinSequenceDataset(X, y, tokenizer=tokenizer, max_length=128)
dataloader = DataLoader(dataset, batch_size=16, shuffle=True)
```

---

### 3. Training & Evaluation Pipeline (`esm_sumo.pipeline`)

**`train_epoch(model, dataloader, optimizer, device)`**

Executes a single training epoch across the provided PyTorch `DataLoader` and returns the average training loss.

```python
import torch
from torch.utils.data import DataLoader
from esm_sumo import load_esm_sumo_model, train_epoch, ProteinSequenceDataset

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model, tokenizer = load_esm_sumo_model(finetuned=False, device=device)

# Single epoch training
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-5)
loss = train_epoch(model, dataloader, optimizer, device)
print("Epoch Loss:", loss)
```

**`evaluate_model(model, dataloader, device)`**

Evaluates model performance and returns a dictionary of classification metrics:

- Average Loss
- Accuracy
- Matthew's Correlation Coefficient (MCC)
- F1-Score
- ROC-AUC Score (`auc_score`)
- Precision-Recall AUC (`aupr`)

```python
from esm_sumo import evaluate_model

metrics = evaluate_model(model, val_dataloader, device)
print(metrics)
```

**`run_kfold_cv(X, y, n_splits, model_repo_id, tokenizer, batch_sizes, learning_rates, weight_decays, epochs, device, max_length=130, seed=42, csv_path="cv_results_progress.csv")`**

Executes Grid Search K-Fold Cross-Validation across hyperparameter combinations (batch sizes, learning rates, weight decays) for the specified number of epochs. Real-time step-by-step fold progress is automatically logged to a CSV file.

```python
import torch
from esm_sumo import run_kfold_cv, load_default_dataset, load_esm_sumo_model

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
X, y = load_default_dataset("128mer")
_, tokenizer = load_esm_sumo_model(finetuned=False)

# Grid Search 5-Fold CV with 3 epochs per fold
cv_results = run_kfold_cv(
    X=X,
    y=y,
    n_splits=5,
    model_repo_id="facebook/esm2_t33_650M_UR50D",
    tokenizer=tokenizer,
    batch_sizes=[16, 32],
    learning_rates=[3e-5, 5e-5],
    weight_decays=[0.01],
    epochs=3,
    device=device,
    csv_path="cv_results_progress.csv"
)

print("CV Execution Complete. Final Averaged Results:", cv_results)
```

---

### 4. Interpretability & Mutational Scanning (`esm_sumo.interpretability`)

**`run_128mer_mutational_scan(sequence, model, tokenizer, device, target_pos_1based=65, batch_size=32)`**

Performs batched single-point in-silico mutagenesis across a 128-mer protein sequence. Evaluates prediction probability changes (ΔP) for all 20 standard amino acid substitutions at every position.

**`plot_mutational_heatmap(deltaP, x_labels, title, output_pdf_path)`**

Generates publication-quality heatmap plots illustrating mutational impact scores across sequence positions.

**`plot_positional_impact(deltaP, x_labels, title, output_pdf_path)`**

Generates positional impact barplots highlighting sequence regions sensitive to mutation.

```python
import torch
from esm_sumo import (
    load_esm_sumo_model, 
    run_128mer_mutational_scan, 
    plot_mutational_heatmap, 
    plot_positional_impact
)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model, tokenizer = load_esm_sumo_model(finetuned=True, device=device)

# Target 128-mer sequence
seq_128mer = (
    "RNLVSLGISLPDLNINSMLEQRREPWSGESEVKIAKNSDGRECIKGVNTGSSYALGSNAEDKPIKKQLGVSFHLHLSELELFPDERVINGCNQVENFINHSSSVSCLQEMSSSVKTPIFNRNDFDDSS"
)

# Execute mutagenesis scan
results = run_128mer_mutational_scan(
    sequence=seq_128mer,
    model=model,
    tokenizer=tokenizer,
    device=device,
    target_pos_1based=65
)

# Export plots to PDF
plot_mutational_heatmap(
    results['deltaP'], 
    results['x_labels'], 
    title="Mutational Impact Heatmap",
    output_pdf_path="./outputs/heatmap.pdf"
)
```
