diff --git a/.github/workflows/github_actions.yml b/.github/workflows/github_actions.yml index ca37b3fa..1bc49788 100644 --- a/.github/workflows/github_actions.yml +++ b/.github/workflows/github_actions.yml @@ -38,7 +38,8 @@ jobs: use-mamba: true - run: | - pip install -e .[test] + # Fresh-interpreter compatibility tests import AF2's real JAX tree API. + pip install -e .[test] 'jax==0.5.3' - run: | pytest -n auto --dist loadfile test/unit @@ -70,7 +71,8 @@ jobs: use-mamba: true - run: | - pip install -e .[test] + # Match the CPU runtime used by the smoke-test matrix. + pip install -e .[test] 'jax==0.5.3' - run: | pytest -n auto --dist loadfile test/unit \ diff --git a/README.md b/README.md index 75436af3..adefa218 100644 --- a/README.md +++ b/README.md @@ -840,6 +840,10 @@ structure_inference_arguments: > **Note**: AlphaPulldown supports: `alphafold2`, `alphafold3`, and `alphalink` backends. +The legacy `--fold_backend=unifold` and `--use_unifold` options are disabled in +this release. The bundled runtime is the AlphaLink2 fork and does not provide +validated native UniFold inference. AlphaLink requires its own model weights. + ### Backend-specific flags You can pass backend CLI switches through `structure_inference_arguments`. Common options are listed below; keep or remove lines based on your needs. diff --git a/alphapulldown/folding_backend/__init__.py b/alphapulldown/folding_backend/__init__.py index a6543421..f84bfd94 100755 --- a/alphapulldown/folding_backend/__init__.py +++ b/alphapulldown/folding_backend/__init__.py @@ -31,7 +31,6 @@ def __init__(self): { "alphafold2": "alphapulldown.folding_backend.alphafold2_backend:AlphaFold2Backend", "alphafold3": "alphapulldown.folding_backend.alphafold3_backend:AlphaFold3Backend", - "unifold": "alphapulldown.folding_backend.unifold_backend:UnifoldBackend", "alphalink": "alphapulldown.folding_backend.alphalink_backend:AlphaLinkBackend", } ) @@ -73,6 +72,9 @@ def available_backends(self) -> List[str]: return sorted(ok) def _load_backend_class(self, backend_name: str) -> Type: + from alphapulldown.prediction.inference_flags import validate_backend_availability + + validate_backend_availability(backend_name) if backend_name not in self._BACKEND_REGISTRY: available = ", ".join(sorted(self._BACKEND_REGISTRY.keys())) raise NotImplementedError( @@ -132,4 +134,3 @@ def change_backend(backend_name: str, **backend_kwargs) -> None: """Change the backend for structure prediction.""" mgr = _get_manager() mgr.change_backend(backend_name, **backend_kwargs) - diff --git a/alphapulldown/folding_backend/alphafold2_backend.py b/alphapulldown/folding_backend/alphafold2_backend.py index 59c5bee6..98bf48ae 100644 --- a/alphapulldown/folding_backend/alphafold2_backend.py +++ b/alphapulldown/folding_backend/alphafold2_backend.py @@ -461,6 +461,36 @@ def _write_processed_template_debug_artifacts( ) +def _msa_depths( + evoformer_config: Mapping, + num_predictions: int, + *, + msa_depth=None, + msa_depth_scan: bool = False, +) -> list[tuple[int, int]]: + """``(num_msa, num_extra_msa)`` for each prediction of one model. + + A fixed ``--msa_depth`` applies to every prediction, with the extra MSA at + about four times the depth, as in the AlphaFold 2 configs. A scan spaces the + depths logarithmically from 16 (32 for the extra MSA) up to the model's own + defaults, which are read from ``evoformer_config`` before anything is changed. + """ + if msa_depth: + num_msa = int(msa_depth) + return [(num_msa, num_msa * 4)] * num_predictions + if not msa_depth_scan: + raise ValueError("either msa_depth or msa_depth_scan must be set") + msa_ranges = np.rint( + np.logspace(np.log10(16), np.log10(evoformer_config["num_msa"]), num_predictions) + ).astype(int) + extra_msa_ranges = np.rint( + np.logspace( + np.log10(32), np.log10(evoformer_config["num_extra_msa"]), num_predictions + ) + ).astype(int) + return [(int(num_msa), int(num_extra)) for num_msa, num_extra in zip(msa_ranges, extra_msa_ranges)] + + class AlphaFold2Backend(FoldingBackend): """ A backend to perform structure prediction using AlphaFold. @@ -496,7 +526,9 @@ def setup( model_names_custom : list, optional A list of strings that specify which models to run, default is None, meaning all 5 models will be used msa_depth : int or None, optional - A specific MSA depth to use, default is None. + A specific MSA depth to use, default is None. With either MSA depth + option every distinct depth gets its own model runner and config; + the model parameters are loaded once per model name and shared. allow_resume : bool, optional If set to True, resumes prediction from partially completed runs, default is True. dropout : bool, optional @@ -556,59 +588,56 @@ def setup( f"Provided model names {model_names_custom} not part of available {model_names + old_model_names}" ) - for model_name in model_names: - model_config = config.model_config(model_name) + def configured_model(name: str, num_msa=None, num_extra_msa=None): + """A fresh model config for one runner. + + ``RunModel`` keeps a reference to the config it is given and traces + the network from it on first use, so every runner needs its own copy. + This used to build one config per model name, register the same + runner once per MSA depth, and rewrite the shared config in place on + each iteration: every registered runner ended up pointing at one + object holding the LAST depth, so an ``--msa_depth_scan`` produced N + predictions at a single depth that differed only by seed. + """ + model_config = config.model_config(name) model_config.model.num_ensemble_eval = num_ensemble model_config["model"].update({"num_recycle": num_cycle}) if dropout: model_config.model.global_config.eval_dropout = True + if num_msa is not None: + model_config["model"]["embeddings_and_evoformer"].update( + {"num_msa": int(num_msa), "num_extra_msa": int(num_extra_msa)} + ) + return model_config + for model_name in model_names: + # The parameters are immutable and large: load them once per model + # name and share them between the runners built from it. model_params = data.get_model_haiku_params( model_name=model_name, data_dir=model_dir ) - model_runner = model.RunModel(model_config, model_params) - - if msa_depth_scan or msa_depth: - embeddings_and_evo = model_config["model"]["embeddings_and_evoformer"] - num_msa = embeddings_and_evo["num_msa"] - num_extra_msa = embeddings_and_evo["num_extra_msa"] - - msa_ranges = np.rint( - np.logspace( - np.log10(16), - np.log10(num_msa), - num_predictions_per_model, - ) - ).astype(int) - extra_msa_ranges = np.rint( - np.logspace( - np.log10(32), - np.log10(num_extra_msa), - num_predictions_per_model, - ) - ).astype(int) - - for i in range(num_predictions_per_model): - logging.debug(f"msa_depth is type : {type(msa_depth)} value: {msa_depth}") - logging.debug(f"msa_depth_scan is type: {type(msa_depth_scan)} value: {msa_depth_scan}") - if msa_depth or msa_depth_scan: - if msa_depth: - num_msa = int(msa_depth) - # approx. 4x the number of msa, as in the AF2 config file - num_extra_msa = int(num_msa * 4) - elif msa_depth_scan: - num_msa = int(msa_ranges[i]) - num_extra_msa = int(extra_msa_ranges[i]) - - # Conversion to int before because num_msa could be None - embeddings_and_evo.update( - {"num_msa": num_msa, "num_extra_msa": num_extra_msa} - ) - - model_runners[f"{model_name}_pred_{i}_msa_{num_msa}"] = model_runner - else: + if not (msa_depth or msa_depth_scan): + model_runner = model.RunModel(configured_model(model_name), model_params) + for i in range(num_predictions_per_model): model_runners[f"{model_name}_pred_{i}"] = model_runner + continue + + depths = _msa_depths( + configured_model(model_name)["model"]["embeddings_and_evoformer"], + num_predictions_per_model, + msa_depth=msa_depth, + msa_depth_scan=msa_depth_scan, + ) + runners_by_depth: Dict[tuple[int, int], Any] = {} + for i, depth in enumerate(depths): + if depth not in runners_by_depth: + runners_by_depth[depth] = model.RunModel( + configured_model(model_name, *depth), model_params + ) + model_runners[f"{model_name}_pred_{i}_msa_{depth[0]}"] = ( + runners_by_depth[depth] + ) return {"model_runners": model_runners} diff --git a/alphapulldown/folding_backend/unifold_backend.py b/alphapulldown/folding_backend/unifold_backend.py index 50755568..05359b43 100644 --- a/alphapulldown/folding_backend/unifold_backend.py +++ b/alphapulldown/folding_backend/unifold_backend.py @@ -1,111 +1,26 @@ -""" Implements structure prediction backend using UniFold. +"""Compatibility entrypoint for the unavailable legacy UniFold backend. - Copyright (c) 2024 European Molecular Biology Laboratory - - Author: Valentin Maurer +The packaged ``unifold`` namespace belongs to AlphaLink2. It lacks the old +inference helpers and adds crosslink layers to the network, so routing native +UniFold checkpoints through it is not a supported substitute for UniFold. """ -from typing import Dict -from alphapulldown.objects import MultimericObject +from alphapulldown.prediction.inference_flags import validate_backend_availability from .folding_backend import FoldingBackend class UnifoldBackend(FoldingBackend): - """ - A backend class for running protein structure predictions using the UniFold model. - """ - @staticmethod - def setup( - model_name: str, - model_dir: str, - output_dir: str, - multimeric_object: MultimericObject, - **kwargs, - ) -> Dict: - """ - Initializes and configures a UniFold model runner. - - Parameters - ---------- - model_name : str - The name of the model to use for prediction. - model_dir : str - The directory where the model files are located. - output_dir : str - The directory where the prediction outputs will be saved. - multimeric_object : MultimericObject - An object containing the description and features of the - multimeric protein to predict. - **kwargs : dict - Additional keyword arguments for model configuration. - - Returns - ------- - Dict - A dictionary containing the model runner, arguments, and configuration. - """ - from unifold.config import model_config - from unifold.inference import config_args, unifold_config_model - - configs = model_config(model_name) - general_args = config_args( - model_dir, target_name=multimeric_object.description, output_dir=output_dir - ) - model_runner = unifold_config_model(general_args) - - return { - "model_runner": model_runner, - "model_args": general_args, - "model_config": configs, - } + """Preserve imports while rejecting legacy calls with an actionable error.""" - def predict( - self, - model_runner, - model_args, - model_config: Dict, - multimeric_object: MultimericObject, - random_seed: int = 42, - **kwargs, - ) -> None: - """ - Predicts the structure of proteins using configured UniFold models. - - Parameters - ---------- - model_runner - The configured model runner for predictions obtained - from :py:meth:`UnifoldBackend.setup`. - model_args - Arguments used for running the UniFold prediction obtained from - from :py:meth:`UnifoldBackend.setup`. - model_config : Dict - Configuration dictionary for the UniFold model obtained from - from :py:meth:`UnifoldBackend.setup`. - multimeric_object : MultimericObject - An object containing the features of the multimeric protein to predict. - random_seed : int, optional - The random seed for prediction reproducibility, default is 42. - **kwargs : dict - Additional keyword arguments for prediction. - """ - from unifold.dataset import process_ap - from unifold.inference import unifold_predict - - processed_features, _ = process_ap( - config=model_config, - features=multimeric_object.feature_dict, - mode="predict", - labels=None, - seed=random_seed, - batch_idx=None, - data_idx=None, - is_distillation=False, - ) - unifold_predict(model_runner, model_args, processed_features) + @staticmethod + def setup(*args, **kwargs): + validate_backend_availability("unifold") - return None + @staticmethod + def predict(*args, **kwargs): + validate_backend_availability("unifold") - def postprocess(**kwargs) -> None: - return None + @staticmethod + def postprocess(*args, **kwargs): + validate_backend_availability("unifold") diff --git a/alphapulldown/objects.py b/alphapulldown/objects.py index 40409060..bd782d98 100644 --- a/alphapulldown/objects.py +++ b/alphapulldown/objects.py @@ -434,6 +434,32 @@ def make_mmseq_features( MonomericObject.zip_msa_files(output_dir) +def _join_template_sequences(slices: List[Any]) -> np.ndarray: + """Join each template's region fragments, in region order, into one sequence. + + ``template_sequence`` holds one string per template, sliced per region like the + template arrays are. Unlike those arrays there is no residue axis to + concatenate along, so the fragments are joined per template here. + """ + counts = {len(fragments) for fragments in slices} + if len(counts) > 1: + raise ValueError( + "Region slices disagree on the number of templates: " + + ", ".join(str(len(fragments)) for fragments in slices) + ) + + def as_bytes(fragment: Any) -> bytes: + if isinstance(fragment, (bytes, bytearray)): + return bytes(fragment) + return str(fragment).encode("utf-8") + + joined = [ + b"".join(as_bytes(fragment) for fragment in fragments) + for fragments in zip(*slices) + ] + return np.array(joined, dtype=object) + + class ChoppedObject(MonomericObject): """A monomeric object chopped into specified regions.""" @@ -606,6 +632,15 @@ def concatenate_sliced_feature_dict( axis=axis_map["template_confidence_scores"], ) + # Template sequences are sliced per region like the template arrays, so + # they have to be joined per template too. They used to sit in the skip + # set below, which kept the FIRST region's fragment only: a multi-region + # chop then carried a template_sequence as long as its first region. + if "template_sequence" in out: + out["template_sequence"] = _join_template_sequences( + [sd["template_sequence"] for sd in slice_dicts] + ) + skip = { "template_domain_names", "template_sequence", "template_sum_probs", "template_release_date", diff --git a/alphapulldown/prediction/fold_preparation.py b/alphapulldown/prediction/fold_preparation.py index c202ce5e..9dfc6e44 100644 --- a/alphapulldown/prediction/fold_preparation.py +++ b/alphapulldown/prediction/fold_preparation.py @@ -24,6 +24,10 @@ # Most filesystems reject paths longer than this. _MAX_PATH_LENGTH = 4096 +# Where --save_features_for_multimeric_object writes the merged features, inside +# the fold's own output directory. +MULTIMERIC_FEATURES_FILENAME = "multimeric_object_features.pkl" + def _interactor_description(interactor: Any) -> str | None: """Name under which this interactor's feature metadata was written, if any.""" @@ -46,6 +50,23 @@ def _output_directory_for(description: str, output_dir: str) -> str: return os.path.join(output_dir, "_and_".join(fragments)) +def _save_multimeric_features(object_to_model: Any, output_dir: str) -> str: + """Pickle the merged multimer features next to the fold's outputs. + + The features are read off the object that was just built. This used to read + ``feature_dict`` off the ``MultimericObject`` *class*, where no such attribute + exists, so ``--save_features_for_multimeric_object`` raised ``AttributeError`` + on every run. It also ran before the output directory was created, and before + ``--use_ap_style`` had chosen the fold's directory, so even a correct lookup + would have failed on a fresh output tree or written one level too high. + """ + destination = os.path.join(output_dir, MULTIMERIC_FEATURES_FILENAME) + with open(destination, "wb") as handle: + pickle.dump(object_to_model.feature_dict, handle) + logging.info("Saved multimeric object features to %s", destination) + return destination + + def _copy_feature_metadata(interactors: List[Any], flags: Any, output_dir: str) -> None: """Copy each interactor's most recent feature metadata, decompressing `.json.xz`. @@ -116,11 +137,6 @@ def prepare_fold( hb_allowance=flags.hb_allowance, plddt_threshold=flags.plddt_threshold, ) - if flags.save_features_for_multimeric_object: - with open( - os.path.join(output_dir, "multimeric_object_features.pkl"), "wb" - ) as handle: - pickle.dump(MultimericObject.feature_dict, handle) else: object_to_model = interactors[0] object_to_model.input_seqs = [object_to_model.sequence] @@ -136,5 +152,8 @@ def prepare_fold( ) os.makedirs(output_dir, exist_ok=True) + if len(interactors) > 1 and flags.save_features_for_multimeric_object: + _save_multimeric_features(object_to_model, output_dir) + _copy_feature_metadata(interactors, flags, output_dir) return object_to_model, output_dir diff --git a/alphapulldown/prediction/inference_flags.py b/alphapulldown/prediction/inference_flags.py index d9e803d1..49b47eaf 100644 --- a/alphapulldown/prediction/inference_flags.py +++ b/alphapulldown/prediction/inference_flags.py @@ -36,6 +36,12 @@ ALPHALINK_EXTRA_FLAGS = frozenset({"crosslinks"}) +UNIFOLD_UNAVAILABLE_REASON = ( + "UniFold is unavailable in this release: the bundled runtime is an AlphaLink " + "fork, not a supported UniFold inference runtime. Use --fold_backend=alphafold2 " + "or --fold_backend=alphafold3; use --fold_backend=alphalink only with AlphaLink weights." +) + AF3_FLAGS = frozenset({ "jax_compilation_cache_dir", "buckets", "flash_attention_implementation", "num_diffusion_samples", "num_seeds", "debug_templates", "debug_msas", @@ -50,8 +56,15 @@ } +def validate_backend_availability(backend_name: str) -> None: + """Reject the legacy UniFold selector before reading inputs or model weights.""" + if backend_name == "unifold": + raise ValueError(UNIFOLD_UNAVAILABLE_REASON) + + def unsupported_flags(backend_name: str, present: Iterable[str]) -> list[str]: """Flag names the named backend does not accept. Unknown backend: nothing to say.""" + validate_backend_availability(backend_name) allowed = FLAGS_BY_BACKEND.get(backend_name) if allowed is None: return [] @@ -106,15 +119,24 @@ def unsupported_flags(backend_name: str, present: Iterable[str]) -> list[str]: } +def model_name_for_backend(flags: Any) -> str: + """The model configuration a backend starts from. + + AlphaFold 2 multimers are switched to ``multimer`` by the caller, per object. + """ + validate_backend_availability(flags.fold_backend) + if flags.fold_backend == "alphalink": + return "multimer_af2_crop" + return "monomer_ptm" + + def model_flags(flags: Any) -> Dict[str, Any]: """Configuration for ``backend.setup`` built from the parsed invocation.""" configuration = { key: getattr(flags, attribute) for key, attribute in _MODEL_FLAG_SOURCES.items() } - configuration["model_name"] = ( - "multimer_af2_crop" if flags.fold_backend == "alphalink" else "monomer_ptm" - ) + configuration["model_name"] = model_name_for_backend(flags) return configuration diff --git a/alphapulldown/prediction/jax_devices.py b/alphapulldown/prediction/jax_devices.py new file mode 100644 index 00000000..ad65055b --- /dev/null +++ b/alphapulldown/prediction/jax_devices.py @@ -0,0 +1,34 @@ +"""Initialise JAX's GPU backend before anything else can claim the device. + +The prediction commands call :func:`initialise_jax_gpu_backend` at import time, +before the folding backends import TensorFlow and OpenMM, so that JAX initialises +CUDA first and its device memory is not pre-empted by a library that merely got +imported earlier. That used to be a bare ``jax.local_devices(backend='gpu')``, +which raises when JAX has no GPU backend, so ``--help`` and the head-node flag +validation the workflow documentation recommends failed on any machine without a +GPU. A missing GPU is reported here, not fatal: whether inference can run +without one is each backend's decision (AlphaFold 3 refuses, AlphaFold 2 falls +back to CPU). +""" + +from __future__ import annotations + +from typing import Any, List + +from absl import logging + + +def initialise_jax_gpu_backend() -> List[Any]: + """Return the GPU devices JAX can see, or an empty list when it sees none.""" + import jax + + try: + devices = list(jax.local_devices(backend="gpu")) + except RuntimeError as exc: + logging.warning( + "JAX found no usable GPU backend (%s); only CPU inference is possible.", + exc, + ) + return [] + logging.info("JAX GPU devices: %s", devices) + return devices diff --git a/alphapulldown/scripts/run_multimer_jobs.py b/alphapulldown/scripts/run_multimer_jobs.py index a8756d7c..2b5cb0cc 100644 --- a/alphapulldown/scripts/run_multimer_jobs.py +++ b/alphapulldown/scripts/run_multimer_jobs.py @@ -10,8 +10,8 @@ from absl import app, logging, flags import os import sys -import jax -gpus = jax.local_devices(backend='gpu') +# Importing the prediction command initialises JAX's GPU backend first, tolerating +# a machine without a GPU; see alphapulldown.prediction.jax_devices. from alphapulldown.scripts.run_structure_prediction import FLAGS from alphapulldown.utils.modelling_setup import parse_fold from alphapulldown.utils.output_paths import derive_af3_job_name_from_json @@ -27,12 +27,10 @@ flags.DEFINE_string("alphalink_weight", None, "Path to AlphaLink neural network weights") flags.DEFINE_string("unifold_param", None, "Path to UniFold neural network weights") flags.DEFINE_boolean("use_unifold", False, - "Whether unifold models are going to be used. Default it False") + "Legacy option. UniFold is unavailable in this release.") flags.DEFINE_boolean("use_alphalink", False, "Whether alphalink models are going to be used. Default it False") -flags.DEFINE_enum("unifold_model_name", "multimer_af2", - ["multimer_af2", "multimer_ft", "multimer", "multimer_af2_v3", "multimer_af2_model45_v3"], - "choose unifold model structure") +# --unifold_model_name is defined by run_structure_prediction and shared here. flags.DEFINE_integer("job_index", None, "index of sequence in the fasta file, starting from 1") flags.DEFINE_boolean("dry_run", False, "Report number of jobs that would be run and exit without running them") @@ -76,6 +74,9 @@ def _resolve_af3_wrapper_output_dir( def main(argv): FLAGS(argv) + from alphapulldown.prediction.inference_flags import validate_backend_availability + + validate_backend_availability("unifold" if FLAGS.use_unifold else FLAGS.fold_backend) protein_lists = FLAGS.protein_lists if FLAGS.mode == "all_vs_all": protein_lists = [FLAGS.protein_lists[0], FLAGS.protein_lists[0]] @@ -99,13 +100,11 @@ def main(argv): script_path = os.path.join(parent_dir, "run_structure_prediction.py") base_command = [sys.executable, script_path] - # Use the specified fold_backend, only override for alphalink/unifold + # Use the specified fold_backend, only override for AlphaLink. fold_backend = FLAGS.fold_backend model_dir = FLAGS.data_dir if FLAGS.use_alphalink: fold_backend, model_dir = "alphalink", FLAGS.alphalink_weight - elif FLAGS.use_unifold: - fold_backend, model_dir = "unifold", FLAGS.unifold_param af3_use_ap_style = FLAGS.use_ap_style if fold_backend == "alphafold3" and FLAGS["use_ap_style"].using_default_value: diff --git a/alphapulldown/scripts/run_structure_prediction.py b/alphapulldown/scripts/run_structure_prediction.py index 3721f282..7fa34f20 100644 --- a/alphapulldown/scripts/run_structure_prediction.py +++ b/alphapulldown/scripts/run_structure_prediction.py @@ -9,8 +9,12 @@ """ import pickle -import jax -gpus = jax.local_devices(backend='gpu') +from alphapulldown.prediction.jax_devices import initialise_jax_gpu_backend + +# Initialise JAX's GPU backend before the folding backends import TensorFlow and +# OpenMM, so JAX claims the device first. Tolerates a machine without a GPU, so +# --help and flag validation work on a login node. +gpus = initialise_jax_gpu_backend() from absl import flags, app import os from os import makedirs @@ -132,6 +136,11 @@ # AlphaLink2 settings flags.DEFINE_string('crosslinks', None, 'Path to crosslink information pickle for AlphaLink.') +# Keep the legacy flag parseable so old invocations get the backend's actionable error. +flags.DEFINE_string( + 'unifold_model_name', 'multimer_af2', + 'Legacy option. UniFold is unavailable in this release.') + # AlphaFold3 settings # JAX inference performance tuning. flags.DEFINE_string( @@ -218,7 +227,7 @@ # Global settings flags.DEFINE_string('protein_delimiter', '+', 'Delimiter for proteins of a single fold.') flags.DEFINE_string('fold_backend', 'alphafold2', - 'Folding backend that should be used for structure prediction.') + 'Folding backend: alphafold2, alphafold3, or alphalink. UniFold is unavailable.') flags.DEFINE_boolean( 'debug_templates', False, 'If set, save backend-specific template debug artifacts. AF3 writes generated' diff --git a/alphapulldown/utils/distogram_parser.py b/alphapulldown/utils/distogram_parser.py index 465f7f8a..68f73db2 100644 --- a/alphapulldown/utils/distogram_parser.py +++ b/alphapulldown/utils/distogram_parser.py @@ -21,28 +21,39 @@ class distogram_parser: def __init__(self): pass + @staticmethod + def select_top_ranked_pickle(directory): + """Path of the result pickle in ``directory`` with the highest ranking confidence. + + Returns ``(path, ranking_confidence)``, or ``(None, 0.0)`` when the directory + holds no pickle with a positive ranking confidence. Only the path and the + score are kept while scanning, so the pickles are not all held in memory. + """ + top_ranked = (None, 0.0) + for fn in sorted(glob.glob(os.path.join(directory, "*.pkl"))): + with open(fn, 'rb') as ifile: + d = pickle.load(ifile) + ranking_confidence = d.get('ranking_confidence', 0) + if ranking_confidence > top_ranked[-1]: + top_ranked = (fn, ranking_confidence) + return top_ranked + def get_contacts(self, directory, distance=8, pbtycutoff=0.8, cross_only=True, verbose=False): """ - selects from datadir a pkl/distogram corresponding to a top-ranked model + selects from directory a pkl/distogram corresponding to a top-ranked model """ - top_ranked_dgram = (None, None, 0.0) - for fn in glob.glob(os.path.join(datadir, "*.pkl")): - with open(fn, 'rb') as ifile: - d=pickle.load(ifile) - if d.get('ranking_confidence',0)>top_ranked_dgram[-1]: - top_ranked_dgram = (fn, d, d.get('ranking_confidence',0)) + top_ranked_fn, top_ranked_confidence = self.select_top_ranked_pickle(directory) - if top_ranked_dgram[0] is None: return [] + if top_ranked_fn is None: return [] if verbose: - print(f"Selected {os.path.basename(top_ranked_dgram[0])} with ranking confidence {top_ranked_drgam[-1]:.2f}") + print(f"Selected {os.path.basename(top_ranked_fn)} with ranking confidence {top_ranked_confidence:.2f}") - d = top_ranked_dgram[1] - - # reparse top ditogram; avoids storing all pickles in memory - with open(fn, 'rb') as ifile: - d=pickle.load(ifile) + # Re-read the SELECTED pickle. This used to re-read whichever file the + # scan visited last, which is the top-ranked one only by coincidence. + with open(top_ranked_fn, 'rb') as ifile: + d = pickle.load(ifile) chain_ids = string.ascii_uppercase asym_id=[] chain_lens = [] @@ -51,8 +62,6 @@ def get_contacts(self, directory, distance=8, pbtycutoff=0.8, cross_only=True, v chain_lens.append(len(_seq)) chain_lens = np.array(chain_lens) - assembly_num_chains = len(d['seqs']) - bin_edges = d['distogram']['bin_edges'] # apply softmax (scipy equivalent) @@ -66,9 +75,12 @@ def get_contacts(self, directory, distance=8, pbtycutoff=0.8, cross_only=True, v distance = np.clip(distance, 3, 20) - bin_idx=np.max(np.where(bin_edges= 3.10`) but GitHub Actions does not exercise The coverage artifact is useful as an audit input, but it does not prove workflow correctness by itself. In particular, `python test/tools/check_function_coverage.py --report-only` highlights functions that were never executed in CI and should be treated as follow-up audit items, not automatic release blockers. diff --git a/test/integration/test_backend_test_isolation.py b/test/integration/test_backend_test_isolation.py new file mode 100644 index 00000000..41400331 --- /dev/null +++ b/test/integration/test_backend_test_isolation.py @@ -0,0 +1,32 @@ +"""Backend stubs must not replace real dependencies in subsequent tests.""" + +import os +import subprocess +import sys +from pathlib import Path + + +def test_af3_backend_stubs_do_not_leak_into_modelcif_tests(): + root = Path(__file__).resolve().parents[2] + env = os.environ.copy() + env.pop("PYTEST_ADDOPTS", None) + result = subprocess.run( + [ + sys.executable, + "-m", + "pytest", + "-q", + "--tb=short", + "-n", "0", + "test/unit/test_alphafold3_backend_helpers.py::" + "test_output_name_helpers_compact_and_normalise_fragments", + "test/unit/test_af3_modelcif.py::" + "test_augment_real_af3_modelcif_preserves_comments_and_is_modelcif_readable", + ], + cwd=root, + env=env, + capture_output=True, + text=True, + timeout=60, + ) + assert result.returncode == 0, result.stdout + result.stderr diff --git a/test/integration/test_pytest_collection.py b/test/integration/test_pytest_collection.py new file mode 100644 index 00000000..87687bc9 --- /dev/null +++ b/test/integration/test_pytest_collection.py @@ -0,0 +1,33 @@ +"""CI must execute every parametrized case, including a failing later case.""" + +import os +from pathlib import Path +import subprocess +import sys + +import pytest + + +@pytest.mark.parametrize("parallel", [False, True]) +def test_repository_hooks_preserve_parameter_identity(tmp_path, parallel): + root = Path(__file__).resolve().parents[2] + (tmp_path / "conftest.py").write_text((root / "conftest.py").read_text()) + (tmp_path / "pytest.ini").write_text("[pytest]\n") + (tmp_path / "test_parameters.py").write_text( + 'import pytest\n' + '@pytest.mark.parametrize("value", [0, 1])\n' + 'def test_value(value):\n' + ' assert value == 0\n' + ) + command = [sys.executable, "-m", "pytest", "-q", "--tb=no"] + if parallel: + command += ["-n", "2", "--dist", "loadfile"] + environment = os.environ.copy() + environment.pop("PYTEST_ADDOPTS", None) + result = subprocess.run( + command, cwd=tmp_path, env=environment, + capture_output=True, text=True, timeout=60, + ) + + assert result.returncode == 1, result.stdout + result.stderr + assert "1 failed, 1 passed" in result.stdout diff --git a/test/unit/test_alphafold2_backend_helpers.py b/test/unit/test_alphafold2_backend_helpers.py index b5a25a42..9b400672 100644 --- a/test/unit/test_alphafold2_backend_helpers.py +++ b/test/unit/test_alphafold2_backend_helpers.py @@ -535,6 +535,11 @@ def test_setup_configures_model_runners_and_validates_custom_names(af2_backend_m assert runner.config["model"]["num_recycle"] == 5 assert runner.config.model.global_config.eval_dropout is True assert runner.params["data_dir"] == "/models" + # Each depth is its own runner with its own config; the two must not share. + deeper = runners["model_1_multimer_pred_1_msa_64"] + assert deeper is not runner + assert runner.config["model"]["embeddings_and_evoformer"]["num_msa"] == 16 + assert deeper.config["model"]["embeddings_and_evoformer"]["num_msa"] == 64 with pytest.raises(Exception, match="Provided model names"): af2_backend_module.AlphaFold2Backend.setup( @@ -1191,3 +1196,103 @@ def test_postprocess_uses_ptm_threshold_for_best_monomer_relaxation( assert (tmp_path / "relaxed_modelA.pdb").read_text(encoding="utf-8") == "RELAXED:mono" assert (tmp_path / "ranked_0.pdb").read_text(encoding="utf-8") == "RELAXED:mono" + + +def test_msa_depth_scan_gives_every_depth_its_own_runner_and_config( + af2_backend_module, monkeypatch +): + """One RunModel was registered under every depth while its single config was + rewritten in place, so a scan ran N predictions at the LAST depth and differed + only by seed. The parameters are still loaded once per model name.""" + data_mod = sys.modules["alphafold.model.data"] + loads = [] + + def counting_params(model_name, data_dir): + loads.append(model_name) + return {"model_name": model_name, "data_dir": data_dir} + + monkeypatch.setattr(data_mod, "get_model_haiku_params", counting_params) + + configured = af2_backend_module.AlphaFold2Backend.setup( + model_name="multimer", + num_cycle=3, + model_dir="/models", + num_predictions_per_model=3, + msa_depth_scan=True, + model_names_custom=["model_1_multimer"], + ) + + runners = configured["model_runners"] + depths = { + name: ( + runner.config["model"]["embeddings_and_evoformer"]["num_msa"], + runner.config["model"]["embeddings_and_evoformer"]["num_extra_msa"], + ) + for name, runner in runners.items() + } + assert depths == { + "model_1_multimer_pred_0_msa_16": (16, 32), + "model_1_multimer_pred_1_msa_32": (32, 91), + "model_1_multimer_pred_2_msa_64": (64, 256), + } + assert len({id(runner) for runner in runners.values()}) == 3 + assert loads == ["model_1_multimer"] + assert all( + runner.params == {"model_name": "model_1_multimer", "data_dir": "/models"} + for runner in runners.values() + ) + + +def test_fixed_msa_depth_shares_one_runner_at_that_depth(af2_backend_module): + configured = af2_backend_module.AlphaFold2Backend.setup( + model_name="multimer", + num_cycle=3, + model_dir="/models", + num_predictions_per_model=2, + msa_depth=48, + model_names_custom=["model_1_multimer"], + ) + + runners = configured["model_runners"] + assert sorted(runners) == [ + "model_1_multimer_pred_0_msa_48", + "model_1_multimer_pred_1_msa_48", + ] + first, second = runners.values() + assert first is second + evoformer = first.config["model"]["embeddings_and_evoformer"] + assert (evoformer["num_msa"], evoformer["num_extra_msa"]) == (48, 192) + + +def test_setup_without_msa_depth_options_shares_one_runner_per_model(af2_backend_module): + configured = af2_backend_module.AlphaFold2Backend.setup( + model_name="multimer", + num_cycle=3, + model_dir="/models", + num_predictions_per_model=2, + ) + + runners = configured["model_runners"] + assert sorted(runners) == [ + "model_1_multimer_v3_pred_0", + "model_1_multimer_v3_pred_1", + "model_2_multimer_v3_pred_0", + "model_2_multimer_v3_pred_1", + ] + assert runners["model_1_multimer_v3_pred_0"] is runners["model_1_multimer_v3_pred_1"] + assert runners["model_1_multimer_v3_pred_0"] is not runners["model_2_multimer_v3_pred_0"] + evoformer = runners["model_1_multimer_v3_pred_0"].config["model"]["embeddings_and_evoformer"] + assert (evoformer["num_msa"], evoformer["num_extra_msa"]) == (64, 256) + + +def test_msa_depths_follow_the_fixed_depth_or_the_logarithmic_scan(af2_backend_module): + evoformer = {"num_msa": 64, "num_extra_msa": 256} + + assert af2_backend_module._msa_depths(evoformer, 2, msa_depth=48) == [(48, 192), (48, 192)] + assert af2_backend_module._msa_depths(evoformer, 3, msa_depth_scan=True) == [ + (16, 32), + (32, 91), + (64, 256), + ] + with pytest.raises(ValueError): + af2_backend_module._msa_depths(evoformer, 2) diff --git a/test/unit/test_alphafold3_backend_helpers.py b/test/unit/test_alphafold3_backend_helpers.py index af267538..cc443fc0 100644 --- a/test/unit/test_alphafold3_backend_helpers.py +++ b/test/unit/test_alphafold3_backend_helpers.py @@ -5,6 +5,7 @@ import types from pathlib import Path from types import SimpleNamespace +from unittest.mock import patch import numpy as np import pytest @@ -31,7 +32,7 @@ def _package(name: str) -> types.ModuleType: return module -def _install_alphafold3_backend_stubs(tmp_path: Path) -> None: +def _install_alphafold3_backend_stubs(tmp_path: Path, monkeypatch) -> None: for module_name in list(sys.modules): if module_name == "alphafold3" or module_name.startswith("alphafold3."): sys.modules.pop(module_name, None) @@ -40,15 +41,21 @@ def _install_alphafold3_backend_stubs(tmp_path: Path) -> None: import jax # type: ignore if not hasattr(jax, "Device"): - jax.Device = type("Device", (), {}) + monkeypatch.setattr(jax, "Device", type("Device", (), {}), raising=False) if not hasattr(jax, "tree_map"): - jax.tree_map = jax.tree_util.tree_map + monkeypatch.setattr(jax, "tree_map", jax.tree_util.tree_map, raising=False) if not hasattr(jax, "device_put"): - jax.device_put = lambda value, device=None: value + monkeypatch.setattr( + jax, "device_put", lambda value, device=None: value, raising=False + ) if not hasattr(jax, "jit"): - jax.jit = lambda func, device=None: func + monkeypatch.setattr( + jax, "jit", lambda func, device=None: func, raising=False + ) if not hasattr(jax, "random"): - jax.random = SimpleNamespace(PRNGKey=lambda seed: seed) + monkeypatch.setattr( + jax, "random", SimpleNamespace(PRNGKey=lambda seed: seed), raising=False + ) except Exception: # pragma: no cover - conftest already installs a stub pass @@ -394,17 +401,20 @@ class Config: @pytest.fixture(scope="module") def af3_backend_module(tmp_path_factory): tmp_path = tmp_path_factory.mktemp("af3_backend_stubs") - _install_alphafold3_backend_stubs(tmp_path) - sys.modules.pop("alphapulldown.folding_backend.alphafold3_backend", None) - spec = importlib.util.spec_from_file_location( - "alphapulldown.folding_backend.alphafold3_backend", - MODULE_PATH, - ) - module = importlib.util.module_from_spec(spec) - sys.modules[spec.name] = module - assert spec.loader is not None - spec.loader.exec_module(module) - return module + # Other modules exercise the real AF3 parser. Restore imports and JAX + # attributes even if loading the backend or one of its tests fails. + with patch.dict(sys.modules), pytest.MonkeyPatch.context() as monkeypatch: + _install_alphafold3_backend_stubs(tmp_path, monkeypatch) + sys.modules.pop("alphapulldown.folding_backend.alphafold3_backend", None) + spec = importlib.util.spec_from_file_location( + "alphapulldown.folding_backend.alphafold3_backend", + MODULE_PATH, + ) + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + assert spec.loader is not None + spec.loader.exec_module(module) + yield module class FakeChainsTable: diff --git a/test/unit/test_distogram_parser.py b/test/unit/test_distogram_parser.py index 4ca0bc07..f5b41163 100644 --- a/test/unit/test_distogram_parser.py +++ b/test/unit/test_distogram_parser.py @@ -1,38 +1,94 @@ -import pickle - -import numpy as np - -import alphapulldown.utils.distogram_parser as distogram_parser_module +"""Contact extraction from AlphaFold 2 distograms. +``get_contacts`` used to reference a module global ``datadir`` that never existed, +so every call raised ``NameError``; and after choosing the top-ranked pickle it +re-read whichever file the scan had visited LAST. +""" -def test_get_contacts_returns_empty_list_when_no_pickles_exist(monkeypatch, tmp_path): - monkeypatch.setattr(distogram_parser_module, "datadir", str(tmp_path), raising=False) +import pickle - parser = distogram_parser_module.distogram_parser() +import numpy as np +import pytest - assert parser.get_contacts("ignored") == [] +from alphapulldown.utils.distogram_parser import distogram_parser -def test_get_contacts_extracts_top_ranked_inter_chain_contact(monkeypatch, tmp_path): - logits = np.full((4, 4, 3), -10.0, dtype=np.float32) - logits[0, 2, 0] = 10.0 - logits[2, 0, 0] = 10.0 - payload = { - "ranking_confidence": 0.9, +def _payload(ranking_confidence, *, contact): + logits = np.full((4, 4, 4), -10.0, dtype=np.float32) + if contact: + logits[0, 2, 0] = 10.0 + logits[2, 0, 0] = 10.0 + return { + "ranking_confidence": ranking_confidence, "seqs": ["AA", "BB"], "distogram": { "bin_edges": np.array([4.0, 8.0, 12.0], dtype=np.float32), "logits": logits, }, } - with open(tmp_path / "result_model.pkl", "wb") as handle: + + +def _write(path, payload): + with open(path, "wb") as handle: pickle.dump(payload, handle) - monkeypatch.setattr(distogram_parser_module, "datadir", str(tmp_path), raising=False) - parser = distogram_parser_module.distogram_parser() - contacts = parser.get_contacts("ignored", distance=9, pbtycutoff=0.5, cross_only=True) + +def test_get_contacts_returns_empty_list_when_no_pickles_exist(tmp_path): + assert distogram_parser().get_contacts(str(tmp_path)) == [] + + +def test_get_contacts_reads_the_directory_it_is_given(tmp_path): + _write(tmp_path / "result_model.pkl", _payload(0.9, contact=True)) + + contacts = distogram_parser().get_contacts( + str(tmp_path), distance=9, pbtycutoff=0.5, cross_only=True + ) assert len(contacts) == 1 assert contacts[0][0] == (1, "A") assert contacts[0][1] == (1, "B") assert contacts[0][2] > 0.99 + + +def test_get_contacts_reads_the_top_ranked_model_not_the_last_file_scanned(tmp_path, capsys): + _write(tmp_path / "result_model_1.pkl", _payload(0.9, contact=True)) + _write(tmp_path / "result_model_2.pkl", _payload(0.5, contact=False)) + + contacts = distogram_parser().get_contacts( + str(tmp_path), distance=9, pbtycutoff=0.5, verbose=True + ) + + assert len(contacts) == 1 + assert "Selected result_model_1.pkl with ranking confidence 0.90" in capsys.readouterr().out + + +def test_select_top_ranked_pickle_needs_a_positive_ranking_confidence(tmp_path): + _write(tmp_path / "result_model_1.pkl", {"ranking_confidence": 0.0}) + _write(tmp_path / "result_model_2.pkl", {"seqs": []}) + + assert distogram_parser.select_top_ranked_pickle(str(tmp_path)) == (None, 0.0) + + +@pytest.mark.parametrize("cutoff, contact_bin, expected", [ + (8.0, 18, True), # The entire 7.625--7.9375 A bin is below the cutoff. + (8.0, 19, False), # The bin crossing the cutoff is not wholly below it. + (7.9375, 18, True), + (3.0, 0, True), +]) +def test_contacts_sum_all_bins_wholly_below_cutoff(tmp_path, cutoff, contact_bin, expected): + # Match real AF2 output: 64 probability bins separated by 63 boundaries. + edges = np.linspace(2.3125, 21.6875, 63) + logits = np.full((2, 2, 64), -30.0) + logits[:, :, contact_bin] = 30.0 + _write(tmp_path / "result_model.pkl", { + "ranking_confidence": 0.9, + "seqs": ["A", "B"], + "distogram": {"bin_edges": edges, "logits": logits}, + }) + + contacts = distogram_parser().get_contacts(str(tmp_path), distance=cutoff) + + assert bool(contacts) is expected + if expected: + assert len(contacts) == 1 + assert contacts[0][2] == pytest.approx(1.0) diff --git a/test/unit/test_inference_flags.py b/test/unit/test_inference_flags.py index 1fb1ab81..0ca42caf 100644 --- a/test/unit/test_inference_flags.py +++ b/test/unit/test_inference_flags.py @@ -86,3 +86,19 @@ def test_postprocess_flags_are_built_from_the_invocation(): built = inference_flags.postprocess_flags(_Flags(use_gpu_relax=False)) assert built["use_gpu_relax"] is False assert built["compress_pickles"] is False + + +def test_unifold_model_flags_fail_with_the_release_limitation(): + with pytest.raises(ValueError, match="UniFold.*unavailable"): + inference_flags.model_flags(_Flags(fold_backend="unifold")) + + +def test_unifold_validation_rejects_the_backend_before_checking_flags(): + with pytest.raises(ValueError, match="UniFold.*unavailable"): + inference_flags.unsupported_flags("unifold", ["unifold_model_name"]) + assert inference_flags.unsupported_flags("alphafold2", ["unifold_model_name"]) == [ + "unifold_model_name" + ] + assert inference_flags.unsupported_flags("alphafold3", ["unifold_model_name"]) == [ + "unifold_model_name" + ] diff --git a/test/unit/test_jax_devices.py b/test/unit/test_jax_devices.py new file mode 100644 index 00000000..c8c6ad12 --- /dev/null +++ b/test/unit/test_jax_devices.py @@ -0,0 +1,34 @@ +"""The import-time GPU probe must report a missing GPU, not raise. + +``jax.local_devices(backend='gpu')`` raises on any machine without a GPU backend, +so the prediction commands could not even print ``--help`` on a login node. +""" + +import sys +import types + +from alphapulldown.prediction import jax_devices + + +def _with_jax(monkeypatch, local_devices): + stub = types.ModuleType("jax") + stub.local_devices = local_devices + monkeypatch.setitem(sys.modules, "jax", stub) + + +def test_reports_the_gpu_devices_jax_sees(monkeypatch): + _with_jax(monkeypatch, lambda backend: [f"{backend}:0", f"{backend}:1"]) + + assert jax_devices.initialise_jax_gpu_backend() == ["gpu:0", "gpu:1"] + + +def test_a_missing_gpu_backend_is_reported_not_fatal(monkeypatch): + def no_gpu_backend(backend): + raise RuntimeError( + "Unknown backend: 'gpu' requested, but no platforms that are instances " + "of gpu are present. Platforms present are: cpu" + ) + + _with_jax(monkeypatch, no_gpu_backend) + + assert jax_devices.initialise_jax_gpu_backend() == [] diff --git a/test/unit/test_light_import_invariant.py b/test/unit/test_light_import_invariant.py index a0cbe491..3e04b92f 100644 --- a/test/unit/test_light_import_invariant.py +++ b/test/unit/test_light_import_invariant.py @@ -29,6 +29,8 @@ "alphapulldown.features.feature_batch", "alphapulldown.scripts._mmseqs2_cli", "alphapulldown.prediction.inference_flags", + # Probes JAX only when called, so importing it costs nothing on a login node. + "alphapulldown.prediction.jax_devices", # Compatibility paths must retain the same lightweight behavior. "alphapulldown.feature_batch", "alphapulldown.af2_feature_finalizer", diff --git a/test/unit/test_objects.py b/test/unit/test_objects.py index 79b5fb3b..285d8fa3 100644 --- a/test/unit/test_objects.py +++ b/test/unit/test_objects.py @@ -1839,3 +1839,46 @@ def test_create_all_chain_features_skips_multimeric_template_postprocessing_when assert multimer.feature_dict["template_sequence"] == ["kept"] assert "multichain_mask" not in multimer.feature_dict + + +def test_concatenate_sliced_feature_dict_joins_template_sequences_across_regions(): + """template_sequence sat in the skip set, so a multi-region chop kept the FIRST + region's fragment only, shorter than the chain the template arrays describe.""" + chopped = ChoppedObject( + "proteinA", "ABCDEFGHIJ", _feature_dict(template_count=2), [(1, 2), (5, 7)] + ) + slice_a = chopped.prepare_individual_sliced_feature_dict(chopped.feature_dict, 1, 2) + slice_b = chopped.prepare_individual_sliced_feature_dict(chopped.feature_dict, 5, 7) + + merged = chopped.concatenate_sliced_feature_dict([slice_a, slice_b]) + + assert merged["template_sequence"].tolist() == [b"ABEFG", b"ABEFG"] + assert merged["template_aatype"].shape == (2, 5, 22) + + +def test_prepare_final_sliced_feature_dict_keeps_template_sequences_at_chain_length(): + chopped = ChoppedObject("proteinA", "ABCDEFGHIJ", _feature_dict(), [(1, 2), (5, 7)]) + + chopped.prepare_final_sliced_feature_dict() + + assert chopped.sequence == "ABEFG" + assert chopped.feature_dict["template_sequence"].tolist() == [b"ABEFG"] + assert chopped.feature_dict["template_sequence"].dtype == object + + +def test_concatenate_sliced_feature_dict_without_templates_keeps_no_template_sequence(): + chopped = ChoppedObject( + "proteinA", "ABCDEFGHIJ", _feature_dict(template_count=0), [(1, 2), (3, 4)] + ) + slice_a = chopped.prepare_individual_sliced_feature_dict(chopped.feature_dict, 1, 2) + slice_b = chopped.prepare_individual_sliced_feature_dict(chopped.feature_dict, 3, 4) + + merged = chopped.concatenate_sliced_feature_dict([slice_a, slice_b]) + + assert merged["template_sequence"].shape == (0,) + assert merged["template_aatype"].shape == (0, 4, 22) + + +def test_join_template_sequences_refuses_slices_with_different_template_counts(): + with pytest.raises(ValueError, match="number of templates"): + objects_mod._join_template_sequences([[b"AB"], [b"CD", b"EF"]]) diff --git a/test/unit/test_prediction_batch.py b/test/unit/test_prediction_batch.py index 3c05f2b6..2190f929 100644 --- a/test/unit/test_prediction_batch.py +++ b/test/unit/test_prediction_batch.py @@ -620,3 +620,63 @@ def test_metadata_is_taken_from_every_feature_directory(tmp_path): copied = sorted(p.name for p in out.iterdir()) assert copied == ["P1_feature_metadata_2025-01-01.json"] + + +def test_prepare_fold_saves_the_merged_features_of_the_object_it_built(tmp_path, monkeypatch): + """`--save_features_for_multimeric_object` never worked. + + It read `feature_dict` off the MultimericObject CLASS (AttributeError on every + run), and did so before the fold directory existed and before --use_ap_style + had chosen it. + """ + from alphapulldown.prediction import fold_preparation + + class FakeMultimer: + def __init__(self, interactors, **_kwargs): + self.description = "_and_".join(i.description for i in interactors) + self.feature_dict = {"msa": [[1, 2], [3, 4]], "asym_id": [1, 2]} + + monkeypatch.setattr(fold_preparation, "MultimericObject", FakeMultimer) + flags = SimpleNamespace( + pair_msa=False, multimeric_template=False, description_file=None, + path_to_mmt=None, threshold_clashes=1000, hb_allowance=0.4, + plddt_threshold=0, save_features_for_multimeric_object=True, + use_ap_style=True, features_directory=[], + ) + interactors = [ + SimpleNamespace(description="A", sequence="AC", skip_msa=False), + SimpleNamespace(description="B", sequence="DE", skip_msa=False), + ] + output_root = tmp_path / "predictions" # created by prepare_fold itself + + fold, output_dir = fold_preparation.prepare_fold(interactors, str(output_root), flags) + + assert Path(output_dir) == output_root / "A_and_B" + saved = Path(output_dir) / fold_preparation.MULTIMERIC_FEATURES_FILENAME + with saved.open("rb") as handle: + assert pickle.load(handle) == fold.feature_dict + + +def test_prepare_fold_does_not_save_features_unless_asked(tmp_path, monkeypatch): + from alphapulldown.prediction import fold_preparation + + class FakeMultimer: + def __init__(self, interactors, **_kwargs): + self.description = "A_and_B" + self.feature_dict = {"msa": []} + + monkeypatch.setattr(fold_preparation, "MultimericObject", FakeMultimer) + flags = SimpleNamespace( + pair_msa=False, multimeric_template=False, description_file=None, + path_to_mmt=None, threshold_clashes=1000, hb_allowance=0.4, + plddt_threshold=0, save_features_for_multimeric_object=False, + use_ap_style=False, features_directory=[], + ) + interactors = [ + SimpleNamespace(description="A", sequence="AC", skip_msa=False), + SimpleNamespace(description="B", sequence="DE", skip_msa=False), + ] + + _, output_dir = fold_preparation.prepare_fold(interactors, str(tmp_path / "out"), flags) + + assert not (Path(output_dir) / fold_preparation.MULTIMERIC_FEATURES_FILENAME).exists() diff --git a/test/unit/test_script_entrypoints.py b/test/unit/test_script_entrypoints.py index cc360435..d2c1be6f 100644 --- a/test/unit/test_script_entrypoints.py +++ b/test/unit/test_script_entrypoints.py @@ -140,7 +140,7 @@ def _set_flag(flags_obj, name, value, *, present=True, using_default_value=False flag.using_default_value = using_default_value -def _load_run_structure_prediction_module(): +def _load_run_structure_prediction_module(jax_local_devices=None): module_name = "test_run_structure_prediction_module" names_to_replace = [ "absl", @@ -179,7 +179,7 @@ def _load_run_structure_prediction_module(): absl_pkg.logging = logging_mod jax_mod = types.ModuleType("jax") - jax_mod.local_devices = lambda backend="gpu": [] + jax_mod.local_devices = jax_local_devices or (lambda backend="gpu": []) class ModelsToRelax(Enum): NONE = "none" @@ -234,6 +234,7 @@ def __init__( self.description = "_and_".join(interactor.description for interactor in interactors) self.input_seqs = [interactor.sequence for interactor in interactors] self.multimeric_mode = True + self.feature_dict = {"merged": [interactor.description for interactor in interactors]} objects_mod.MonomericObject = MonomericObject objects_mod.ChoppedObject = ChoppedObject @@ -363,6 +364,7 @@ def _load_run_multimer_jobs_module(): "debug_templates": False, "debug_msas": False, "job_index": None, + "unifold_model_name": "multimer_af2", } for name, default in shared_flag_defaults.items(): flags_mod.FLAGS.define(name, default) @@ -459,9 +461,30 @@ def test_importing_shared_prediction_flags_does_not_require_single_job_outputs( def test_single_prediction_initializes_jax_before_importing_backend(): source = RUN_STRUCTURE_PREDICTION_PATH.read_text(encoding="utf-8") - assert source.index("gpus = jax.local_devices(backend='gpu')") < source.index( - "from alphapulldown.folding_backend import backend" + probe = source.index("gpus = initialise_jax_gpu_backend()") + assert probe < source.index("from alphapulldown.folding_backend import backend") + assert probe < source.index("from alphapulldown.objects import") + + +def _raise_no_gpu_backend(backend="gpu"): + raise RuntimeError( + "Unknown backend: 'gpu' requested, but no platforms that are instances of " + "gpu are present. Platforms present are: cpu" + ) + + +def test_single_prediction_imports_without_a_gpu_backend(): + """The import-time probe raised on any machine without a GPU, so --help and + the head-node flag validation the workflow recommends both failed.""" + module, saved_modules = _load_run_structure_prediction_module( + jax_local_devices=_raise_no_gpu_backend ) + try: + assert module.gpus == [] + assert "unifold_model_name" in module.FLAGS + finally: + sys.modules.pop(module.__name__, None) + _restore_modules(saved_modules) def test_validate_flags_for_af3_allows_modelcif_conversion( @@ -691,10 +714,13 @@ def test_pre_modelling_setup_saves_multimer_features_and_builds_unique_ap_style_ assert isinstance(returned_object, run_structure_prediction_module.MultimericObject) assert returned_output_dir.endswith("protA_and_protB") + # The features of the object that was built, written into the fold's own + # directory once it exists. The old code read them off the class and wrote + # one level up, before the directory was created. assert dumped == [ ( - run_structure_prediction_module.MultimericObject.feature_dict, - str(tmp_path / "outputs" / "multimeric_object_features.pkl"), + {"merged": ["protA", "protB"]}, + str(tmp_path / "outputs" / "protA_and_protB" / "multimeric_object_features.pkl"), ) ] @@ -1380,3 +1406,60 @@ def test_run_multimer_jobs_forwards_relax_best_score_threshold( assert len(calls) == 1 assert "--relax_best_score_threshold" in calls[0] assert calls[0][calls[0].index("--relax_best_score_threshold") + 1] == "0.6" + + +@pytest.mark.parametrize("legacy_switch", [True, False]) +def test_run_multimer_jobs_rejects_unifold_before_reading_inputs( + run_multimer_jobs_module, monkeypatch, legacy_switch, +): + monkeypatch.setattr( + run_multimer_jobs_module, "generate_fold_specifications", + lambda **kwargs: pytest.fail("must reject UniFold before reading inputs"), + ) + _set_flag(run_multimer_jobs_module.FLAGS, "use_unifold", legacy_switch) + _set_flag(run_multimer_jobs_module.FLAGS, "fold_backend", + "alphafold2" if legacy_switch else "unifold") + with pytest.raises(ValueError, match="UniFold.*unavailable"): + run_multimer_jobs_module.main(["prog"]) + + +def test_run_multimer_jobs_does_not_forward_the_unifold_model_name_to_other_backends( + run_multimer_jobs_module, + monkeypatch, +): + calls = [] + monkeypatch.setattr( + run_multimer_jobs_module.subprocess, + "run", + lambda command, check, env: calls.append(command), + ) + run_multimer_jobs_module.generate_fold_specifications = ( + lambda input_files, delimiter, exclude_permutations: ["A,B"] + ) + + _set_flag(run_multimer_jobs_module.FLAGS, "mode", "custom") + _set_flag(run_multimer_jobs_module.FLAGS, "protein_lists", ["proteins.txt"]) + _set_flag(run_multimer_jobs_module.FLAGS, "dry_run", False) + _set_flag(run_multimer_jobs_module.FLAGS, "fold_backend", "alphafold2") + _set_flag(run_multimer_jobs_module.FLAGS, "output_path", "/tmp/output") + _set_flag(run_multimer_jobs_module.FLAGS, "data_dir", "/tmp/models") + _set_flag(run_multimer_jobs_module.FLAGS, "monomer_objects_dir", ["/tmp/features"]) + _set_flag(run_multimer_jobs_module.FLAGS, "use_unifold", False) + + run_multimer_jobs_module.main(["prog"]) + + assert len(calls) == 1 + assert "--unifold_model_name" not in calls[0] + assert calls[0][calls[0].index("--fold_backend") + 1] == "alphafold2" + + +def test_main_rejects_unifold_before_reading_inputs( + run_structure_prediction_module, monkeypatch, tmp_path, +): + _set_flag(run_structure_prediction_module.FLAGS, "fold_backend", "unifold") + monkeypatch.setattr( + run_structure_prediction_module, "parse_fold", + lambda *args: pytest.fail("must reject UniFold before reading inputs"), + ) + with pytest.raises(ValueError, match="UniFold.*unavailable"): + run_structure_prediction_module.main([]) diff --git a/test/unit/test_unifold_backend.py b/test/unit/test_unifold_backend.py index 2ca5c94a..2b6e83ad 100644 --- a/test/unit/test_unifold_backend.py +++ b/test/unit/test_unifold_backend.py @@ -1,161 +1,53 @@ -import importlib.util -import sys -import types +"""The unavailable legacy backend must fail before importing a model runtime.""" + from pathlib import Path from types import SimpleNamespace +import pytest -MODULE_PATH = ( - Path(__file__).resolve().parents[2] - / "alphapulldown" - / "folding_backend" - / "unifold_backend.py" +from alphapulldown.folding_backend import FoldingBackendManager +from alphapulldown.folding_backend.unifold_backend import UnifoldBackend +from alphapulldown.prediction.prediction_batch import ( + PredictionBatch, PredictionJob, PreparedPredictionAdapter, ) -def _restore_modules(saved_modules: dict[str, types.ModuleType | None]) -> None: - for name, module in saved_modules.items(): - if module is None: - sys.modules.pop(name, None) - else: - sys.modules[name] = module - - -def _install_unifold_backend_stubs() -> dict[str, types.ModuleType | None]: - names_to_replace = [ - "alphapulldown.objects", - "unifold", - "unifold.config", - "unifold.inference", - "unifold.dataset", - ] - saved_modules = {name: sys.modules.get(name) for name in names_to_replace} - - objects_mod = types.ModuleType("alphapulldown.objects") - objects_mod.MultimericObject = type("MultimericObject", (), {}) - - unifold_pkg = types.ModuleType("unifold") - config_mod = types.ModuleType("unifold.config") - config_mod.model_config = lambda model_name: {"model_name": model_name} - - inference_mod = types.ModuleType("unifold.inference") - inference_mod.calls = [] - inference_mod.config_args = ( - lambda model_dir, target_name, output_dir: { - "model_dir": model_dir, - "target_name": target_name, - "output_dir": output_dir, - } - ) - inference_mod.unifold_config_model = lambda general_args: {"runner_args": general_args} - inference_mod.unifold_predict = ( - lambda model_runner, model_args, processed_features: inference_mod.calls.append( - (model_runner, model_args, processed_features) - ) - ) - - dataset_mod = types.ModuleType("unifold.dataset") - dataset_mod.process_ap = ( - lambda config, features, mode, labels, seed, batch_idx, data_idx, is_distillation: ( - { - "processed_features": features, - "seed": seed, - "mode": mode, - "config": config, - }, - None, - ) - ) - - modules = { - "alphapulldown.objects": objects_mod, - "unifold": unifold_pkg, - "unifold.config": config_mod, - "unifold.inference": inference_mod, - "unifold.dataset": dataset_mod, - } - for name, module in modules.items(): - sys.modules[name] = module - - unifold_pkg.config = config_mod - unifold_pkg.inference = inference_mod - unifold_pkg.dataset = dataset_mod - - return saved_modules - - -def _load_unifold_backend_module(): - saved_modules = _install_unifold_backend_stubs() - sys.modules.pop("alphapulldown.folding_backend.unifold_backend", None) - spec = importlib.util.spec_from_file_location( - "alphapulldown.folding_backend.unifold_backend", - MODULE_PATH, +@pytest.mark.parametrize("model_name", [ + "multimer_af2", "multimer_ft", "multimer", "multimer_af2_v3", + "multimer_af2_model45_v3", +]) +def test_legacy_model_choices_fail_with_actionable_error(model_name, monkeypatch): + import sys + monkeypatch.setitem(sys.modules, "unifold", None) + with pytest.raises(ValueError, match="UniFold.*unavailable.*AlphaLink"): + UnifoldBackend.setup(model_dir="/missing/weights", model_name=model_name) + + +def test_direct_prediction_rejects_unifold_before_creating_outputs(tmp_path): + output = tmp_path / "fold" + with pytest.raises(ValueError, match="UniFold.*unavailable"): + list(UnifoldBackend.predict( + objects_to_model=[{"object": SimpleNamespace(), "output_dir": str(output)}], + model_dir="/missing/weights", model_config={}, + )) + assert not output.exists() + + +def test_prediction_adapter_reports_unavailable_backend(tmp_path): + output = tmp_path / "fold" + adapter = PreparedPredictionAdapter( + backend=FoldingBackendManager(), fold_backend="unifold", + objects_to_model=[{"object": SimpleNamespace(), "output_dir": str(output)}], + model_flags={"model_dir": "/missing/weights", "model_name": "multimer_af2"}, + postprocess_flags={}, random_seed=7, ) - module = importlib.util.module_from_spec(spec) - sys.modules[spec.name] = module - assert spec.loader is not None - try: - spec.loader.exec_module(module) - return module, saved_modules - except Exception: - sys.modules.pop(spec.name, None) - _restore_modules(saved_modules) - raise - - -def test_unifold_setup_predict_and_postprocess(): - module, saved_modules = _load_unifold_backend_module() - try: - multimeric_object = SimpleNamespace(description="complex", feature_dict={"msa": [1, 2]}) - - configured = module.UnifoldBackend.setup( - model_name="multimer", - model_dir="/models", - output_dir="/output", - multimeric_object=multimeric_object, - ) - assert configured == { - "model_runner": { - "runner_args": { - "model_dir": "/models", - "target_name": "complex", - "output_dir": "/output", - } - }, - "model_args": { - "model_dir": "/models", - "target_name": "complex", - "output_dir": "/output", - }, - "model_config": {"model_name": "multimer"}, - } + # Backend setup is a batch-level rejection, not a recoverable fold failure. + with pytest.raises(ValueError, match="UniFold.*unavailable"): + PredictionBatch((PredictionJob("legacy", "", Path(output)),)).run(adapter) + assert not output.exists() - backend = module.UnifoldBackend() - assert ( - backend.predict( - model_runner="runner", - model_args={"arg": 1}, - model_config={"cfg": 2}, - multimeric_object=multimeric_object, - random_seed=11, - ) - is None - ) - inference_mod = sys.modules["unifold.inference"] - assert inference_mod.calls == [ - ( - "runner", - {"arg": 1}, - { - "processed_features": {"msa": [1, 2]}, - "seed": 11, - "mode": "predict", - "config": {"cfg": 2}, - }, - ) - ] - assert module.UnifoldBackend.postprocess() is None - finally: - sys.modules.pop("alphapulldown.folding_backend.unifold_backend", None) - _restore_modules(saved_modules) +def test_unifold_is_not_advertised_as_available(monkeypatch): + import alphapulldown.folding_backend as manager_module + monkeypatch.setattr(manager_module, "_try_import", lambda *args: object) + assert "unifold" not in FoldingBackendManager().available_backends()