Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 4 additions & 2 deletions .github/workflows/github_actions.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 \
Expand Down
4 changes: 4 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
5 changes: 3 additions & 2 deletions alphapulldown/folding_backend/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
}
)
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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)

117 changes: 73 additions & 44 deletions alphapulldown/folding_backend/alphafold2_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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}

Expand Down
115 changes: 15 additions & 100 deletions alphapulldown/folding_backend/unifold_backend.py
Original file line number Diff line number Diff line change
@@ -1,111 +1,26 @@
""" Implements structure prediction backend using UniFold.
"""Compatibility entrypoint for the unavailable legacy UniFold backend.

Copyright (c) 2024 European Molecular Biology Laboratory

Author: Valentin Maurer <valentin.maurer@embl-hamburg.de>
The packaged ``unifold`` namespace belongs to AlphaLink2. It lacks the old
inference helpers and adds crosslink layers to the network, so routing native
UniFold checkpoints through it is not a supported substitute for UniFold.
"""
from typing import Dict

from alphapulldown.objects import MultimericObject
from alphapulldown.prediction.inference_flags import validate_backend_availability

from .folding_backend import FoldingBackend


class UnifoldBackend(FoldingBackend):
"""
A backend class for running protein structure predictions using the UniFold model.
"""
@staticmethod
def setup(
model_name: str,
model_dir: str,
output_dir: str,
multimeric_object: MultimericObject,
**kwargs,
) -> Dict:
"""
Initializes and configures a UniFold model runner.

Parameters
----------
model_name : str
The name of the model to use for prediction.
model_dir : str
The directory where the model files are located.
output_dir : str
The directory where the prediction outputs will be saved.
multimeric_object : MultimericObject
An object containing the description and features of the
multimeric protein to predict.
**kwargs : dict
Additional keyword arguments for model configuration.

Returns
-------
Dict
A dictionary containing the model runner, arguments, and configuration.
"""
from unifold.config import model_config
from unifold.inference import config_args, unifold_config_model

configs = model_config(model_name)
general_args = config_args(
model_dir, target_name=multimeric_object.description, output_dir=output_dir
)
model_runner = unifold_config_model(general_args)

return {
"model_runner": model_runner,
"model_args": general_args,
"model_config": configs,
}
"""Preserve imports while rejecting legacy calls with an actionable error."""

def predict(
self,
model_runner,
model_args,
model_config: Dict,
multimeric_object: MultimericObject,
random_seed: int = 42,
**kwargs,
) -> None:
"""
Predicts the structure of proteins using configured UniFold models.

Parameters
----------
model_runner
The configured model runner for predictions obtained
from :py:meth:`UnifoldBackend.setup`.
model_args
Arguments used for running the UniFold prediction obtained from
from :py:meth:`UnifoldBackend.setup`.
model_config : Dict
Configuration dictionary for the UniFold model obtained from
from :py:meth:`UnifoldBackend.setup`.
multimeric_object : MultimericObject
An object containing the features of the multimeric protein to predict.
random_seed : int, optional
The random seed for prediction reproducibility, default is 42.
**kwargs : dict
Additional keyword arguments for prediction.
"""
from unifold.dataset import process_ap
from unifold.inference import unifold_predict

processed_features, _ = process_ap(
config=model_config,
features=multimeric_object.feature_dict,
mode="predict",
labels=None,
seed=random_seed,
batch_idx=None,
data_idx=None,
is_distillation=False,
)
unifold_predict(model_runner, model_args, processed_features)
@staticmethod
def setup(*args, **kwargs):
validate_backend_availability("unifold")

return None
@staticmethod
def predict(*args, **kwargs):
validate_backend_availability("unifold")

def postprocess(**kwargs) -> None:
return None
@staticmethod
def postprocess(*args, **kwargs):
validate_backend_availability("unifold")
35 changes: 35 additions & 0 deletions alphapulldown/objects.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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",
Expand Down
Loading
Loading