# How to Use Guided Generation with PocketConditioning Constraints for Binder Design

> Design protein binders with specific pocket contact requirements using Biohub/esm's guided generation and PocketConditioning constraints. Learn this derivative-free decoding engine.

- Repository: [Biohub/esm](https://github.com/Biohub/esm)
- Tags: how-to-guide
- Published: 2026-05-30

---

**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`](https://github.com/Biohub/esm/blob/main/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`](https://github.com/Biohub/esm/blob/main/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`](https://github.com/Biohub/esm/blob/main/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`](https://github.com/Biohub/esm/blob/main/esm/utils/structure/input_builder.py) – Defines `PocketConditioning` and input serialization logic.
- [`esm/sdk/experimental/guided_generation.py`](https://github.com/Biohub/esm/blob/main/esm/sdk/experimental/guided_generation.py) – Contains the `GuidedDecodingScoringFunction` abstract base class.
- [`esm/sdk/experimental/constrained_generation.py`](https://github.com/Biohub/esm/blob/main/esm/sdk/experimental/constrained_generation.py) – Implements `ESM3GuidedDecodingWithConstraints` with MDMM optimization.
- [`esm/models/esmfold2/prepare_input.py`](https://github.com/Biohub/esm/blob/main/esm/models/esmfold2/prepare_input.py) – Prepares model tensors; note that `pocket_feature = torch.zeros` (lines 26-28) indicates pocket constraints are enforced purely through scoring functions.
- [`esm/sdk/api.py`](https://github.com/Biohub/esm/blob/main/esm/sdk/api.py) – Provides `ESMProtein` and `ESMProteinTensor` data 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:

```python
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:

```python
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:

```python
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:

```python
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:

```python

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

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

1. **Encode** the initial masked protein into `ESMProteinTensor` using `client.encode`.
2. **Unmask** a subset of positions via `randomly_unmask_positions` at each decoding step.
3. **Score** all sampled candidates using your `GuidedDecodingScoringFunction` and evaluate constraints via `_score_and_constraints`.
4. **Update** Lagrange multipliers using `c.update_lambda` to drive the optimization toward feasible designs.
5. **Select** the best-scoring candidate as the new state for the next iteration.
6. **Decode** the final structure deterministically using `client.decode` after 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.py`](https://github.com/Biohub/esm/blob/main/esm/utils/structure/input_builder.py) defines binder chains and required contacts.
- **GuidedDecodingScoringFunction** provides the interface for custom reward functions evaluating `ESMProtein` objects.
- **ESM3GuidedDecodingWithConstraints** in [`esm/sdk/experimental/constrained_generation.py`](https://github.com/Biohub/esm/blob/main/esm/sdk/experimental/constrained_generation.py) uses 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`](https://github.com/Biohub/esm/blob/main/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`](https://github.com/Biohub/esm/blob/main/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.