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

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 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 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, 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, 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) 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:


# 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 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.


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

# 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:

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, which utilizes load_local_model in esm/pretrained.py for checkpoint management.
  • Tokenize protein sequences with EsmSequenceTokenizer (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.

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 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 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. 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.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →