# How to Fine-Tune ESMC Models on Custom Protein Datasets: A Complete Guide

> Master fine-tuning ESMC models on custom protein datasets. This guide details loading checkpoints, tokenizing sequences, and optimizing loss for advanced protein analysis.

- Repository: [Biohub/esm](https://github.com/Biohub/esm)
- Tags: how-to-guide
- Published: 2026-05-30

---

**Fine-tuning ESMC models on custom protein datasets requires loading a pretrained checkpoint via `ESMC.from_pretrained()`, tokenizing FASTA sequences with `EsmSequenceTokenizer`, and optimizing cross-entropy loss against the sequence logits produced by the `RegressionHead` while iterating through a standard PyTorch training loop.**

The [Biohub/esm](https://github.com/Biohub/esm) repository implements ESM-C (Evolutionary Scale Modeling – Contrastive) as a modular transformer-based protein language model. Because the `ESMC` class inherits from both `nn.Module` and `ESMCInferenceClient`, it supports both inference and fine-tuning on custom datasets without architectural modifications.

## Understanding the ESM-C Architecture

Before fine-tuning, it helps to understand the three core components defined in the source code. The model architecture in [`esm/models/esmc.py`](https://github.com/Biohub/esm/blob/main/esm/models/esmc.py) consists of:

- **Embedding layer** – Maps a 64-token alphabet (amino acids plus special tokens such as `<cls>`, `<sep>`, and `<mask>`) to a configurable hidden dimension (`d_model`).
- **Transformer stack** – Implemented in [`esm/layers/transformer_stack.py`](https://github.com/Biohub/esm/blob/main/esm/layers/transformer_stack.py), this stacks `n_layers` of multi-head self-attention blocks. When available, setting `use_flash_attn=True` switches to memory-efficient flash-attention kernels without changing the model logic.
- **Regression head** – Located in [`esm/layers/regression_head.py`](https://github.com/Biohub/esm/blob/main/esm/layers/regression_head.py), this linear layer projects final hidden states back to the 64-token vocabulary, producing **sequence logits** for next-token prediction.

The `forward()` method (line 123 in [`esm/models/esmc.py`](https://github.com/Biohub/esm/blob/main/esm/models/esmc.py)) returns an `ESMCOutput` dataclass containing `sequence_logits`, optional `embeddings`, and attention tensors, making it compatible with standard PyTorch loss functions.

## Prerequisites and Optional Optimizations

While the base model runs on CPU, fine-tuning benefits significantly from GPU acceleration. Install optional dependencies to enable performance optimizations:

```bash

# Optional: Install flash-attention for memory-efficient training

pip install flash-attn

```

Flash attention reduces memory overhead during the self-attention computation in `TransformerStack`, allowing larger batch sizes or longer sequences on the same hardware.

## Step-by-Step Fine-Tuning Workflow

Fine-tuning follows the standard PyTorch supervised learning pattern. The high-level steps involve:

1. **Load a pretrained checkpoint** – `ESMC.from_pretrained()` calls `load_local_model` from [`esm/pretrained.py`](https://github.com/Biohub/esm/blob/main/esm/pretrained.py) to fetch weights (e.g., `ESMC_600M`) and initialize the model on the appropriate device.
2. **Tokenize raw sequences** – `EsmSequenceTokenizer.encode()` converts FASTA strings into padded `torch.LongTensor` objects, handling special token insertion automatically.
3. **Create a PyTorch Dataset** – Return tuples of `(input_ids, target_ids)` where targets are the next-token predictions shifted by one position.
4. **Define loss and optimizer** – Use `nn.CrossEntropyLoss(ignore_index=tokenizer.pad_id)` to ignore padding tokens, paired with `AdamW` and a learning rate scheduler.
5. **Execute the training loop** – Call `model(input_ids)` to obtain logits, compute loss, and run `backward()` and `optimizer.step()`.
6. **Save the fine-tuned weights** – Use `torch.save(model.state_dict(), path)` to export the adapted model for downstream inference.

## Complete Fine-Tuning Implementation

The following end-to-end example demonstrates how to fine-tune ESMC on a custom FASTA file. This script handles tokenization, batching, and the training loop while respecting the model's expected input format.

```python

# ----------------------------------------------------------------------

# 1️⃣ Imports

# ----------------------------------------------------------------------

import torch
from torch import nn
from torch.utils.data import Dataset, DataLoader
from esm.models.esmc import ESMC
from esm.tokenization.sequence_tokenizer import EsmSequenceTokenizer
from esm.utils.sampling import _BatchedESMProteinTensor

# ----------------------------------------------------------------------

# 2️⃣ Custom Dataset (FASTA → token IDs)

# ----------------------------------------------------------------------

class ProteinFastaDataset(Dataset):
    def __init__(self, fasta_paths, tokenizer):
        self.seqs = []
        for path in fasta_paths:
            with open(path) as f:
                seq = ""
                for line in f:
                    if line.startswith(">"):
                        continue
                    seq += line.strip()
                self.seqs.append(seq)
        self.tokenizer = tokenizer

    def __len__(self):
        return len(self.seqs)

    def __getitem__(self, idx):
        seq = self.seqs[idx]
        token_ids = self.tokenizer.encode(seq, add_special_tokens=True)
        input_ids = token_ids[:-1]          # all but last token

        target_ids = token_ids[1:]          # next-token supervision

        return torch.tensor(input_ids, dtype=torch.long), torch.tensor(target_ids, dtype=torch.long)

# ----------------------------------------------------------------------

# 3️⃣ Instantiate tokenizer and model

# ----------------------------------------------------------------------

tokenizer = EsmSequenceTokenizer()
model = ESMC.from_pretrained(use_flash_attn=True)  # loads ESMC_600M by default

model.train()
model = model.to(torch.device("cuda" if torch.cuda.is_available() else "cpu"))

# ----------------------------------------------------------------------

# 4️⃣ DataLoader with padding

# ----------------------------------------------------------------------

train_dataset = ProteinFastaDataset(["./data/train.fasta"], tokenizer)
train_loader = DataLoader(
    train_dataset, 
    batch_size=8, 
    shuffle=True, 
    collate_fn=lambda b: torch.nn.utils.rnn.pad_sequence(b, batch_first=True, padding_value=tokenizer.pad_id)
)

# ----------------------------------------------------------------------

# 5️⃣ Optimizer & loss

# ----------------------------------------------------------------------

optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01)
criterion = nn.CrossEntropyLoss(ignore_index=tokenizer.pad_id)

# ----------------------------------------------------------------------

# 6️⃣ Training loop

# ----------------------------------------------------------------------

for epoch in range(5):
    total_loss = 0.0
    for input_ids, target_ids in train_loader:
        input_ids = input_ids.to(model.device)
        target_ids = target_ids.to(model.device)

        # Forward pass returns ESMCOutput

        out = model(input_ids)
        logits = out.sequence_logits  # shape (B, L, vocab)

        # Compute cross-entropy

        loss = criterion(logits.view(-1, logits.size(-1)), target_ids.view(-1))

        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        total_loss += loss.item()
    print(f"Epoch {epoch+1}: avg loss = {total_loss/len(train_loader):.4f}")

# ----------------------------------------------------------------------

# 7️⃣ Save fine-tuned checkpoint

# ----------------------------------------------------------------------

torch.save(model.state_dict(), "esmc_finetuned.pt")

```

## Inference After Fine-Tuning

Once fine-tuning completes, you can decode generated sequences using the `decode` method inherited from `ESMCInferenceClient` defined in [`esm/sdk/api.py`](https://github.com/Biohub/esm/blob/main/esm/sdk/api.py):

```python
model.eval()

# Example: decode a random sample (replace with actual generated tokens)

sample = torch.randint(high=tokenizer.vocab_size, size=(1, 128), device=model.device)
protein_tensor = _BatchedESMProteinTensor(sample, tokenizer)
decoded = model.decode(protein_tensor)  # returns list of strings

print(decoded[0])

```

## Summary

- **Load** pretrained weights using `ESMC.from_pretrained()` from [`esm/models/esmc.py`](https://github.com/Biohub/esm/blob/main/esm/models/esmc.py), which utilizes `load_local_model` in [`esm/pretrained.py`](https://github.com/Biohub/esm/blob/main/esm/pretrained.py) for checkpoint management.
- **Tokenize** protein sequences with `EsmSequenceTokenizer` ([`esm/tokenization/sequence_tokenizer.py`](https://github.com/Biohub/esm/blob/main/esm/tokenization/sequence_tokenizer.py)) to handle the 64-token alphabet and special tokens.
- **Optimize** using `nn.CrossEntropyLoss(ignore_index=tokenizer.pad_id)` against the `sequence_logits` output from the `RegressionHead`.
- **Accelerate** training by setting `use_flash_attn=True` when instantiating the model, leveraging the flash-attention kernels in `TransformerStack` without code changes.
- **Save** adapted weights with standard PyTorch `state_dict()` methods or Forge SDK helpers in [`esm/sdk/forge.py`](https://github.com/Biohub/esm/blob/main/esm/sdk/forge.py).

## Frequently Asked Questions

### What learning rate works best for fine-tuning ESMC models?

A learning rate of **1e-4 to 5e-5** with the `AdamW` optimizer typically yields stable convergence. The example in [`esm/models/esmc.py`](https://github.com/Biohub/esm/blob/main/esm/models/esmc.py) uses standard transformer training dynamics, so applying a linear warm-up over the first 10% of steps before cosine decay often improves downstream task performance.

### Can I freeze the transformer layers and only train the regression head?

Yes. Because `ESMC` inherits from `nn.Module`, you can freeze the transformer parameters by setting `param.requires_grad = False` for all parameters except those in `model.regression_head`. This parameter-efficient approach is useful when your dataset is small, as the `RegressionHead` in [`esm/layers/regression_head.py`](https://github.com/Biohub/esm/blob/main/esm/layers/regression_head.py) contains minimal parameters compared to the full `TransformerStack`.

### How does flash attention affect fine-tuning?

Flash attention reduces memory consumption and increases throughput during the self-attention computation in [`esm/layers/transformer_stack.py`](https://github.com/Biohub/esm/blob/main/esm/layers/transformer_stack.py). When you instantiate the model with `use_flash_attn=True`, the code automatically dispatches to optimized kernels if installed, allowing larger batch sizes or longer protein sequences on the same GPU hardware.

### How do I handle variable-length sequences in the DataLoader?

Use `torch.nn.utils.rnn.pad_sequence` with `batch_first=True` and set `padding_value=tokenizer.pad_id` as the `collate_fn` in your `DataLoader`. Your loss function should then ignore this padding index via `ignore_index=tokenizer.pad_id` in `nn.CrossEntropyLoss` to ensure padding tokens do not contribute to the gradient updates.