Skip to main content

Overview

The run_alphafold.py script is the main entry point for running AlphaFold 3 structure predictions. It orchestrates the complete prediction pipeline including data processing, MSA generation, template search, and model inference.

Main Functions

make_model_config

Creates a model configuration with customizable parameters.
str
default:"triton"
Flash attention implementation to use. Options: 'triton', 'cudnn', or 'xla'. Triton is fastest and requires Ampere GPUs or later.
int
default:"5"
Number of diffusion samples to generate per seed.
int
default:"10"
Number of recycling iterations during inference.
bool
default:"False"
Whether to return the final trunk single and pair embeddings. Embeddings are large float16 arrays: num_tokens * 384 + num_tokens * num_tokens * 128.
bool
default:"False"
Whether to return the final distogram. Distogram is a large float16 array: num_tokens * num_tokens * 64.
model.Model.Config
Configured model instance ready for inference.

predict_structure

Runs the full inference pipeline to predict structures for each seed.
folding_input.Input
required
The input containing chains, sequences, MSAs, and templates.
ModelRunner
required
The model runner instance for executing predictions.
Sequence[int] | None
Token bucket sizes for compilation caching. If None, calculates appropriate bucket from token count.
datetime.date | None
Maximum date for using CCD model coordinates as fallback.
int | None
Maximum iterations for RDKit conformer search.
bool
default:"True"
Whether to deduplicate unpaired MSA against paired MSA.
Sequence[ResultsForSeed]
List of results for each seed, containing inference results and full fold input.

process_fold_input

Runs data pipeline and/or inference on a single fold input.
folding_input.Input
required
Fold input to process.
pipeline.DataPipelineConfig | None
required
Data pipeline config to use. If None, skip the data pipeline.
ModelRunner | None
required
Model runner to use. If None, skip inference.
str
required
Output directory to write results to.
bool
default:"False"
If True, use existing output directory even if non-empty. If False, create timestamped directory.
bool
default:"False"
If True, compress large output files (mmCIF and confidences JSON) using zstandard.

ModelRunner Class

Helper class to run structure prediction stages.

Constructor

model.Model.Config
required
Model configuration.
jax.Device
required
JAX device to run inference on (e.g., GPU).
pathlib.Path
required
Path to directory containing model parameters.

Methods

run_inference

Computes a forward pass of the model on a featurised example.

extract_inference_results

Extracts inference results from model outputs.

extract_embeddings

Extracts single and pair embeddings from model outputs.

extract_distogram

Extracts distogram from model outputs.

ResultsForSeed

Dataclass storing inference results for a single seed.
int
The random seed used to generate the samples.
Sequence[model.InferenceResult]
The inference results, one per diffusion sample.
folding_input.Input
The fold input including MSA and templates from data pipeline.
dict[str, np.ndarray] | None
The final trunk single and pair embeddings, if requested.
np.ndarray | None
The token distance histogram, if requested.

Command Line Flags

Input/Output

  • --json_path: Path to input JSON file
  • --input_dir: Path to directory containing input JSON files
  • --output_dir: Path to output directory (required)
  • --model_dir: Path to model directory (default: ~/models)

Pipeline Control

  • --run_data_pipeline: Whether to run data pipeline (default: True)
  • --run_inference: Whether to run inference (default: True)

Database Paths

  • --db_dir: Database directory path (can specify multiple)
  • --small_bfd_database_path: Small BFD database path
  • --mgnify_database_path: Mgnify database path
  • --uniref90_database_path: UniRef90 database path
  • --uniprot_cluster_annot_database_path: UniProt database path
  • --ntrna_database_path: NT-RNA database path
  • --rfam_database_path: Rfam database path
  • --rna_central_database_path: RNAcentral database path
  • --pdb_database_path: PDB mmCIF files directory
  • --seqres_database_path: PDB sequence database path

Performance Tuning

  • --num_recycles: Number of recycles (default: 10)
  • --num_diffusion_samples: Number of diffusion samples (default: 5)
  • --num_seeds: Number of seeds to generate
  • --gpu_device: GPU device index (default: 0)
  • --flash_attention_implementation: Flash attention type: triton, cudnn, or xla (default: triton)
  • --buckets: Token bucket sizes for compilation caching

Output Control

  • --save_embeddings: Save final embeddings (default: False)
  • --save_distogram: Save distogram (default: False)
  • --compress_large_output_files: Compress output files (default: False)
  • --force_output_dir: Use existing output directory (default: False)

Usage Example

Command Line Usage

See Also