From de711b80bf573be6cec01a890b8f80d66d2a7923 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 09:17:43 +0000 Subject: [PATCH 1/4] fix: AF2 MSA-depth runners, saved multimer features, UniFold contract, GPU-less import - AlphaFold 2 setup registered one RunModel under every MSA depth while rewriting its single config in place, so --msa_depth_scan predicted N times at the last depth, differing only by seed. Each distinct depth now gets its own config and runner; parameters are still loaded once per model name. - --save_features_for_multimeric_object read feature_dict off the MultimericObject class (AttributeError on every run) and ran before the fold directory existed. It now pickles the built object's features into the fold's own directory once that exists. - The UniFold backend required output_dir and multimeric_object in setup() and took one object in an instance-method predict(), so no adapter could call it. It now follows the backend contract; --unifold_model_name is defined by run_structure_prediction, forwarded by run_multimer_jobs, validated for the unifold backend, and kept for multimers instead of being replaced by the AlphaFold 2 "multimer" preset. - The import-time jax.local_devices(backend='gpu') probe raised on any machine without a GPU, so --help and head-node flag validation failed. It is now a tolerant helper (alphapulldown.prediction.jax_devices) that still runs before the backends import TensorFlow and OpenMM. - ChoppedObject kept only the first region's template_sequence fragment when concatenating regions; fragments are now joined per template, so the sequence matches the chopped chain length. - distogram_parser.get_contacts referenced an undefined `datadir` and re-read whichever pickle the scan visited last rather than the top-ranked one. - requires-python is 3.10: the package uses slots dataclasses and zip(strict=True), and CI runs 3.10 and 3.11. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01Mkya8Y3utHEPbVPc2RbJu3 --- .../folding_backend/alphafold2_backend.py | 117 ++++++----- .../folding_backend/unifold_backend.py | 151 +++++++++------ alphapulldown/objects.py | 35 ++++ alphapulldown/prediction/fold_preparation.py | 29 ++- alphapulldown/prediction/inference_flags.py | 23 ++- alphapulldown/prediction/jax_devices.py | 34 ++++ alphapulldown/scripts/run_multimer_jobs.py | 9 +- .../scripts/run_structure_prediction.py | 18 +- alphapulldown/utils/distogram_parser.py | 43 +++-- pyproject.toml | 5 +- test/RELEASE_READINESS.md | 2 +- test/unit/test_alphafold2_backend_helpers.py | 105 ++++++++++ test/unit/test_distogram_parser.py | 66 +++++-- test/unit/test_inference_flags.py | 22 +++ test/unit/test_jax_devices.py | 34 ++++ test/unit/test_light_import_invariant.py | 2 + test/unit/test_objects.py | 43 +++++ test/unit/test_prediction_batch.py | 60 ++++++ test/unit/test_script_entrypoints.py | 161 +++++++++++++++- test/unit/test_unifold_backend.py | 182 ++++++++++-------- 20 files changed, 896 insertions(+), 245 deletions(-) create mode 100644 alphapulldown/prediction/jax_devices.py create mode 100644 test/unit/test_jax_devices.py diff --git a/alphapulldown/folding_backend/alphafold2_backend.py b/alphapulldown/folding_backend/alphafold2_backend.py index 59c5bee68..98bf48aec 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 50755568f..2c2da7174 100644 --- a/alphapulldown/folding_backend/unifold_backend.py +++ b/alphapulldown/folding_backend/unifold_backend.py @@ -4,108 +4,141 @@ Author: Valentin Maurer """ -from typing import Dict +from __future__ import annotations -from alphapulldown.objects import MultimericObject +import os +from typing import Any, Dict, Iterator, List + +from absl import logging from .folding_backend import FoldingBackend +# The model configurations UniFold ships; the default is what run_multimer_jobs +# always offered through --unifold_model_name. +UNIFOLD_MODEL_NAMES = ( + "multimer_af2", + "multimer_ft", + "multimer", + "multimer_af2_v3", + "multimer_af2_model45_v3", +) +DEFAULT_UNIFOLD_MODEL_NAME = "multimer_af2" + + class UnifoldBackend(FoldingBackend): """ A backend class for running protein structure predictions using the UniFold model. + + It follows the same contract as the other backends, which is what the + prediction adapters drive: ``setup(**model_flags)`` builds a session, + ``predict(**session, objects_to_model=..., random_seed=..., **model_flags)`` + yields one record per fold, and ``postprocess`` receives the record. The + previous version required ``output_dir`` and ``multimeric_object`` in + ``setup`` and took a single object in an instance-method ``predict``, so no + caller in the package could invoke it. """ + @staticmethod def setup( - model_name: str, model_dir: str, - output_dir: str, - multimeric_object: MultimericObject, + model_name: str = DEFAULT_UNIFOLD_MODEL_NAME, **kwargs, ) -> Dict: """ - Initializes and configures a UniFold model runner. + Resolve the UniFold model configuration for this session. 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. + The directory where the UniFold parameters are located. + model_name : str + The UniFold model configuration (``--unifold_model_name``). **kwargs : dict - Additional keyword arguments for model configuration. + Additional keyword arguments for model configuration. Ignored. Returns ------- Dict - A dictionary containing the model runner, arguments, and configuration. + The model configuration under ``model_config``. The weights are + loaded per fold in :py:meth:`UnifoldBackend.predict`, because + UniFold ties its runner to the target name and output directory. """ 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, - } + if model_name not in UNIFOLD_MODEL_NAMES: + raise ValueError( + f"Unknown UniFold model {model_name!r}; choose one of " + f"{', '.join(UNIFOLD_MODEL_NAMES)}" + ) + return {"model_config": model_config(model_name)} + @staticmethod def predict( - self, - model_runner, - model_args, + objects_to_model: List[Dict[str, Any]], model_config: Dict, - multimeric_object: MultimericObject, + model_dir: str, random_seed: int = 42, **kwargs, - ) -> None: + ) -> Iterator[Dict[str, Any]]: """ - Predicts the structure of proteins using configured UniFold models. + Predicts the structure of each object with the configured UniFold model. 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`. + objects_to_model : List[Dict[str, Any]] + One ``{"object": ..., "output_dir": ...}`` record per fold; the object + carries ``description`` and ``feature_dict``. 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. + Configuration obtained from :py:meth:`UnifoldBackend.setup`. + model_dir : str + The directory where the UniFold parameters are located. random_seed : int, optional The random seed for prediction reproducibility, default is 42. **kwargs : dict - Additional keyword arguments for prediction. + Additional keyword arguments for prediction. Ignored. + + Yields + ------ + Dict + ``object``, ``prediction_results`` and ``output_dir`` for each fold, + in order. UniFold writes its structures itself, so the results are + an empty mapping. """ 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) + from unifold.inference import config_args, unifold_config_model, unifold_predict - return None + for entry in objects_to_model: + object_to_model = entry["object"] + output_dir = entry["output_dir"] + os.makedirs(output_dir, exist_ok=True) + logging.info( + "Now running UniFold prediction on %s", object_to_model.description + ) + general_args = config_args( + model_dir, + target_name=object_to_model.description, + output_dir=output_dir, + ) + model_runner = unifold_config_model(general_args) + processed_features, _ = process_ap( + config=model_config, + features=object_to_model.feature_dict, + mode="predict", + labels=None, + seed=random_seed, + batch_idx=None, + data_idx=None, + is_distillation=False, + ) + unifold_predict(model_runner, general_args, processed_features) + yield { + "object": object_to_model, + "prediction_results": {}, + "output_dir": output_dir, + } + @staticmethod def postprocess(**kwargs) -> None: + """UniFold writes its own outputs; nothing to post-process.""" return None diff --git a/alphapulldown/objects.py b/alphapulldown/objects.py index 404090603..bd782d982 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 c202ce5ea..9dfc6e446 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 d9e803d19..713446d7e 100644 --- a/alphapulldown/prediction/inference_flags.py +++ b/alphapulldown/prediction/inference_flags.py @@ -36,6 +36,10 @@ ALPHALINK_EXTRA_FLAGS = frozenset({"crosslinks"}) +# UniFold picks its model configuration by name; it serves monomers and multimers alike. +UNIFOLD_EXTRA_FLAGS = frozenset({"unifold_model_name"}) +DEFAULT_UNIFOLD_MODEL_NAME = "multimer_af2" + AF3_FLAGS = frozenset({ "jax_compilation_cache_dir", "buckets", "flash_attention_implementation", "num_diffusion_samples", "num_seeds", "debug_templates", "debug_msas", @@ -46,6 +50,7 @@ FLAGS_BY_BACKEND: Mapping[str, frozenset] = { "alphafold2": COMMON_FLAGS | AF2_LIKE_FLAGS, "alphalink": COMMON_FLAGS | AF2_LIKE_FLAGS | ALPHALINK_EXTRA_FLAGS, + "unifold": COMMON_FLAGS | AF2_LIKE_FLAGS | UNIFOLD_EXTRA_FLAGS, "alphafold3": COMMON_FLAGS | AF3_FLAGS, } @@ -106,15 +111,27 @@ 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. + UniFold has one model for monomers and multimers, chosen by + ``--unifold_model_name``, so the caller must leave it alone. + """ + if flags.fold_backend == "alphalink": + return "multimer_af2_crop" + if flags.fold_backend == "unifold": + return getattr(flags, "unifold_model_name", None) or DEFAULT_UNIFOLD_MODEL_NAME + 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 000000000..ad65055bd --- /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 a8756d7c3..20f1854d2 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 @@ -30,9 +30,7 @@ "Whether unifold models are going to be used. Default it False") 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") @@ -153,6 +151,7 @@ def main(argv): "--msa_depth": FLAGS.msa_depth, "--crosslinks": FLAGS.crosslinks, "--fold_backend": fold_backend, + "--unifold_model_name": FLAGS.unifold_model_name if FLAGS.use_unifold else None, "--description_file": FLAGS.description_file, "--path_to_mmt": FLAGS.path_to_mmt, "--compress_result_pickles": FLAGS.compress_result_pickles, diff --git a/alphapulldown/scripts/run_structure_prediction.py b/alphapulldown/scripts/run_structure_prediction.py index 3721f2823..3e2c02bb5 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,12 @@ # AlphaLink2 settings flags.DEFINE_string('crosslinks', None, 'Path to crosslink information pickle for AlphaLink.') +# UniFold settings +flags.DEFINE_enum( + 'unifold_model_name', 'multimer_af2', + ['multimer_af2', 'multimer_ft', 'multimer', 'multimer_af2_v3', 'multimer_af2_model45_v3'], + 'UniFold model configuration used with --fold_backend=unifold.') + # AlphaFold3 settings # JAX inference performance tuning. flags.DEFINE_string( @@ -369,8 +379,10 @@ def main(argv): json_output_dir = real_out # Flags for THIS object, not for whichever fold happens to come last. + # UniFold's model is chosen by --unifold_model_name and serves monomers + # and multimers alike, so it keeps the name the flags resolved. object_model_flags = default_model_flags.copy() - if isinstance(obj, MultimericObject): + if isinstance(obj, MultimericObject) and FLAGS.fold_backend != "unifold": object_model_flags.update({ "model_name": "multimer", "msa_depth_scan": FLAGS.msa_depth_scan, diff --git a/alphapulldown/utils/distogram_parser.py b/alphapulldown/utils/distogram_parser.py index 465f7f8ad..e5346e86f 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}") - - d = top_ranked_dgram[1] + print(f"Selected {os.path.basename(top_ranked_fn)} with ranking confidence {top_ranked_confidence:.2f}") - # 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) @@ -97,6 +106,6 @@ def get_contacts(self, directory, distance=8, pbtycutoff=0.8, cross_only=True, v if __name__=="__main__": do=distogram_parser() - contacts=do.get_contacts(datadir='.', verbose=0) + contacts=do.get_contacts(directory='.', verbose=0) diff --git a/pyproject.toml b/pyproject.toml index ca94e76d0..9ab4a52ff 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,7 +6,7 @@ build-backend = "setuptools.build_meta" name = "alphapulldown" description = "Pipeline allows massive screening using alphafold" readme = { file = "README.md", content-type = "text/markdown" } -requires-python = ">=3.8" +requires-python = ">=3.10" license = "MIT" license-files = ["LICENSE"] authors = [ @@ -14,6 +14,9 @@ authors = [ ] classifiers = [ "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.10", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", "Operating System :: OS Independent", ] dynamic = ["version"] diff --git a/test/RELEASE_READINESS.md b/test/RELEASE_READINESS.md index ef3f78a6f..d1d3361b5 100644 --- a/test/RELEASE_READINESS.md +++ b/test/RELEASE_READINESS.md @@ -45,7 +45,7 @@ The following areas are only partially protected, optional, or report-only: - `test/alphalink` workflows beyond CPU-safe helper tests - legacy scenarios still parked under `test/outdated` - analysis-pipeline utilities and some deeper ModelCIF internals -- Python `3.8`, which is still advertised in packaging but is not exercised by GitHub Actions +- Python `3.12`, which packaging allows (`requires-python >= 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/unit/test_alphafold2_backend_helpers.py b/test/unit/test_alphafold2_backend_helpers.py index b5a25a425..9b4006721 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_distogram_parser.py b/test/unit/test_distogram_parser.py index 4ca0bc07e..0918b4c2f 100644 --- a/test/unit/test_distogram_parser.py +++ b/test/unit/test_distogram_parser.py @@ -1,38 +1,68 @@ -import pickle - -import numpy as np +"""Contact extraction from AlphaFold 2 distograms. -import alphapulldown.utils.distogram_parser as distogram_parser_module +``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. +""" +import pickle -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) - - parser = distogram_parser_module.distogram_parser() +import numpy as np - 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): +def _payload(ranking_confidence, *, contact): 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, + 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) diff --git a/test/unit/test_inference_flags.py b/test/unit/test_inference_flags.py index 1fb1ab81e..d8c74bd8f 100644 --- a/test/unit/test_inference_flags.py +++ b/test/unit/test_inference_flags.py @@ -86,3 +86,25 @@ 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_model_flags_name_the_unifold_model_from_its_own_flag(): + """UniFold has one model for monomers and multimers, chosen by --unifold_model_name.""" + chosen = inference_flags.model_flags( + _Flags(fold_backend="unifold", unifold_model_name="multimer_ft") + ) + assert chosen["model_name"] == "multimer_ft" + default = inference_flags.model_flags(_Flags(fold_backend="unifold")) + assert default["model_name"] == inference_flags.DEFAULT_UNIFOLD_MODEL_NAME == "multimer_af2" + + +def test_unifold_accepts_its_model_name_flag_and_the_af2_like_flags(): + assert inference_flags.unsupported_flags( + "unifold", ["unifold_model_name", "allow_resume", "num_cycle"] + ) == [] + 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 000000000..c8c6ad12b --- /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 a0cbe4911..3e04b92fa 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 79b5fb3b3..285d8fa3f 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 3c05f2b66..2190f929c 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 cc360435a..4c76bf21c 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,11 +461,32 @@ 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( run_structure_prediction_module, ): @@ -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,126 @@ 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" + + +def test_run_multimer_jobs_forwards_the_unifold_backend_and_model_name( + 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", True) + _set_flag(run_multimer_jobs_module.FLAGS, "unifold_param", "/tmp/unifold-params") + _set_flag(run_multimer_jobs_module.FLAGS, "unifold_model_name", "multimer_ft") + + run_multimer_jobs_module.main(["prog"]) + + assert len(calls) == 1 + command = calls[0] + assert command[command.index("--fold_backend") + 1] == "unifold" + assert command[command.index("--data_directory") + 1] == "/tmp/unifold-params" + assert command[command.index("--unifold_model_name") + 1] == "multimer_ft" + + +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_keeps_the_unifold_model_name_for_multimers( + run_structure_prediction_module, + monkeypatch, + tmp_path, +): + """The multimer override to model_name "multimer" is an AlphaFold 2 preset; UniFold + has one model for monomers and multimers, chosen by --unifold_model_name.""" + captured_calls = [] + protein_obj = run_structure_prediction_module.MonomericObject("protA", "AC") + multimer_obj = run_structure_prediction_module.MultimericObject( + [protein_obj, protein_obj], + pair_msa=True, + multimeric_template=False, + multimeric_template_meta_data=None, + multimeric_template_dir=None, + ) + + _set_flag(run_structure_prediction_module.FLAGS, "fold_backend", "unifold") + _set_flag(run_structure_prediction_module.FLAGS, "unifold_model_name", "multimer_ft") + _set_flag(run_structure_prediction_module.FLAGS, "input", ["job1"]) + _set_flag( + run_structure_prediction_module.FLAGS, + "output_directory", + [str(tmp_path / "shared-output")], + ) + _set_flag( + run_structure_prediction_module.FLAGS, + "features_directory", + [str(tmp_path / "features")], + ) + _set_flag(run_structure_prediction_module.FLAGS, "protein_delimiter", "+") + _set_flag(run_structure_prediction_module.FLAGS, "data_directory", "/unifold-params") + + monkeypatch.setattr(run_structure_prediction_module, "parse_fold", lambda *args: [["parsed"]]) + monkeypatch.setattr(run_structure_prediction_module, "create_custom_info", lambda parsed: "data") + monkeypatch.setattr( + run_structure_prediction_module, + "create_interactors", + lambda data, features_directory: [[protein_obj, protein_obj]], + ) + monkeypatch.setattr( + run_structure_prediction_module, + "pre_modelling_setup", + lambda prot_objs, output_dir: (multimer_obj, str(tmp_path / "protein")), + ) + monkeypatch.setattr( + run_structure_prediction_module, + "predict_structure", + lambda **kwargs: captured_calls.append(kwargs), + ) + + run_structure_prediction_module.main([]) + + assert len(captured_calls) == 1 + assert captured_calls[0]["fold_backend"] == "unifold" + assert captured_calls[0]["model_flags"]["model_name"] == "multimer_ft" + assert captured_calls[0]["model_flags"]["model_dir"] == "/unifold-params" diff --git a/test/unit/test_unifold_backend.py b/test/unit/test_unifold_backend.py index 2ca5c94a4..0d8d0fb5e 100644 --- a/test/unit/test_unifold_backend.py +++ b/test/unit/test_unifold_backend.py @@ -1,9 +1,18 @@ +"""The UniFold backend must be drivable through the prediction adapters. + +Its ``setup()`` used to require ``output_dir`` and ``multimeric_object`` and its +``predict()`` took one object as an instance method, so neither adapter could +call it: ``--fold_backend=unifold`` failed on the first ``setup(**model_flags)``. +""" + import importlib.util import sys import types from pathlib import Path from types import SimpleNamespace +import pytest + MODULE_PATH = ( Path(__file__).resolve().parents[2] @@ -11,6 +20,7 @@ / "folding_backend" / "unifold_backend.py" ) +MODULE_NAME = "alphapulldown.folding_backend.unifold_backend" def _restore_modules(saved_modules: dict[str, types.ModuleType | None]) -> None: @@ -21,20 +31,12 @@ def _restore_modules(saved_modules: dict[str, types.ModuleType | None]) -> None: 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", - ] +def _install_unifold_stubs() -> dict[str, types.ModuleType | None]: + names_to_replace = ["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") + unifold_pkg.__path__ = [] # type: ignore[attr-defined] config_mod = types.ModuleType("unifold.config") config_mod.model_config = lambda model_name: {"model_name": model_name} @@ -57,18 +59,12 @@ def _install_unifold_backend_stubs() -> dict[str, types.ModuleType | None]: 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, - }, + {"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, @@ -76,86 +72,106 @@ def _install_unifold_backend_stubs() -> dict[str, types.ModuleType | None]: } 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.fixture +def unifold_backend_module(): + saved_modules = _install_unifold_stubs() + sys.modules.pop(MODULE_NAME, None) + spec = importlib.util.spec_from_file_location(MODULE_NAME, MODULE_PATH) 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) + yield module + finally: + sys.modules.pop(MODULE_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]}) +def _unifold_flags(**overrides): + defaults = dict( + fold_backend="unifold", unifold_model_name="multimer_ft", num_cycle=3, + data_directory="/weights", num_predictions_per_model=1, crosslinks=None, + desired_num_res=None, desired_num_msa=None, skip_templates=False, + allow_resume=True, num_diffusion_samples=5, num_recycles=10, + save_embeddings=False, save_distogram=False, + flash_attention_implementation="triton", buckets=["256"], + jax_compilation_cache_dir=None, features_directory=["/features"], + num_seeds=None, debug_templates=False, debug_msas=False, dropout=False, + ) + defaults.update(overrides) + return SimpleNamespace(**defaults) - 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 = 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 +def test_unifold_backend_runs_through_the_prediction_adapters(unifold_backend_module, tmp_path): + from alphapulldown.folding_backend import FoldingBackendManager + from alphapulldown.prediction.inference_flags import model_flags + from alphapulldown.prediction.prediction_batch import ( + PredictionBatch, + PredictionJob, + PreparedPredictionAdapter, + ) + + fold = SimpleNamespace(description="A_and_B", feature_dict={"msa": [1, 2]}) + output_dir = tmp_path / "A_and_B" # created by the backend + adapter = PreparedPredictionAdapter( + backend=FoldingBackendManager(), + fold_backend="unifold", + objects_to_model=[{"object": fold, "output_dir": str(output_dir)}], + model_flags=model_flags(_unifold_flags()), + postprocess_flags={}, + random_seed=7, + ) + + summary = PredictionBatch((PredictionJob("legacy", "", output_dir),)).run(adapter) + + assert summary.failures == () + assert summary.completed_job_ids == ("legacy",) + assert output_dir.is_dir() + general_args = { + "model_dir": "/weights", + "target_name": "A_and_B", + "output_dir": str(output_dir), + } + assert sys.modules["unifold.inference"].calls == [ + ( + {"runner_args": general_args}, + general_args, + { + "processed_features": {"msa": [1, 2]}, + "seed": 7, + "mode": "predict", + "config": {"model_name": "multimer_ft"}, + }, ) + ] - 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_predict_handles_every_object_in_order(unifold_backend_module, tmp_path): + backend = unifold_backend_module.UnifoldBackend + session = backend.setup(model_dir="/weights", model_name="multimer_af2") + objects = [ + {"object": SimpleNamespace(description=name, feature_dict={"n": index}), + "output_dir": str(tmp_path / name)} + for index, name in enumerate(("first", "second")) + ] + + records = list(backend.predict(objects, random_seed=3, model_dir="/weights", **session)) + + assert [record["object"].description for record in records] == ["first", "second"] + assert [record["output_dir"] for record in records] == [o["output_dir"] for o in objects] + assert all(record["prediction_results"] == {} for record in records) + targets = [call[1]["target_name"] for call in sys.modules["unifold.inference"].calls] + assert targets == ["first", "second"] + assert backend.postprocess(prediction_results={}, output_dir=str(tmp_path)) is None + + +def test_unifold_setup_rejects_an_alphafold_model_name(unifold_backend_module): + with pytest.raises(ValueError, match="multimer_af2"): + unifold_backend_module.UnifoldBackend.setup(model_dir="/weights", model_name="monomer_ptm") From ba2514751fbbe905dbdadd591a837016e41245f6 Mon Sep 17 00:00:00 2001 From: Dima Molodenskiy Date: Thu, 1 Oct 2026 12:37:32 +0200 Subject: [PATCH 2/4] fix: disable unsupported UniFold and close release test gaps Reject legacy UniFold selectors before input processing because the bundled AlphaLink fork lacks a compatible native UniFold inference API and network. Keep legacy flags/imports parseable for an actionable error and document the release limitation without advertising unsupported model configurations. Preserve pytest parameter IDs: the previous docstring rewrite made xdist loadfile silently omit cases. A subprocess regression now verifies that a failing second parameter fails both serial and parallel execution. Include every complete AF2 distogram bin at or below the distance cutoff, with regressions matching the real 64-bin/63-edge output representation. Validation: 748 passed, 12 skipped in both serial and xdist CPU suites; identical 760 executed test identities including skips. Fresh MMseqs features generated; AF2 depth-scan and AF3 species-pairing GPU checks submitted separately. --- README.md | 4 + alphapulldown/folding_backend/__init__.py | 5 +- .../folding_backend/unifold_backend.py | 142 ++----------- alphapulldown/prediction/inference_flags.py | 21 +- alphapulldown/scripts/run_multimer_jobs.py | 10 +- .../scripts/run_structure_prediction.py | 13 +- alphapulldown/utils/distogram_parser.py | 10 +- conftest.py | 12 -- test/README.md | 4 + test/RELEASE_READINESS.md | 12 ++ test/integration/test_pytest_collection.py | 33 +++ test/unit/test_distogram_parser.py | 28 ++- test/unit/test_inference_flags.py | 18 +- test/unit/test_script_entrypoints.py | 98 ++------- test/unit/test_unifold_backend.py | 196 ++++-------------- 15 files changed, 182 insertions(+), 424 deletions(-) create mode 100644 test/integration/test_pytest_collection.py diff --git a/README.md b/README.md index 75436af3e..adefa218e 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 a65434219..f84bfd947 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/unifold_backend.py b/alphapulldown/folding_backend/unifold_backend.py index 2c2da7174..05359b430 100644 --- a/alphapulldown/folding_backend/unifold_backend.py +++ b/alphapulldown/folding_backend/unifold_backend.py @@ -1,144 +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 __future__ import annotations - -import os -from typing import Any, Dict, Iterator, List -from absl import logging +from alphapulldown.prediction.inference_flags import validate_backend_availability from .folding_backend import FoldingBackend -# The model configurations UniFold ships; the default is what run_multimer_jobs -# always offered through --unifold_model_name. -UNIFOLD_MODEL_NAMES = ( - "multimer_af2", - "multimer_ft", - "multimer", - "multimer_af2_v3", - "multimer_af2_model45_v3", -) -DEFAULT_UNIFOLD_MODEL_NAME = "multimer_af2" - - class UnifoldBackend(FoldingBackend): - """ - A backend class for running protein structure predictions using the UniFold model. - - It follows the same contract as the other backends, which is what the - prediction adapters drive: ``setup(**model_flags)`` builds a session, - ``predict(**session, objects_to_model=..., random_seed=..., **model_flags)`` - yields one record per fold, and ``postprocess`` receives the record. The - previous version required ``output_dir`` and ``multimeric_object`` in - ``setup`` and took a single object in an instance-method ``predict``, so no - caller in the package could invoke it. - """ + """Preserve imports while rejecting legacy calls with an actionable error.""" @staticmethod - def setup( - model_dir: str, - model_name: str = DEFAULT_UNIFOLD_MODEL_NAME, - **kwargs, - ) -> Dict: - """ - Resolve the UniFold model configuration for this session. - - Parameters - ---------- - model_dir : str - The directory where the UniFold parameters are located. - model_name : str - The UniFold model configuration (``--unifold_model_name``). - **kwargs : dict - Additional keyword arguments for model configuration. Ignored. - - Returns - ------- - Dict - The model configuration under ``model_config``. The weights are - loaded per fold in :py:meth:`UnifoldBackend.predict`, because - UniFold ties its runner to the target name and output directory. - """ - from unifold.config import model_config - - if model_name not in UNIFOLD_MODEL_NAMES: - raise ValueError( - f"Unknown UniFold model {model_name!r}; choose one of " - f"{', '.join(UNIFOLD_MODEL_NAMES)}" - ) - return {"model_config": model_config(model_name)} + def setup(*args, **kwargs): + validate_backend_availability("unifold") @staticmethod - def predict( - objects_to_model: List[Dict[str, Any]], - model_config: Dict, - model_dir: str, - random_seed: int = 42, - **kwargs, - ) -> Iterator[Dict[str, Any]]: - """ - Predicts the structure of each object with the configured UniFold model. - - Parameters - ---------- - objects_to_model : List[Dict[str, Any]] - One ``{"object": ..., "output_dir": ...}`` record per fold; the object - carries ``description`` and ``feature_dict``. - model_config : Dict - Configuration obtained from :py:meth:`UnifoldBackend.setup`. - model_dir : str - The directory where the UniFold parameters are located. - random_seed : int, optional - The random seed for prediction reproducibility, default is 42. - **kwargs : dict - Additional keyword arguments for prediction. Ignored. - - Yields - ------ - Dict - ``object``, ``prediction_results`` and ``output_dir`` for each fold, - in order. UniFold writes its structures itself, so the results are - an empty mapping. - """ - from unifold.dataset import process_ap - from unifold.inference import config_args, unifold_config_model, unifold_predict - - for entry in objects_to_model: - object_to_model = entry["object"] - output_dir = entry["output_dir"] - os.makedirs(output_dir, exist_ok=True) - logging.info( - "Now running UniFold prediction on %s", object_to_model.description - ) - general_args = config_args( - model_dir, - target_name=object_to_model.description, - output_dir=output_dir, - ) - model_runner = unifold_config_model(general_args) - processed_features, _ = process_ap( - config=model_config, - features=object_to_model.feature_dict, - mode="predict", - labels=None, - seed=random_seed, - batch_idx=None, - data_idx=None, - is_distillation=False, - ) - unifold_predict(model_runner, general_args, processed_features) - yield { - "object": object_to_model, - "prediction_results": {}, - "output_dir": output_dir, - } + def predict(*args, **kwargs): + validate_backend_availability("unifold") @staticmethod - def postprocess(**kwargs) -> None: - """UniFold writes its own outputs; nothing to post-process.""" - return None + def postprocess(*args, **kwargs): + validate_backend_availability("unifold") diff --git a/alphapulldown/prediction/inference_flags.py b/alphapulldown/prediction/inference_flags.py index 713446d7e..49b47eaf7 100644 --- a/alphapulldown/prediction/inference_flags.py +++ b/alphapulldown/prediction/inference_flags.py @@ -36,9 +36,11 @@ ALPHALINK_EXTRA_FLAGS = frozenset({"crosslinks"}) -# UniFold picks its model configuration by name; it serves monomers and multimers alike. -UNIFOLD_EXTRA_FLAGS = frozenset({"unifold_model_name"}) -DEFAULT_UNIFOLD_MODEL_NAME = "multimer_af2" +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", @@ -50,13 +52,19 @@ FLAGS_BY_BACKEND: Mapping[str, frozenset] = { "alphafold2": COMMON_FLAGS | AF2_LIKE_FLAGS, "alphalink": COMMON_FLAGS | AF2_LIKE_FLAGS | ALPHALINK_EXTRA_FLAGS, - "unifold": COMMON_FLAGS | AF2_LIKE_FLAGS | UNIFOLD_EXTRA_FLAGS, "alphafold3": COMMON_FLAGS | AF3_FLAGS, } +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 [] @@ -115,13 +123,10 @@ 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. - UniFold has one model for monomers and multimers, chosen by - ``--unifold_model_name``, so the caller must leave it alone. """ + validate_backend_availability(flags.fold_backend) if flags.fold_backend == "alphalink": return "multimer_af2_crop" - if flags.fold_backend == "unifold": - return getattr(flags, "unifold_model_name", None) or DEFAULT_UNIFOLD_MODEL_NAME return "monomer_ptm" diff --git a/alphapulldown/scripts/run_multimer_jobs.py b/alphapulldown/scripts/run_multimer_jobs.py index 20f1854d2..2b5cb0cc2 100644 --- a/alphapulldown/scripts/run_multimer_jobs.py +++ b/alphapulldown/scripts/run_multimer_jobs.py @@ -27,7 +27,7 @@ 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") # --unifold_model_name is defined by run_structure_prediction and shared here. @@ -74,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]] @@ -97,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: @@ -151,7 +152,6 @@ def main(argv): "--msa_depth": FLAGS.msa_depth, "--crosslinks": FLAGS.crosslinks, "--fold_backend": fold_backend, - "--unifold_model_name": FLAGS.unifold_model_name if FLAGS.use_unifold else None, "--description_file": FLAGS.description_file, "--path_to_mmt": FLAGS.path_to_mmt, "--compress_result_pickles": FLAGS.compress_result_pickles, diff --git a/alphapulldown/scripts/run_structure_prediction.py b/alphapulldown/scripts/run_structure_prediction.py index 3e2c02bb5..7fa34f205 100644 --- a/alphapulldown/scripts/run_structure_prediction.py +++ b/alphapulldown/scripts/run_structure_prediction.py @@ -136,11 +136,10 @@ # AlphaLink2 settings flags.DEFINE_string('crosslinks', None, 'Path to crosslink information pickle for AlphaLink.') -# UniFold settings -flags.DEFINE_enum( +# Keep the legacy flag parseable so old invocations get the backend's actionable error. +flags.DEFINE_string( 'unifold_model_name', 'multimer_af2', - ['multimer_af2', 'multimer_ft', 'multimer', 'multimer_af2_v3', 'multimer_af2_model45_v3'], - 'UniFold model configuration used with --fold_backend=unifold.') + 'Legacy option. UniFold is unavailable in this release.') # AlphaFold3 settings # JAX inference performance tuning. @@ -228,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' @@ -379,10 +378,8 @@ def main(argv): json_output_dir = real_out # Flags for THIS object, not for whichever fold happens to come last. - # UniFold's model is chosen by --unifold_model_name and serves monomers - # and multimers alike, so it keeps the name the flags resolved. object_model_flags = default_model_flags.copy() - if isinstance(obj, MultimericObject) and FLAGS.fold_backend != "unifold": + if isinstance(obj, MultimericObject): object_model_flags.update({ "model_name": "multimer", "msa_depth_scan": FLAGS.msa_depth_scan, diff --git a/alphapulldown/utils/distogram_parser.py b/alphapulldown/utils/distogram_parser.py index e5346e86f..68f73db20 100644 --- a/alphapulldown/utils/distogram_parser.py +++ b/alphapulldown/utils/distogram_parser.py @@ -75,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 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_stubs() -> dict[str, types.ModuleType | None]: - names_to_replace = ["unifold", "unifold.config", "unifold.inference", "unifold.dataset"] - saved_modules = {name: sys.modules.get(name) for name in names_to_replace} - - unifold_pkg = types.ModuleType("unifold") - unifold_pkg.__path__ = [] # type: ignore[attr-defined] - 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 = { - "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 - -@pytest.fixture -def unifold_backend_module(): - saved_modules = _install_unifold_stubs() - sys.modules.pop(MODULE_NAME, None) - spec = importlib.util.spec_from_file_location(MODULE_NAME, MODULE_PATH) - module = importlib.util.module_from_spec(spec) - sys.modules[spec.name] = module - assert spec.loader is not None - try: - spec.loader.exec_module(module) - yield module - finally: - sys.modules.pop(MODULE_NAME, None) - _restore_modules(saved_modules) +@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 _unifold_flags(**overrides): - defaults = dict( - fold_backend="unifold", unifold_model_name="multimer_ft", num_cycle=3, - data_directory="/weights", num_predictions_per_model=1, crosslinks=None, - desired_num_res=None, desired_num_msa=None, skip_templates=False, - allow_resume=True, num_diffusion_samples=5, num_recycles=10, - save_embeddings=False, save_distogram=False, - flash_attention_implementation="triton", buckets=["256"], - jax_compilation_cache_dir=None, features_directory=["/features"], - num_seeds=None, debug_templates=False, debug_msas=False, dropout=False, - ) - defaults.update(overrides) - return SimpleNamespace(**defaults) +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_unifold_backend_runs_through_the_prediction_adapters(unifold_backend_module, tmp_path): - from alphapulldown.folding_backend import FoldingBackendManager - from alphapulldown.prediction.inference_flags import model_flags - from alphapulldown.prediction.prediction_batch import ( - PredictionBatch, - PredictionJob, - PreparedPredictionAdapter, - ) - fold = SimpleNamespace(description="A_and_B", feature_dict={"msa": [1, 2]}) - output_dir = tmp_path / "A_and_B" # created by the backend +def test_prediction_adapter_reports_unavailable_backend(tmp_path): + output = tmp_path / "fold" adapter = PreparedPredictionAdapter( - backend=FoldingBackendManager(), - fold_backend="unifold", - objects_to_model=[{"object": fold, "output_dir": str(output_dir)}], - model_flags=model_flags(_unifold_flags()), - postprocess_flags={}, - random_seed=7, + 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, ) - - summary = PredictionBatch((PredictionJob("legacy", "", output_dir),)).run(adapter) - - assert summary.failures == () - assert summary.completed_job_ids == ("legacy",) - assert output_dir.is_dir() - general_args = { - "model_dir": "/weights", - "target_name": "A_and_B", - "output_dir": str(output_dir), - } - assert sys.modules["unifold.inference"].calls == [ - ( - {"runner_args": general_args}, - general_args, - { - "processed_features": {"msa": [1, 2]}, - "seed": 7, - "mode": "predict", - "config": {"model_name": "multimer_ft"}, - }, - ) - ] - - -def test_unifold_predict_handles_every_object_in_order(unifold_backend_module, tmp_path): - backend = unifold_backend_module.UnifoldBackend - session = backend.setup(model_dir="/weights", model_name="multimer_af2") - objects = [ - {"object": SimpleNamespace(description=name, feature_dict={"n": index}), - "output_dir": str(tmp_path / name)} - for index, name in enumerate(("first", "second")) - ] - - records = list(backend.predict(objects, random_seed=3, model_dir="/weights", **session)) - - assert [record["object"].description for record in records] == ["first", "second"] - assert [record["output_dir"] for record in records] == [o["output_dir"] for o in objects] - assert all(record["prediction_results"] == {} for record in records) - targets = [call[1]["target_name"] for call in sys.modules["unifold.inference"].calls] - assert targets == ["first", "second"] - assert backend.postprocess(prediction_results={}, output_dir=str(tmp_path)) is None + # 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() -def test_unifold_setup_rejects_an_alphafold_model_name(unifold_backend_module): - with pytest.raises(ValueError, match="multimer_af2"): - unifold_backend_module.UnifoldBackend.setup(model_dir="/weights", model_name="monomer_ptm") +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() From f8063c090a813cdcc584685fe2868f18287ddb71 Mon Sep 17 00:00:00 2001 From: Dima Molodenskiy Date: Thu, 1 Oct 2026 12:45:57 +0200 Subject: [PATCH 3/4] Install CPU JAX for fresh-interpreter compatibility tests Canonical pytest IDs exposed two previously omitted import-order cases. Their subprocesses intentionally bypass pytest stubs and need the real AF2 JAX tree API. Install the supported CPU JAX version in smoke and coverage jobs without enabling CUDA dependencies. --- .github/workflows/github_actions.yml | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/.github/workflows/github_actions.yml b/.github/workflows/github_actions.yml index ca37b3faa..1bc497880 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 \ From 2128f446715967dff347d09293b7cb2f16a08a77 Mon Sep 17 00:00:00 2001 From: Dima Molodenskiy Date: Thu, 1 Oct 2026 12:55:11 +0200 Subject: [PATCH 4/4] Restore AF3 test stubs before subsequent test modules Coverage exposed an order-dependent fake mmCIF parser leaking from the backend fixture into real ModelCIF checks. Scope module and JAX attribute replacements to the fixture, including setup failures. Add a serial subprocess regression for the reproducing module order. The original probe failed; the corrected probe and all 43 related tests pass. --- .../test_backend_test_isolation.py | 32 ++++++++++++++ test/unit/test_alphafold3_backend_helpers.py | 44 ++++++++++++------- 2 files changed, 59 insertions(+), 17 deletions(-) create mode 100644 test/integration/test_backend_test_isolation.py diff --git a/test/integration/test_backend_test_isolation.py b/test/integration/test_backend_test_isolation.py new file mode 100644 index 000000000..41400331d --- /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/unit/test_alphafold3_backend_helpers.py b/test/unit/test_alphafold3_backend_helpers.py index af2675384..cc443fc00 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: