Team Ai
Modelpublic

OneScience-Group/AlphaFold3

sourceHugging Facecc-by-nc-sa-4.0updated 2mo agoView on Hugging Face
3likes254downloads
run_alphafold.py1195 linesDownload Raw Back to scripts
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