ase-calculator-kit 0.3.2__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- ase_calculator_kit/__init__.py +50 -0
- ase_calculator_kit/backends/__init__.py +24 -0
- ase_calculator_kit/backends/base.py +25 -0
- ase_calculator_kit/backends/dft/__init__.py +11 -0
- ase_calculator_kit/backends/dft/espresso.py +56 -0
- ase_calculator_kit/backends/dft/vasp.py +46 -0
- ase_calculator_kit/backends/mlip/__init__.py +17 -0
- ase_calculator_kit/backends/mlip/chgnet.py +73 -0
- ase_calculator_kit/backends/mlip/fairchem.py +95 -0
- ase_calculator_kit/backends/mlip/mattersim.py +86 -0
- ase_calculator_kit/backends/mlip/nequip.py +110 -0
- ase_calculator_kit/backends/mlip/sevennet.py +153 -0
- ase_calculator_kit/config.py +98 -0
- ase_calculator_kit/device.py +63 -0
- ase_calculator_kit/dispersion.py +193 -0
- ase_calculator_kit/errors.py +39 -0
- ase_calculator_kit/factory.py +113 -0
- ase_calculator_kit/py.typed +0 -0
- ase_calculator_kit/registry.py +37 -0
- ase_calculator_kit-0.3.2.dist-info/METADATA +502 -0
- ase_calculator_kit-0.3.2.dist-info/RECORD +24 -0
- ase_calculator_kit-0.3.2.dist-info/WHEEL +5 -0
- ase_calculator_kit-0.3.2.dist-info/licenses/LICENSE +21 -0
- ase_calculator_kit-0.3.2.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
"""ase-calculator-kit: a unified ASE calculator factory.
|
|
2
|
+
|
|
3
|
+
Call supported MLIP and DFT ASE calculators from one public factory::
|
|
4
|
+
|
|
5
|
+
from ase_calculator_kit import get_calculator
|
|
6
|
+
atoms.calc = get_calculator("uma", task="omat")
|
|
7
|
+
atoms.calc = get_calculator("vasp", config="examples/dft/vasp_pbe_static.yaml")
|
|
8
|
+
energy = atoms.get_potential_energy()
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from __future__ import annotations
|
|
12
|
+
|
|
13
|
+
from importlib.metadata import PackageNotFoundError, version
|
|
14
|
+
|
|
15
|
+
from .config import resolve_calculator_config
|
|
16
|
+
from .dispersion import DispersionPolicy, get_dispersion_policy
|
|
17
|
+
from .errors import CalculatorKitError, DispersionError, MissingDependencyError
|
|
18
|
+
from .factory import (
|
|
19
|
+
attach_calculator,
|
|
20
|
+
available_calculators,
|
|
21
|
+
available_dft_calculators,
|
|
22
|
+
available_mlip_models,
|
|
23
|
+
available_models,
|
|
24
|
+
get_calculator,
|
|
25
|
+
get_dft_calculator,
|
|
26
|
+
get_mlip_calculator,
|
|
27
|
+
)
|
|
28
|
+
|
|
29
|
+
try:
|
|
30
|
+
__version__ = version("ase-calculator-kit")
|
|
31
|
+
except PackageNotFoundError: # not installed (e.g. running from a source tree)
|
|
32
|
+
__version__ = "0.0.0"
|
|
33
|
+
|
|
34
|
+
__all__ = [
|
|
35
|
+
"get_calculator",
|
|
36
|
+
"get_mlip_calculator",
|
|
37
|
+
"get_dft_calculator",
|
|
38
|
+
"attach_calculator",
|
|
39
|
+
"available_calculators",
|
|
40
|
+
"available_models",
|
|
41
|
+
"available_mlip_models",
|
|
42
|
+
"available_dft_calculators",
|
|
43
|
+
"resolve_calculator_config",
|
|
44
|
+
"DispersionPolicy",
|
|
45
|
+
"get_dispersion_policy",
|
|
46
|
+
"CalculatorKitError",
|
|
47
|
+
"MissingDependencyError",
|
|
48
|
+
"DispersionError",
|
|
49
|
+
"__version__",
|
|
50
|
+
]
|
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
"""Calculator backends."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from .base import BaseBackend
|
|
6
|
+
from .dft import EspressoBackend, VaspBackend
|
|
7
|
+
from .mlip import (
|
|
8
|
+
CHGNetBackend,
|
|
9
|
+
FairChemBackend,
|
|
10
|
+
MatterSimBackend,
|
|
11
|
+
NequIPBackend,
|
|
12
|
+
SevenNetBackend,
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
__all__ = [
|
|
16
|
+
"BaseBackend",
|
|
17
|
+
"CHGNetBackend",
|
|
18
|
+
"EspressoBackend",
|
|
19
|
+
"FairChemBackend",
|
|
20
|
+
"MatterSimBackend",
|
|
21
|
+
"NequIPBackend",
|
|
22
|
+
"SevenNetBackend",
|
|
23
|
+
"VaspBackend",
|
|
24
|
+
]
|
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
"""Backend base class."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from abc import ABC, abstractmethod
|
|
6
|
+
|
|
7
|
+
from ase.calculators.calculator import Calculator
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class BaseBackend(ABC):
|
|
11
|
+
"""Abstract base for a calculator backend.
|
|
12
|
+
|
|
13
|
+
Subclasses lazily import their underlying package inside
|
|
14
|
+
:meth:`create_calculator` and raise
|
|
15
|
+
:class:`ase_calculator_kit.errors.MissingDependencyError` if it is not
|
|
16
|
+
installed.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
#: Canonical backend name (e.g. ``"chgnet"`` or ``"vasp"``).
|
|
20
|
+
name: str
|
|
21
|
+
|
|
22
|
+
@abstractmethod
|
|
23
|
+
def create_calculator(self, **kwargs) -> Calculator:
|
|
24
|
+
"""Build and return a fresh ASE calculator."""
|
|
25
|
+
raise NotImplementedError
|
|
@@ -0,0 +1,56 @@
|
|
|
1
|
+
"""Quantum ESPRESSO backend using ASE's Espresso calculator."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
from ase.calculators.calculator import Calculator
|
|
9
|
+
|
|
10
|
+
from ...config import resolve_calculator_config, write_resolved_config_file
|
|
11
|
+
from ..base import BaseBackend
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class EspressoBackend(BaseBackend):
|
|
15
|
+
name = "qe"
|
|
16
|
+
|
|
17
|
+
def create_calculator(
|
|
18
|
+
self,
|
|
19
|
+
*,
|
|
20
|
+
config: str | Path | dict[str, Any],
|
|
21
|
+
overrides: dict[str, Any] | None = None,
|
|
22
|
+
write_resolved_config: bool = False,
|
|
23
|
+
) -> Calculator:
|
|
24
|
+
"""Create an ASE :class:`ase.calculators.espresso.Espresso` from config."""
|
|
25
|
+
from ase.calculators.espresso import Espresso, EspressoProfile
|
|
26
|
+
|
|
27
|
+
resolved = resolve_calculator_config(
|
|
28
|
+
"qe",
|
|
29
|
+
config=config,
|
|
30
|
+
overrides=overrides,
|
|
31
|
+
)
|
|
32
|
+
profile_cfg = resolved.get("profile", {})
|
|
33
|
+
if "command" not in profile_cfg:
|
|
34
|
+
raise ValueError("QE config requires profile.command.")
|
|
35
|
+
if "pseudo_dir" not in profile_cfg:
|
|
36
|
+
raise ValueError("QE config requires profile.pseudo_dir.")
|
|
37
|
+
|
|
38
|
+
pseudopotentials = resolved.get("pseudopotentials")
|
|
39
|
+
if not pseudopotentials:
|
|
40
|
+
raise ValueError("QE config requires pseudopotentials.")
|
|
41
|
+
|
|
42
|
+
parameters = resolved.get("parameters", {})
|
|
43
|
+
directory = resolved.get("directory", ".")
|
|
44
|
+
if write_resolved_config:
|
|
45
|
+
write_resolved_config_file(resolved, directory)
|
|
46
|
+
|
|
47
|
+
profile = EspressoProfile(
|
|
48
|
+
command=profile_cfg["command"],
|
|
49
|
+
pseudo_dir=profile_cfg["pseudo_dir"],
|
|
50
|
+
)
|
|
51
|
+
return Espresso(
|
|
52
|
+
profile=profile,
|
|
53
|
+
directory=directory,
|
|
54
|
+
pseudopotentials=pseudopotentials,
|
|
55
|
+
**parameters,
|
|
56
|
+
)
|
|
@@ -0,0 +1,46 @@
|
|
|
1
|
+
"""VASP backend using ASE's Vasp calculator."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
from ase.calculators.calculator import Calculator
|
|
9
|
+
|
|
10
|
+
from ...config import resolve_calculator_config, write_resolved_config_file
|
|
11
|
+
from ..base import BaseBackend
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class VaspBackend(BaseBackend):
|
|
15
|
+
name = "vasp"
|
|
16
|
+
|
|
17
|
+
def create_calculator(
|
|
18
|
+
self,
|
|
19
|
+
*,
|
|
20
|
+
config: str | Path | dict[str, Any],
|
|
21
|
+
overrides: dict[str, Any] | None = None,
|
|
22
|
+
write_resolved_config: bool = False,
|
|
23
|
+
) -> Calculator:
|
|
24
|
+
"""Create an ASE :class:`ase.calculators.vasp.Vasp` from config."""
|
|
25
|
+
from ase.calculators.vasp import Vasp
|
|
26
|
+
|
|
27
|
+
resolved = resolve_calculator_config(
|
|
28
|
+
"vasp",
|
|
29
|
+
config=config,
|
|
30
|
+
overrides=overrides,
|
|
31
|
+
)
|
|
32
|
+
profile = resolved.get("profile", {})
|
|
33
|
+
if "command" not in profile:
|
|
34
|
+
raise ValueError("VASP config requires profile.command.")
|
|
35
|
+
|
|
36
|
+
parameters = resolved.get("parameters", {})
|
|
37
|
+
directory = resolved.get("directory", ".")
|
|
38
|
+
if write_resolved_config:
|
|
39
|
+
write_resolved_config_file(resolved, directory)
|
|
40
|
+
|
|
41
|
+
return Vasp(
|
|
42
|
+
command=profile["command"],
|
|
43
|
+
directory=directory,
|
|
44
|
+
txt=profile.get("txt", "vasp.out"),
|
|
45
|
+
**parameters,
|
|
46
|
+
)
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
"""Machine-learning interatomic potential backends."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from .chgnet import CHGNetBackend
|
|
6
|
+
from .fairchem import FairChemBackend
|
|
7
|
+
from .mattersim import MatterSimBackend
|
|
8
|
+
from .nequip import NequIPBackend
|
|
9
|
+
from .sevennet import SevenNetBackend
|
|
10
|
+
|
|
11
|
+
__all__ = [
|
|
12
|
+
"CHGNetBackend",
|
|
13
|
+
"FairChemBackend",
|
|
14
|
+
"MatterSimBackend",
|
|
15
|
+
"NequIPBackend",
|
|
16
|
+
"SevenNetBackend",
|
|
17
|
+
]
|
|
@@ -0,0 +1,73 @@
|
|
|
1
|
+
"""CHGNet backend (https://github.com/CederGroupHub/chgnet)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from ase.calculators.calculator import Calculator
|
|
6
|
+
|
|
7
|
+
from ...device import resolve_device
|
|
8
|
+
from ...dispersion import precheck_dispersion_xc, wrap_with_d3
|
|
9
|
+
from ...errors import MissingDependencyError
|
|
10
|
+
from ..base import BaseBackend
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class CHGNetBackend(BaseBackend):
|
|
14
|
+
name = "chgnet"
|
|
15
|
+
|
|
16
|
+
def create_calculator(
|
|
17
|
+
self,
|
|
18
|
+
*,
|
|
19
|
+
device: str = "auto",
|
|
20
|
+
model: str | None = None,
|
|
21
|
+
checkpoint: str | None = None,
|
|
22
|
+
dispersion: bool = False,
|
|
23
|
+
dispersion_xc: str | None = None,
|
|
24
|
+
**kwargs,
|
|
25
|
+
) -> Calculator:
|
|
26
|
+
"""Create a :class:`chgnet.model.dynamics.CHGNetCalculator`.
|
|
27
|
+
|
|
28
|
+
CHGNet is best for inorganic crystalline materials and fast
|
|
29
|
+
pre-relaxation. Be careful with isolated atoms and chemistry far from
|
|
30
|
+
its training domain.
|
|
31
|
+
|
|
32
|
+
Parameters
|
|
33
|
+
----------
|
|
34
|
+
device:
|
|
35
|
+
``"auto"`` (cuda > mps > cpu), or ``"cuda"`` / ``"mps"`` / ``"cpu"``.
|
|
36
|
+
CHGNet supports Apple Silicon ``"mps"``.
|
|
37
|
+
model:
|
|
38
|
+
Optional pretrained model name passed to ``CHGNet.load(model_name=...)``.
|
|
39
|
+
When omitted, CHGNet's bundled default model is used.
|
|
40
|
+
checkpoint:
|
|
41
|
+
Optional path to a ``.pth`` checkpoint, loaded via
|
|
42
|
+
``CHGNetCalculator.from_file``.
|
|
43
|
+
dispersion, dispersion_xc:
|
|
44
|
+
Add a Grimme-D3(BJ) correction (CHGNet is PBE+U, so ``xc="pbe"`` by
|
|
45
|
+
default). See ``docs/models.md`` for the per-model policy.
|
|
46
|
+
"""
|
|
47
|
+
# Validate the dispersion policy before loading the model (fail fast).
|
|
48
|
+
d3_xc = precheck_dispersion_xc(
|
|
49
|
+
self.name, model or "default",
|
|
50
|
+
dispersion=dispersion, dispersion_xc=dispersion_xc,
|
|
51
|
+
)
|
|
52
|
+
use_device = resolve_device(device, allow_mps=True)
|
|
53
|
+
|
|
54
|
+
try:
|
|
55
|
+
from chgnet.model.dynamics import CHGNetCalculator
|
|
56
|
+
except ImportError as exc: # pragma: no cover - exercised via tests with mocks
|
|
57
|
+
raise MissingDependencyError("CHGNet") from exc
|
|
58
|
+
|
|
59
|
+
if checkpoint is not None:
|
|
60
|
+
bare = CHGNetCalculator.from_file(
|
|
61
|
+
checkpoint, use_device=use_device, **kwargs
|
|
62
|
+
)
|
|
63
|
+
elif model is not None:
|
|
64
|
+
from chgnet.model.model import CHGNet
|
|
65
|
+
|
|
66
|
+
loaded = CHGNet.load(model_name=model)
|
|
67
|
+
bare = CHGNetCalculator(model=loaded, use_device=use_device, **kwargs)
|
|
68
|
+
else:
|
|
69
|
+
bare = CHGNetCalculator(use_device=use_device, **kwargs)
|
|
70
|
+
|
|
71
|
+
if d3_xc is not None:
|
|
72
|
+
return wrap_with_d3(bare, xc=d3_xc, device=use_device)
|
|
73
|
+
return bare
|
|
@@ -0,0 +1,95 @@
|
|
|
1
|
+
"""fairchem / UMA backend (https://github.com/facebookresearch/fairchem)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from ase.calculators.calculator import Calculator
|
|
6
|
+
|
|
7
|
+
from ...device import resolve_device
|
|
8
|
+
from ...dispersion import precheck_dispersion_xc, wrap_with_d3
|
|
9
|
+
from ...errors import MissingDependencyError
|
|
10
|
+
from ..base import BaseBackend
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class FairChemBackend(BaseBackend):
|
|
14
|
+
name = "uma"
|
|
15
|
+
|
|
16
|
+
def create_calculator(
|
|
17
|
+
self,
|
|
18
|
+
*,
|
|
19
|
+
device: str = "auto",
|
|
20
|
+
model: str = "uma-s-1p2",
|
|
21
|
+
task: str = "omat",
|
|
22
|
+
dispersion: bool = False,
|
|
23
|
+
dispersion_xc: str | None = None,
|
|
24
|
+
**kwargs,
|
|
25
|
+
) -> Calculator:
|
|
26
|
+
"""Create a :class:`fairchem.core.FAIRChemCalculator` for a UMA model.
|
|
27
|
+
|
|
28
|
+
UMA checkpoints are gated on Hugging Face. If creation fails with an
|
|
29
|
+
authorization error, request access to the model repository and run
|
|
30
|
+
``huggingface-cli login``.
|
|
31
|
+
|
|
32
|
+
Parameters
|
|
33
|
+
----------
|
|
34
|
+
device:
|
|
35
|
+
``"auto"`` (cuda > cpu) or explicit ``"cuda"`` / ``"cpu"``. Apple
|
|
36
|
+
Silicon ``"mps"`` is not supported: fairchem-core's predict unit
|
|
37
|
+
asserts ``device in {"cpu", "cuda"}``, so ``"mps"`` is rejected
|
|
38
|
+
before this wrapper runs.
|
|
39
|
+
model:
|
|
40
|
+
UMA model name. Defaults to ``"uma-s-1p2"``.
|
|
41
|
+
task:
|
|
42
|
+
The ``task_name`` selecting the domain-specific head. A single UMA
|
|
43
|
+
model serves many domains; pick the task matching your system:
|
|
44
|
+
|
|
45
|
+
======= ==================================================
|
|
46
|
+
``task`` Use for
|
|
47
|
+
======= ==================================================
|
|
48
|
+
``omat`` Inorganic bulk/materials, stress, cell optimization
|
|
49
|
+
``omol`` Molecules and polymers
|
|
50
|
+
``oc20`` Catalyst surfaces and adsorption
|
|
51
|
+
``oc22`` Oxide catalysis
|
|
52
|
+
``oc25`` Electrochemistry / solid-liquid interfaces
|
|
53
|
+
``odac`` MOFs and direct air capture
|
|
54
|
+
``omc`` Molecular crystals
|
|
55
|
+
======= ==================================================
|
|
56
|
+
|
|
57
|
+
For the molecular task (``omol``), set ``atoms.info["charge"]`` (total
|
|
58
|
+
charge) and ``atoms.info["spin"]`` (spin multiplicity, ``2S+1``)
|
|
59
|
+
*before* computing::
|
|
60
|
+
|
|
61
|
+
atoms.info["charge"] = -1
|
|
62
|
+
atoms.info["spin"] = 2
|
|
63
|
+
atoms.calc = get_calculator("uma", task="omol")
|
|
64
|
+
|
|
65
|
+
This is not optional in practice, only in form: fairchem does **not**
|
|
66
|
+
raise when they are missing. It logs a warning, writes
|
|
67
|
+
``charge=0`` / ``spin=1`` into ``atoms.info`` (mutating the object you
|
|
68
|
+
passed in), and returns a neutral closed-shell result. An anion or a
|
|
69
|
+
radical therefore comes back silently wrong unless both keys are set.
|
|
70
|
+
``charge`` and ``spin`` are read only by the ``omol`` head; other
|
|
71
|
+
tasks ignore them.
|
|
72
|
+
dispersion, dispersion_xc:
|
|
73
|
+
Add a Grimme-D3(BJ) correction. The D3 ``xc`` depends on the task's
|
|
74
|
+
DFT level (e.g. ``omat``→pbe, ``oc20``→rpbe). Rejected for the tasks
|
|
75
|
+
whose reference data already accounts for dispersion: ``oc25``
|
|
76
|
+
(RPBE+D3), ``omol`` (ωB97M-V nonlocal), ``odac`` (PBE-D3) and ``omc``
|
|
77
|
+
(PBE+D3). See ``docs/models.md``.
|
|
78
|
+
"""
|
|
79
|
+
# Validate the dispersion policy before loading the model (fail fast).
|
|
80
|
+
d3_xc = precheck_dispersion_xc(
|
|
81
|
+
self.name, task, dispersion=dispersion, dispersion_xc=dispersion_xc,
|
|
82
|
+
)
|
|
83
|
+
resolved_device = resolve_device(device)
|
|
84
|
+
|
|
85
|
+
try:
|
|
86
|
+
from fairchem.core import FAIRChemCalculator, pretrained_mlip
|
|
87
|
+
except ImportError as exc: # pragma: no cover - exercised via tests with mocks
|
|
88
|
+
raise MissingDependencyError("fairchem-core") from exc
|
|
89
|
+
|
|
90
|
+
predictor = pretrained_mlip.get_predict_unit(model, device=resolved_device)
|
|
91
|
+
bare = FAIRChemCalculator(predictor, task_name=task, **kwargs)
|
|
92
|
+
|
|
93
|
+
if d3_xc is not None:
|
|
94
|
+
return wrap_with_d3(bare, xc=d3_xc, device=resolved_device)
|
|
95
|
+
return bare
|
|
@@ -0,0 +1,86 @@
|
|
|
1
|
+
"""MatterSim backend (https://github.com/microsoft/mattersim)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from ase.calculators.calculator import Calculator
|
|
6
|
+
|
|
7
|
+
from ...device import resolve_device
|
|
8
|
+
from ...dispersion import precheck_dispersion_xc, wrap_with_d3
|
|
9
|
+
from ...errors import MissingDependencyError
|
|
10
|
+
from ..base import BaseBackend
|
|
11
|
+
|
|
12
|
+
#: Friendly model keys -> checkpoint file names shipped with MatterSim.
|
|
13
|
+
_MODEL_TO_LOAD_PATH = {
|
|
14
|
+
"1M": "MatterSim-v1.0.0-1M.pth",
|
|
15
|
+
"5M": "MatterSim-v1.0.0-5M.pth",
|
|
16
|
+
"MatterSim-v1.0.0-1M": "MatterSim-v1.0.0-1M.pth",
|
|
17
|
+
"MatterSim-v1.0.0-5M": "MatterSim-v1.0.0-5M.pth",
|
|
18
|
+
}
|
|
19
|
+
_DEFAULT_MODELS = {"1M", "MatterSim-v1.0.0-1M"}
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class MatterSimBackend(BaseBackend):
|
|
23
|
+
name = "mattersim"
|
|
24
|
+
|
|
25
|
+
def create_calculator(
|
|
26
|
+
self,
|
|
27
|
+
*,
|
|
28
|
+
device: str = "auto",
|
|
29
|
+
model: str = "1M",
|
|
30
|
+
load_path: str | None = None,
|
|
31
|
+
dispersion: bool = False,
|
|
32
|
+
dispersion_xc: str | None = None,
|
|
33
|
+
**kwargs,
|
|
34
|
+
) -> Calculator:
|
|
35
|
+
"""Create a :class:`mattersim.forcefield.MatterSimCalculator`.
|
|
36
|
+
|
|
37
|
+
Parameters
|
|
38
|
+
----------
|
|
39
|
+
device:
|
|
40
|
+
``"auto"`` (cuda > mps > cpu), or explicit ``"cuda"`` / ``"mps"`` /
|
|
41
|
+
``"cpu"``. MatterSim supports Apple Silicon ``"mps"``.
|
|
42
|
+
model:
|
|
43
|
+
``"1M"`` (default, fast screening) or ``"5M"`` (more accurate).
|
|
44
|
+
Keep the checkpoint fixed across a campaign.
|
|
45
|
+
load_path:
|
|
46
|
+
Explicit checkpoint path; overrides ``model``. When ``model="1M"``
|
|
47
|
+
and no ``load_path`` is given, MatterSim's default 1M checkpoint is
|
|
48
|
+
used (no ``load_path`` passed).
|
|
49
|
+
dispersion, dispersion_xc:
|
|
50
|
+
Add a Grimme-D3(BJ) correction (MatterSim is PBE, so ``xc="pbe"`` by
|
|
51
|
+
default). See ``docs/models.md``.
|
|
52
|
+
"""
|
|
53
|
+
if model in {"1M", "MatterSim-v1.0.0-1M"}:
|
|
54
|
+
key = "1M"
|
|
55
|
+
elif model in {"5M", "MatterSim-v1.0.0-5M"}:
|
|
56
|
+
key = "5M"
|
|
57
|
+
else:
|
|
58
|
+
key = "default"
|
|
59
|
+
# Validate the dispersion policy before loading the model (fail fast).
|
|
60
|
+
d3_xc = precheck_dispersion_xc(
|
|
61
|
+
self.name, key, dispersion=dispersion, dispersion_xc=dispersion_xc,
|
|
62
|
+
)
|
|
63
|
+
resolved_device = resolve_device(device, allow_mps=True)
|
|
64
|
+
|
|
65
|
+
try:
|
|
66
|
+
from mattersim.forcefield import MatterSimCalculator
|
|
67
|
+
except ImportError as exc: # pragma: no cover - exercised via tests with mocks
|
|
68
|
+
raise MissingDependencyError("MatterSim") from exc
|
|
69
|
+
|
|
70
|
+
params: dict = {"device": resolved_device}
|
|
71
|
+
|
|
72
|
+
resolved = load_path or _MODEL_TO_LOAD_PATH.get(model)
|
|
73
|
+
if load_path is not None or model not in _DEFAULT_MODELS:
|
|
74
|
+
if resolved is None:
|
|
75
|
+
raise ValueError(
|
|
76
|
+
f"Unknown MatterSim model '{model}'. Use '1M', '5M', or pass "
|
|
77
|
+
"an explicit load_path."
|
|
78
|
+
)
|
|
79
|
+
params["load_path"] = resolved
|
|
80
|
+
|
|
81
|
+
params.update(kwargs)
|
|
82
|
+
bare = MatterSimCalculator(**params)
|
|
83
|
+
|
|
84
|
+
if d3_xc is not None:
|
|
85
|
+
return wrap_with_d3(bare, xc=d3_xc, device=resolved_device)
|
|
86
|
+
return bare
|
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
"""NequIP OAM backend (https://www.nequip.net/)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
|
|
7
|
+
from ase.calculators.calculator import Calculator
|
|
8
|
+
|
|
9
|
+
from ...device import resolve_device
|
|
10
|
+
from ...dispersion import precheck_dispersion_xc, wrap_with_d3
|
|
11
|
+
from ...errors import MissingDependencyError
|
|
12
|
+
from ..base import BaseBackend
|
|
13
|
+
|
|
14
|
+
_MODEL_IDS = {
|
|
15
|
+
"S": "mir-group/NequIP-OAM-S:0.1",
|
|
16
|
+
"M": "mir-group/NequIP-OAM-M:0.1",
|
|
17
|
+
"L": "mir-group/NequIP-OAM-L:0.1",
|
|
18
|
+
"XL": "mir-group/NequIP-OAM-XL:0.1",
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def _normalized_model_size(model: str) -> str:
|
|
23
|
+
normalized = model.upper()
|
|
24
|
+
if normalized not in _MODEL_IDS:
|
|
25
|
+
valid = ", ".join(_MODEL_IDS)
|
|
26
|
+
raise ValueError(f"Unknown NequIP OAM model '{model}'. Use one of: {valid}.")
|
|
27
|
+
return normalized
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _load_target(model: str, model_path: str | Path | None) -> str:
|
|
31
|
+
"""Return either an explicit local model or NequIP's pinned OAM identifier."""
|
|
32
|
+
return str(model_path) if model_path is not None else f"nequip.net:{_MODEL_IDS[model]}"
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class NequIPBackend(BaseBackend):
|
|
36
|
+
name = "nequip"
|
|
37
|
+
|
|
38
|
+
def create_calculator(
|
|
39
|
+
self,
|
|
40
|
+
*,
|
|
41
|
+
device: str = "auto",
|
|
42
|
+
model: str = "L",
|
|
43
|
+
model_path: str | Path | None = None,
|
|
44
|
+
chemical_species_to_atom_type_map: dict[str, str] | bool | None = True,
|
|
45
|
+
compile_mode: str = "eager",
|
|
46
|
+
model_name: str = "sole_model",
|
|
47
|
+
neighborlist_backend: str = "matscipy",
|
|
48
|
+
allow_tf32: bool = False,
|
|
49
|
+
dispersion: bool = False,
|
|
50
|
+
dispersion_xc: str | None = None,
|
|
51
|
+
**kwargs,
|
|
52
|
+
) -> Calculator:
|
|
53
|
+
"""Create a NequIP ASE calculator for one of the OAM foundation models.
|
|
54
|
+
|
|
55
|
+
Parameters
|
|
56
|
+
----------
|
|
57
|
+
device:
|
|
58
|
+
``"auto"`` (cuda > cpu) or explicit ``"cuda"`` / ``"cpu"``. Apple
|
|
59
|
+
Silicon ``"mps"`` is intentionally not enabled for NequIP OAM:
|
|
60
|
+
local testing fails with ``Cannot convert a MPS Tensor to float64``
|
|
61
|
+
because the packaged OAM models use float64 buffers and PyTorch MPS
|
|
62
|
+
does not support float64.
|
|
63
|
+
model:
|
|
64
|
+
OAM model size: ``"S"``, ``"M"``, ``"L"`` (default), or ``"XL"``.
|
|
65
|
+
Case-insensitive.
|
|
66
|
+
model_path:
|
|
67
|
+
Optional local ``.nequip.zip`` or checkpoint path. When omitted, the
|
|
68
|
+
selected OAM model is loaded through NequIP's ``nequip.net:`` loader
|
|
69
|
+
and cached by NequIP.
|
|
70
|
+
chemical_species_to_atom_type_map:
|
|
71
|
+
Passed to NequIP's ASE integration. ``True`` selects the identity
|
|
72
|
+
map and silences the upstream fallback warning, which is correct for
|
|
73
|
+
OAM model type names.
|
|
74
|
+
compile_mode, model_name, neighborlist_backend, allow_tf32:
|
|
75
|
+
Forwarded to NequIP's saved-model ASE loader.
|
|
76
|
+
dispersion, dispersion_xc:
|
|
77
|
+
Optionally add Grimme-D3(BJ). OAM uses PBE-level reference data, so
|
|
78
|
+
the default D3 parameters are PBE; an explicit ``dispersion_xc``
|
|
79
|
+
overrides that verified default.
|
|
80
|
+
"""
|
|
81
|
+
normalized_model = _normalized_model_size(model)
|
|
82
|
+
|
|
83
|
+
# Validate the dispersion policy before loading the model (fail fast).
|
|
84
|
+
d3_xc = precheck_dispersion_xc(
|
|
85
|
+
self.name,
|
|
86
|
+
normalized_model,
|
|
87
|
+
dispersion=dispersion,
|
|
88
|
+
dispersion_xc=dispersion_xc,
|
|
89
|
+
)
|
|
90
|
+
resolved_device = resolve_device(device)
|
|
91
|
+
|
|
92
|
+
try:
|
|
93
|
+
from nequip.integrations.ase import NequIPCalculator
|
|
94
|
+
except ImportError as exc: # pragma: no cover - exercised via tests with mocks
|
|
95
|
+
raise MissingDependencyError("NequIP") from exc
|
|
96
|
+
|
|
97
|
+
bare = NequIPCalculator._from_saved_model(
|
|
98
|
+
_load_target(normalized_model, model_path),
|
|
99
|
+
device=resolved_device,
|
|
100
|
+
chemical_species_to_atom_type_map=chemical_species_to_atom_type_map,
|
|
101
|
+
allow_tf32=allow_tf32,
|
|
102
|
+
model_name=model_name,
|
|
103
|
+
compile_mode=compile_mode,
|
|
104
|
+
neighborlist_backend=neighborlist_backend,
|
|
105
|
+
**kwargs,
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
if d3_xc is not None:
|
|
109
|
+
return wrap_with_d3(bare, xc=d3_xc, device=resolved_device)
|
|
110
|
+
return bare
|