How to Use Guided Generation with PocketConditioning Constraints for Binder Design
The Biohub/esm repository provides a derivative-free guided decoding engine that combines PocketConditioning with MDMM-based constrained optimization to design protein binders satisfying specific pocket contact requirements.
The ESM3 architecture supports advanced protein engineering workflows through its experimental constrained generation API. By leveraging guided generation with PocketConditioning constraints, researchers can direct the model to create binders that form specific contacts with target pockets, using only a scoring function without explicit token-level conditioning.
Understanding the Three-Layer Architecture
The pocket conditioning workflow consists of three conceptual layers that bridge constraint definition with structure generation.
Define the Pocket Constraints
First, create a PocketConditioning object to specify which chain serves as the binder and which target residues must make contact. According to the Biohub/esm source code, this dataclass resides in esm/utils/structure/input_builder.py (lines 58-62).
Create the Scoring Function
Next, implement a subclass of GuidedDecodingScoringFunction to evaluate generated proteins. This abstract base class is defined in esm/sdk/experimental/guided_generation.py (lines 22-25). Your implementation receives an ESMProtein and returns a scalar reward based on contact satisfaction, pTM scores, or custom geometric metrics.
Execute Constrained Decoding
Finally, instantiate ESM3GuidedDecodingWithConstraints from esm/sdk/experimental/constrained_generation.py (lines 94-119) to perform the iterative optimization. This class uses the Modified Differential Method of Multipliers (MDMM) to enforce constraints while maximizing your reward function.
Key Implementation Files
esm/utils/structure/input_builder.py– DefinesPocketConditioningand input serialization logic.esm/sdk/experimental/guided_generation.py– Contains theGuidedDecodingScoringFunctionabstract base class.esm/sdk/experimental/constrained_generation.py– ImplementsESM3GuidedDecodingWithConstraintswith MDMM optimization.esm/models/esmfold2/prepare_input.py– Prepares model tensors; note thatpocket_feature = torch.zeros(lines 26-28) indicates pocket constraints are enforced purely through scoring functions.esm/sdk/api.py– ProvidesESMProteinandESMProteinTensordata structures.
Step-by-Step Implementation
The following workflow demonstrates how to design a binder for chain B residues 0 and 2.
Step 1: Configure PocketConditioning
Instantiate the pocket definition specifying the binder chain and contact residues:
from esm.utils.structure.input_builder import PocketConditioning
pocket = PocketConditioning(
binder_chain_id="A",
contacts=[("B", 0), ("B", 2)],
)
Step 2: Implement the Scoring Function
Subclass GuidedDecodingScoringFunction to count satisfied contacts within an 8 Å threshold:
from esm.sdk.experimental import GuidedDecodingScoringFunction
class ContactScorer(GuidedDecodingScoringFunction):
def __call__(self, protein) -> float:
complex = protein.to_protein_complex()
binder = complex.get_chain_by_id("A")
target = complex.get_chain_by_id("B")
dmat = binder.cbeta_contacts(distance_threshold=8.0)
satisfied = sum(
1 for chain_id, res_idx in pocket.contacts
if dmat[:, target[res_idx].residue_index].min() < 8.0
)
return float(satisfied)
Step 3: Set Up Constraints
Wrap the scorer in a GenerationConstraint requiring at least two contacts:
from esm.sdk.experimental import GenerationConstraint, ConstraintType
contact_constraint = GenerationConstraint(
scoring_function=ContactScorer(),
value=2.0,
constraint_type=ConstraintType.GREATER_EQUAL,
)
Step 4: Run Guided Generation
Execute the constrained decoding loop:
from esm.sdk.experimental import ESM3GuidedDecodingWithConstraints
from esm.sdk.api import ESM3InferenceClient
client = ESM3InferenceClient.from_pretrained("esm3_t33_650M_UR50D")
guided = ESM3GuidedDecodingWithConstraints(
client=client,
scoring_function=ContactScorer(),
constraints=[contact_constraint],
damping=10.0,
learning_rate=1.0,
)
final_protein = guided.guided_generate(
protein=client.encode(structure_input),
num_decoding_steps=10,
num_samples_per_step=5,
track="structure",
)
Complete Binder Design Example
Here is the complete, runnable implementation combining all components:
# ------------------------------------------------------------
# 1️⃣ Imports
# ------------------------------------------------------------
from esm.sdk.api import ESM3InferenceClient, ESMProtein, SamplingConfig
from esm.sdk.experimental import (
ESM3GuidedDecodingWithConstraints,
GuidedDecodingScoringFunction,
GenerationConstraint,
ConstraintType,
)
from esm.utils.structure.input_builder import (
StructurePredictionInput,
ProteinInput,
PocketConditioning,
)
import torch
import numpy as np
# ------------------------------------------------------------
# 2️⃣ Define a pocket (binder = chain A, contacts on chain B)
# ------------------------------------------------------------
pocket = PocketConditioning(
binder_chain_id="A",
contacts=[("B", 0), ("B", 2)], # target residues on chain B
)
# ------------------------------------------------------------
# 3️⃣ Build the input structure (binder + target)
# ------------------------------------------------------------
binder_seq = "MKTLLILTCLVAVALAR..." # placeholder sequence for the binder
target_seq = "GVALV...,B" # include chain‑break token "|" if needed
inp = StructurePredictionInput(
sequences=[
ProteinInput(id="binder", sequence=binder_seq),
ProteinInput(id="target", sequence=target_seq),
],
pocket=pocket,
)
# ------------------------------------------------------------
# 4️⃣ Initialise the model client (uses the public inference endpoint)
# ------------------------------------------------------------
client = ESM3InferenceClient.from_pretrained("esm3_t33_650M_UR50D")
# ------------------------------------------------------------
# 5️⃣ Scoring function: maximise number of satisfied contacts
# ------------------------------------------------------------
class ContactScorer(GuidedDecodingScoringFunction):
def __call__(self, protein: ESMProtein) -> float:
# Convert the full prediction to a ProteinComplex
complex = protein.to_protein_complex()
binder = complex.get_chain_by_id(pocket.binder_chain_id)
target = complex.get_chain_by_id("B")
# Cβ‑Cβ distance matrix (Å)
dmat = binder.cbeta_contacts(distance_threshold=8.0)
# Count how many of the specified contacts are within 8 Å
satisfied = 0
for chain_id, res_idx in pocket.contacts:
# res_idx is 0‑based; get the corresponding residue in the target
target_res = target[res_idx]
# Find minimal distance between this target residue and any binder residue
min_dist = dmat[:, target_res.residue_index].min()
if min_dist < 8.0:
satisfied += 1
# Return a positive reward (higher is better)
return float(satisfied)
# ------------------------------------------------------------
# 6️⃣ Wrap the contact scorer in a constraint (≥ 2 contacts required)
# ------------------------------------------------------------
contact_constraint = GenerationConstraint(
scoring_function=ContactScorer(),
value=2.0, # we want at least 2 contacts
constraint_type=ConstraintType.GREATER_EQUAL,
)
# ------------------------------------------------------------
# 7️⃣ Set up constrained guided decoding
# ------------------------------------------------------------
guided = ESM3GuidedDecodingWithConstraints(
client=client,
scoring_function=ContactScorer(), # overall reward (could be same as constraint)
constraints=[contact_constraint],
damping=10.0,
learning_rate=1.0,
)
# ------------------------------------------------------------
# 8️⃣ Run the optimisation
# ------------------------------------------------------------
final_protein = guided.guided_generate(
protein=client.encode(inp), # the initial masked input is created automatically
num_decoding_steps=10,
num_samples_per_step=5,
denoised_prediction_temperature=0.0,
track="structure",
verbose=True,
)
# ------------------------------------------------------------
# 9️⃣ Inspect the result
# ------------------------------------------------------------
print("Generated sequence :", final_protein.sequence)
print("Number of satisfied contacts :", ContactScorer()(final_protein))
final_protein.to_pdb("binder_design_output.pdb")
How the Constrained Decoding Loop Works
The ESM3GuidedDecodingWithConstraints class implements an iterative unmasking strategy:
- Encode the initial masked protein into
ESMProteinTensorusingclient.encode. - Unmask a subset of positions via
randomly_unmask_positionsat each decoding step. - Score all sampled candidates using your
GuidedDecodingScoringFunctionand evaluate constraints via_score_and_constraints. - Update Lagrange multipliers using
c.update_lambdato drive the optimization toward feasible designs. - Select the best-scoring candidate as the new state for the next iteration.
- Decode the final structure deterministically using
client.decodeafter the last step.
Because the pocket feature is currently a placeholder (pocket_feature = torch.zeros), the constraint is enforced purely through the scoring function. This design allows you to plug in any bespoke metric—such as geometric distances, interface pTM, or Rosetta energy terms—while retaining the model's internal structure prediction capabilities.
Summary
- The PocketConditioning dataclass in
esm/utils/structure/input_builder.pydefines binder chains and required contacts. - GuidedDecodingScoringFunction provides the interface for custom reward functions evaluating
ESMProteinobjects. - ESM3GuidedDecodingWithConstraints in
esm/sdk/experimental/constrained_generation.pyuses MDMM to enforce constraints while optimizing your scoring function. - Pocket constraints are enforced purely through scoring due to the placeholder implementation in
esm/models/esmfold2/prepare_input.py, enabling flexible custom metrics without token-level conditioning. - The workflow supports multiple constraints and arbitrary reward functions, including contact-based, energy-based, or confidence-based objectives.
Frequently Asked Questions
How does the MDMM optimizer enforce pocket constraints during generation?
The Modified Differential Method of Multipliers updates Lagrange multipliers via c.update_lambda after each scoring iteration, adjusting the optimization landscape to penalize designs that violate the specified contact requirements while rewarding those that satisfy them.
Can I use multiple scoring functions simultaneously for binder design?
Yes. Pass a list of GenerationConstraint objects to the constraints parameter of ESM3GuidedDecodingWithConstraints. The MDMM machinery handles each constraint independently, allowing you to combine contact requirements with secondary objectives like pTM scores or stability metrics.
Why is the pocket feature implemented as a zero tensor in the model inputs?
As implemented in esm/models/esmfold2/prepare_input.py (lines 26-28), pocket_feature = torch.zeros(...) indicates that pocket conditioning is not performed at the token embedding level. Instead, constraints are enforced purely through the derivative-free guided decoding scoring function, providing flexibility to define arbitrary geometric or energetic constraints without requiring model retraining.
What is the difference between the scoring function passed to the constructor versus the constraint?
The scoring_function parameter in ESM3GuidedDecodingWithConstraints serves as the primary reward to maximize, while GenerationConstraint objects enforce hard requirements (e.g., minimum number of contacts). They can use the same underlying logic or different metrics, allowing the optimizer to balance reward maximization against constraint satisfaction.
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 →