How to Compute and Visualize Predicted Aligned Error (PAE) Matrices in the ESM Repository

Predicted aligned error (PAE) matrices quantify inter-residue distance confidence in protein structure predictions by computing expected alignment errors from softmax probability distributions over 64 distance bins ranging from 0 to 31 Å.

The Biohub/esm repository implements predicted aligned error estimation to evaluate confidence in generated protein structures. Understanding how these matrices are computed and visualized enables researchers to distinguish high-confidence domains from uncertain regions in AI-designed proteins.

How PAE Matrices Are Computed

Bin Definition and Distance Ranges

In esm/utils/structure/predicted_aligned_error.py, the _pae_bins function constructs 64 linearly-spaced distance bins ranging from 0 Å to a configurable max_bin parameter (default 31.0 Å). According to lines 20-28 of the source, the final bin edge is shifted by half a step to ensure the upper boundary is properly represented.

Masking Valid Residue Pairs

Before computing probabilities, the _compute_pae_masks function (lines 32-34) generates a square boolean mask of shape L × L (where L is sequence length). This mask is True only for residue pairs where both positions contain valid amino acids according to the aa_mask tensor, effectively excluding gaps and padding from the error calculation.

Converting Logits to Expected Error

The compute_predicted_aligned_error function implements the core statistical transformation in lines 43-48:

  1. Mask application: Invalid positions in the raw logits (shape [L, L, num_bins]) are set to torch.finfo(logits.dtype).min (lines 45-47)
  2. Softmax normalization: A softmax operation over the bin dimension converts logits to probability distributions per residue pair
  3. Expectation calculation: The expected error for each pair is computed as the probability-weighted sum of bin centers, yielding a final matrix of shape [L, L] expressed in Ångströms

Integration in the VQ-VAE Decoder

The StructureTokenDecoder class in esm/models/vqvae.py (lines 410-417) extracts PAE logits from the model's pairwise head and invokes compute_predicted_aligned_error. The resulting matrix is stored under the predicted_aligned_error key in the decoder output dictionary, making it available for downstream processing.

SDK Propagation

When using the ESM SDK with the include_pae=True flag, the Forge class in esm/sdk/forge.py (lines 149-166) automatically retrieves the computed matrix and populates the pae field of the returned ESMProtein object. If PAE computation is not requested, this field remains None to conserve memory.

Visualizing PAE Matrices

While the repository does not ship a dedicated plotting widget, standard Python visualization tools effectively render these confidence maps. Low-error regions appear in cool colors (blue/green), while high-error regions appear warm (yellow/red), with the diagonal always near 0 Å since residues are perfectly aligned with themselves.

import matplotlib.pyplot as plt
import torch
from esm.utils.structure.predicted_aligned_error import compute_predicted_aligned_error

# pae_logits: raw model output of shape [L, L, num_bins]

# aa_mask: 1-D boolean tensor where True marks valid residues

pae_matrix = compute_predicted_aligned_error(
    logits=pae_logits,
    aa_mask=aa_mask,
    max_bin=31.0,
)

# Convert to NumPy for visualization

pae_np = pae_matrix.detach().cpu().numpy()

plt.figure(figsize=(6, 5))
im = plt.imshow(pae_np, cmap='viridis', origin='lower')
plt.title('Predicted Aligned Error (PAE) Matrix')
plt.xlabel('Residue index')
plt.ylabel('Residue index')
cbar = plt.colorbar(im)
cbar.set_label('PAE (Å)')
plt.show()

Practical Usage with the ESM SDK

To generate structures with confidence analysis, use the Forge interface with the include_pae parameter enabled. The matrix is then accessible via the pae attribute of the returned protein object.

from esm.sdk import Forge

# Initialize the SDK

forge = Forge()

# Request generation with PAE computation

result = forge.generate(
    prompt="Design a stable 100-residue protein",
    include_pae=True,  # Critical flag for PAE matrix

    num_samples=1,
)

protein = result.samples[0]  # ESMProtein object

pae = protein.pae            # torch.Tensor of shape [L, L] or None

# Visualize if available

if pae is not None:
    import matplotlib.pyplot as plt
    plt.imshow(pae.squeeze().cpu().numpy(), cmap='magma')
    plt.title('Predicted Aligned Error')
    plt.colorbar(label='Expected Error (Å)')
    plt.show()
else:
    print("PAE not returned")

Key Implementation Files

The end-to-end PAE computation flows through these critical source files:

Summary

  • PAE matrices represent expected distance alignment errors between residue pairs, computed as probability-weighted sums over 64 discrete bins from 0–31 Å
  • Computation involves masking invalid residues, applying softmax to logits, and calculating expectations using compute_predicted_aligned_error in esm/utils/structure/predicted_aligned_error.py
  • Model integration occurs in StructureTokenDecoder (esm/models/vqvae.py), which extracts pairwise logits and converts them to error estimates
  • SDK access requires setting include_pae=True in forge.generate() calls, returning the matrix via the ESMProtein.pae attribute
  • Visualization uses standard matplotlib heatmaps, where cool colors indicate high-confidence (low error) regions and warm colors indicate uncertain (high error) regions

Frequently Asked Questions

What does a predicted aligned error matrix represent?

A predicted aligned error matrix quantifies the model's confidence in the relative spatial positioning of every residue pair in a protein structure. Each cell (i, j) contains the expected distance error (in Ångströms) between residues i and j if their relative positions were aligned to the predicted structure, with lower values indicating higher confidence in the predicted inter-residue distance.

Why does the ESM repository use 64 bins for PAE calculation?

The implementation in esm/utils/structure/predicted_aligned_error.py defines 64 bins to balance computational efficiency with sufficient resolution for error discrimination. These bins span 0 to 31 Å (default), providing approximately 0.5 Å resolution per bin, which captures meaningful structural variations while keeping the softmax computation over the bin dimension numerically stable and memory-efficient.

How do I interpret high versus low PAE values in the visualization?

In a PAE heatmap, values near 0 Å (typically shown in blue/green) indicate the model is highly confident that the corresponding residue pair is positioned correctly relative to each other. Values above 15–20 Å (yellow/red regions) indicate substantial uncertainty, often occurring between structurally distant domains or flexible loops. The diagonal is always near 0 Å since a residue is trivially aligned with itself.

Can PAE matrices be computed for partial protein structures?

Yes, the masking mechanism in _compute_pae_masks automatically handles partial structures by filtering the computation to only valid residues where aa_mask == True. When visualizing partial predictions, ensure you align the matrix indices with the actual residue positions in your sequence, as the output matrix dimensions correspond to the full sequence length with masked positions excluded from error calculations.

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 →