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
18 changes: 16 additions & 2 deletions src/Auto3D/cli/commands/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,20 @@
from Auto3D.cli.console import console
from Auto3D.cli.errors import handle_error
from Auto3D.exceptions import ConfigurationError
from Auto3D.models.policy import ANI_ELEMENTS
from Auto3D.models.species import format_elements

#: What ANI2x and ANI2xt accept, quoted from the gate that enforces it rather
#: than retyped -- see ``format_elements``. Rendering at import pulls rdkit in
#: with this module; that is a knowing choice and costs nothing, since rdkit is a
#: required dependency and the CLI already imports torch.
#:
#: The four AIMNet2 entries below stay literals on purpose. Their source of truth
#: is each model file's own ``implemented_species`` metadata, and deriving them
#: here would mean loading four NNPs to print a help table.
#: ``tests/test_element_sets.py`` pins them to that metadata in the slow tier
#: instead.
_ANI_ELEMENT_STRING = format_elements(ANI_ELEMENTS)


def check_dependency_status(name: str) -> tuple[bool, str]:
Expand Down Expand Up @@ -165,7 +179,7 @@ def execute_models_list() -> None:
"ANI2X": {
"name": "ANI-2x",
"description": "Accurate neural network potential for organic molecules.",
"elements": "H, C, N, O, F, S, Cl",
"elements": _ANI_ELEMENT_STRING,
"speed": "8-model ensemble (not benchmarked here)",
"accuracy": "Excellent for covered elements",
"reference": "https://github.com/aiqm/torchani",
Expand All @@ -178,7 +192,7 @@ def execute_models_list() -> None:
"ANI2XT": {
"name": "ANI-2xt",
"description": "Extended ANI-2x with improved torsion handling.",
"elements": "H, C, N, O, F, S, Cl",
"elements": _ANI_ELEMENT_STRING,
"speed": "Single model (not benchmarked here)",
"accuracy": "Good for conformer generation",
"reference": "https://github.com/aiqm/torchani",
Expand Down
13 changes: 13 additions & 0 deletions src/Auto3D/models/policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,11 +16,24 @@

from Auto3D.constants import BUILTIN_ANI_MODELS
from Auto3D.exceptions import ConfigurationError, GPUError
from Auto3D.models.species import ANI2XT_INDEX

#: Elements ANI2x/ANI2xt were trained on. AIMNET (and any aimnet registry
#: model) and a custom NNP path are not restricted to this set.
ANI_ELEMENTS = frozenset({1, 6, 7, 8, 9, 16, 17})

# The gate above and ANI2xt's network-index table must cover the same elements.
# Asserted rather than derived, deliberately: they are the same seven numbers but
# not the same fact -- this set is what ANI2x AND ANI2xt were trained on, while
# ANI2XT_INDEX is one engine's 0-based index order. Defining either in terms of
# the other would record a provenance that is not true and would stop being a
# check. Same construction as model_factory's BUILTIN_ANI_MODELS assert.
assert frozenset(ANI2XT_INDEX) == ANI_ELEMENTS, (
"ANI_ELEMENTS (what check_engine_supports_molecules admits) and "
"ANI2XT_INDEX (what to_ani2xt_species can remap) have drifted apart: "
f"{sorted(ANI_ELEMENTS ^ frozenset(ANI2XT_INDEX))}"
)


def check_gpu_requested(use_gpu: bool) -> None:
"""Raise if GPU was requested but no CUDA device is visible.
Expand Down
38 changes: 35 additions & 3 deletions src/Auto3D/models/species.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,13 +32,42 @@

from __future__ import annotations

from collections.abc import Sequence
from collections.abc import Iterable, Sequence

# Atomic number -> ANI2xt network index. The order matches the ModuleList in
# models/ani2xt.py; changing one without the other misroutes elements.
ANI2XT_INDEX: dict[int, int] = {1: 0, 6: 1, 7: 2, 8: 3, 9: 4, 16: 5, 17: 6}

__all__ = ["ANI2XT_INDEX", "to_ani2xt_species"]
__all__ = ["ANI2XT_INDEX", "format_elements", "to_ani2xt_species"]


def format_elements(atomic_numbers: Iterable[int]) -> str:
"""Render an element set as ``"H, C, N, O, F, S, Cl"``.

Ordered by **atomic number**, which is the order every hand-written copy of
every element string in this package already used -- so routing them through
here changes no user-visible output. Sorting by symbol instead would silently
rewrite all of them.

One renderer because there were five hand-maintained copies of the ANI set
across three layers: the numeric gate in :mod:`Auto3D.models.policy`, the keys
of :data:`ANI2XT_INDEX` below, the message in :func:`to_ani2xt_species`, and
two entries in the CLI's ``ENGINE_INFO``. They agreed by hand, which is the
same arrangement the engine registry replaced for engine *names*.

Args:
atomic_numbers: Atomic numbers, in any order and with any duplicates.

Returns:
Comma-separated element symbols, ascending by atomic number.
"""
# Deferred exactly as in to_ani2xt_species below, and for the same reason:
# this module is a lookup table imported by the padder, by ``ASE/`` and by
# ``cli/``, none of which otherwise need rdkit. Only rendering does.
from rdkit import Chem

table = Chem.GetPeriodicTable()
return ", ".join(table.GetElementSymbol(int(z)) for z in sorted(set(atomic_numbers)))


def to_ani2xt_species(atomic_numbers: Sequence[int]) -> list[int]:
Expand Down Expand Up @@ -67,8 +96,11 @@ def to_ani2xt_species(atomic_numbers: Sequence[int]) -> list[int]:
from rdkit import Chem

symbol = Chem.GetPeriodicTable().GetElementSymbol(int(atomic_num))
# The supported set is rendered from the table this loop indexes,
# not retyped beside it: the message and the check it explains
# cannot disagree about which elements ANI2xt accepts.
raise ValueError(
f"Element Z={atomic_num} ({symbol}) is not supported by "
f"ANI2xt (supported: H, C, N, O, F, S, Cl)."
f"ANI2xt (supported: {format_elements(ANI2XT_INDEX)})."
) from None
return converted
133 changes: 133 additions & 0 deletions tests/test_element_sets.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,133 @@
# tests/test_element_sets.py
"""Each engine's element set, and the one place it is written down.

The ANI element set appeared five times across three layers -- a numeric
``frozenset`` in ``models/policy.py`` (the gate), the keys of ``ANI2XT_INDEX`` in
``models/species.py`` (the remap), a hand-written symbol string in that module's
error message, and two more copies of that string in the CLI's ``ENGINE_INFO``.
All five agreed, by hand, with nothing connecting them. That is the same shape as
the three parallel engine-name lists the registry collapsed: correct today,
correct only until someone edits one of them.

AIMNet2's sets are different in kind and are handled differently here. The
authoritative source is the model file's own ``implemented_species`` metadata,
which ``AIMNet2Calculator`` enforces at call time -- so Auto3D does not get to
define them, only to quote them correctly. The slow test at the bottom is what
makes "correctly" checkable.
"""

from __future__ import annotations

import pytest

from Auto3D.models.policy import ANI_ELEMENTS
from Auto3D.models.species import ANI2XT_INDEX, format_elements

#: The string every one of these sites has always shown. Written out rather than
#: computed, so that a change to the renderer has something to be wrong against.
ANI_ELEMENT_STRING = "H, C, N, O, F, S, Cl"


def test_format_elements_orders_by_atomic_number():
"""Not sorted by symbol, and not in set-iteration order.

Atomic number is the order every existing string already used, which is what
lets this renderer replace all of them without changing a single character of
user-visible output.
"""
assert format_elements({8, 1, 6}) == "H, C, O"
assert format_elements({53, 35, 46, 34}) == "Se, Br, Pd, I"


def test_the_ani_element_string_is_the_one_it_has_always_been():
"""The behavior lock. If this fails, the renderer changed the output."""
assert format_elements(ANI_ELEMENTS) == ANI_ELEMENT_STRING


def test_ani2xt_index_covers_exactly_the_ani_element_set():
"""The gate and the remap agree -- asserted, not assumed.

Deliberately an equality check rather than deriving one from the other. They
are the same seven numbers but not the same fact: ``ANI_ELEMENTS`` is what
ANI2x and ANI2xt were *trained* on, ``ANI2XT_INDEX`` is one engine's 0-based
network index order. Defining either in terms of the other would record a
provenance that is not true, and would quietly stop being a check.
"""
assert ANI_ELEMENTS == frozenset(ANI2XT_INDEX)


def test_unsupported_element_message_names_the_supported_set(caplog):
"""The remap's rejection message renders from the set it enforces."""
from Auto3D.models.species import to_ani2xt_species

with pytest.raises(ValueError) as exc:
to_ani2xt_species([1, 6, 5]) # boron: in AIMNet2's set, not ANI's

message = str(exc.value)
assert "Z=5" in message and "(B)" in message
assert ANI_ELEMENT_STRING in message


def test_engine_info_ani_entries_render_from_the_element_set():
"""The CLI quotes the gate rather than restating it.

Two entries carried this string as a literal. A retrained ANI with a
different element set would have moved the gate and left the CLI advertising
the old one -- and ``auto3d models info`` is where a user checks precisely
this before choosing an engine.
"""
from Auto3D.cli.commands.models import ENGINE_INFO

for name in ("ANI2X", "ANI2XT"):
assert ENGINE_INFO[name]["elements"] == format_elements(ANI_ELEMENTS), (
f"{name}'s advertised element set no longer matches the set "
f"check_engine_supports_molecules actually enforces"
)


@pytest.mark.slow
@pytest.mark.parametrize(
("info_key", "registry_name"),
[
("AIMNET", "aimnet2"),
("AIMNET2-2025", "aimnet2-2025"),
("AIMNET2-NSE", "aimnet2-nse"),
("AIMNET2-PD", "aimnet2-pd"),
],
)
def test_aimnet_engine_info_elements_match_the_model_metadata(info_key, registry_name):
"""What ``auto3d models info`` advertises is what the model file declares.

These four strings stay literals in ``ENGINE_INFO`` -- unlike the ANI pair
above -- because deriving them would mean loading four NNPs at CLI import.
This test is the alternative: it loads them once, in the slow tier, and pins
the literals to ground truth. ``AIMNet2Calculator`` reads the same
``implemented_species`` to reject out-of-set atomic numbers at call time, so
a mismatch here is the CLI promising chemistry the engine will refuse.

**If this goes red, the literal is stale -- do not delete the test.** It fails
exactly when aimnet ships a model whose element set changed, which is the one
moment the advertised string needs updating and the one moment nothing else
would say so.

CPU-only and constructed directly rather than through ``create_model``: the
subject is what the aimnet package declares, not how Auto3D wraps it, and
reaching into an adapter for it would be the abstraction leak this codebase
just finished removing.
"""
import torch

from Auto3D.cli.commands.models import ENGINE_INFO

aimnet_calculators = pytest.importorskip("aimnet.calculators")

calc = aimnet_calculators.AIMNet2Calculator(registry_name, device=torch.device("cpu"))
implemented = (calc.metadata or {}).get("implemented_species")
assert implemented is not None, (
f"{registry_name} declares no implemented_species, so aimnet's own "
f"element validation is a silent no-op for it and this test cannot "
f"check the advertised set"
)
numbers = implemented.tolist() if hasattr(implemented, "tolist") else list(implemented)

assert ENGINE_INFO[info_key]["elements"] == format_elements(numbers)
Loading