Skip to main content

Overview

The model.py module defines the core AlphaFold 3 model architecture, including the trunk (Evoformer), diffusion head, confidence head, and distogram head. It handles the complete forward pass and generates structure predictions.

Model Class

Main model class that orchestrates the full prediction pipeline.
Model.Config
required
Model configuration including Evoformer, heads, and global settings.
str
default:"diffuser"
Module name for Haiku.

Configuration

Model.Config

evoformer_network.Evoformer.Config
Configuration for the Evoformer trunk (embedding module).
model_config.GlobalConfig
Global configuration including precision and attention settings.
Model.HeadsConfig
Configuration for diffusion, confidence, and distogram heads.
int
default:"10"
Number of recycling iterations through the trunk.
bool
default:"False"
Whether to return single and pair embeddings in output.
bool
default:"False"
Whether to return distogram in output.

Model.HeadsConfig

Forward Pass

features.BatchDict
required
Input batch containing featurized sequences, MSAs, and templates.
jax.Array | None
Random key for JAX operations. If None, uses hk.next_rng_key().
ModelResult
Dictionary containing diffusion samples, confidence metrics, and optionally embeddings/distogram.

ModelResult Structure

The model returns a dictionary with the following keys:
dict
Contains atom_positions array with predicted atom coordinates.
dict
Distance histogram predictions and contact probabilities.
np.ndarray
Per-atom predicted local distance difference test (pLDDT) scores.
np.ndarray
Full predicted aligned error (PAE) matrix [num_samples, num_tokens, num_tokens].
np.ndarray
Full predicted distance error (PDE) matrix.
np.ndarray
TM-score adjusted PAE for global structure assessment.
np.ndarray
TM-score adjusted PAE for interface assessment.
jnp.ndarray
Single embeddings if return_embeddings=True [num_tokens, 384].
jnp.ndarray
Pair embeddings if return_embeddings=True [num_tokens, num_tokens, 128].

Class Methods

get_inference_result

Processes model outputs and computes inference-time metrics.
features.BatchDict
required
Data batch used for model inference.
ModelResult
required
Output dict from the model’s forward pass.
str
default:""
Target name to be saved within structure.
InferenceResult
Yields one InferenceResult per diffusion sample containing predicted structure and confidence metrics.

InferenceResult Class

Postprocessed model result containing the predicted structure and all confidence metrics.
structure.Structure
Predicted protein structure with atom coordinates and B-factors.
Mapping[str, float | int | np.ndarray]
Large numerical arrays like full PAE, PDE, and contact probabilities.
Mapping[str, float | int | np.ndarray]
Confidence metrics and summary statistics (see Metadata Fields below).
Mapping[str, Any]
Additional debugging information.
bytes
Model identifier from parameters.

Metadata Fields

The metadata dictionary contains the following confidence scores:
float
Primary ranking score combining pTM, ipTM, disorder, and clash penalties.
float
Predicted TM-score (pTM) measuring overall structure quality.
float
Interface predicted TM-score (ipTM) for multi-chain complexes.
float
Weighted average: 0.8 * ipTM + 0.2 * pTM.
float
Ranking confidence (equals ipTM for multi-chain, pTM for single chain).
float
Alternative ranking metric based on PAE.
float
Average predicted distance error across structure.
float
Fraction of structure predicted to be disordered.
bool
Whether structure has atomic clashes.
np.ndarray
Mean PDE between chain pairs [num_chains, num_chains].
np.ndarray
Minimum PDE between chain pairs [num_chains, num_chains].
np.ndarray
Minimum PAE between chain pairs [num_chains, num_chains].
np.ndarray
Interface pTM between chain pairs [num_chains, num_chains].
float
Average PDE within chains (intra-chain contacts).
float
Average PDE between chains (inter-chain contacts).
np.ndarray
Per-chain PAE scores [num_chains].
np.ndarray
Cross-chain PAE scores [num_chains].
np.ndarray
Per-chain ipTM scores [num_chains].
np.ndarray
Cross-chain ipTM scores [num_chains].
list[str]
Chain IDs for each token.
np.ndarray
Residue IDs for each token.

Numerical Data Fields

The numerical_data dictionary contains large arrays:
np.ndarray
Full predicted distance error matrix [num_tokens, num_tokens].
np.ndarray
Full predicted aligned error matrix [num_tokens, num_tokens].
np.ndarray
Contact probability matrix [num_tokens, num_tokens].

Helper Functions

get_predicted_structure

Creates the predicted structure from model outputs.
ModelResult
required
Model output in model-specific layout.
feat_batch.Batch
required
Model input batch for layout conversion.
structure.Structure
Predicted structure with atom coordinates and B-factors.

create_target_feat_embedding

Creates target feature embedding for the Evoformer.
feat_batch.Batch
required
Input batch data.
evoformer_network.Evoformer.Config
required
Evoformer configuration.
model_config.GlobalConfig
required
Global model configuration.
jnp.ndarray
Target feature embeddings [num_tokens, feature_dim].

Usage Examples

Basic Model Inference

Processing Results

Accessing Confidence Metrics

Multi-Sample Analysis

With Embeddings

Architecture Overview

The Model consists of:
  1. Evoformer (Trunk): Processes MSA and creates single/pair embeddings through multiple recycling iterations
  2. Diffusion Head: Generates atom coordinates through denoising diffusion process
  3. Confidence Head: Predicts pLDDT, PAE, and PDE confidence metrics
  4. Distogram Head: Predicts distance histograms and contact probabilities

Forward Pass Flow

Performance Considerations

  • Recycling: More recycles (10-20) improve quality but increase compute time
  • Diffusion Samples: More samples (5-10) provide better coverage but are slower
  • Embeddings: Enabling embeddings significantly increases memory usage
  • Flash Attention: Use triton or cudnn for best performance on Ampere+ GPUs

See Also