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, orfunction) - num_steps: Number of iterative unmasking rounds to perform
- schedule: Token unmasking schedule (
cosineorlinear) - strategy: Position selection strategy (
randomorentropy) - 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 (typicallyTruefor 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:
- Stacks input tensors into a batched
ESMProteinTensor - Runs forward passes on the entire batch each step via
_batch_forward - Samples logits using the configured SamplingConfig for the current track
- Updates only positions selected by the unmasking mask returned from
_get_iterative_sampling_mask_for_prompt_and_step - Repeats until
num_stepsis 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
esm/sdk/api.py: Definitions forGenerationConfig,SamplingConfig, andSamplingTrackConfigdata classesesm/utils/generation.py: Core iterative sampling loop includingiterative_sampling_tokensand_batch_forwardesm/utils/sampling.py: Helper functions includingget_default_sampling_configand the low-levelsample_logitsimplementationcookbook/snippets/esm3.py: End-to-end examples demonstrating encoding, generation, and decoding workflowsesm/widgets/views/generation.py: UI-side configuration logic showing how fields are exposed in the web interface
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()fromesm/utils/sampling.pyand modify specificSamplingTrackConfiginstances before passing to low-level helpers - Set
invalid_idsto exclude special tokens like padding, and usetopk_logprobsto retain probability distributions for analysis - The core generation logic resides in
esm/utils/generation.pywithin theiterative_sampling_tokensfunction
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →