How to Configure SamplingConfig for Iterative Protein Generation in ESM3

To configure SamplingConfig for iterative protein generation, instantiate a GenerationConfig to control the high-level unmasking schedule, then either let the SDK auto-build a per-track SamplingConfig or manually construct SamplingTrackConfig objects with custom temperature, top-p, and invalid_ids parameters.

The Biohub/esm repository provides ESM3, a generative protein language model that creates sequences, structures, and functional annotations through iterative unmasking. Properly configuring the SamplingConfig is essential for controlling the stochastic behavior during this iterative protein generation process.

Understanding the Two-Layer Configuration Architecture

The ESM3 SDK implements a hierarchical configuration system where GenerationConfig handles high-level generation parameters while SamplingConfig manages low-level per-track sampling behavior.

GenerationConfig (High-Level Controls)

Defined in esm/sdk/api.py, the GenerationConfig class governs the overall generation process through these key fields:

  • track: Specifies which modality to generate (sequence, structure, secondary_structure, sasa, or function)
  • num_steps: Number of iterative unmasking rounds to perform
  • schedule: Token unmasking schedule (cosine or linear)
  • strategy: Position selection strategy (random or entropy)
  • temperature and top_p: Global sampling parameters applied across all steps

SamplingConfig and SamplingTrackConfig (Low-Level Controls)

Also defined in esm/sdk/api.py, the SamplingConfig class contains per-track SamplingTrackConfig objects that directly control logits sampling:

  • temperature: Controls randomness; lower values generate more deterministic outputs (typical range: 0.7–1.2)
  • top_p: Nucleus sampling threshold; higher values allow broader distributions (typical range: 0.8–1.0)
  • only_sample_masked_tokens: Restricts sampling to <mask> positions (typically True for iterative generation)
  • invalid_ids: Token IDs to exclude from sampling (e.g., [tokenizer.pad_token_id])
  • topk_logprobs: Number of top log-probabilities to retain in output for analysis (0–5)

Automatic vs Manual SamplingConfig Construction

Automatic Configuration via Client Methods

When calling client.generate() or client.batch_generate(), the SDK automatically constructs a SamplingConfig from your GenerationConfig. As implemented in the source code, the library builds a per-track SamplingTrackConfig as follows:

track_cfg = SamplingTrackConfig(
    temperature=gen_cfg.temperature,
    top_p=gen_cfg.top_p,
    only_sample_masked_tokens=True,
    invalid_ids=gen_cfg.invalid_ids,
    topk_logprobs=0,
)
sampling_cfg = SamplingConfig(**{gen_cfg.track: track_cfg})

Manual Configuration for Fine-Grained Control

For advanced use cases requiring different parameters per track, manually construct a SamplingConfig using the helper in esm/utils/sampling.py:

from esm.utils.sampling import get_default_sampling_config
from esm.tokenization import get_tokenizers

tokenizers = get_tokenizers()
sampling_cfg = get_default_sampling_config(tokenizers)

This creates a SamplingConfig with sensible defaults for all tracks (temperature=1.0, top-p=1.0) and automatically populates invalid_ids based on tokenizer specifications.

Core Iterative Generation Loop

The iterative sampling process resides in esm/utils/generation.py within the iterative_sampling_tokens function. According to the source code, the loop executes the following operations:

  1. Stacks input tensors into a batched ESMProteinTensor
  2. Runs forward passes on the entire batch each step via _batch_forward
  3. Samples logits using the configured SamplingConfig for the current track
  4. Updates only positions selected by the unmasking mask returned from _get_iterative_sampling_mask_for_prompt_and_step
  5. Repeats until num_steps is reached or the mask is empty

Practical Configuration Examples

Basic Iterative Generation with Auto-Configuration

from esm.sdk.api import GenerationConfig
from esm.sdk.forge import ESM3InferenceClient, ESMProtein

client = ESM3InferenceClient(model="esm3")
protein = ESMProtein(sequence="MKTIIALSYIFCLVFADYKDDDDK")

gen_cfg = GenerationConfig(
    track="sequence",
    num_steps=30,
    temperature=0.8,
    top_p=0.9,
    schedule="cosine",
    strategy="random",
    temperature_annealing=True,
)

generated = client.generate(protein, gen_cfg)
print(generated.sequence)

Custom Per-Track SamplingConfig

from esm.sdk.api import SamplingConfig, SamplingTrackConfig
from esm.utils.sampling import get_default_sampling_config
from esm.tokenization import get_tokenizers
from esm.sdk.forge import ESM3InferenceClient

client = ESM3InferenceClient(model="esm3")
tokenizers = get_tokenizers()

# Get defaults for all tracks

sampling_cfg = get_default_sampling_config(tokenizers)

# Override sequence track specifically

sampling_cfg.sequence = SamplingTrackConfig(
    temperature=0.6,
    top_p=0.95,
    only_sample_masked_tokens=True,
    invalid_ids=[tokenizers.sequence.pad_token_id],
    topk_logprobs=5,
)

# Encode protein for low-level API usage

protein = client.encode(ESMProtein(sequence="MKTIIALSYIFCLVFADYKDDDDK"))

Low-Level Raw Tensor Access

raw_outputs = client.iterative_sampling_raw(
    client,
    [ESMProtein(sequence="MKTIIALSYIFCLVFADYKDDDDK")],
    [GenerationConfig(track="structure", num_steps=15, temperature=1.0)],
)

Key Configuration Parameters

Parameter Effect Typical Values
temperature Controls sampling randomness; lower values increase determinism 0.7–1.2
top_p Nucleus sampling threshold; higher values allow broader distributions 0.8–1.0
only_sample_masked_tokens Restricts sampling to <mask> positions during iterative generation True
invalid_ids Excludes specific token IDs (padding, BOS, EOS) from sampling [tokenizer.pad_token_id]
topk_logprobs Retains top-k log-probabilities in output for downstream analysis 0–5

Key Source Files

Summary

  • GenerationConfig controls the high-level iterative unmasking process including track selection, number of steps, and global temperature
  • SamplingConfig manages per-track stochastic parameters and is automatically constructed by client.generate() from your GenerationConfig settings
  • For custom behavior, use get_default_sampling_config() from esm/utils/sampling.py and modify specific SamplingTrackConfig instances before passing to low-level helpers
  • Set invalid_ids to exclude special tokens like padding, and use topk_logprobs to retain probability distributions for analysis
  • The core generation logic resides in esm/utils/generation.py within the iterative_sampling_tokens function

Frequently Asked Questions

What is the difference between GenerationConfig and SamplingConfig?

GenerationConfig acts as the high-level controller that specifies which track to generate, how many unmasking steps to perform, and the global sampling strategy. SamplingConfig contains the low-level per-track parameters found in esm/sdk/api.py that directly control logits sampling, including temperature, top-p values, and token exclusions.

How do I prevent the model from generating padding tokens?

Set the invalid_ids parameter in your SamplingTrackConfig to include the pad token ID from your specific tokenizer, such as invalid_ids=[tokenizers.sequence.pad_token_id]. The helper function get_default_sampling_config() automatically configures these exclusions based on tokenizer specifications.

Can I use different sampling temperatures for different tracks?

Yes. While the standard client.generate() method uses a single global temperature, you can manually construct a SamplingConfig with different SamplingTrackConfig instances for each track. This allows you to set temperature=0.8 for sequence generation while using temperature=1.2 for structure generation within the same model execution.

When should I use iterative_sampling_raw instead of client.generate?

Use iterative_sampling_raw when you need access to intermediate tensors, logits, or raw model outputs for debugging or custom post-processing. The standard client.generate() method handles encoding, generation, and decoding automatically, while _raw variants return the underlying data structures without automatic decoding to ESMProtein objects.

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 →