OneScience-Group/AlphaFold3
3254
1 2 3"""AlphaFold 3 structure prediction script.4 5AlphaFold 3 source code is licensed under CC BY-NC-SA 4.0. To view a copy of6this license, visit https://creativecommons.org/licenses/by-nc-sa/4.0/7 8To request access to the AlphaFold 3 model parameters, follow the process set9out at https://github.com/google-deepmind/alphafold3. You may only use these10if received directly from Google. Use is subject to terms of use available at11https://github.com/google-deepmind/alphafold3/blob/main/WEIGHTS_TERMS_OF_USE.md12"""13 14from collections.abc import Callable, Sequence15import csv16import dataclasses17import datetime18import functools19import multiprocessing20import os21import pathlib22import shutil23import string24import sys25import textwrap26import time27import typing28from typing import overload29 30_PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[1]31if str(_PROJECT_ROOT) not in sys.path:32 sys.path.insert(0, str(_PROJECT_ROOT))33 34from absl import app35from absl import flags36from flax_model.alphafold3.common import folding_input37from flax_model.alphafold3.common import resources38from flax_model.alphafold3.constants import chemical_components39import flax_model.alphafold3.cpp as af3_cpp40from flax_model.alphafold3.data import featurisation41from flax_model.alphafold3.data import pipeline42from flax_model.alphafold3.jax.attention import attention43from flax_model.alphafold3.model import features44from flax_model.alphafold3.model import model45from flax_model.alphafold3.model import params46from flax_model.alphafold3.model import post_processing47from flax_model.alphafold3.model.components import utils48import haiku as hk49import jax50from jax import numpy as jnp51import numpy as np52 53 54_HOME_DIR = pathlib.Path(os.environ.get('HOME', _PROJECT_ROOT))55_MODELS_ROOT = pathlib.Path(56 os.environ.get('ONESCIENCE_MODELS_DIR', _PROJECT_ROOT / 'weight')57)58_DATASETS_ROOT = pathlib.Path(59 os.environ.get('ONESCIENCE_DATASETS_DIR', _HOME_DIR)60) / 'alphafold3'61_DEFAULT_MODEL_DIR = pathlib.Path(62 os.environ.get('ALPHAFOLD3_MODEL_DIR', _MODELS_ROOT / 'AlphaFold3')63)64_DEFAULT_DB_DIR = pathlib.Path(65 os.environ.get('ALPHAFOLD3_DB_DIR', _DATASETS_ROOT / 'public_databases')66)67_DEFAULT_MMSEQS_DB_DIR = pathlib.Path(68 os.environ.get('ALPHAFOLD3_MMSEQS_DB_DIR', _DATASETS_ROOT / 'mmseqsDB')69)70 71# Input and output paths.72_JSON_PATH = flags.DEFINE_string(73 'json_path',74 None,75 'Path to the input JSON file.',76)77_INPUT_DIR = flags.DEFINE_string(78 'input_dir',79 None,80 'Path to the directory containing input JSON files.',81)82_OUTPUT_DIR = flags.DEFINE_string(83 'output_dir',84 None,85 'Path to a directory where the results will be saved.',86)87MODEL_DIR = flags.DEFINE_string(88 'model_dir',89 _DEFAULT_MODEL_DIR.as_posix(),90 'Path to the model to use for inference.',91)92 93# Control which stages to run.94_RUN_DATA_PIPELINE = flags.DEFINE_bool(95 'run_data_pipeline',96 True,97 'Whether to run the data pipeline on the fold inputs.',98)99_RUN_INFERENCE = flags.DEFINE_bool(100 'run_inference',101 True,102 'Whether to run inference on the fold inputs.',103)104 105 106_DEFAULT_MMSEQS_OPTIONS ='--num-iterations 1 --db-load-mode 2 -a --max-seqs 10000 --prefilter-mode 1'107_DEFAULT_R2MSA_OPTIONS ='--filter-msa 1 --filter-min-enable 1000 --diff 3000 --qid 0.0,0.2,0.4,0.6,0.8,1.0 --qsc 0 --max-seq-id 0.95'108_AF3_DIR = _PROJECT_ROOT / 'flax_model' / 'alphafold3'109_HMMER_BIN_DIR = _AF3_DIR / '_tools' / 'hmmer' / 'bin'110 111 112def _resolve_af3_tool(binary_name: str) -> str | None:113 local_binary = _HMMER_BIN_DIR / binary_name114 if local_binary.exists():115 return str(local_binary)116 return shutil.which(binary_name)117 118 119def _available_cpu_count() -> int:120 if hasattr(os, 'sched_getaffinity'):121 return len(os.sched_getaffinity(0))122 return multiprocessing.cpu_count()123 124 125_USE_MMSEQS = flags.DEFINE_bool(126 'use_mmseqs',127 False,128 'Whether to use mmseqs for protein MSA search',129)130_USE_MMSEQS_GPU = flags.DEFINE_bool(131 'use_mmseqs_gpu',132 False,133 'Whether to use mmseqs GPU for protein MSA search',134)135_MMSEQS_OPTIONS = flags.DEFINE_string(136 'mmseqs_options',137 _DEFAULT_MMSEQS_OPTIONS,138 'mmseqs serach options',139)140_R2MSA_OPTIONS = flags.DEFINE_string(141 'result2msa_options',142 _DEFAULT_R2MSA_OPTIONS,143 'mmseqs result2msa options',144)145 146 147# Binary paths.148_JACKHMMER_BINARY_PATH = flags.DEFINE_string(149 'jackhmmer_binary_path',150 _resolve_af3_tool('jackhmmer'),151 'Path to the Jackhmmer binary.',152)153_NHMMER_BINARY_PATH = flags.DEFINE_string(154 'nhmmer_binary_path',155 _resolve_af3_tool('nhmmer'),156 'Path to the Nhmmer binary.',157)158_HMMALIGN_BINARY_PATH = flags.DEFINE_string(159 'hmmalign_binary_path',160 _resolve_af3_tool('hmmalign'),161 'Path to the Hmmalign binary.',162)163_HMMSEARCH_BINARY_PATH = flags.DEFINE_string(164 'hmmsearch_binary_path',165 _resolve_af3_tool('hmmsearch'),166 'Path to the Hmmsearch binary.',167)168_HMMBUILD_BINARY_PATH = flags.DEFINE_string(169 'hmmbuild_binary_path',170 _resolve_af3_tool('hmmbuild'),171 'Path to the Hmmbuild binary.',172)173_MMSEQS_BINARY_PATH = flags.DEFINE_string(174 'mmseqs_binary_path',175 shutil.which('mmseqs'),176 'Path to the mmseqs binary.',177)178 179# Database paths.180DB_DIR = flags.DEFINE_multi_string(181 'db_dir',182 (_DEFAULT_DB_DIR.as_posix(),),183 'Path to the directory containing the databases. Can be specified multiple'184 ' times to search multiple directories in order.',185)186MMSEQS_DB_DIR = flags.DEFINE_multi_string(187 'mmseqs_db_dir',188 (_DEFAULT_MMSEQS_DB_DIR.as_posix(),),189 'Path to the directory containing the mmseqs databases. Can be specified multiple'190 ' times to search multiple directories in order.', 191)192 193_SMALL_BFD_DATABASE_PATH = flags.DEFINE_string(194 'small_bfd_database_path',195 '${DB_DIR}/bfd-first_non_consensus_sequences.fasta',196 'Small BFD database path, used for protein MSA search.',197)198_SMALL_BFD_Z_VALUE = flags.DEFINE_integer(199 'small_bfd_z_value',200 None,201 'The Z-value representing the database size in number of sequences for'202 ' E-value calculation. Must be set for sharded databases.',203 lower_bound=0,204)205_MGNIFY_DATABASE_PATH = flags.DEFINE_string(206 'mgnify_database_path',207 '${DB_DIR}/mgy_clusters_2022_05.fa',208 'Mgnify database path, used for protein MSA search.',209)210_MGNIFY_Z_VALUE = flags.DEFINE_integer(211 'mgnify_z_value',212 None,213 'The Z-value representing the database size in number of sequences for'214 ' E-value calculation. Must be set for sharded databases.',215 lower_bound=0,216)217_UNIPROT_CLUSTER_ANNOT_DATABASE_PATH = flags.DEFINE_string(218 'uniprot_cluster_annot_database_path',219 '${DB_DIR}/uniprot_all_2021_04.fa',220 'UniProt database path, used for protein paired MSA search.',221)222_UNIPROT_CLUSTER_ANNOT_Z_VALUE = flags.DEFINE_integer(223 'uniprot_cluster_annot_z_value',224 None,225 'The Z-value representing the database size in number of sequences for'226 ' E-value calculation. Must be set for sharded databases.',227 lower_bound=0,228)229_UNIREF90_DATABASE_PATH = flags.DEFINE_string(230 'uniref90_database_path',231 '${DB_DIR}/uniref90_2022_05.fa',232 'UniRef90 database path, used for MSA search. The MSA obtained by '233 'searching it is used to construct the profile for template search.',234)235_UNIREF90_Z_VALUE = flags.DEFINE_integer(236 'uniref90_z_value',237 None,238 'The Z-value representing the database size in number of sequences for'239 ' E-value calculation. Must be set for sharded databases.',240 lower_bound=0,241)242_NTRNA_DATABASE_PATH = flags.DEFINE_string(243 'ntrna_database_path',244 '${DB_DIR}/nt_rna_2023_02_23_clust_seq_id_90_cov_80_rep_seq.fasta',245 'NT-RNA database path, used for RNA MSA search.',246)247_NTRNA_Z_VALUE = flags.DEFINE_float(248 'ntrna_z_value',249 None,250 'The Z-value representing the database size in megabases for E-value'251 ' calculation. Must be set for sharded databases.',252 lower_bound=0.0,253)254_RFAM_DATABASE_PATH = flags.DEFINE_string(255 'rfam_database_path',256 '${DB_DIR}/rfam_14_9_clust_seq_id_90_cov_80_rep_seq.fasta',257 'Rfam database path, used for RNA MSA search.',258)259_RFAM_Z_VALUE = flags.DEFINE_float(260 'rfam_z_value',261 None,262 'The Z-value representing the database size in megabases for E-value'263 ' calculation. Must be set for sharded databases.',264 lower_bound=0.0,265)266_RNA_CENTRAL_DATABASE_PATH = flags.DEFINE_string(267 'rna_central_database_path',268 '${DB_DIR}/rnacentral_active_seq_id_90_cov_80_linclust.fasta',269 'RNAcentral database path, used for RNA MSA search.',270)271_RNA_CENTRAL_Z_VALUE = flags.DEFINE_float(272 'rna_central_z_value',273 None,274 'The Z-value representing the database size in megabases for E-value'275 ' calculation. Must be set for sharded databases.',276 lower_bound=0.0,277)278_PDB_DATABASE_PATH = flags.DEFINE_string(279 'pdb_database_path',280 '${DB_DIR}/mmcif_files',281 'PDB database directory with mmCIF files path, used for template search.',282)283_SEQRES_DATABASE_PATH = flags.DEFINE_string(284 'seqres_database_path',285 '${DB_DIR}/pdb_seqres_2022_09_28.fasta',286 'PDB sequence database path, used for template search.',287)288 289# MMSEQS Database paths.290_MMSEQS_SMALL_BFD_DATABASE_PATH = flags.DEFINE_string(291 'mmseqs_small_bfd_database_path',292 '${MMSEQS_DB_DIR}/small_bfd_db',293 'Small BFD database path, used for protein MSA search.',294) 295_MMSEQS_MGNIFY_DATABASE_PATH = flags.DEFINE_string(296 'mmseqs_mgnify_database_path',297 '${MMSEQS_DB_DIR}/mgnify_db',298 'Mgnify database path, used for protein MSA search.',299)300_MMSEQS_UNIPROT_CLUSTER_ANNOT_DATABASE_PATH = flags.DEFINE_string(301 'mmseqs_uniprot_cluster_annot_database_path',302 '${MMSEQS_DB_DIR}/uniprot_cluster_annot_db',303 'UniProt database path, used for protein paired MSA search.',304)305_MMSEQS_UNIREF90_DATABASE_PATH = flags.DEFINE_string(306 'mmseqs_uniref90_database_path',307 '${MMSEQS_DB_DIR}/uniref90_db',308 'UniRef90 database path, used for MSA search. ',309)310 311_JACKHMMER_MAX_THREADS = flags.DEFINE_integer(312 'jackhmmer_max_threads',313 None,314 'Maximum number of threads used when running sharded databases. If unset,'315 ' defaults to None (no limit).',316 lower_bound=1,317)318# Number of CPUs to use for MSA tools.319_JACKHMMER_N_CPU = flags.DEFINE_integer(320 'jackhmmer_n_cpu',321 # Unfortunately, os.process_cpu_count() is only available in Python 3.13+.322 min(_available_cpu_count(), 8),323 'Number of CPUs to use for Jackhmmer. Defaults to min(cpu_count, 8). Going'324 ' above 8 CPUs provides very little additional speedup.',325 lower_bound=0,326)327_JACKHMMER_MAX_PARALLEL_SHARDS = flags.DEFINE_integer(328 'jackhmmer_max_parallel_shards',329 None,330 'Maximum number of shards to search against in parallel. If unset, one'331 ' Jackhmmer instance will be run per shard. Only applicable if the'332 ' database is sharded.',333 lower_bound=1,334)335_NHMMER_N_CPU = flags.DEFINE_integer(336 'nhmmer_n_cpu',337 # Unfortunately, os.process_cpu_count() is only available in Python 3.13+.338 min(_available_cpu_count(), 8),339 'Number of CPUs to use for Nhmmer. Defaults to min(cpu_count, 8). Going'340 ' above 8 CPUs provides very little additional speedup.',341 lower_bound=0,342)343_NHMMER_MAX_PARALLEL_SHARDS = flags.DEFINE_integer(344 'nhmmer_max_parallel_shards',345 None,346 'Maximum number of shards to search against in parallel. If unset, one'347 ' Nhmmer instance will be run per shard. Only applicable if the'348 ' database is sharded.',349 lower_bound=1,350)351_NHMMER_MAX_THREADS = flags.DEFINE_integer(352 'nhmmer_max_threads',353 None,354 'Maximum number of threads used when running sharded databases. If unset,'355 ' defaults to None (no limit).',356 lower_bound=1,357)358# Data pipeline configuration.359_RESOLVE_MSA_OVERLAPS = flags.DEFINE_bool(360 'resolve_msa_overlaps',361 True,362 'Whether to deduplicate unpaired MSA against paired MSA. The default'363 ' behaviour matches the method described in the AlphaFold 3 paper. Set this'364 ' to false if providing custom paired MSA using the unpaired MSA field to'365 ' keep it exactly as is as deduplication against the paired MSA could break'366 ' the manually crafted pairing between MSA sequences.',367)368_MMSEQS_N_CPU = flags.DEFINE_integer(369 'mmseqs_n_cpu',370 min(multiprocessing.cpu_count(), 8),371 'Number of CPUs to use for MMseqs. Default to min(cpu_count, 8). Going'372 ' beyond 8 CPUs provides very little additional speedup.',373)374 375# Template search configuration.376_MAX_TEMPLATE_DATE = flags.DEFINE_string(377 'max_template_date',378 '2021-09-30', # By default, use the date from the AlphaFold 3 paper.379 'Maximum template release date to consider. Format: YYYY-MM-DD. All'380 ' templates released after this date will be ignored. Controls also whether'381 ' to allow use of model coordinates for a chemical component from the CCD'382 ' if RDKit conformer generation fails and the component does not have ideal'383 ' coordinates set. Only for components that have been released before this'384 ' date the model coordinates can be used as a fallback.',385)386 387_CONFORMER_MAX_ITERATIONS = flags.DEFINE_integer(388 'conformer_max_iterations',389 None, # Default to RDKit default parameters value.390 'Optional override for maximum number of iterations to run for RDKit '391 'conformer search.',392 lower_bound=0,393)394 395# JAX inference performance tuning.396_JAX_COMPILATION_CACHE_DIR = flags.DEFINE_string(397 'jax_compilation_cache_dir',398 None,399 'Path to a directory for the JAX compilation cache.',400)401_GPU_DEVICE = flags.DEFINE_integer(402 'gpu_device',403 0,404 'Optional override for the GPU device to use for inference, uses zero-based'405 ' indexing. Defaults to the 0th GPU on the system. Useful on multi-GPU'406 ' systems to pin each run to a specific GPU. Note that if GPUs are already'407 ' pre-filtered by the environment (e.g. by using CUDA_VISIBLE_DEVICES),'408 ' this flag refers to the GPU index after the filtering has been done.',409)410_BUCKETS = flags.DEFINE_list(411 'buckets',412 # pyformat: disable413 ['256', '512', '768', '1024', '1280', '1536', '2048', '2560', '3072',414 '3584', '4096', '4608', '5120'],415 # pyformat: enable416 'Strictly increasing order of token sizes for which to cache compilations.'417 ' For any input with more tokens than the largest bucket size, a new bucket'418 ' is created for exactly that number of tokens.',419)420_FLASH_ATTENTION_IMPLEMENTATION = flags.DEFINE_enum(421 'flash_attention_implementation',422 default='triton',423 enum_values=['triton', 'cudnn', 'xla', 'cutlass',],424 help=(425 "Flash attention implementation to use. 'triton' and 'cudnn' uses a"426 ' Triton and cuDNN flash attention implementation, respectively. The'427 ' Triton kernel is fastest and has been tested more thoroughly. The'428 " Triton and cuDNN kernels require Ampere GPUs or later. 'xla' uses an"429 ' XLA attention implementation (no flash attention) and is portable'430 ' across GPU devices.'431 ),432)433_NUM_RECYCLES = flags.DEFINE_integer(434 'num_recycles',435 10,436 'Number of recycles to use during inference.',437 lower_bound=1,438)439_NUM_DIFFUSION_SAMPLES = flags.DEFINE_integer(440 'num_diffusion_samples',441 5,442 'Number of diffusion samples to generate.',443 lower_bound=1,444)445_NUM_SEEDS = flags.DEFINE_integer(446 'num_seeds',447 None,448 'Number of seeds to use for inference. If set, only a single seed must be'449 ' provided in the input JSON. AlphaFold 3 will then generate random seeds'450 ' in sequence, starting from the single seed specified in the input JSON.'451 ' The full input JSON produced by AlphaFold 3 will include the generated'452 ' random seeds. If not set, AlphaFold 3 will use the seeds as provided in'453 ' the input JSON.',454 lower_bound=1,455)456 457# Output controls.458_SAVE_EMBEDDINGS = flags.DEFINE_bool(459 'save_embeddings',460 False,461 'Whether to save the final trunk single and pair embeddings in the output.'462 ' Note that the embeddings are large float16 arrays: num_tokens * 384'463 ' + num_tokens * num_tokens * 128.',464)465_SAVE_DISTOGRAM = flags.DEFINE_bool(466 'save_distogram',467 False,468 'Whether to save the final distogram in the output. Note that the distogram'469 ' is a large float16 array: num_tokens * num_tokens * 64.',470)471_FORCE_OUTPUT_DIR = flags.DEFINE_bool(472 'force_output_dir',473 False,474 'Whether to force the output directory to be used even if it already exists'475 ' and is non-empty. Useful to set this to True to run the data pipeline and'476 ' the inference separately, but use the same output directory.',477)478 479 480def make_model_config(481 *,482 flash_attention_implementation: attention.Implementation = 'triton',483 num_diffusion_samples: int = 5,484 num_recycles: int = 10,485 return_embeddings: bool = False,486 return_distogram: bool = False,487) -> model.Model.Config:488 """Returns a model config with some defaults overridden."""489 config = model.Model.Config()490 config.global_config.flash_attention_implementation = (491 flash_attention_implementation492 )493 config.heads.diffusion.eval.num_samples = num_diffusion_samples494 config.num_recycles = num_recycles495 config.return_embeddings = return_embeddings496 config.return_distogram = return_distogram497 return config498 499 500class ModelRunner:501 """Helper class to run structure prediction stages."""502 503 def __init__(504 self,505 config: model.Model.Config,506 device: jax.Device,507 model_dir: pathlib.Path,508 ):509 self._model_config = config510 self._device = device511 self._model_dir = model_dir512 513 @functools.cached_property514 def model_params(self) -> hk.Params:515 """Loads model parameters from the model directory."""516 return params.get_model_haiku_params(model_dir=self._model_dir)517 518 @functools.cached_property519 def _model(520 self,521 ) -> Callable[[jnp.ndarray, features.BatchDict], model.ModelResult]:522 """Loads model parameters and returns a jitted model forward pass."""523 524 @hk.transform525 def forward_fn(batch):526 return model.Model(self._model_config)(batch)527 528 return functools.partial(529 jax.jit(forward_fn.apply, device=self._device), self.model_params530 )531 532 def run_inference(533 self, featurised_example: features.BatchDict, rng_key: jnp.ndarray534 ) -> model.ModelResult:535 """Computes a forward pass of the model on a featurised example."""536 featurised_example = jax.device_put(537 jax.tree_util.tree_map(538 jnp.asarray, utils.remove_invalidly_typed_feats(featurised_example)539 ),540 self._device,541 )542 543 result = self._model(rng_key, featurised_example)544 result = jax.tree.map(np.asarray, result)545 result = jax.tree.map(546 lambda x: x.astype(jnp.float32) if x.dtype == jnp.bfloat16 else x,547 result,548 )549 result = dict(result)550 identifier = self.model_params['__meta__']['__identifier__'].tobytes()551 result['__identifier__'] = identifier552 return result553 554 def extract_inference_results(555 self,556 batch: features.BatchDict,557 result: model.ModelResult,558 target_name: str,559 ) -> list[model.InferenceResult]:560 """Extracts inference results from model outputs."""561 return list(562 model.Model.get_inference_result(563 batch=batch, result=result, target_name=target_name564 )565 )566 567 def extract_embeddings(568 self, result: model.ModelResult, num_tokens: int569 ) -> dict[str, np.ndarray] | None:570 """Extracts embeddings from model outputs."""571 embeddings = {}572 if 'single_embeddings' in result:573 embeddings['single_embeddings'] = result['single_embeddings'][574 :num_tokens575 ].astype(np.float16)576 if 'pair_embeddings' in result:577 embeddings['pair_embeddings'] = result['pair_embeddings'][578 :num_tokens, :num_tokens579 ].astype(np.float16)580 return embeddings or None581 582 def extract_distogram(583 self, result: model.ModelResult, num_tokens: int584 ) -> np.ndarray | None:585 """Extracts distogram from model outputs."""586 if 'distogram' not in result['distogram']:587 return None588 distogram = result['distogram']['distogram'][:num_tokens, :num_tokens, :]589 return distogram590 591 592@dataclasses.dataclass(frozen=True, slots=True, kw_only=True)593class ResultsForSeed:594 """Stores the inference results (diffusion samples) for a single seed.595 596 Attributes:597 seed: The seed used to generate the samples.598 inference_results: The inference results, one per sample.599 full_fold_input: The fold input that must also include the results of600 running the data pipeline - MSA and templates.601 embeddings: The final trunk single and pair embeddings, if requested.602 distogram: The token distance histogram, if requested.603 """604 605 seed: int606 inference_results: Sequence[model.InferenceResult]607 full_fold_input: folding_input.Input608 embeddings: dict[str, np.ndarray] | None = None609 distogram: np.ndarray | None = None610 611 612def predict_structure(613 fold_input: folding_input.Input,614 model_runner: ModelRunner,615 buckets: Sequence[int] | None = None,616 ref_max_modified_date: datetime.date | None = None,617 conformer_max_iterations: int | None = None,618 resolve_msa_overlaps: bool = True,619) -> Sequence[ResultsForSeed]:620 """Runs the full inference pipeline to predict structures for each seed."""621 622 print(f'Featurising data with {len(fold_input.rng_seeds)} seed(s)...')623 featurisation_start_time = time.time()624 ccd = chemical_components.Ccd(user_ccd=fold_input.user_ccd)625 featurised_examples = featurisation.featurise_input(626 fold_input=fold_input,627 buckets=buckets,628 ccd=ccd,629 verbose=True,630 ref_max_modified_date=ref_max_modified_date,631 conformer_max_iterations=conformer_max_iterations,632 resolve_msa_overlaps=resolve_msa_overlaps,633 )634 print(635 f'Featurising data with {len(fold_input.rng_seeds)} seed(s) took'636 f' {time.time() - featurisation_start_time:.2f} seconds.'637 )638 print(639 'Running model inference and extracting output structure samples with'640 f' {len(fold_input.rng_seeds)} seed(s)...'641 )642 all_inference_start_time = time.time()643 all_inference_results = []644 for seed, example in zip(fold_input.rng_seeds, featurised_examples):645 print(f'Running model inference with seed {seed}...')646 inference_start_time = time.time()647 rng_key = jax.random.PRNGKey(seed)648 result = model_runner.run_inference(example, rng_key)649 print(650 f'Running model inference with seed {seed} took'651 f' {time.time() - inference_start_time:.2f} seconds.'652 )653 print(f'Extracting inference results with seed {seed}...')654 extract_structures = time.time()655 inference_results = model_runner.extract_inference_results(656 batch=example, result=result, target_name=fold_input.name657 )658 num_tokens = len(inference_results[0].metadata['token_chain_ids'])659 embeddings = model_runner.extract_embeddings(660 result=result, num_tokens=num_tokens661 )662 distogram = model_runner.extract_distogram(663 result=result, num_tokens=num_tokens664 )665 print(666 f'Extracting {len(inference_results)} inference samples with'667 f' seed {seed} took {time.time() - extract_structures:.2f} seconds.'668 )669 670 all_inference_results.append(671 ResultsForSeed(672 seed=seed,673 inference_results=inference_results,674 full_fold_input=fold_input,675 embeddings=embeddings,676 distogram=distogram,677 )678 )679 print(680 'Running model inference and extracting output structures with'681 f' {len(fold_input.rng_seeds)} seed(s) took'682 f' {time.time() - all_inference_start_time:.2f} seconds.'683 )684 return all_inference_results685 686 687def write_fold_input_json(688 fold_input: folding_input.Input,689 output_dir: os.PathLike[str] | str,690) -> None:691 """Writes the input JSON to the output directory."""692 os.makedirs(output_dir, exist_ok=True)693 path = os.path.join(output_dir, f'{fold_input.sanitised_name()}_data.json')694 print(f'Writing model input JSON to {path}')695 with open(path, 'wt') as f:696 f.write(fold_input.to_json())697 698 699def write_outputs(700 all_inference_results: Sequence[ResultsForSeed],701 output_dir: os.PathLike[str] | str,702 job_name: str,703) -> None:704 """Writes outputs to the specified output directory."""705 ranking_scores = []706 max_ranking_score = None707 max_ranking_result = None708 try:709 output_terms = (710 pathlib.Path(af3_cpp.__file__).parent / 'OUTPUT_TERMS_OF_USE.md'711 ).read_text()712 except FileNotFoundError:713 output_terms = None714 os.makedirs(output_dir, exist_ok=True)715 for results_for_seed in all_inference_results:716 seed = results_for_seed.seed717 for sample_idx, result in enumerate(results_for_seed.inference_results):718 sample_dir = os.path.join(output_dir, f'seed-{seed}_sample-{sample_idx}')719 os.makedirs(sample_dir, exist_ok=True)720 post_processing.write_output(721 inference_result=result,722 output_dir=sample_dir,723 name=f'{job_name}_seed-{seed}_sample-{sample_idx}',724 )725 ranking_score = float(result.metadata['ranking_score'])726 ranking_scores.append((seed, sample_idx, ranking_score))727 if max_ranking_score is None or ranking_score > max_ranking_score:728 max_ranking_score = ranking_score729 max_ranking_result = result730 731 if embeddings := results_for_seed.embeddings:732 embeddings_dir = os.path.join(output_dir, f'seed-{seed}_embeddings')733 os.makedirs(embeddings_dir, exist_ok=True)734 post_processing.write_embeddings(735 embeddings=embeddings,736 output_dir=embeddings_dir,737 name=f'{job_name}_seed-{seed}',738 )739 740 if (distogram := results_for_seed.distogram) is not None:741 distogram_dir = os.path.join(output_dir, f'seed-{seed}_distogram')742 os.makedirs(distogram_dir, exist_ok=True)743 distogram_path = os.path.join(744 distogram_dir, f'{job_name}_seed-{seed}_distogram.npz'745 )746 with open(distogram_path, 'wb') as f:747 np.savez_compressed(f, distogram=distogram.astype(np.float16))748 749 if max_ranking_result is not None: # True iff ranking_scores non-empty.750 post_processing.write_output(751 inference_result=max_ranking_result,752 output_dir=output_dir,753 # The output terms of use are the same for all seeds/samples.754 terms_of_use=output_terms,755 name=job_name,756 )757 # Save csv of ranking scores with seeds and sample indices, to allow easier758 # comparison of ranking scores across different runs.759 with open(760 os.path.join(output_dir, f'{job_name}_ranking_scores.csv'), 'wt'761 ) as f:762 writer = csv.writer(f)763 writer.writerow(['seed', 'sample', 'ranking_score'])764 writer.writerows(ranking_scores)765 766 767def replace_db_dir(path_with_db_dir: str, db_dirs: Sequence[str]) -> str:768 """Replaces the DB_DIR placeholder in a path with the given DB_DIR."""769 template = string.Template(path_with_db_dir)770 if 'DB_DIR' in template.get_identifiers():771 for db_dir in db_dirs:772 path = template.substitute(DB_DIR=db_dir)773 if os.path.exists(path):774 return path775 raise FileNotFoundError(776 f'{path_with_db_dir} with ${{DB_DIR}} not found in any of {db_dirs}.'777 )778 if not os.path.exists(path_with_db_dir):779 raise FileNotFoundError(f'{path_with_db_dir} does not exist.')780 return path_with_db_dir781 782 783def replace_mmseqs_db_dir(784 path_with_db_dir: str, 785 db_dirs: Sequence[str],786 mmseqs_db_dirs: Sequence[str]787) -> str:788 """Replaces the MMSEQS_DB_DIR placeholder in a path with the given MMSEQS_DB_DIR.789 790 Args:791 path_with_db_dir: Path containing MMSEQS_DB_DIR placeholder792 db_dirs: List of database directories793 mmseqs_db_dirs: List of mmseqs database directories794 use_gpu: Whether to use GPU version of mmseqs databases795 796 Returns:797 The expanded path if found, otherwise raises FileNotFoundError798 """799 template = string.Template(path_with_db_dir)800 801 is_jackhmmer_db = any(802 path_with_db_dir.endswith(db) for db in [803 'bfd-first_non_consensus_sequences.fasta',804 'mgy_clusters_2022_05.fa',805 'uniprot_all_2021_04.fa',806 'uniref90_2022_05.fa'807 ]808 )809 810 if 'MMSEQS_DB_DIR' in template.get_identifiers():811 db_suffixes = [812 'small_bfd_db',813 'mgnify_db',814 'uniprot_cluster_annot_db',815 'uniref90_db'816 ]817 818 is_mmseqs_db = any(819 path_with_db_dir.endswith(suffix) for suffix in db_suffixes820 )821 822 for mmseqs_db_dir in mmseqs_db_dirs:823 path = template.substitute(MMSEQS_DB_DIR=mmseqs_db_dir)824 if is_mmseqs_db:825 if os.path.exists(path):826 return path827 else:828 return path829 830 if is_mmseqs_db:831 raise FileNotFoundError(832 f'{path_with_db_dir} with ${{MMSEQS_DB_DIR}} not found in any of {mmseqs_db_dirs}.'833 )834 return template.substitute(MMSEQS_DB_DIR=mmseqs_db_dirs[0])835 836 if 'DB_DIR' in template.get_identifiers():837 for db_dir in db_dirs:838 path = template.substitute(DB_DIR=db_dir)839 if is_jackhmmer_db:840 return path841 if os.path.exists(path):842 return path843 if is_jackhmmer_db:844 return template.substitute(DB_DIR=db_dirs[0])845 raise FileNotFoundError(846 f'{path_with_db_dir} with ${{DB_DIR}} not found in any of {db_dirs}.'847 )848 if not is_jackhmmer_db and not os.path.exists(path_with_db_dir):849 raise FileNotFoundError(f'{path_with_db_dir} does not exist.')850 return path_with_db_dir851 852 853@overload854def process_fold_input(855 fold_input: folding_input.Input,856 data_pipeline_config: pipeline.DataPipelineConfig | None,857 model_runner: None,858 output_dir: os.PathLike[str] | str,859 buckets: Sequence[int] | None = None,860 ref_max_modified_date: datetime.date | None = None,861 conformer_max_iterations: int | None = None,862 resolve_msa_overlaps: bool = True,863 force_output_dir: bool = False,864) -> folding_input.Input:865 ...866 867 868@overload869def process_fold_input(870 fold_input: folding_input.Input,871 data_pipeline_config: pipeline.DataPipelineConfig | None,872 model_runner: ModelRunner,873 output_dir: os.PathLike[str] | str,874 buckets: Sequence[int] | None = None,875 ref_max_modified_date: datetime.date | None = None,876 conformer_max_iterations: int | None = None,877 resolve_msa_overlaps: bool = True,878 force_output_dir: bool = False,879) -> Sequence[ResultsForSeed]:880 ...881 882 883def process_fold_input(884 fold_input: folding_input.Input,885 data_pipeline_config: pipeline.DataPipelineConfig | None,886 model_runner: ModelRunner | None,887 output_dir: os.PathLike[str] | str,888 buckets: Sequence[int] | None = None,889 ref_max_modified_date: datetime.date | None = None,890 conformer_max_iterations: int | None = None,891 resolve_msa_overlaps: bool = True,892 force_output_dir: bool = False,893) -> folding_input.Input | Sequence[ResultsForSeed]:894 """Runs data pipeline and/or inference on a single fold input.895 896 Args:897 fold_input: Fold input to process.898 data_pipeline_config: Data pipeline config to use. If None, skip the data899 pipeline.900 model_runner: Model runner to use. If None, skip inference.901 output_dir: Output directory to write to.902 buckets: Bucket sizes to pad the data to, to avoid excessive re-compilation903 of the model. If None, calculate the appropriate bucket size from the904 number of tokens. If not None, must be a sequence of at least one integer,905 in strictly increasing order. Will raise an error if the number of tokens906 is more than the largest bucket size.907 ref_max_modified_date: Optional maximum date that controls whether to allow908 use of model coordinates for a chemical component from the CCD if RDKit909 conformer generation fails and the component does not have ideal910 coordinates set. Only for components that have been released before this911 date the model coordinates can be used as a fallback.912 conformer_max_iterations: Optional override for maximum number of iterations913 to run for RDKit conformer search.914 resolve_msa_overlaps: Whether to deduplicate unpaired MSA against paired915 MSA. The default behaviour matches the method described in the AlphaFold 3916 paper. Set this to false if providing custom paired MSA using the unpaired917 MSA field to keep it exactly as is as deduplication against the paired MSA918 could break the manually crafted pairing between MSA sequences.919 force_output_dir: If True, do not create a new output directory even if the920 existing one is non-empty. Instead use the existing output directory and921 potentially overwrite existing files. If False, create a new timestamped922 output directory instead if the existing one is non-empty.923 924 Returns:925 The processed fold input, or the inference results for each seed.926 927 Raises:928 ValueError: If the fold input has no chains.929 """930 print(f'\nRunning fold job {fold_input.name}...')931 932 if not fold_input.chains:933 raise ValueError('Fold input has no chains.')934 935 if (936 not force_output_dir937 and os.path.exists(output_dir)938 and os.listdir(output_dir)939 ):940 new_output_dir = (941 f'{output_dir}_{datetime.datetime.now().strftime("%Y%m%d_%H%M%S")}'942 )943 print(944 f'Output will be written in {new_output_dir} since {output_dir} is'945 ' non-empty.'946 )947 output_dir = new_output_dir948 else:949 print(f'Output will be written in {output_dir}')950 951 # Create model loader callback function952 def load_model_callback():953 if model_runner is not None:954 _ = model_runner.model_params955 956 if data_pipeline_config is None:957 print('Skipping data pipeline...')958 # Load model immediately when skipping pipeline959 if model_runner is not None:960 print('Loading model parameters (no pipeline to wait for)...')961 load_model_callback()962 else: 963 print('Running data pipeline...') 964 fold_input = pipeline.DataPipeline(data_pipeline_config, load_model_callback).process(fold_input) 965 966 write_fold_input_json(fold_input, output_dir)967 if model_runner is None:968 print('Skipping model inference...')969 output = fold_input970 else:971 print(972 f'Predicting 3D structure for {fold_input.name} with'973 f' {len(fold_input.rng_seeds)} seed(s)...'974 )975 all_inference_results = predict_structure(976 fold_input=fold_input,977 model_runner=model_runner,978 buckets=buckets,979 ref_max_modified_date=ref_max_modified_date,980 conformer_max_iterations=conformer_max_iterations,981 resolve_msa_overlaps=resolve_msa_overlaps,982 )983 print(f'Writing outputs with {len(fold_input.rng_seeds)} seed(s)...')984 write_outputs(985 all_inference_results=all_inference_results,986 output_dir=output_dir,987 job_name=fold_input.sanitised_name(),988 )989 output = all_inference_results990 991 print(f'Fold job {fold_input.name} done, output written to {output_dir}\n')992 return output993 994 995def main(_):996 if _JAX_COMPILATION_CACHE_DIR.value is not None:997 jax.config.update(998 'jax_compilation_cache_dir', _JAX_COMPILATION_CACHE_DIR.value999 )1000 1001 if _JSON_PATH.value is None == _INPUT_DIR.value is None:1002 raise ValueError(1003 'Exactly one of --json_path or --input_dir must be specified.'1004 )1005 1006 if not _RUN_INFERENCE.value and not _RUN_DATA_PIPELINE.value:1007 raise ValueError(1008 'At least one of --run_inference or --run_data_pipeline must be'1009 ' set to true.'1010 )1011 1012 if _INPUT_DIR.value is not None:1013 fold_inputs = folding_input.load_fold_inputs_from_dir(1014 pathlib.Path(_INPUT_DIR.value)1015 )1016 elif _JSON_PATH.value is not None:1017 fold_inputs = folding_input.load_fold_inputs_from_path(1018 pathlib.Path(_JSON_PATH.value)1019 )1020 else:1021 raise AssertionError(1022 'Exactly one of --json_path or --input_dir must be specified.'1023 )1024 1025 # Make sure we can create the output directory before running anything.1026 try:1027 os.makedirs(_OUTPUT_DIR.value, exist_ok=True)1028 except OSError as e:1029 print(f'Failed to create output directory {_OUTPUT_DIR.value}: {e}')1030 raise1031 1032 # if _RUN_INFERENCE.value:1033 # # Fail early on incompatible devices, but only if we're running inference.1034 # gpu_devices = jax.local_devices(backend='gpu')1035 # if gpu_devices:1036 # compute_capability = float(1037 # gpu_devices[_GPU_DEVICE.value].compute_capability1038 # )1039 # if compute_capability < 6.0:1040 # raise ValueError(1041 # 'AlphaFold 3 requires at least GPU compute capability 6.0 (see'1042 # ' https://developer.nvidia.com/cuda-gpus).'1043 # )1044 # elif 7.0 <= compute_capability < 8.0:1045 # xla_flags = os.environ.get('XLA_FLAGS')1046 # required_flag = '--xla_disable_hlo_passes=custom-kernel-fusion-rewriter'1047 # if not xla_flags or required_flag not in xla_flags:1048 # raise ValueError(1049 # 'For devices with GPU compute capability 7.x (see'1050 # ' https://developer.nvidia.com/cuda-gpus) the ENV XLA_FLAGS must'1051 # f' include "{required_flag}".'1052 # )1053 # if _FLASH_ATTENTION_IMPLEMENTATION.value != 'xla':1054 # raise ValueError(1055 # 'For devices with GPU compute capability 7.x (see'1056 # ' https://developer.nvidia.com/cuda-gpus) the'1057 # ' --flash_attention_implementation must be set to "xla".'1058 # )1059 1060 notice = textwrap.wrap(1061 'Running AlphaFold 3. Please note that standard AlphaFold 3 model'1062 ' parameters are only available under terms of use provided at'1063 ' https://github.com/google-deepmind/alphafold3/blob/main/WEIGHTS_TERMS_OF_USE.md.'1064 ' If you do not agree to these terms and are using AlphaFold 3 derived'1065 ' model parameters, cancel execution of AlphaFold 3 inference with'1066 ' CTRL-C, and do not use the model parameters.',1067 break_long_words=False,1068 break_on_hyphens=False,1069 width=80,1070 )1071 print('\n' + '\n'.join(notice) + '\n')1072 1073 max_template_date = datetime.date.fromisoformat(_MAX_TEMPLATE_DATE.value)1074 if _RUN_DATA_PIPELINE.value:1075 if not _USE_MMSEQS.value:1076 expand_path = lambda x: replace_db_dir(x, DB_DIR.value)1077 data_pipeline_config = pipeline.DataPipelineConfig(1078 jackhmmer_binary_path=_JACKHMMER_BINARY_PATH.value,1079 nhmmer_binary_path=_NHMMER_BINARY_PATH.value,1080 hmmalign_binary_path=_HMMALIGN_BINARY_PATH.value,1081 hmmsearch_binary_path=_HMMSEARCH_BINARY_PATH.value,1082 hmmbuild_binary_path=_HMMBUILD_BINARY_PATH.value,1083 mmseqs_binary_path=_MMSEQS_BINARY_PATH.value,1084 small_bfd_database_path=expand_path(_SMALL_BFD_DATABASE_PATH.value),1085 small_bfd_z_value=_SMALL_BFD_Z_VALUE.value,1086 mgnify_database_path=expand_path(_MGNIFY_DATABASE_PATH.value),1087 mgnify_z_value=_MGNIFY_Z_VALUE.value,1088 uniprot_cluster_annot_database_path=expand_path(1089 _UNIPROT_CLUSTER_ANNOT_DATABASE_PATH.value1090 ),1091 uniprot_cluster_annot_z_value=_UNIPROT_CLUSTER_ANNOT_Z_VALUE.value,1092 uniref90_database_path=expand_path(_UNIREF90_DATABASE_PATH.value),1093 uniref90_z_value=_UNIREF90_Z_VALUE.value,1094 ntrna_database_path=expand_path(_NTRNA_DATABASE_PATH.value),1095 ntrna_z_value=_NTRNA_Z_VALUE.value,1096 rfam_database_path=expand_path(_RFAM_DATABASE_PATH.value),1097 rfam_z_value=_RFAM_Z_VALUE.value,1098 rna_central_database_path=expand_path(_RNA_CENTRAL_DATABASE_PATH.value),1099 rna_central_z_value=_RNA_CENTRAL_Z_VALUE.value,1100 pdb_database_path=expand_path(_PDB_DATABASE_PATH.value),1101 seqres_database_path=expand_path(_SEQRES_DATABASE_PATH.value),1102 jackhmmer_n_cpu=_JACKHMMER_N_CPU.value,1103 jackhmmer_max_parallel_shards=_JACKHMMER_MAX_PARALLEL_SHARDS.value,1104 jackhmmer_max_threads=_JACKHMMER_MAX_THREADS.value,1105 nhmmer_n_cpu=_NHMMER_N_CPU.value,1106 nhmmer_max_parallel_shards=_NHMMER_MAX_PARALLEL_SHARDS.value,1107 nhmmer_max_threads=_NHMMER_MAX_THREADS.value,1108 max_template_date=max_template_date,1109 use_mmseqs=_USE_MMSEQS.value,1110 mmseqs_options=_MMSEQS_OPTIONS.value,1111 result2msa_options=_R2MSA_OPTIONS.value,1112 )1113 else:1114 expand_path = lambda x: replace_mmseqs_db_dir(x, DB_DIR.value, MMSEQS_DB_DIR.value)1115 data_pipeline_config = pipeline.DataPipelineConfig(1116 jackhmmer_binary_path=_JACKHMMER_BINARY_PATH.value,1117 nhmmer_binary_path=_NHMMER_BINARY_PATH.value,1118 hmmalign_binary_path=_HMMALIGN_BINARY_PATH.value,1119 hmmsearch_binary_path=_HMMSEARCH_BINARY_PATH.value,1120 hmmbuild_binary_path=_HMMBUILD_BINARY_PATH.value,1121 mmseqs_binary_path=_MMSEQS_BINARY_PATH.value,1122 small_bfd_database_path=expand_path(_MMSEQS_SMALL_BFD_DATABASE_PATH.value),1123 mgnify_database_path=expand_path(_MMSEQS_MGNIFY_DATABASE_PATH.value),1124 uniprot_cluster_annot_database_path=expand_path(1125 _MMSEQS_UNIPROT_CLUSTER_ANNOT_DATABASE_PATH.value1126 ),1127 uniref90_database_path=expand_path(_MMSEQS_UNIREF90_DATABASE_PATH.value),1128 ntrna_database_path=expand_path(_NTRNA_DATABASE_PATH.value),1129 rfam_database_path=expand_path(_RFAM_DATABASE_PATH.value),1130 rna_central_database_path=expand_path(_RNA_CENTRAL_DATABASE_PATH.value),1131 pdb_database_path=expand_path(_PDB_DATABASE_PATH.value),1132 seqres_database_path=expand_path(_SEQRES_DATABASE_PATH.value),1133 mmseqs_n_cpu=_MMSEQS_N_CPU.value,1134 nhmmer_n_cpu=_NHMMER_N_CPU.value,1135 max_template_date=max_template_date,1136 use_mmseqs=_USE_MMSEQS.value,1137 use_mmseqs_gpu=_USE_MMSEQS_GPU.value,1138 mmseqs_options=_MMSEQS_OPTIONS.value,1139 result2msa_options=_R2MSA_OPTIONS.value,1140 ) 1141 else:1142 data_pipeline_config = None1143 1144 if _RUN_INFERENCE.value:1145 devices = jax.local_devices(backend='gpu')1146 print(1147 f'Found local devices: {devices}, using device {_GPU_DEVICE.value}:'1148 f' {devices[_GPU_DEVICE.value]}'1149 )1150 1151 print('Building model from scratch...')1152 model_runner = ModelRunner(1153 config=make_model_config(1154 flash_attention_implementation=typing.cast(1155 attention.Implementation, _FLASH_ATTENTION_IMPLEMENTATION.value1156 ),1157 num_diffusion_samples=_NUM_DIFFUSION_SAMPLES.value,1158 num_recycles=_NUM_RECYCLES.value,1159 return_embeddings=_SAVE_EMBEDDINGS.value,1160 return_distogram=_SAVE_DISTOGRAM.value,1161 ),1162 device=devices[_GPU_DEVICE.value],1163 model_dir=pathlib.Path(MODEL_DIR.value),1164 )1165 # Check we can load the model parameters before launching anything.1166 print('Checking that model parameters can be loaded...')1167 _ = model_runner.model_params1168 else:1169 model_runner = None1170 1171 num_fold_inputs = 01172 for fold_input in fold_inputs:1173 if _NUM_SEEDS.value is not None:1174 print(f'Expanding fold job {fold_input.name} to {_NUM_SEEDS.value} seeds')1175 fold_input = fold_input.with_multiple_seeds(_NUM_SEEDS.value)1176 process_fold_input(1177 fold_input=fold_input,1178 data_pipeline_config=data_pipeline_config,1179 model_runner=model_runner,1180 output_dir=os.path.join(_OUTPUT_DIR.value, fold_input.sanitised_name()),1181 buckets=tuple(int(bucket) for bucket in _BUCKETS.value),1182 ref_max_modified_date=max_template_date,1183 conformer_max_iterations=_CONFORMER_MAX_ITERATIONS.value,1184 resolve_msa_overlaps=_RESOLVE_MSA_OVERLAPS.value,1185 force_output_dir=_FORCE_OUTPUT_DIR.value,1186 )1187 num_fold_inputs += 11188 1189 print(f'Done running {num_fold_inputs} fold jobs.')1190 1191 1192if __name__ == '__main__':1193 flags.mark_flags_as_required(['output_dir'])1194 app.run(main)1195 