Overview
Themodel.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
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
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
Themetadata 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
Thenumerical_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
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
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:- Evoformer (Trunk): Processes MSA and creates single/pair embeddings through multiple recycling iterations
- Diffusion Head: Generates atom coordinates through denoising diffusion process
- Confidence Head: Predicts pLDDT, PAE, and PDE confidence metrics
- 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
tritonorcudnnfor best performance on Ampere+ GPUs
See Also
- run_alphafold.py - Main prediction script
- Input Dataclass - Input format specification
- DataPipeline - MSA and template processing