How ESMFold2 Diffusion Sampling Generates Protein Structures
ESMFold2 generates 3D protein structures by running a diffusion-based generative process that iteratively denoises atom coordinates through multiple loops, guided by configurable noise schedules and selective masking.
The Biohub/esm repository implements ESMFold2 as a state-of-the-art protein structure prediction model that replaces traditional deterministic folding with a stochastic generative approach. This article examines how ESMFold2 diffusion sampling works by analyzing the actual source code implementation, including the specific functions in processor.py and noise_schedules.py that control the denoising pipeline.
The Architecture of ESMFold2 Diffusion Sampling
ESMFold2 generates structures by treating atomic coordinates as a field that evolves through a series of noisy states toward a final prediction. The process is implemented in esm/models/esmfold2/processor.py, where the ESMFold2InputBuilder class orchestrates the diffusion loop.
The core workflow involves three distinct phases: input preparation, where sequence and geometry tensors are constructed; the diffusion loop, where noise is iteratively added and removed; and post-processing, where raw coordinates are converted into structured outputs with confidence metrics.
Step-by-Step: How Diffusion Sampling Works
Input Preparation and Tokenization
The process begins in ESMFold2InputBuilder.prepare_input, located in esm/models/esmfold2/processor.py at lines 45-81. This method tokenizes the protein sequence and constructs geometry-aware tensors including atom masks and reference encodings.
When a seed is supplied, this preparation step becomes fully deterministic, ensuring reproducible structure generation even when random SMILES-derived conformers are involved.
Configuring the Diffusion Loop
The model accepts three critical hyper-parameters that control the sampling behavior:
- num_loops: Defines how many full denoising cycles are performed, with each loop restarting the noise schedule.
- num_sampling_steps: Specifies the number of intermediate diffusion timesteps within each loop.
- num_diffusion_samples: Determines how many independent noisy trajectories are generated in parallel, expanding the batch dimension to
Bm = B × num_diffusion_samples.
Noise Scheduling and Injection
The amount of noise present at each timestep is governed by schedules stored in esm.utils.noise_schedules.NOISE_SCHEDULE_REGISTRY, implemented in esm/utils/noise_schedules.py at lines 28-34. Available schedules include cosine, linear, and square-root decay curves.
During the forward pass, the model repeatedly injects Gaussian noise into the current atom-coordinate field and predicts a denoised update based on the selected schedule.
Selective Coordinate Updates via Masking
At every diffusion step, the library determines which atoms to update using esm.utils.sampling.get_sampling_mask, found in esm/utils/sampling.py at lines 22-38. This function removes special tokens such as BOS and EOS, then identifies masked positions by checking for mask tokens (torch.inf).
The model outputs sample_atom_coords with shape [Bm, L, 3], representing the predicted coordinates for the batch. After each diffusion step, newly predicted coordinates replace old ones only at the masked positions, while unmasked coordinates remain unchanged.
Structure Decoding and Confidence Scoring
Once the diffusion loops complete, ESMFold2InputBuilder.decode processes the final coordinates. Located in esm/models/esmfold2/processor.py at lines 225-252 and 241-280, this method constructs a MolecularComplexResult object.
The result includes per-residue confidence metrics (pLDDT), PTM/iptm scores, and optional PAEs or distograms, providing a complete structural prediction with quality estimates.
Practical Implementation: Running ESMFold2 Structure Prediction
The following example demonstrates how to invoke the diffusion sampling pipeline using the Biohub/esm API:
from esm.models.esmfold2.processor import ESMFold2InputBuilder
from esm.pretrained import load_local_model
import torch
# Build the model
model = load_local_model("esmfold2", device=torch.device("cuda"))
# Create input specification
from esm.models.esmfold2.types import StructurePredictionInput, ProteinInput
inp = StructurePredictionInput(
sequences=[ProteinInput(id=["0"], sequence="MKTIIALSYIFCLVFA")]
)
# Execute diffusion sampling
builder = ESMFold2InputBuilder()
results = builder.fold(
model,
inp,
num_loops=3,
num_sampling_steps=200,
num_diffusion_samples=2,
seed=42,
)
# Process results
for i, res in enumerate(results):
print(f"Sample {i} – pLDDT mean: {res.plddt.mean():.2f}")
coords = res.complex.coordinates # Shape: [L, 3]
Key implementation details from this example:
- The
builder.foldmethod forwards prepared tensors through the complete diffusion loop. - Setting
num_diffusion_samples=2triggers parallel sampling, generating multiple structural hypotheses. - The
seedparameter ensures reproducibility for both SMILES-based conformer generation and diffusion noise.
Key Source Files in the ESMFold2 Pipeline
Understanding ESMFold2 diffusion sampling requires familiarity with these specific modules:
esm/models/esmfold2/processor.py: ContainsESMFold2InputBuildermethods (prepare_input,fold,decode) that drive the diffusion loop and coordinate decoding.esm/utils/noise_schedules.py: Implements selectable noise decay curves including cosine and linear schedules used by the sampler.esm/utils/sampling.py: Providesget_sampling_maskand other utilities for masking tokens and computing per-track metadata.esm/models/esmfold2/types.py: Defines data classes for prediction requests includingStructurePredictionInputandProteinInput.
Summary
- ESMFold2 diffusion sampling generates structures by iteratively denoising atom coordinates through multiple loops controlled by configurable noise schedules.
- The process uses three key parameters:
num_loops,num_sampling_steps, andnum_diffusion_samplesto control sampling depth and diversity. - Input preparation occurs in
ESMFold2InputBuilder.prepare_input(processor.py lines 45-81), which tokenizes sequences and builds geometry-aware tensors. - Noise schedules from
esm/utils/noise_schedules.py(lines 28-34) determine how quickly noise decays during the denoising process. - The
get_sampling_maskfunction (sampling.py lines 22-38) ensures only masked atom positions are updated at each diffusion step. - Final decoding in
ESMFold2InputBuilder.decode(processor.py lines 241-280) producesMolecularComplexResultobjects containing coordinates and confidence metrics like pLDDT.
Frequently Asked Questions
What is the difference between num_loops and num_sampling_steps in ESMFold2?
num_loops controls how many complete denoising cycles are executed, with each loop restarting the noise schedule from maximum noise. num_sampling_steps defines the granularity within each loop, specifying how many intermediate timesteps occur between the fully noised and fully denoised states. Higher values for both parameters increase computational cost but may improve structural quality.
How does ESMFold2 ensure deterministic sampling when requested?
Determinism is achieved by supplying a seed parameter to ESMFold2InputBuilder.fold, which initializes random number generators for both the SMILES-derived conformer generation and the Gaussian noise injection during diffusion. When a seed is provided, the prepare_input method produces identical initial conditions across runs.
What noise schedules are available in ESMFold2, and how do they differ?
The NOISE_SCHEDULE_REGISTRY in esm/utils/noise_schedules.py provides multiple decay curves including cosine, linear, and square-root schedules. Cosine schedules provide smoother transitions near the end of the diffusion process, while linear schedules offer constant decay rates. The choice of schedule affects how quickly the model transitions from high-noise to low-noise states.
How does the sampling mask determine which atoms to update during diffusion?
The get_sampling_mask function (sampling.py lines 22-38) identifies valid positions by removing special tokens (BOS/EOS) and checking for mask tokens (torch.inf). At each diffusion step, only coordinates at these masked positions are updated with the model's predictions, while unmasked positions retain their previous values. This selective updating ensures the diffusion process respects prior structural constraints when available.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →