From 411ad7b05be32b1a9e60d12cfb22033f6878ff20 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Tue, 12 May 2026 13:47:03 +0200 Subject: [PATCH 01/43] Add an empty metatomic-core python package, re-exporting metatomic-torch --- pyproject.toml | 14 +- python/metatomic_core/AUTHORS | 1 + .../CMakeLists.txt} | 0 python/metatomic_core/LICENSE | 1 + python/metatomic_core/MANIFEST.in | 6 + python/metatomic_core/metatomic/__init__.py | 0 python/metatomic_core/metatomic/torch.py | 14 ++ python/metatomic_core/pyproject.toml | 54 +++++++ python/metatomic_core/setup.py | 146 ++++++++++++++++++ python/metatomic_torch/CMakeLists.txt | 15 +- python/metatomic_torch/MANIFEST.in | 2 +- python/metatomic_torch/README.rst | 6 +- .../torch => metatomic_torch}/__init__.py | 8 + .../torch => metatomic_torch}/_c_lib.py | 0 .../torch => metatomic_torch}/_extensions.py | 0 .../ase_calculator.py | 0 .../data/dftd3_parameters.npz | Bin .../torch => metatomic_torch}/dftd3.py | 0 .../documentation.py | 0 .../torch => metatomic_torch}/heat_flux.py | 2 +- .../torch => metatomic_torch}/model.py | 0 .../torch => metatomic_torch}/o3/__init__.py | 0 .../o3/_tranformations.py | 0 .../torch => metatomic_torch}/o3/_wigner.py | 0 .../serialization.py | 0 .../systems_to_torch.py | 0 .../torch => metatomic_torch}/utils.py | 0 .../torch => metatomic_torch}/version.py | 0 python/metatomic_torch/pyproject.toml | 4 - python/metatomic_torch/setup.py | 8 +- scripts/clean-python.sh | 9 ++ setup.py | 20 ++- tox.ini | 22 ++- 33 files changed, 303 insertions(+), 29 deletions(-) create mode 120000 python/metatomic_core/AUTHORS rename python/{metatomic_torch/metatomic/__init__.py => metatomic_core/CMakeLists.txt} (100%) create mode 120000 python/metatomic_core/LICENSE create mode 100644 python/metatomic_core/MANIFEST.in create mode 100644 python/metatomic_core/metatomic/__init__.py create mode 100644 python/metatomic_core/metatomic/torch.py create mode 100644 python/metatomic_core/pyproject.toml create mode 100644 python/metatomic_core/setup.py rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/__init__.py (92%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/_c_lib.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/_extensions.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/ase_calculator.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/data/dftd3_parameters.npz (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/dftd3.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/documentation.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/heat_flux.py (99%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/model.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/o3/__init__.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/o3/_tranformations.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/o3/_wigner.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/serialization.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/systems_to_torch.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/utils.py (100%) rename python/metatomic_torch/{metatomic/torch => metatomic_torch}/version.py (100%) diff --git a/pyproject.toml b/pyproject.toml index 88dc392b9..2db00795e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -63,12 +63,16 @@ filterwarnings = [ "ignore:ast.NameConstant is deprecated and will be removed in Python 3.14:DeprecationWarning", # TorchScript deprecation warnings "ignore:`torch.jit.script` is deprecated. Please switch to `torch.compile` or `torch.export`:DeprecationWarning", + "ignore:`torch.jit.script_method` is deprecated. Please switch to `torch.compile` or `torch.export`:DeprecationWarning", "ignore:`torch.jit.save` is deprecated. Please switch to `torch.export`:DeprecationWarning", - "ignore:.*vesin.metatomic was only tested with metatomic.torch >=0.1.3,<0.2.*:UserWarning", "ignore:`torch.jit.load` is deprecated. Please switch to `torch.export`.:DeprecationWarning", "ignore:`torch.jit.script` is not supported in Python 3.14+:DeprecationWarning", + "ignore:`torch.jit.script_method` is not supported in Python 3.14+:DeprecationWarning", "ignore:`torch.jit.save` is not supported in Python 3.14+:DeprecationWarning", - # deprecation warning from warp/nvalchemi + # vesin and metatomic warning + "ignore:.*vesin.metatomic was only tested with metatomic.torch >=0.1.3,<0.2.*:UserWarning", + # Warnings from warp (dependency of nvalchemi) + "ignore:.*Structure will use memory layout compatible with MSVC:DeprecationWarning", "ignore:warp.config.quiet is deprecated:DeprecationWarning", ] @@ -95,6 +99,8 @@ docstring-code-format = true [tool.uv.pip] reinstall-package = [ - "metatomic-torch", - "metatomic-torchsim", + "metatomic_core", + "metatomic_torch", + "metatomic_torchsim", + "metatomic_ase", ] diff --git a/python/metatomic_core/AUTHORS b/python/metatomic_core/AUTHORS new file mode 120000 index 000000000..f04b7e8a2 --- /dev/null +++ b/python/metatomic_core/AUTHORS @@ -0,0 +1 @@ +../../AUTHORS \ No newline at end of file diff --git a/python/metatomic_torch/metatomic/__init__.py b/python/metatomic_core/CMakeLists.txt similarity index 100% rename from python/metatomic_torch/metatomic/__init__.py rename to python/metatomic_core/CMakeLists.txt diff --git a/python/metatomic_core/LICENSE b/python/metatomic_core/LICENSE new file mode 120000 index 000000000..30cff7403 --- /dev/null +++ b/python/metatomic_core/LICENSE @@ -0,0 +1 @@ +../../LICENSE \ No newline at end of file diff --git a/python/metatomic_core/MANIFEST.in b/python/metatomic_core/MANIFEST.in new file mode 100644 index 000000000..02404051b --- /dev/null +++ b/python/metatomic_core/MANIFEST.in @@ -0,0 +1,6 @@ +include pyproject.toml +include CMakeLists.txt +include AUTHORS +include LICENSE + +include git_version_info diff --git a/python/metatomic_core/metatomic/__init__.py b/python/metatomic_core/metatomic/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/metatomic_core/metatomic/torch.py b/python/metatomic_core/metatomic/torch.py new file mode 100644 index 000000000..060e7bccf --- /dev/null +++ b/python/metatomic_core/metatomic/torch.py @@ -0,0 +1,14 @@ +import sys + + +try: + import metatomic_torch +except ImportError as e: + raise ImportError( + "metatomic-torch is required to use the metatomic.torch module. " + "Please install it with `pip install metatomic-torch` or using " + "your favorite Python package manager." + ) from e + +# metatomic.torch is registered as an alias in metatomic_torch's __init__.py +assert sys.modules["metatomic.torch"] is metatomic_torch diff --git a/python/metatomic_core/pyproject.toml b/python/metatomic_core/pyproject.toml new file mode 100644 index 000000000..9107f6805 --- /dev/null +++ b/python/metatomic_core/pyproject.toml @@ -0,0 +1,54 @@ +[project] +name = "metatomic-core" +dynamic = ["version", "authors", "dependencies"] +requires-python = ">=3.10" + +# readme = "TODO" +license = "BSD-3-Clause" +description = "Interface between atomistic machine learning models and simulation tools" + +keywords = ["machine learning", "molecular modeling"] +classifiers = [ + "Development Status :: 4 - Beta", + "Intended Audience :: Science/Research", + "Operating System :: POSIX", + "Operating System :: MacOS :: MacOS X", + "Operating System :: Microsoft :: Windows", + "Programming Language :: Python", + "Programming Language :: Python :: 3", + "Topic :: Scientific/Engineering", + "Topic :: Scientific/Engineering :: Bio-Informatics", + "Topic :: Scientific/Engineering :: Chemistry", + "Topic :: Scientific/Engineering :: Physics", + "Topic :: Software Development :: Libraries", + "Topic :: Software Development :: Libraries :: Python Modules", +] + +[project.urls] +homepage = "https://docs.metatensor.org/metatomic/" +documentation = "https://docs.metatensor.org/metatomic/" +repository = "https://github.com/metatensor/metatomic" +# changelog = "TODO" + +### ======================================================================== ### +[build-system] +requires = [ + "setuptools >=77", + "packaging >=26", + "cmake", + "metatensor-core >=0.2.0,<0.3", +] + +build-backend = "setuptools.build_meta" + + +[tool.setuptools] +zip-safe = false + +### ======================================================================== ### +[tool.pytest.ini_options] +python_files = ["*.py"] +testpaths = ["tests"] +filterwarnings = [ + "error", +] diff --git a/python/metatomic_core/setup.py b/python/metatomic_core/setup.py new file mode 100644 index 000000000..f2654eef3 --- /dev/null +++ b/python/metatomic_core/setup.py @@ -0,0 +1,146 @@ +import os +import subprocess +import sys + +import packaging.version +from setuptools import setup +from setuptools.command.bdist_egg import bdist_egg +from setuptools.command.sdist import sdist + + +ROOT = os.path.realpath(os.path.dirname(__file__)) + +METATOMIC_CORE_VERSION = "0.1.0" + +METATOMIC_BUILD_TYPE = os.environ.get("METATOMIC_BUILD_TYPE", "release") +if METATOMIC_BUILD_TYPE not in ["debug", "release"]: + raise Exception( + f"invalid build type passed: '{METATOMIC_BUILD_TYPE}', " + "expected 'debug' or 'release'" + ) + + +class bdist_egg_disabled(bdist_egg): + """Disabled version of bdist_egg + + Prevents setup.py install performing setuptools' default easy_install, + which it should never ever do. + """ + + def run(self): + sys.exit( + "Aborting implicit building of eggs.\nUse `pip install .` or " + "`python -m build --wheel . && pip install dist/metatomic_torch-*.whl` " + "to install from source." + ) + + +class sdist_generate_data(sdist): + """ + Create a sdist with an additional generated files: + - `git_version_info` + """ + + def run(self): + n_commits, git_hash = git_version_info() + with open("git_version_info", "w") as fd: + fd.write(f"{n_commits}\n{git_hash}\n") + + # run original sdist + super().run() + + os.unlink("git_version_info") + + +def git_version_info(): + """ + If git is available and we are building from a checkout, get the number of commits + since the last tag & full hash of the code. Otherwise, this always returns (0, ""). + """ + TAG_PREFIX = "metatomic-v" + + if os.path.exists("git_version_info"): + # we are building from a sdist, without git available, but the git + # version was recorded in the `git_version_info` file + with open("git_version_info") as fd: + n_commits = int(fd.readline().strip()) + git_hash = fd.readline().strip() + else: + script = os.path.join(ROOT, "..", "..", "scripts", "git-version-info.py") + assert os.path.exists(script) + + output = subprocess.run( + [sys.executable, script, TAG_PREFIX], + stderr=subprocess.PIPE, + stdout=subprocess.PIPE, + encoding="utf8", + ) + + if output.returncode != 0: + raise Exception( + "failed to get git version info.\n" + f"stdout: {output.stdout}\n" + f"stderr: {output.stderr}\n" + ) + elif output.stderr: + print(output.stderr, file=sys.stderr) + n_commits = 0 + git_hash = "" + else: + lines = output.stdout.splitlines() + n_commits = int(lines[0].strip()) + git_hash = lines[1].strip() + + return n_commits, git_hash + + +def create_version_number(version): + version = packaging.version.parse(version) + + n_commits, git_hash = git_version_info() + + if n_commits != 0: + # if we have commits since the last tag, this mean we are in a pre-release of + # the next version. So we increase either the minor version number or the + # release candidate number (if we are closing up on a release) + if version.pre is not None: + assert version.pre[0] == "rc" + pre = ("rc", version.pre[1] + 1) + release = version.release + else: + major, minor, _ = version.release + release = (major, minor + 1, 0) + pre = None + + version = version.__replace__( + release=release, + pre=pre, + dev=n_commits, + local=git_hash, + ) + + return str(version) + + +if __name__ == "__main__": + with open(os.path.join(ROOT, "AUTHORS")) as fd: + authors = fd.read().splitlines() + + if authors[0].startswith(".."): + # handle "raw" symlink files (on Windows or from full repo tarball) + with open(os.path.join(ROOT, authors[0])) as fd: + authors = fd.read().splitlines() + + install_requires = [ + "metatensor-core >=0.2.0,<0.3", + ] + + setup( + version=create_version_number(METATOMIC_CORE_VERSION), + author=", ".join(authors), + install_requires=install_requires, + cmdclass={ + "bdist_egg": bdist_egg if "bdist_egg" in sys.argv else bdist_egg_disabled, + "sdist": sdist_generate_data, + }, + ) diff --git a/python/metatomic_torch/CMakeLists.txt b/python/metatomic_torch/CMakeLists.txt index 3578cd11f..74702d3ac 100644 --- a/python/metatomic_torch/CMakeLists.txt +++ b/python/metatomic_torch/CMakeLists.txt @@ -63,6 +63,9 @@ else() add_subdirectory("${METATOMIC_TORCH_SOURCE_DIR}" metatomic-torch) + if (CMAKE_VERSION VERSION_LESS "3.25") + set(LINUX $) + endif() if (LINUX OR APPLE) if (LINUX) @@ -74,12 +77,12 @@ else() set(metatomic_install_rpath "${CMAKE_INSTALL_RPATH}") # when loading the libraries from a Python installation: - # - $ORIGIN/../../../../torch/lib is where libtorch.so will be - # - $ORIGIN/../../../../metatensor/lib is where libmetatensor.so will be - # - $ORIGIN/../../../../metatensor/torch/torch-x.y/lib is where libmetatensor_torch.so will be - set(metatomic_install_rpath "${metatomic_install_rpath};${rpath_origin}/../../../../torch/lib") - set(metatomic_install_rpath "${metatomic_install_rpath};${rpath_origin}/../../../../metatensor/lib") - set(metatomic_install_rpath "${metatomic_install_rpath};${rpath_origin}/../../../../metatensor/torch/torch-${Torch_VERSION_MAJOR}.${Torch_VERSION_MINOR}/lib") + # - $ORIGIN/../../../torch/lib is where libtorch.so will be + # - $ORIGIN/../../../metatensor/lib is where libmetatensor.so will be + # - $ORIGIN/../../../metatensor_torch/torch-${Torch_VERSION_MAJOR}.${Torch_VERSION_MINOR}/lib is where libmetatensor_torch.so will be + set(metatomic_install_rpath "${metatomic_install_rpath};${rpath_origin}/../../../torch/lib") + set(metatomic_install_rpath "${metatomic_install_rpath};${rpath_origin}/../../../metatensor/lib") + set(metatomic_install_rpath "${metatomic_install_rpath};${rpath_origin}/../../../metatensor_torch/torch-${Torch_VERSION_MAJOR}.${Torch_VERSION_MINOR}/lib") set_target_properties( metatomic_torch PROPERTIES INSTALL_RPATH "${metatomic_install_rpath}" diff --git a/python/metatomic_torch/MANIFEST.in b/python/metatomic_torch/MANIFEST.in index 5f3ae9425..eb7359c7c 100644 --- a/python/metatomic_torch/MANIFEST.in +++ b/python/metatomic_torch/MANIFEST.in @@ -5,7 +5,7 @@ include LICENSE include git_version_info -include metatomic-torch-*.tar.gz +include metatomic-torch-cxx-*.tar.gz recursive-include build-backend *.py diff --git a/python/metatomic_torch/README.rst b/python/metatomic_torch/README.rst index f06f2b8af..994fda75e 100644 --- a/python/metatomic_torch/README.rst +++ b/python/metatomic_torch/README.rst @@ -1,4 +1,4 @@ -metatensor-torch -================ +metatomic-torch +=============== -This package contains the TorchScript bindings to the core API of metatensor. +This package contains the TorchScript bindings to the core API of metatomic. diff --git a/python/metatomic_torch/metatomic/torch/__init__.py b/python/metatomic_torch/metatomic_torch/__init__.py similarity index 92% rename from python/metatomic_torch/metatomic/torch/__init__.py rename to python/metatomic_torch/metatomic_torch/__init__.py index 06a9ae9c5..ce03634b5 100644 --- a/python/metatomic_torch/metatomic/torch/__init__.py +++ b/python/metatomic_torch/metatomic_torch/__init__.py @@ -1,8 +1,11 @@ import os +import sys from typing import TYPE_CHECKING import torch +import metatomic + from ._c_lib import _load_library from .version import __version__ # noqa: F401 @@ -68,3 +71,8 @@ save_buffer, ) from .systems_to_torch import systems_to_torch # noqa: F401 + + +sys.modules["metatomic.torch"] = sys.modules[__name__] +if not hasattr(metatomic, "torch"): + metatomic.torch = sys.modules[__name__] diff --git a/python/metatomic_torch/metatomic/torch/_c_lib.py b/python/metatomic_torch/metatomic_torch/_c_lib.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/_c_lib.py rename to python/metatomic_torch/metatomic_torch/_c_lib.py diff --git a/python/metatomic_torch/metatomic/torch/_extensions.py b/python/metatomic_torch/metatomic_torch/_extensions.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/_extensions.py rename to python/metatomic_torch/metatomic_torch/_extensions.py diff --git a/python/metatomic_torch/metatomic/torch/ase_calculator.py b/python/metatomic_torch/metatomic_torch/ase_calculator.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/ase_calculator.py rename to python/metatomic_torch/metatomic_torch/ase_calculator.py diff --git a/python/metatomic_torch/metatomic/torch/data/dftd3_parameters.npz b/python/metatomic_torch/metatomic_torch/data/dftd3_parameters.npz similarity index 100% rename from python/metatomic_torch/metatomic/torch/data/dftd3_parameters.npz rename to python/metatomic_torch/metatomic_torch/data/dftd3_parameters.npz diff --git a/python/metatomic_torch/metatomic/torch/dftd3.py b/python/metatomic_torch/metatomic_torch/dftd3.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/dftd3.py rename to python/metatomic_torch/metatomic_torch/dftd3.py diff --git a/python/metatomic_torch/metatomic/torch/documentation.py b/python/metatomic_torch/metatomic_torch/documentation.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/documentation.py rename to python/metatomic_torch/metatomic_torch/documentation.py diff --git a/python/metatomic_torch/metatomic/torch/heat_flux.py b/python/metatomic_torch/metatomic_torch/heat_flux.py similarity index 99% rename from python/metatomic_torch/metatomic/torch/heat_flux.py rename to python/metatomic_torch/metatomic_torch/heat_flux.py index 4de0828e5..167149b06 100644 --- a/python/metatomic_torch/metatomic/torch/heat_flux.py +++ b/python/metatomic_torch/metatomic_torch/heat_flux.py @@ -4,7 +4,7 @@ from metatensor.torch import Labels, TensorBlock, TensorMap from vesin.metatomic import NeighborList -from metatomic.torch import ( +from . import ( AtomisticModel, ModelCapabilities, ModelOutput, diff --git a/python/metatomic_torch/metatomic/torch/model.py b/python/metatomic_torch/metatomic_torch/model.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/model.py rename to python/metatomic_torch/metatomic_torch/model.py diff --git a/python/metatomic_torch/metatomic/torch/o3/__init__.py b/python/metatomic_torch/metatomic_torch/o3/__init__.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/o3/__init__.py rename to python/metatomic_torch/metatomic_torch/o3/__init__.py diff --git a/python/metatomic_torch/metatomic/torch/o3/_tranformations.py b/python/metatomic_torch/metatomic_torch/o3/_tranformations.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/o3/_tranformations.py rename to python/metatomic_torch/metatomic_torch/o3/_tranformations.py diff --git a/python/metatomic_torch/metatomic/torch/o3/_wigner.py b/python/metatomic_torch/metatomic_torch/o3/_wigner.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/o3/_wigner.py rename to python/metatomic_torch/metatomic_torch/o3/_wigner.py diff --git a/python/metatomic_torch/metatomic/torch/serialization.py b/python/metatomic_torch/metatomic_torch/serialization.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/serialization.py rename to python/metatomic_torch/metatomic_torch/serialization.py diff --git a/python/metatomic_torch/metatomic/torch/systems_to_torch.py b/python/metatomic_torch/metatomic_torch/systems_to_torch.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/systems_to_torch.py rename to python/metatomic_torch/metatomic_torch/systems_to_torch.py diff --git a/python/metatomic_torch/metatomic/torch/utils.py b/python/metatomic_torch/metatomic_torch/utils.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/utils.py rename to python/metatomic_torch/metatomic_torch/utils.py diff --git a/python/metatomic_torch/metatomic/torch/version.py b/python/metatomic_torch/metatomic_torch/version.py similarity index 100% rename from python/metatomic_torch/metatomic/torch/version.py rename to python/metatomic_torch/metatomic_torch/version.py diff --git a/python/metatomic_torch/pyproject.toml b/python/metatomic_torch/pyproject.toml index 2d3c34368..a5ff0280a 100644 --- a/python/metatomic_torch/pyproject.toml +++ b/python/metatomic_torch/pyproject.toml @@ -48,10 +48,6 @@ backend-path = ["build-backend"] [tool.setuptools] zip-safe = false -[tool.setuptools.packages.find] -include = ["metatomic*"] -namespaces = true - ### ======================================================================== ### [tool.pytest.ini_options] python_files = ["*.py"] diff --git a/python/metatomic_torch/setup.py b/python/metatomic_torch/setup.py index bfc81072c..f2a9613ed 100644 --- a/python/metatomic_torch/setup.py +++ b/python/metatomic_torch/setup.py @@ -24,6 +24,7 @@ METATOMIC_TORCH_SRC = os.path.realpath( os.path.join(ROOT, "..", "..", "metatomic-torch") ) +METATOMIC_CORE = os.path.realpath(os.path.join(ROOT, "..", "metatomic_core")) METATOMIC_ASE = os.path.realpath(os.path.join(ROOT, "..", "metatomic_ase")) @@ -50,7 +51,7 @@ def run(self): source_dir = ROOT build_dir = os.path.join(ROOT, "build", "cmake-build") - install_dir = os.path.join(os.path.realpath(self.build_lib), "metatomic/torch") + install_dir = os.path.join(os.path.realpath(self.build_lib), "metatomic_torch") os.makedirs(build_dir, exist_ok=True) @@ -326,11 +327,14 @@ def create_version_number(version): # when packaging a sdist for release, we should never use local dependencies METATOMIC_NO_LOCAL_DEPS = os.environ.get("METATOMIC_NO_LOCAL_DEPS", "0") == "1" - if not METATOMIC_NO_LOCAL_DEPS and os.path.exists(METATOMIC_ASE): + if not METATOMIC_NO_LOCAL_DEPS and os.path.exists(METATOMIC_CORE): + assert os.path.exists(METATOMIC_ASE) # we are building from a git checkout or full repo archive + install_requires.append(f"metatomic-core @ file://{METATOMIC_CORE}") install_requires.append(f"metatomic-ase @ file://{METATOMIC_ASE}") else: # we are building from a sdist/installing from a wheel + install_requires.append("metatomic-core >=0.1.0,<0.2.0") install_requires.append("metatomic-ase >=0.1.1,<0.2.0") setup( diff --git a/scripts/clean-python.sh b/scripts/clean-python.sh index ba6a9e9f5..81e69b26a 100755 --- a/scripts/clean-python.sh +++ b/scripts/clean-python.sh @@ -14,9 +14,18 @@ rm -rf docs/build rm -rf docs/src/examples rm -rf docs/src/sg_execution_times.rst +rm -rf python/metatomic_core/dist +rm -rf python/metatomic_core/build + rm -rf python/metatomic_torch/dist rm -rf python/metatomic_torch/build +rm -rf python/metatomic_ase/dist +rm -rf python/metatomic_ase/build + +rm -rf python/metatomic_torchsim/dist +rm -rf python/metatomic_torchsim/build + find . -name "*.egg-info" -exec rm -rf "{}" + find . -name "__pycache__" -exec rm -rf "{}" + find . -name ".coverage" -exec rm -rf "{}" + diff --git a/setup.py b/setup.py index ced9f7146..2124530b5 100644 --- a/setup.py +++ b/setup.py @@ -4,29 +4,39 @@ ROOT = os.path.realpath(os.path.dirname(__file__)) +METATOMIC_CORE = os.path.join(ROOT, "python", "metatomic_core") METATOMIC_TORCH = os.path.join(ROOT, "python", "metatomic_torch") +METATOMIC_ASE = os.path.join(ROOT, "python", "metatomic_ase") METATOMIC_TORCHSIM = os.path.join(ROOT, "python", "metatomic_torchsim") if __name__ == "__main__": extras_require = {} + install_requires = [] # when packaging a sdist for release, we should never use local dependencies METATOMIC_NO_LOCAL_DEPS = os.environ.get("METATOMIC_NO_LOCAL_DEPS", "0") == "1" - if not METATOMIC_NO_LOCAL_DEPS and os.path.exists(METATOMIC_TORCH): + if not METATOMIC_NO_LOCAL_DEPS and os.path.exists(METATOMIC_CORE): + assert os.path.exists(METATOMIC_TORCH) + assert os.path.exists(METATOMIC_ASE) + assert os.path.exists(METATOMIC_TORCHSIM) + # we are building from a git checkout + install_requires.append(f"metatomic-core @ file://{METATOMIC_CORE}") extras_require["torch"] = f"metatomic-torch @ file://{METATOMIC_TORCH}" + extras_require["ase"] = f"metatomic-ase @ file://{METATOMIC_ASE}" + extras_require["torchsim"] = f"metatomic-torchsim @ file://{METATOMIC_TORCHSIM}" else: # we are building from a sdist/installing from a wheel - extras_require["torch"] = "metatomic-torch" + install_requires.append("metatomic-core") - if not METATOMIC_NO_LOCAL_DEPS and os.path.exists(METATOMIC_TORCHSIM): - extras_require["torchsim"] = f"metatomic-torchsim @ file://{METATOMIC_TORCHSIM}" - else: + extras_require["torch"] = "metatomic-torch" + extras_require["ase"] = "metatomic-ase" extras_require["torchsim"] = "metatomic-torchsim" setup( author=", ".join(open(os.path.join(ROOT, "AUTHORS")).read().splitlines()), + install_requires=install_requires, extras_require=extras_require, ) diff --git a/tox.ini b/tox.ini index 919350ccd..0165f865a 100644 --- a/tox.ini +++ b/tox.ini @@ -38,6 +38,7 @@ packaging_deps = testing_deps = pytest pytest-cov + pytest-custom_exit_code metatomic_deps = metatensor-torch >=0.10.0,<0.11 @@ -134,6 +135,7 @@ deps = changedir = python/metatomic_torch commands = + pip install {[testenv]build_single_wheel} ../metatomic_core pip install {[testenv]build_single_wheel} . pip install {[testenv]build_single_wheel} ../metatomic_ase @@ -158,12 +160,23 @@ deps = vesin >=0.6.0,<0.7 ase + torch-sim-atomistic + +setenv = + # ignore the fact that metatensor.torch.operations was loaded from a file + # not in `metatensor/torch/operations` + PY_IGNORE_IMPORTMISMATCH = 1 commands = + pip install {[testenv]build_single_wheel} python/metatomic_core pip install {[testenv]build_single_wheel} python/metatomic_torch pip install {[testenv]build_single_wheel} python/metatomic_ase + pip install {[testenv]build_single_wheel} python/metatomic_torchsim - pytest --doctest-modules --pyargs metatomic + pytest --suppress-no-test-exit-code --doctest-modules --pyargs metatomic + pytest --suppress-no-test-exit-code --doctest-modules --pyargs metatomic_torch + pytest --suppress-no-test-exit-code --doctest-modules --pyargs metatomic_ase + pytest --suppress-no-test-exit-code --doctest-modules --pyargs metatomic_torchsim ################################################################################ @@ -193,8 +206,9 @@ deps = changedir = python/metatomic_ase commands = - pip install {[testenv]build_single_wheel} . + pip install {[testenv]build_single_wheel} ../metatomic_core pip install {[testenv]build_single_wheel} ../metatomic_torch + pip install {[testenv]build_single_wheel} . # use the reference LJ implementation for tests {[testenv]install_lj_tests} @@ -224,8 +238,9 @@ deps = changedir = python/metatomic_torchsim commands = - pip install {[testenv]build_single_wheel} . + pip install {[testenv]build_single_wheel} ../metatomic_core pip install {[testenv]build_single_wheel} ../metatomic_torch + pip install {[testenv]build_single_wheel} . # use the reference LJ implementation for tests {[testenv]install_lj_tests} @@ -294,6 +309,7 @@ deps = chemiscope commands = + pip install {[testenv]build_single_wheel} python/metatomic_core pip install {[testenv]build_single_wheel} python/metatomic_torch pip install {[testenv]build_single_wheel} python/metatomic_ase pip install {[testenv]build_single_wheel} python/metatomic_torchsim From 450547d6a83df5f44d8e881dd0130da44effee5e Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Thu, 21 May 2026 10:26:30 +0200 Subject: [PATCH 02/43] Use pathlib for all path manipulations --- python/metatomic_ase/setup.py | 21 ++++++------ python/metatomic_core/setup.py | 15 +++++---- python/metatomic_torch/setup.py | 53 ++++++++++++++---------------- python/metatomic_torchsim/setup.py | 19 ++++++----- setup.py | 29 ++++++++-------- 5 files changed, 70 insertions(+), 67 deletions(-) diff --git a/python/metatomic_ase/setup.py b/python/metatomic_ase/setup.py index 836dc3b5d..4752f46d1 100644 --- a/python/metatomic_ase/setup.py +++ b/python/metatomic_ase/setup.py @@ -1,4 +1,5 @@ import os +import pathlib import subprocess import sys @@ -8,8 +9,8 @@ from setuptools.command.sdist import sdist -ROOT = os.path.realpath(os.path.dirname(__file__)) -METATOMIC_TORCH = os.path.realpath(os.path.join(ROOT, "..", "metatomic_torch")) +ROOT = pathlib.Path(__file__).parent.resolve() +METATOMIC_TORCH = (ROOT / ".." / "metatomic_torch").resolve() METATOMIC_ASE_VERSION = "0.1.2" @@ -53,15 +54,15 @@ def git_version_info(): """ TAG_PREFIX = "metatomic-ase-v" - if os.path.exists("git_version_info"): + if (ROOT / "git_version_info").exists(): # we are building from a sdist, without git available, but the git # version was recorded in the `git_version_info` file - with open("git_version_info") as fd: + with open(ROOT / "git_version_info") as fd: n_commits = int(fd.readline().strip()) git_hash = fd.readline().strip() else: - script = os.path.join(ROOT, "..", "..", "scripts", "git-version-info.py") - assert os.path.exists(script) + script = (ROOT / ".." / ".." / "scripts" / "git-version-info.py").resolve() + assert script.exists() output = subprocess.run( [sys.executable, script, TAG_PREFIX], @@ -127,19 +128,19 @@ def create_version_number(version): # when packaging a sdist for release, we should never use local dependencies METATOMIC_NO_LOCAL_DEPS = os.environ.get("METATOMIC_NO_LOCAL_DEPS", "0") == "1" - if not METATOMIC_NO_LOCAL_DEPS and os.path.exists(METATOMIC_TORCH): + if not METATOMIC_NO_LOCAL_DEPS and METATOMIC_TORCH.exists(): # we are building from a git checkout or full repo archive - install_requires.append(f"metatomic-torch @ file://{METATOMIC_TORCH}") + install_requires.append(f"metatomic-torch @ {METATOMIC_TORCH.as_uri()}") else: # we are building from a sdist/installing from a wheel install_requires.append("metatomic-torch >=0.1.12,<0.2") - with open(os.path.join(ROOT, "AUTHORS")) as fd: + with open(ROOT / "AUTHORS") as fd: authors = fd.read().splitlines() if authors[0].startswith(".."): # handle "raw" symlink files (on Windows or from full repo tarball) - with open(os.path.join(ROOT, authors[0])) as fd: + with open(ROOT / authors[0]) as fd: authors = fd.read().splitlines() setup( diff --git a/python/metatomic_core/setup.py b/python/metatomic_core/setup.py index f2654eef3..35a2ef16c 100644 --- a/python/metatomic_core/setup.py +++ b/python/metatomic_core/setup.py @@ -1,4 +1,5 @@ import os +import pathlib import subprocess import sys @@ -8,7 +9,7 @@ from setuptools.command.sdist import sdist -ROOT = os.path.realpath(os.path.dirname(__file__)) +ROOT = pathlib.Path(__file__).parent.resolve() METATOMIC_CORE_VERSION = "0.1.0" @@ -59,15 +60,15 @@ def git_version_info(): """ TAG_PREFIX = "metatomic-v" - if os.path.exists("git_version_info"): + if (ROOT / "git_version_info").exists(): # we are building from a sdist, without git available, but the git # version was recorded in the `git_version_info` file - with open("git_version_info") as fd: + with open(ROOT / "git_version_info") as fd: n_commits = int(fd.readline().strip()) git_hash = fd.readline().strip() else: - script = os.path.join(ROOT, "..", "..", "scripts", "git-version-info.py") - assert os.path.exists(script) + script = (ROOT / ".." / ".." / "scripts" / "git-version-info.py").resolve() + assert script.exists() output = subprocess.run( [sys.executable, script, TAG_PREFIX], @@ -123,12 +124,12 @@ def create_version_number(version): if __name__ == "__main__": - with open(os.path.join(ROOT, "AUTHORS")) as fd: + with open(ROOT / "AUTHORS") as fd: authors = fd.read().splitlines() if authors[0].startswith(".."): # handle "raw" symlink files (on Windows or from full repo tarball) - with open(os.path.join(ROOT, authors[0])) as fd: + with open(ROOT / authors[0]) as fd: authors = fd.read().splitlines() install_requires = [ diff --git a/python/metatomic_torch/setup.py b/python/metatomic_torch/setup.py index f2a9613ed..58f8e3125 100644 --- a/python/metatomic_torch/setup.py +++ b/python/metatomic_torch/setup.py @@ -1,5 +1,6 @@ import glob import os +import pathlib import subprocess import sys @@ -12,7 +13,7 @@ from setuptools.command.sdist import sdist -ROOT = os.path.realpath(os.path.dirname(__file__)) +ROOT = pathlib.Path(__file__).parent.resolve() METATOMIC_BUILD_TYPE = os.environ.get("METATOMIC_BUILD_TYPE", "release") if METATOMIC_BUILD_TYPE not in ["debug", "release"]: @@ -21,11 +22,9 @@ "expected 'debug' or 'release'" ) -METATOMIC_TORCH_SRC = os.path.realpath( - os.path.join(ROOT, "..", "..", "metatomic-torch") -) -METATOMIC_CORE = os.path.realpath(os.path.join(ROOT, "..", "metatomic_core")) -METATOMIC_ASE = os.path.realpath(os.path.join(ROOT, "..", "metatomic_ase")) +METATOMIC_TORCH_SRC = (ROOT / ".." / ".." / "metatomic-torch").resolve() +METATOMIC_CORE = (ROOT / ".." / "metatomic_core").resolve() +METATOMIC_ASE = (ROOT / ".." / "metatomic_ase").resolve() class universal_wheel(bdist_wheel): @@ -50,10 +49,10 @@ def run(self): import torch source_dir = ROOT - build_dir = os.path.join(ROOT, "build", "cmake-build") - install_dir = os.path.join(os.path.realpath(self.build_lib), "metatomic_torch") + build_dir = ROOT / "build" / "cmake-build" + install_dir = pathlib.Path(self.build_lib).resolve() / "metatomic_torch" - os.makedirs(build_dir, exist_ok=True) + build_dir.mkdir(parents=True, exist_ok=True) # Tell CMake where to find metatensor, metatensor_torch, and torch cmake_prefix_path = [ @@ -66,9 +65,7 @@ def run(self): # compile the code. This allows having multiple version of this shared library # inside the wheel; and dynamically pick the right one. torch_major, torch_minor, *_ = torch.__version__.split(".") - cmake_install_prefix = os.path.join( - install_dir, f"torch-{torch_major}.{torch_minor}" - ) + cmake_install_prefix = install_dir / f"torch-{torch_major}.{torch_minor}" use_external_lib = os.environ.get( "METATOMIC_TORCH_PYTHON_USE_EXTERNAL_LIB", "OFF" @@ -142,8 +139,8 @@ def run(self): def generate_cxx_tar(): - script = os.path.join(ROOT, "..", "..", "scripts", "package-torch.sh") - assert os.path.exists(script) + script = (ROOT / ".." / ".." / "scripts" / "package-torch.sh").resolve() + assert script.exists() try: output = subprocess.run( @@ -180,15 +177,15 @@ def git_version_info(): """ TAG_PREFIX = "metatomic-torch-v" - if os.path.exists("git_version_info"): + if (ROOT / "git_version_info").exists(): # we are building from a sdist, without git available, but the git # version was recorded in the `git_version_info` file - with open("git_version_info") as fd: + with open(ROOT / "git_version_info") as fd: n_commits = int(fd.readline().strip()) git_hash = fd.readline().strip() else: - script = os.path.join(ROOT, "..", "..", "scripts", "git-version-info.py") - assert os.path.exists(script) + script = (ROOT / ".." / ".." / "scripts" / "git-version-info.py").resolve() + assert script.exists() output = subprocess.run( [sys.executable, script, TAG_PREFIX], @@ -275,10 +272,10 @@ def create_version_number(version): # End of Windows/MKL/PIP hack - if not os.path.exists(METATOMIC_TORCH_SRC): + if not METATOMIC_TORCH_SRC.exists(): # we are building from a sdist, which should include metatomic-torch C++ # sources as a tarball - tarballs = glob.glob(os.path.join(ROOT, "metatomic-torch-cxx-*.tar.gz")) + tarballs = glob.glob(ROOT / "metatomic-torch-cxx-*.tar.gz") if not len(tarballs) == 1: raise RuntimeError( @@ -286,7 +283,7 @@ def create_version_number(version): "metatomic-torch C++ sources" ) - METATOMIC_TORCH_SRC = os.path.realpath(tarballs[0]) + METATOMIC_TORCH_SRC = pathlib.Path(tarballs[0]).resolve() subprocess.run( ["cmake", "-E", "tar", "xf", METATOMIC_TORCH_SRC], cwd=ROOT, @@ -295,15 +292,15 @@ def create_version_number(version): METATOMIC_TORCH_SRC = ".".join(METATOMIC_TORCH_SRC.split(".")[:-2]) - with open(os.path.join(METATOMIC_TORCH_SRC, "VERSION")) as fd: + with open(METATOMIC_TORCH_SRC / "VERSION") as fd: METATOMIC_TORCH_VERSION = fd.read().strip() - with open(os.path.join(ROOT, "AUTHORS")) as fd: + with open(ROOT / "AUTHORS") as fd: authors = fd.read().splitlines() if authors[0].startswith(".."): # handle "raw" symlink files (on Windows or from full repo tarball) - with open(os.path.join(ROOT, authors[0])) as fd: + with open(ROOT / authors[0]) as fd: authors = fd.read().splitlines() try: @@ -327,11 +324,11 @@ def create_version_number(version): # when packaging a sdist for release, we should never use local dependencies METATOMIC_NO_LOCAL_DEPS = os.environ.get("METATOMIC_NO_LOCAL_DEPS", "0") == "1" - if not METATOMIC_NO_LOCAL_DEPS and os.path.exists(METATOMIC_CORE): - assert os.path.exists(METATOMIC_ASE) + if not METATOMIC_NO_LOCAL_DEPS and METATOMIC_CORE.exists(): + assert METATOMIC_ASE.exists() # we are building from a git checkout or full repo archive - install_requires.append(f"metatomic-core @ file://{METATOMIC_CORE}") - install_requires.append(f"metatomic-ase @ file://{METATOMIC_ASE}") + install_requires.append(f"metatomic-core @ {METATOMIC_CORE.as_uri()}") + install_requires.append(f"metatomic-ase @ {METATOMIC_ASE.as_uri()}") else: # we are building from a sdist/installing from a wheel install_requires.append("metatomic-core >=0.1.0,<0.2.0") diff --git a/python/metatomic_torchsim/setup.py b/python/metatomic_torchsim/setup.py index 505982e57..55ea95df8 100644 --- a/python/metatomic_torchsim/setup.py +++ b/python/metatomic_torchsim/setup.py @@ -1,4 +1,5 @@ import os +import pathlib import subprocess import sys @@ -7,8 +8,8 @@ from setuptools.command.sdist import sdist -ROOT = os.path.realpath(os.path.dirname(__file__)) -METATOMIC_TORCH = os.path.realpath(os.path.join(ROOT, "..", "metatomic_torch")) +ROOT = pathlib.Path(__file__).parent.resolve() +METATOMIC_TORCH = (ROOT / ".." / "metatomic_torch").resolve() METATOMIC_TORCHSIM_VERSION = "0.1.4" @@ -38,15 +39,15 @@ def git_version_info(): """ TAG_PREFIX = "metatomic-torchsim-v" - if os.path.exists("git_version_info"): + if (ROOT / "git_version_info").exists(): # we are building from a sdist, without git available, but the git # version was recorded in the `git_version_info` file - with open("git_version_info") as fd: + with open(ROOT / "git_version_info") as fd: n_commits = int(fd.readline().strip()) git_hash = fd.readline().strip() else: - script = os.path.join(ROOT, "..", "..", "scripts", "git-version-info.py") - assert os.path.exists(script) + script = (ROOT / ".." / ".." / "scripts" / "git-version-info.py").resolve() + assert script.exists() output = subprocess.run( [sys.executable, script, TAG_PREFIX], @@ -102,7 +103,7 @@ def create_version_number(version): if __name__ == "__main__": - with open(os.path.join(ROOT, "AUTHORS")) as fd: + with open(ROOT / "AUTHORS") as fd: authors = fd.read().splitlines() install_requires = [ @@ -113,9 +114,9 @@ def create_version_number(version): # when packaging a sdist for release, we should never use local dependencies METATOMIC_NO_LOCAL_DEPS = os.environ.get("METATOMIC_NO_LOCAL_DEPS", "0") == "1" - if not METATOMIC_NO_LOCAL_DEPS and os.path.exists(METATOMIC_TORCH): + if not METATOMIC_NO_LOCAL_DEPS and METATOMIC_TORCH.exists(): # we are building from a git checkout or full repo archive - install_requires.append(f"metatomic-torch @ file://{METATOMIC_TORCH}") + install_requires.append(f"metatomic-torch @ {METATOMIC_TORCH.as_uri()}") else: # we are building from a sdist/installing from a wheel install_requires.append("metatomic-torch >=0.1.12,<0.2") diff --git a/setup.py b/setup.py index 2124530b5..69699d06e 100644 --- a/setup.py +++ b/setup.py @@ -1,13 +1,14 @@ import os +import pathlib from setuptools import setup -ROOT = os.path.realpath(os.path.dirname(__file__)) -METATOMIC_CORE = os.path.join(ROOT, "python", "metatomic_core") -METATOMIC_TORCH = os.path.join(ROOT, "python", "metatomic_torch") -METATOMIC_ASE = os.path.join(ROOT, "python", "metatomic_ase") -METATOMIC_TORCHSIM = os.path.join(ROOT, "python", "metatomic_torchsim") +ROOT = pathlib.Path(__file__).parent.resolve() +METATOMIC_CORE = (ROOT / "python" / "metatomic_core").resolve() +METATOMIC_TORCH = (ROOT / "python" / "metatomic_torch").resolve() +METATOMIC_ASE = (ROOT / "python" / "metatomic_ase").resolve() +METATOMIC_TORCHSIM = (ROOT / "python" / "metatomic_torchsim").resolve() if __name__ == "__main__": @@ -17,16 +18,18 @@ # when packaging a sdist for release, we should never use local dependencies METATOMIC_NO_LOCAL_DEPS = os.environ.get("METATOMIC_NO_LOCAL_DEPS", "0") == "1" - if not METATOMIC_NO_LOCAL_DEPS and os.path.exists(METATOMIC_CORE): - assert os.path.exists(METATOMIC_TORCH) - assert os.path.exists(METATOMIC_ASE) - assert os.path.exists(METATOMIC_TORCHSIM) + if not METATOMIC_NO_LOCAL_DEPS and METATOMIC_CORE.exists(): + assert METATOMIC_TORCH.exists() + assert METATOMIC_ASE.exists() + assert METATOMIC_TORCHSIM.exists() # we are building from a git checkout - install_requires.append(f"metatomic-core @ file://{METATOMIC_CORE}") - extras_require["torch"] = f"metatomic-torch @ file://{METATOMIC_TORCH}" - extras_require["ase"] = f"metatomic-ase @ file://{METATOMIC_ASE}" - extras_require["torchsim"] = f"metatomic-torchsim @ file://{METATOMIC_TORCHSIM}" + install_requires.append(f"metatomic-core @ {METATOMIC_CORE.as_uri()}") + extras_require["torch"] = f"metatomic-torch @ {METATOMIC_TORCH.as_uri()}" + extras_require["ase"] = f"metatomic-ase @ {METATOMIC_ASE.as_uri()}" + extras_require["torchsim"] = ( + f"metatomic-torchsim @ {METATOMIC_TORCHSIM.as_uri()}" + ) else: # we are building from a sdist/installing from a wheel install_requires.append("metatomic-core") From e1ad98afd7f2dab129de4ac3e8c5abc4a150aec2 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Tue, 12 May 2026 15:26:39 +0200 Subject: [PATCH 03/43] Switch main test runner from tox to cargo --- .github/workflows/python-tests.yml | 96 +++++ .github/workflows/torch-tests.yml | 96 ++--- .gitignore | 3 + CONTRIBUTING.rst | 80 +++- Cargo.toml | 7 + docs/src/devdoc/get-started.rst | 6 + docs/src/devdoc/index.rst | 26 ++ docs/src/index.rst | 1 + metatomic-torch/Cargo.toml | 13 + metatomic-torch/lib.rs | 1 + metatomic-torch/tests/CMakeLists.txt | 6 +- metatomic-torch/tests/check-torch-install.rs | 207 ++++++++++ metatomic-torch/tests/run-torch-tests.rs | 47 +++ metatomic-torch/tests/utils/mod.rs | 413 +++++++++++++++++++ python/Cargo.toml | 12 + python/lib.rs | 1 + python/tests/run-python-tests.rs | 23 ++ tox.ini | 71 ---- 18 files changed, 974 insertions(+), 135 deletions(-) create mode 100644 .github/workflows/python-tests.yml create mode 100644 Cargo.toml create mode 100644 docs/src/devdoc/get-started.rst create mode 100644 docs/src/devdoc/index.rst create mode 100644 metatomic-torch/Cargo.toml create mode 100644 metatomic-torch/lib.rs create mode 100644 metatomic-torch/tests/check-torch-install.rs create mode 100644 metatomic-torch/tests/run-torch-tests.rs create mode 100644 metatomic-torch/tests/utils/mod.rs create mode 100644 python/Cargo.toml create mode 100644 python/lib.rs create mode 100644 python/tests/run-python-tests.rs diff --git a/.github/workflows/python-tests.yml b/.github/workflows/python-tests.yml new file mode 100644 index 000000000..282f915d4 --- /dev/null +++ b/.github/workflows/python-tests.yml @@ -0,0 +1,96 @@ +name: Python tests + +on: + push: + branches: [main] + pull_request: + # Check all PR + +concurrency: + group: python-tests-${{ github.ref }} + cancel-in-progress: ${{ github.ref != 'refs/heads/main' }} + +jobs: + python-tests: + runs-on: ${{ matrix.os }} + name: ${{ matrix.os }} / Python ${{ matrix.python-version }} / Torch ${{ matrix.torch-version }} + strategy: + matrix: + include: + - os: ubuntu-24.04 + python-version: "3.10" + torch-version: "2.3" + numpy-version-pin: "<2.0" + # Do not run docs-tests with python 3.10 since torch-sim-atomistic + # is not available for this version of python + tox-envs: lint,torch-tests + - os: ubuntu-24.04 + python-version: "3.10" + torch-version: "2.13" + # See above + tox-envs: lint,torch-tests + - os: ubuntu-24.04 + # TorchScript is no longer supported in Python 3.14 + # so we keep a test with 3.13 to make sure this doesn't break + python-version: "3.13" + torch-version: "2.13" + tox-envs: lint,torch-tests,docs-tests + - os: ubuntu-24.04 + python-version: "3.14" + torch-version: "2.13" + tox-envs: lint,torch-tests,docs-tests + - os: macos-15 + python-version: "3.14" + torch-version: "2.13" + tox-envs: lint,torch-tests,docs-tests + - os: windows-2022 + python-version: "3.14" + torch-version: "2.13" + tox-envs: lint,torch-tests,docs-tests + steps: + - uses: actions/checkout@v7 + with: + fetch-depth: 0 + + - name: setup Python + uses: actions/setup-python@v6 + with: + python-version: ${{ matrix.python-version }} + + - name: setup rust + uses: dtolnay/rust-toolchain@master + with: + toolchain: stable + + - name: Cache Rust dependencies + uses: Leafwing-Studios/cargo-cache@v2.6.1 + with: + sweep-cache: true + + - name: Setup sccache + if: ${{ !env.ACT }} + uses: mozilla-actions/sccache-action@v0.0.10 + with: + version: "v0.10.0" + + - name: setup MSVC command prompt + uses: ilammy/msvc-dev-cmd@v1 + + - name: Setup sccache environnement variables + if: ${{ !env.ACT }} + run: | + echo "SCCACHE_GHA_ENABLED=true" >> $GITHUB_ENV + echo "RUSTC_WRAPPER=sccache" >> $GITHUB_ENV + echo "CMAKE_C_COMPILER_LAUNCHER=sccache" >> $GITHUB_ENV + echo "CMAKE_CXX_COMPILER_LAUNCHER=sccache" >> $GITHUB_ENV + + - name: install tests dependencies + run: | + python -m pip install --upgrade pip + python -m pip install tox coverage + + - name: run tests + run: tox -e ${{ matrix.tox-envs }} + env: + PIP_EXTRA_INDEX_URL: https://download.pytorch.org/whl/cpu + METATOMIC_TESTS_TORCH_VERSION: ${{ matrix.torch-version }} diff --git a/.github/workflows/torch-tests.yml b/.github/workflows/torch-tests.yml index e769aa350..0a3819575 100644 --- a/.github/workflows/torch-tests.yml +++ b/.github/workflows/torch-tests.yml @@ -13,81 +13,85 @@ concurrency: jobs: tests: runs-on: ${{ matrix.os }} - name: ${{ matrix.os }} / Python ${{ matrix.python-version }} / Torch ${{ matrix.torch-version }} + name: ${{ matrix.os }} / Torch ${{ matrix.torch-version }}${{ matrix.extra-name }} + container: ${{ matrix.container }} strategy: matrix: include: - os: ubuntu-24.04 - python-version: "3.10" - torch-version: "2.3" - - os: ubuntu-24.04 - python-version: "3.10" torch-version: "2.13" - - os: ubuntu-24.04 - # Keep a building with Python 3.13 since TorchScript is deprecated - # in Python 3.14 - python-version: "3.13" - torch-version: "2.13" - - os: ubuntu-24.04 python-version: "3.14" - torch-version: "2.13" + cargo-test-flags: --release + do-valgrind: true + + # check the build on a stock Ubuntu 22.04, which uses cmake 3.22 + - os: ubuntu-24.04 + container: ubuntu:22.04 + extra-name: ", cmake 3.22" + torch-version: "2.3" + cargo-test-flags: "" + - os: macos-15 - python-version: "3.14" torch-version: "2.13" - - os: windows-2022 python-version: "3.14" + cargo-test-flags: --release + + - os: windows-2022 torch-version: "2.13" + python-version: "3.14" + cargo-test-flags: --release steps: + - name: install dependencies in container + if: matrix.container == 'ubuntu:22.04' + run: | + apt update + apt install -y software-properties-common + add-apt-repository ppa:deadsnakes/ppa + apt install -y cmake make gcc g++ git curl python3.10 python3.10-venv + + update-alternatives --install /usr/local/bin/python python /usr/bin/python3.10 1 + - uses: actions/checkout@v7 with: fetch-depth: 0 - - name: setup Python - uses: actions/setup-python@v6 + - name: Configure git safe directory + if: matrix.container == 'ubuntu:22.04' + run: git config --global --add safe.directory /__w/metatomic/metatomic + + - name: setup rust + uses: dtolnay/rust-toolchain@master with: - python-version: ${{ matrix.python-version }} + toolchain: stable + + - name: Cache Rust dependencies + uses: Leafwing-Studios/cargo-cache@v2.6.1 + with: + sweep-cache: true + + - name: install valgrind + if: matrix.do-valgrind + run: | + sudo apt-get install -y valgrind - name: Setup sccache + if: ${{ !env.ACT }} uses: mozilla-actions/sccache-action@v0.0.10 with: version: "v0.10.0" - - name: setup MSVC command prompt - uses: ilammy/msvc-dev-cmd@v1 - - name: Setup sccache environnement variables + if: ${{ !env.ACT }} run: | echo "SCCACHE_GHA_ENABLED=true" >> $GITHUB_ENV echo "RUSTC_WRAPPER=sccache" >> $GITHUB_ENV echo "CMAKE_C_COMPILER_LAUNCHER=sccache" >> $GITHUB_ENV echo "CMAKE_CXX_COMPILER_LAUNCHER=sccache" >> $GITHUB_ENV - - name: install tests dependencies - run: | - python -m pip install --upgrade pip - python -m pip install tox coverage - - - name: run Python tests - run: tox -e lint,torch-tests,docs-tests + - name: run TorchScript C++ tests + run: cargo test --package metatomic-torch ${{ matrix.cargo-test-flags }} env: + # Use the CPU only version of torch when building/running the code PIP_EXTRA_INDEX_URL: https://download.pytorch.org/whl/cpu METATOMIC_TESTS_TORCH_VERSION: ${{ matrix.torch-version }} - - - name: run C++ tests - run: tox -e torch-tests-cxx,torch-install-tests-cxx - env: - PIP_EXTRA_INDEX_URL: https://download.pytorch.org/whl/cpu - METATOMIC_TESTS_TORCH_VERSION: ${{ matrix.torch-version }} - - - name: combine Python coverage files - shell: bash - run: | - coverage combine .tox/*/.coverage - coverage xml - - - name: upload to codecov.io - uses: codecov/codecov-action@v7 - with: - fail_ci_if_error: true - files: coverage.xml - token: ${{ secrets.CODECOV_TOKEN }} + CXXFLAGS: ${{ matrix.cxx-flags }} diff --git a/.gitignore b/.gitignore index ab865aa23..265263ff8 100644 --- a/.gitignore +++ b/.gitignore @@ -7,3 +7,6 @@ build/ htmlcov/ .coverage* coverage.xml + +Cargo.lock +target/ diff --git a/CONTRIBUTING.rst b/CONTRIBUTING.rst index 50c8dc986..ff0e53af3 100644 --- a/CONTRIBUTING.rst +++ b/CONTRIBUTING.rst @@ -16,6 +16,10 @@ on metatomic: - **git**: the software we use for version control of the source code. See https://git-scm.com/downloads for installation instructions. +- **the rust compiler**: you will need both ``rustc`` (the compiler) and + ``cargo`` (associated build tool). You can install both using `rustup`_, or + use a version provided by your operating system. We need at least Rust version + 1.74 to build metatomic. - **Python**: you can install ``Python`` and ``pip`` on your operating system. We require a Python version of at least 3.9. - **tox**: a Python test runner, see https://tox.readthedocs.io/en/latest/. You @@ -28,17 +32,21 @@ not have to interact with them directly: - **a C++ compiler** we need a compiler supporting C++11. GCC >= 7, clang >= 5 and MSVC >= 19 should all work, although MSVC is not yet tested continuously. +.. _rustup: https://rustup.rs +.. _`cargo` : https://doc.rust-lang.org/cargo/ +.. _tox: https://tox.readthedocs.io/en/latest + .. admonition:: Optional tools Depending on which part of the code you are working on, you might experience a - lot of time spent re-compiling code, even if you did not directly change them. - For faster builds (and in turn faster tests), you can use compiler cache, like - `sccache`_ or the classic `ccache`_ to reduce the recompilation of unchanged - source code. To do this, you should install and configure one of these tools - (we suggest ``sccache`` since it also supports Rust), and then configure - ``cmake`` and ``cargo`` to use them by setting environnement variables. On - Linux and macOS, you should set the following (look up how to do set - environment variable with your shell): + lot of time spend re-compiling Rust or C++ code, even if you did not change + them. If you'd like faster builds (and in turn faster tests), you can use + `sccache`_ or the classic `ccache`_ to only re-run the compiler if the + corresponding source code changed. To do this, you should install and configure + one of these tools (we suggest sccache since it also supports Rust), and then + configure cmake and cargo to use them by setting environnement variables. On + Linux and macOS, you should set the following (look up how to do set environment + variable with your shell): .. code-block:: bash @@ -88,32 +96,70 @@ changes: Running tests ------------- -The continuous integration pipeline is based on `tox`_. You can run all tests +The continuous integration pipeline is based on `cargo`_. You can run all tests with: .. code-block:: bash cd - tox + cargo test # or cargo test --release to run tests in release mode -These are exactly the same tests that will be performed online in our Github CI +These are exactly the same tests that will be performed online in our GitHub CI workflows. You can also run only a subset of tests with one of these commands: +- ``cargo test`` runs everything + +- ``cargo test --package=metatomic-torch`` to run the C++ TorchScript tests only; + + - ``cargo test --test=run-torch-tests`` will run the unit tests for the + TorchScript C++ extension; + - ``cargo test --test=check-torch-install`` will build the C++ TorchScript + extension, install it and then try to build a basic project depending on + this extension with CMake; + +- ``cargo test --package=metatomic-python`` (or ``tox`` directly, see below) to + run Python tests only; +- ``cargo test --lib`` to run unit tests; +- ``cargo test --doc`` to run documentation tests; +- ``cargo bench --test`` compiles and run the benchmarks once, to quickly ensure + they still work. + +You can add some flags to any of above commands to further refine which tests +should run: + +- ``--release`` to run tests in release mode (default is to run tests in debug mode) +- ``-- `` to only run tests whose name contains filter, for example ``cargo test -- system`` + +Also, you can run individual Python tests using `tox`_ if you wish to run a +subset of Python tests, for example: + .. code-block:: bash tox -e lint # check files for formatting errors tox -e torch-tests # unit tests for metatomic-torch, in Python - tox -e torch-tests-cxx # unit tests for metatomic-torch, in C++ - tox -e torch-install-tests-cxx # testing that the C++ code is a valid CMake package + tox -e ase-tests # unit tests for metatomic-ase, in Python + tox -e torchsim-tests # unit tests for metatomic-torchsim, in Python tox -e docs-tests # doctests (checking inline examples) for all packages - tox -e lint # code style tox -e format # format all files -The last command ``tox -e format`` will use ``tox`` to do actual formatting -instead of just checking it, you can use this to automatically fix some of the -issues detected by ``tox -e lint``. +The last command ``tox -e format`` will use tox to do actual formatting instead +of just checking it, you can use to automatically fix some of the issues +detected by ``tox -e lint``. + +You can run only a subset of the tests with ``tox -e tests -- ``, +replacing ```` with the path to the files you want to test, e.g. +``tox -e tests -- python/tests/operations/abs.py``. + +To get the release build for ``tox`` runs, set the environment variable. + +.. code-block:: bash + + METATOMIC_BUILD_TYPE="release" tox -e torch-tests + +This corresponds to running ``cargo test --package-metatensor-python --release`` +but on the subset of interest. You can run only a subset of the tests with ``tox -e torch-tests -- ``, replacing ```` with the path to the files you diff --git a/Cargo.toml b/Cargo.toml new file mode 100644 index 000000000..5256b9601 --- /dev/null +++ b/Cargo.toml @@ -0,0 +1,7 @@ +[workspace] +resolver = "2" + +members = [ + "metatomic-torch", + "python", +] diff --git a/docs/src/devdoc/get-started.rst b/docs/src/devdoc/get-started.rst new file mode 100644 index 000000000..4c19e4ef6 --- /dev/null +++ b/docs/src/devdoc/get-started.rst @@ -0,0 +1,6 @@ +.. _devdoc-get-started: + +Getting started +=============== + +.. include:: ../../../CONTRIBUTING.rst diff --git a/docs/src/devdoc/index.rst b/docs/src/devdoc/index.rst new file mode 100644 index 000000000..43755fdf6 --- /dev/null +++ b/docs/src/devdoc/index.rst @@ -0,0 +1,26 @@ +.. _devdoc: + +Developer documentation +####################### + +This developer documentation contains the following sections: + +1. :ref:`devdoc-get-started` explains how you can start developing code and + documentation; + +.. toctree:: + :maxdepth: 2 + + get-started + +Development team +---------------- + +Metatensor is developed in the `COSMO laboratory`_ at `EPFL`_, and made +available under the `BSD 3-clauses license `_. We welcome +contributions from anyone, feel free to contact us if you need some help working +with the code! + +.. _COSMO laboratory: https://www.epfl.ch/labs/cosmo/ +.. _EPFL: https://www.epfl.ch/ +.. _LICENSE: https://github.com/metatensor/metatensor/blob/main/LICENSE diff --git a/docs/src/index.rst b/docs/src/index.rst index d94c6ded2..170c25c19 100644 --- a/docs/src/index.rst +++ b/docs/src/index.rst @@ -96,4 +96,5 @@ existing trained models, look into the metatrain_ project instead. quantities/index engines/index examples/index + devdoc/index cite diff --git a/metatomic-torch/Cargo.toml b/metatomic-torch/Cargo.toml new file mode 100644 index 000000000..3809a5a99 --- /dev/null +++ b/metatomic-torch/Cargo.toml @@ -0,0 +1,13 @@ +[package] +name = "metatomic-torch" +version = "0.0.0" +edition = "2021" +publish = false +rust-version = "1.74" + +[lib] +path = "lib.rs" + +[dev-dependencies] +lazy_static = "1" +which = "8" diff --git a/metatomic-torch/lib.rs b/metatomic-torch/lib.rs new file mode 100644 index 000000000..59bc69bb6 --- /dev/null +++ b/metatomic-torch/lib.rs @@ -0,0 +1 @@ +// empty lib.rs, this crate only exists to run TorchScript C++ tests with cargo diff --git a/metatomic-torch/tests/CMakeLists.txt b/metatomic-torch/tests/CMakeLists.txt index 89a3db0f2..8a64a4f33 100644 --- a/metatomic-torch/tests/CMakeLists.txt +++ b/metatomic-torch/tests/CMakeLists.txt @@ -14,9 +14,11 @@ if (VALGRIND) "--leak-check=full" "--show-leak-kinds=definite,indirect,possible" "--track-origins=yes" "--gen-suppressions=all" "--suppressions=${CMAKE_CURRENT_SOURCE_DIR}/valgrind.supp" ) + set(USING_VALGRIND ON) endif() else() set(TEST_COMMAND "") + set(USING_VALGRIND OFF) endif() @@ -46,7 +48,9 @@ foreach(_file_ ${ALL_TESTS}) ) # stop tests if they run for more than 30s - set_tests_properties(torch-${_name_} PROPERTIES TIMEOUT 30) + if (NOT USING_VALGRIND) + set_tests_properties(torch-${_name_} PROPERTIES TIMEOUT 30) + endif() if(WIN32) # We need to set the path to allow access to torch.dll diff --git a/metatomic-torch/tests/check-torch-install.rs b/metatomic-torch/tests/check-torch-install.rs new file mode 100644 index 000000000..8883d916e --- /dev/null +++ b/metatomic-torch/tests/check-torch-install.rs @@ -0,0 +1,207 @@ +use std::path::PathBuf; +use std::sync::Mutex; + +mod utils; + +lazy_static::lazy_static! { + // Make sure only one of the tests below run at the time, since they both + // try to modify the same files + static ref LOCK: Mutex<()> = Mutex::new(()); +} + +/// Check that metatomic-torch can be built and installed with cmake, and that +/// the installed version can be used from another cmake project with +/// `find_package` +#[test] +fn check_torch_install() { + let _guard = match LOCK.lock() { + Ok(guard) => guard, + Err(_) => { + panic!("another test failed, stopping") + } + }; + + const CARGO_TARGET_TMPDIR: &str = env!("CARGO_TARGET_TMPDIR"); + let cargo_manifest_dir = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); + + // ====================================================================== // + // build and install metatensor-torch with cmake + let mut build_dir = PathBuf::from(CARGO_TARGET_TMPDIR); + build_dir.push("torch-install"); + build_dir.push("cmake-find-package"); + std::fs::create_dir_all(&build_dir).expect("failed to create build dir"); + + + let deps_dir = build_dir.join("deps"); + + let torch_dep = deps_dir.join("virtualenv"); + std::fs::create_dir_all(&torch_dep).expect("failed to create virtualenv dir"); + let python = utils::create_python_venv(torch_dep); + let pytorch_cmake_prefix = utils::setup_torch_pip(&python); + let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python); + let metatensor_torch_cmake_prefix = utils::setup_metatensor_torch_pip(&python); + + // configure cmake for metatomic-torch + let metatomic_torch_dep = deps_dir.join("metatomic-torch"); + + let cmake_options = vec![ + format!( + "-DCMAKE_PREFIX_PATH={};{};{}", + pytorch_cmake_prefix.display(), + metatensor_cmake_prefix.display(), + metatensor_torch_cmake_prefix.display() + ), + // The two properties below handle the RPATH for metatomic_torch, + // setting it in such a way that we can always load libmetatensor.so and + // libtorch.so from the location they are found at when compiling + // metatomic-torch. See + // https://gitlab.kitware.com/cmake/community/-/wikis/doc/cmake/RPATH-handling + // for more information on CMake RPATH handling + "-DCMAKE_BUILD_WITH_INSTALL_RPATH=ON".into(), + "-DCMAKE_INSTALL_RPATH_USE_LINK_PATH=ON".into(), + ]; + + let install_prefix = utils::setup_metatomic_torch_cmake( + &cargo_manifest_dir, + &metatomic_torch_dep, + cmake_options, + ); + + // ====================================================================== // + // // try to use the installed metatomic-torch from cmake + let mut source_dir = PathBuf::from(&cargo_manifest_dir); + source_dir.extend(["tests", "cmake-project"]); + + // configure cmake for the test cmake project + let mut cmake_config = utils::cmake_config(&source_dir, &build_dir); + cmake_config.arg(format!( + "-DCMAKE_PREFIX_PATH={};{};{};{}", + metatensor_cmake_prefix.display(), + pytorch_cmake_prefix.display(), + metatensor_torch_cmake_prefix.display(), + install_prefix.display(), + )); + + utils::run_command(cmake_config, "cmake configuration"); + + // build the code, linking to metatomic-torch + let cmake_build = utils::cmake_build(&build_dir); + utils::run_command(cmake_build, "cmake build"); + + // run the executables + let ctest = utils::ctest(&build_dir); + utils::run_command(ctest, "ctest"); +} + +/// Same as above, but using pre-built metatensor-torch from the Python wheel, +/// instead of building it from source with cmake. +#[test] +fn check_python_install() { + let _guard = match LOCK.lock() { + Ok(guard) => guard, + Err(_) => { + panic!("another test failed, stopping") + } + }; + + const CARGO_TARGET_TMPDIR: &str = env!("CARGO_TARGET_TMPDIR"); + + // ====================================================================== // + // build and install metatensor and metatensor-torch with pip + let mut build_dir = PathBuf::from(CARGO_TARGET_TMPDIR); + build_dir.push("torch-install"); + build_dir.push("python-wheels"); + std::fs::create_dir_all(&build_dir).expect("failed to create build dir"); + + let mut venv_dir = build_dir.clone(); + venv_dir.push("virtualenv"); + + let python_exe = utils::create_python_venv(venv_dir); + + let cargo_manifest_dir = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); + let pytorch_cmake_prefix = utils::setup_torch_pip(&python_exe); + let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python_exe); + let metatensor_torch_cmake_prefix = utils::setup_metatensor_torch_pip(&python_exe); + + let python_source_dir = cargo_manifest_dir.parent().unwrap().join("python").join("metatomic_torch"); + let metatomic_torch_cmake_prefix = utils::setup_metatomic_torch_pip(&python_exe, &python_source_dir); + + // ====================================================================== // + // try to use the installed metatensor-torch from cmake + let mut source_dir = PathBuf::from(&cargo_manifest_dir); + source_dir.extend(["tests", "cmake-project"]); + + // configure cmake for the test cmake project + let mut cmake_config = utils::cmake_config(&source_dir, &build_dir); + cmake_config.arg(format!( + "-DCMAKE_PREFIX_PATH={};{};{};{}", + pytorch_cmake_prefix.display(), + metatensor_cmake_prefix.display(), + metatensor_torch_cmake_prefix.display(), + metatomic_torch_cmake_prefix.display(), + )); + + utils::run_command(cmake_config, "cmake configuration"); + + // build the code, linking to metatensor-torch + let cmake_build = utils::cmake_build(&build_dir); + utils::run_command(cmake_build, "cmake build"); + + // run the executables + let ctest = utils::ctest(&build_dir); + utils::run_command(ctest, "ctest"); +} + +/// Same test as above, but building metatomic-torch in the same +/// CMake project (i.e. using add_subdirectory instead of find_package) +#[test] +fn check_cmake_subdirectory() { + let _guard = match LOCK.lock() { + Ok(guard) => guard, + Err(_) => { + panic!("another test failed, stopping") + } + }; + + const CARGO_TARGET_TMPDIR: &str = env!("CARGO_TARGET_TMPDIR"); + + // install torch + let mut build_dir = PathBuf::from(CARGO_TARGET_TMPDIR); + build_dir.push("torch-install"); + build_dir.push("cmake-subdirectory"); + std::fs::create_dir_all(&build_dir).expect("failed to create build dir"); + + let deps_dir = build_dir.join("deps"); + + let torch_dep = deps_dir.join("virtualenv"); + std::fs::create_dir_all(&torch_dep).expect("failed to create virtualenv dir"); + let python = utils::create_python_venv(torch_dep); + let pytorch_cmake_prefix = utils::setup_torch_pip(&python); + let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python); + let metatensor_torch_cmake_prefix = utils::setup_metatensor_torch_pip(&python); + + // ====================================================================== // + let cargo_manifest_dir = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); + let mut source_dir = PathBuf::from(&cargo_manifest_dir); + source_dir.extend(["tests", "cmake-project"]); + + // configure cmake for the test cmake project + let mut cmake_config = utils::cmake_config(&source_dir, &build_dir); + cmake_config.arg(format!( + "-DCMAKE_PREFIX_PATH={};{};{}", + pytorch_cmake_prefix.display(), + metatensor_cmake_prefix.display(), + metatensor_torch_cmake_prefix.display() + )); + cmake_config.arg("-DUSE_CMAKE_SUBDIRECTORY=ON"); + + utils::run_command(cmake_config, "cmake configuration"); + + // build the code, linking to metatomic-torch + let cmake_build = utils::cmake_build(&build_dir); + utils::run_command(cmake_build, "cmake build"); + + // run the executables + let ctest = utils::ctest(&build_dir); + utils::run_command(ctest, "ctest"); +} diff --git a/metatomic-torch/tests/run-torch-tests.rs b/metatomic-torch/tests/run-torch-tests.rs new file mode 100644 index 000000000..93772f0a6 --- /dev/null +++ b/metatomic-torch/tests/run-torch-tests.rs @@ -0,0 +1,47 @@ +use std::path::PathBuf; + +mod utils; + +#[test] +fn run_torch_tests() { + const CARGO_TARGET_TMPDIR: &str = env!("CARGO_TARGET_TMPDIR"); + let cargo_manifest_dir = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); + + // ====================================================================== // + // setup dependencies for the torch tests + + let mut build_dir = PathBuf::from(CARGO_TARGET_TMPDIR); + build_dir.push("torch-tests"); + let deps_dir = build_dir.join("deps"); + + let torch_dep = deps_dir.join("virtualenv"); + std::fs::create_dir_all(&torch_dep).expect("failed to create virtualenv dir"); + let python_exe = utils::create_python_venv(torch_dep); + let pytorch_cmake_prefix = utils::setup_torch_pip(&python_exe); + let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python_exe); + let metatensor_torch_cmake_prefix = utils::setup_metatensor_torch_pip(&python_exe); + + // ====================================================================== // + // build the metatomic-torch C++ tests and run them + let source_dir = cargo_manifest_dir; + + // configure cmake for the tests + let mut cmake_config = utils::cmake_config(&source_dir, &build_dir); + cmake_config.arg("-DMETATOMIC_TORCH_TESTS=ON"); + cmake_config.arg(format!( + "-DCMAKE_PREFIX_PATH={};{};{}", + pytorch_cmake_prefix.display(), + metatensor_cmake_prefix.display(), + metatensor_torch_cmake_prefix.display() + )); + + utils::run_command(cmake_config, "cmake configuration"); + + // build the tests + let cmake_build = utils::cmake_build(&build_dir); + utils::run_command(cmake_build, "cmake build"); + + // run the tests + let ctest = utils::ctest(&build_dir); + utils::run_command(ctest, "ctest"); +} diff --git a/metatomic-torch/tests/utils/mod.rs b/metatomic-torch/tests/utils/mod.rs new file mode 100644 index 000000000..66890380e --- /dev/null +++ b/metatomic-torch/tests/utils/mod.rs @@ -0,0 +1,413 @@ +#![allow(dead_code)] +#![allow(clippy::needless_return)] + +use std::io::{Read, Write}; +use std::path::{Path, PathBuf}; +use std::process::{Command, Stdio}; + +fn build_type() -> &'static str { + // assume that debug assertion means that we are building the code in + // debug mode, even if that could be not true in some cases + if cfg!(debug_assertions) { + "debug" + } else { + "release" + } +} + +fn append_flags(existing: Option, extra: &str) -> String { + match existing { + Some(flags) if !flags.trim().is_empty() => format!("{flags} {extra}"), + _ => extra.into(), + } +} + +pub fn cmake_config(source_dir: &Path, build_dir: &Path) -> Command { + let cmake = which::which("cmake").expect("could not find cmake"); + + let mut cmake_config = Command::new(cmake); + cmake_config.current_dir(build_dir); + cmake_config.arg(source_dir); + cmake_config.arg("--no-warn-unused-cli"); + cmake_config.arg(format!("-DCMAKE_BUILD_TYPE={}", build_type())); + + // the cargo executable currently running + let cargo_exe = std::env::var("CARGO").expect("CARGO env var is not set"); + cmake_config.arg(format!("-DCARGO_EXE={}", cargo_exe)); + + if std::env::var_os("CARGO_LLVM_COV").is_some() { + let coverage_compile_flags = "-fprofile-instr-generate -fcoverage-mapping"; + let coverage_link_flags = "-fprofile-instr-generate"; + + let c_flags = append_flags(std::env::var("CFLAGS").ok(), coverage_compile_flags); + let cxx_flags = append_flags(std::env::var("CXXFLAGS").ok(), coverage_compile_flags); + let exe_linker_flags = + append_flags(std::env::var("LDFLAGS").ok(), coverage_link_flags); + + cmake_config.arg(format!("-DCMAKE_C_FLAGS={c_flags}")); + cmake_config.arg(format!("-DCMAKE_CXX_FLAGS={cxx_flags}")); + cmake_config.arg(format!("-DCMAKE_EXE_LINKER_FLAGS={exe_linker_flags}")); + cmake_config.arg(format!("-DCMAKE_SHARED_LINKER_FLAGS={exe_linker_flags}")); + } + + return cmake_config; +} + +pub fn cmake_build(build_dir: &Path) -> Command { + let cmake = which::which("cmake").expect("could not find cmake"); + + let mut cmake_build = Command::new(cmake); + cmake_build.current_dir(build_dir); + cmake_build.arg("--build"); + cmake_build.arg("."); + cmake_build.arg("--parallel"); + cmake_build.arg("--config"); + cmake_build.arg(build_type()); + + return cmake_build; +} + + +pub fn ctest(build_dir: &Path) -> Command { + let ctest = which::which("ctest").expect("could not find ctest"); + + let mut ctest = Command::new(ctest); + ctest.current_dir(build_dir); + ctest.arg("--output-on-failure"); + ctest.arg("--build-config"); + ctest.arg(build_type()); + + return ctest +} + +/// Find the path to the uv binary, or None if not present +fn find_uv() -> Option { + which::which("uv").ok() +} + +/// Find the path to the `python`or `python3` binary on the user system +fn find_python() -> PathBuf { + if let Ok(python) = which::which("python") { + let output = Command::new(&python) + .arg("-c") + .arg("import sys; print(sys.version_info.major)") + .output() + .expect("could not run python"); + + if output.status.success() { + let stdout = String::from_utf8_lossy(&output.stdout); + + if stdout.trim() == "3" { + // we found Python 3 + return python; + } + } + } + + // try python3 + let python = which::which("python3").expect("failed to run `which python3`"); + let output = Command::new(&python) + .arg("-c") + .arg("import sys; print(sys.version_info.major)") + .output() + .expect("could not run python"); + + if output.status.success() { + let stdout = String::from_utf8_lossy(&output.stdout); + if stdout.trim() == "3" { + // we found Python 3 + return python; + } + } + + panic!("could not find Python 3") +} + +/// Helper: get python executable path inside a venv +fn python_in_venv(venv_dir: &Path) -> PathBuf { + let mut python = venv_dir.to_path_buf(); + if cfg!(target_os = "windows") { + python.extend(["Scripts", "python.exe"]); + } else { + python.extend(["bin", "python"]); + } + python +} + +/// Create a fresh Python virtualenv using uv if available, else fallback to +/// `python -m venv`, and return the path to the python executable in the venv +pub fn create_python_venv(build_dir: PathBuf) -> PathBuf { + if let Some(uv_bin) = find_uv() { + let mut cmd = Command::new(&uv_bin); + cmd.arg("venv"); + cmd.arg("--clear"); + cmd.arg(&build_dir); + + run_command(cmd, "uv venv creation"); + } else { + let mut cmd = Command::new(find_python()); + cmd.arg("-m"); + cmd.arg("venv"); + cmd.arg(&build_dir); + + run_command(cmd, "python to create virtualenv with `venv`"); + + // update pip in case the system uses a very old one + let python = python_in_venv(&build_dir); + let mut cmd = Command::new(&python); + cmd.arg("-m"); + cmd.arg("pip"); + cmd.arg("install"); + cmd.arg("--upgrade"); + cmd.arg("pip"); + + run_command(cmd, "pip upgrade in virtualenv"); + } + + python_in_venv(&build_dir) +} + +#[derive(Default)] +pub struct PipInstallOptions { + pub upgrade: bool, + pub no_deps: bool, + pub no_build_isolation: bool, +} + +/// Install a package with pip (uses uv if present, else falls back to python) +fn pip_install( + python: &Path, + packages: &[&str], + options: PipInstallOptions, +) { + if let Some(uv_bin) = find_uv() { + let mut cmd = Command::new(&uv_bin); + cmd.arg("pip").arg("install").arg("--python").arg(python); + + // follow the same behavior as pip when there are multiple indexes + cmd.arg("--index-strategy"); + cmd.arg("unsafe-best-match"); + + if options.upgrade { + cmd.arg("--upgrade"); + } + if options.no_deps { + cmd.arg("--no-deps"); + } + if options.no_build_isolation { + cmd.arg("--no-build-isolation"); + // uv doesn't support --check-build-dependencies + } + + for package in packages { + cmd.arg(package); + } + + run_command(cmd, "uv pip install"); + } else { + let mut cmd = Command::new(python); + cmd.arg("-m").arg("pip").arg("install"); + if options.upgrade { + cmd.arg("--upgrade"); + } + if options.no_deps { + cmd.arg("--no-deps"); + } + if options.no_build_isolation { + // If pip, add both supported options + cmd.arg("--no-build-isolation"); + cmd.arg("--check-build-dependencies"); + } + + for package in packages { + cmd.arg(package); + } + + run_command(cmd, "pip install"); + } +} + +/// Download PyTorch in a Python virtualenv, and return the +/// CMAKE_PREFIX_PATH for the corresponding libtorch +pub fn setup_torch_pip(python: &Path) -> PathBuf { + let torch_version = std::env::var("METATOMIC_TESTS_TORCH_VERSION").unwrap_or("2.13".into()); + pip_install( + python, + &[&format!("torch=={}.*", torch_version)], + PipInstallOptions { upgrade: true, no_deps: false, no_build_isolation: false } + ); + + let mut cmd = Command::new(python); + cmd.arg("-c"); + cmd.arg("import torch; print(torch.utils.cmake_prefix_path)"); + + let output = run_command(cmd, "python to get torch cmake prefix"); + + let stdout = String::from_utf8_lossy(&output.stdout); + let prefix = PathBuf::from(stdout.trim()); + if !prefix.exists() { + panic!("'torch.utils.cmake_prefix' at '{}' does not exist", prefix.display()); + } + + return prefix; +} + +/// Install metatensor in a Python virtualenv with pip, and return the +/// CMAKE_PREFIX_PATH for the installed libmetatensor. +pub fn setup_metatensor_pip(python: &Path) -> PathBuf { + pip_install(python, &["metatensor-core >=0.2.0,<0.3"], PipInstallOptions::default()); + + let mut cmd = Command::new(python); + cmd.arg("-c"); + cmd.arg("import metatensor; print(metatensor.utils.cmake_prefix_path)"); + + let output = run_command(cmd, "python to get metatensor cmake prefix"); + + let stdout = String::from_utf8_lossy(&output.stdout); + let prefix = PathBuf::from(stdout.trim()); + if !prefix.exists() { + panic!("'metatensor.utils.cmake_prefix' at '{}' does not exist", prefix.display()); + } + + return prefix; +} + +/// Install metatensor-torch in a Python virtualenv with pip, and return the +/// CMAKE_PREFIX_PATH for the installed libmetatensor_torch. +pub fn setup_metatensor_torch_pip(python: &Path) -> PathBuf { + pip_install(python, &["metatensor-torch >=0.9.0,<0.10"], PipInstallOptions::default()); + + let mut cmd = Command::new(python); + cmd.arg("-c"); + cmd.arg("import metatensor.torch; print(metatensor.torch.utils.cmake_prefix_path)"); + + let output = run_command(cmd, "python to get metatensor_torch cmake prefix"); + + let stdout = String::from_utf8_lossy(&output.stdout); + let prefix = PathBuf::from(stdout.trim()); + if !prefix.exists() { + panic!("'metatensor.torch.utils.cmake_prefix' at '{}' does not exist", prefix.display()); + } + + return prefix; +} + +/// Build metatomic-torch located in `source_dir` inside `build_dir`, and return +/// the installation prefix. +pub fn setup_metatomic_torch_cmake(source_dir: &Path, build_dir: &Path, cmake_args: Vec) -> PathBuf { + std::fs::create_dir_all(build_dir).expect("failed to create metatomic build dir"); + + // configure cmake for metatomic-torch + let mut cmake_config = cmake_config(source_dir, build_dir); + + let install_prefix = build_dir.join("usr"); + cmake_config.arg(format!("-DCMAKE_INSTALL_PREFIX={}", install_prefix.display())); + + // Add any additional cmake arguments + for arg in cmake_args { + cmake_config.arg(arg); + } + + run_command(cmake_config, "cmake configuration for metatomic_torch"); + + // build and install metatomic-torch + let mut cmake_build = cmake_build(build_dir); + cmake_build.arg("--target"); + cmake_build.arg("install"); + + run_command(cmake_build, "cmake build for metatomic_torch"); + + install_prefix +} + + +/// Install metatomic-torch in a Python virtualenv with pip, and return the +/// CMAKE_PREFIX_PATH for the installed libmetatomic_torch. +pub fn setup_metatomic_torch_pip(python: &Path, source_dir: &Path) -> PathBuf { + // build dependencies + pip_install(python, &["setuptools>=77", "packaging>=23", "cmake"], PipInstallOptions::default()); + // runtime dependencies which are not just metatensor and metatensor-torch + pip_install(python, &["wigners"], PipInstallOptions::default()); + + pip_install( + python, + &[&source_dir.display().to_string()], + PipInstallOptions { + upgrade: true, + no_deps: false, + no_build_isolation: true + } + ); + + let mut cmd = Command::new(python); + cmd.arg("-c"); + cmd.arg("import metatomic.torch; print(metatomic.torch.utils.cmake_prefix_path)"); + + let output = run_command(cmd, "python to get metatomic_torch cmake prefix"); + + let stdout = String::from_utf8_lossy(&output.stdout); + let prefix = PathBuf::from(stdout.trim()); + if !prefix.exists() { + panic!("'metatomic.torch.utils.cmake_prefix' at '{}' does not exist", prefix.display()); + } + + return prefix; +} + + +pub fn run_command(mut command: Command, context: &str) -> std::process::Output { + write!(std::io::stdout().lock(), "\n\n[Running] {:?}\n\n", command).unwrap(); + + let mut child = command + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn().unwrap_or_else(|_| panic!("failed to spawn {}", context)); + + let mut child_stdout = child.stdout.take().expect("missing stdout"); + let mut child_stderr = child.stderr.take().expect("missing stderr"); + + let out_handle = std::thread::spawn(move || -> std::io::Result> { + let mut buf = [0u8; 8192]; + let mut captured = Vec::new(); + let mut sink = std::io::stdout().lock(); + loop { + let n = child_stdout.read(&mut buf)?; + if n == 0 { + break; + } + sink.write_all(&buf[..n])?; + sink.flush()?; + captured.extend_from_slice(&buf[..n]); + } + Ok(captured) + }); + + let err_handle = std::thread::spawn(move || -> std::io::Result> { + let mut buf = [0u8; 8192]; + let mut captured = Vec::new(); + let mut sink = std::io::stderr().lock(); + loop { + let n = child_stderr.read(&mut buf)?; + if n == 0 { + break; + } + sink.write_all(&buf[..n])?; + sink.flush()?; + captured.extend_from_slice(&buf[..n]); + } + Ok(captured) + }); + + let status = child.wait().unwrap_or_else(|_| panic!("failed to run {}", context)); + let stdout = String::from_utf8_lossy(&out_handle.join().unwrap().unwrap()).into_owned(); + let stderr = String::from_utf8_lossy(&err_handle.join().unwrap().unwrap()).into_owned(); + + if !status.success() { + panic!( + "{} failed, status: {}\nstderr:\n\n{}\nstdout:\n\n{}\n", + context, status, stderr, stdout + ); + } + + return std::process::Output { status, stdout: stdout.into_bytes(), stderr: stderr.into_bytes() }; +} diff --git a/python/Cargo.toml b/python/Cargo.toml new file mode 100644 index 000000000..2ca54178e --- /dev/null +++ b/python/Cargo.toml @@ -0,0 +1,12 @@ +[package] +name = "metatomic-python" +version = "0.0.0" +edition = "2021" +publish = false +rust-version = "1.74" + +[lib] +path = "lib.rs" + +[dev-dependencies] +which = "8" diff --git a/python/lib.rs b/python/lib.rs new file mode 100644 index 000000000..5ef74bad8 --- /dev/null +++ b/python/lib.rs @@ -0,0 +1 @@ +// empty lib.rs, this crate only exists to run Python tests with cargo diff --git a/python/tests/run-python-tests.rs b/python/tests/run-python-tests.rs new file mode 100644 index 000000000..8d52a6f83 --- /dev/null +++ b/python/tests/run-python-tests.rs @@ -0,0 +1,23 @@ +use std::path::PathBuf; +use std::process::Command; + +#[test] +fn run_python_tests() { + let tox = which::which("tox").expect("could not find tox"); + + let mut root = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); + root.pop(); + + let mut tox = Command::new(tox); + tox.arg("--"); + if cfg!(debug_assertions) { + // assume that debug assertions means that we are building the code + // in debug mode, even if optimizations could be enabled + tox.env("METATOMIC_BUILD_TYPE", "debug"); + } else { + tox.env("METATOMIC_BUILD_TYPE", "release"); + } + tox.current_dir(&root); + let status = tox.status().expect("failed to run tox"); + assert!(status.success()); +} diff --git a/tox.ini b/tox.ini index 0165f865a..11184a1b9 100644 --- a/tox.ini +++ b/tox.ini @@ -6,8 +6,6 @@ requires = tox >=4.39 # `tox` in the command-line without anything else envlist = lint - torch-tests-cxx - torch-install-tests-cxx torch-tests docs-tests ase-tests @@ -46,75 +44,6 @@ metatomic_deps = wigners >=0.4.0 -################################################################################ -##### C++ tests setup ##### -################################################################################ - -[testenv:torch-tests-cxx] -description = Run the C++ tests for metatomic-torch -deps = - cmake - {[testenv]metatomic_deps} - torch=={env:METATOMIC_TESTS_TORCH_VERSION:2.13}.* - -commands = - # configure cmake - cmake -B {env_dir}/build metatomic-torch \ - -DCMAKE_BUILD_TYPE=Debug \ - -DCMAKE_EXPORT_COMPILE_COMMANDS=ON \ - -DCMAKE_PREFIX_PATH={env_site_packages_dir}/metatensor/;\ - {env_site_packages_dir}/torch/;\ - {env_site_packages_dir}/metatensor_torch/torch-{env:METATOMIC_TESTS_TORCH_VERSION:2.13}/ \ - -DMETATOMIC_TORCH_TESTS=ON - - # build code with cmake - cmake --build {env_dir}/build --config Debug --parallel - - # run all tests - ctest --test-dir {env_dir}/build --build-config Debug --output-on-failure - -[testenv:torch-install-tests-cxx] -description = Run the C++ tests for metatomic-torch -deps = - cmake - {[testenv]metatomic_deps} - torch=={env:METATOMIC_TESTS_TORCH_VERSION:2.13}.* - -commands = - # configure, build and install metatomic-torch - cmake -B {env_dir}/build-metatomic-torch metatomic-torch \ - -DCMAKE_BUILD_TYPE=Debug \ - -DCMAKE_INSTALL_PREFIX={env_dir}/usr/ \ - -DCMAKE_PREFIX_PATH={env_site_packages_dir}/metatensor/;\ - {env_site_packages_dir}/torch/;\ - {env_site_packages_dir}/metatensor_torch/torch-{env:METATOMIC_TESTS_TORCH_VERSION:2.13}/ \ - -DCMAKE_BUILD_WITH_INSTALL_RPATH=ON \ - -DCMAKE_INSTALL_RPATH_USE_LINK_PATH=ON - cmake --build {env_dir}/build-metatomic-torch --config Debug --parallel --target install - - # try to use the installed metatomic-torch from another CMake project - cmake -B {env_dir}/build-find-package metatomic-torch/tests/cmake-project \ - -DCMAKE_BUILD_TYPE=Debug \ - -DCMAKE_PREFIX_PATH={env_site_packages_dir}/metatensor/;\ - {env_site_packages_dir}/torch/;\ - {env_site_packages_dir}/metatensor_torch/torch-{env:METATOMIC_TESTS_TORCH_VERSION:2.13}/;\ - {env_dir}/usr/ \ - -DUSE_CMAKE_SUBDIRECTORY=OFF - - cmake --build {env_dir}/build-find-package --config Debug --parallel - ctest --test-dir {env_dir}/build-find-package --build-config Debug --output-on-failure - - # Same, but using metatomic-torch as a CMake subdirectory - cmake -B {env_dir}/build-subdirectory metatomic-torch/tests/cmake-project \ - -DCMAKE_BUILD_TYPE=Debug \ - -DCMAKE_PREFIX_PATH={env_site_packages_dir}/metatensor/;\ - {env_site_packages_dir}/torch/;\ - {env_site_packages_dir}/metatensor_torch/torch-{env:METATOMIC_TESTS_TORCH_VERSION:2.13}/ \ - -DUSE_CMAKE_SUBDIRECTORY=ON - - cmake --build {env_dir}/build-subdirectory --config Debug --parallel - ctest --test-dir {env_dir}/build-subdirectory --build-config Debug --output-on-failure - ################################################################################ ##### Python tests setup ##### ################################################################################ From 0fc3f8874c6101e6e6da22fecf1d4fba841b9242 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Wed, 13 May 2026 14:45:08 +0200 Subject: [PATCH 04/43] Scaffold a new metatomic-core package --- .github/workflows/build-wheels.yml | 11 +- .github/workflows/torch-tests.yml | 7 + CONTRIBUTING.rst | 6 + Cargo.toml | 1 + metatomic-core/CHANGELOG.md | 18 + metatomic-core/CMakeLists.txt | 506 ++++++++++++++++++ metatomic-core/Cargo.toml | 26 + metatomic-core/Clippy.toml | 1 + metatomic-core/build.rs | 48 ++ metatomic-core/cmake/dev-versions.cmake | 91 ++++ .../cmake/metatomic-config.in.cmake | 91 ++++ metatomic-core/cmake/tempdir.cmake | 51 ++ metatomic-core/include/metatomic.h | 32 ++ metatomic-core/include/metatomic.hpp | 2 + metatomic-core/include/metatomic/model.hpp | 7 + metatomic-core/include/metatomic/system.hpp | 7 + metatomic-core/src/c_api/mod.rs | 18 + metatomic-core/src/lib.rs | 13 + metatomic-core/tests/CMakeLists.txt | 86 +++ metatomic-core/tests/check-cxx-install.rs | 64 +++ .../tests/cmake-project/CMakeLists.txt | 84 +++ metatomic-core/tests/cmake-project/README.md | 3 + metatomic-core/tests/cmake-project/src/main.c | 8 + .../tests/cmake-project/src/main.cpp | 9 + .../tests/external/.gitattributes | 0 .../tests/external/CMakeLists.txt | 0 .../tests/external/catch/catch.cpp | 0 .../tests/external/catch/catch.hpp | 0 metatomic-core/tests/misc.cpp | 15 + metatomic-core/tests/run-cxx-tests.rs | 40 ++ metatomic-core/tests/utils/mod.rs | 473 ++++++++++++++++ metatomic-torch/tests/CMakeLists.txt | 3 +- metatomic-torch/tests/check-torch-install.rs | 10 +- metatomic-torch/tests/utils/mod.rs | 414 +------------- .../metatomic_torch/build-backend/backend.py | 17 +- 35 files changed, 1742 insertions(+), 420 deletions(-) create mode 100644 metatomic-core/CHANGELOG.md create mode 100644 metatomic-core/CMakeLists.txt create mode 100644 metatomic-core/Cargo.toml create mode 100644 metatomic-core/Clippy.toml create mode 100644 metatomic-core/build.rs create mode 100644 metatomic-core/cmake/dev-versions.cmake create mode 100644 metatomic-core/cmake/metatomic-config.in.cmake create mode 100644 metatomic-core/cmake/tempdir.cmake create mode 100644 metatomic-core/include/metatomic.h create mode 100644 metatomic-core/include/metatomic.hpp create mode 100644 metatomic-core/include/metatomic/model.hpp create mode 100644 metatomic-core/include/metatomic/system.hpp create mode 100644 metatomic-core/src/c_api/mod.rs create mode 100644 metatomic-core/src/lib.rs create mode 100644 metatomic-core/tests/CMakeLists.txt create mode 100644 metatomic-core/tests/check-cxx-install.rs create mode 100644 metatomic-core/tests/cmake-project/CMakeLists.txt create mode 100644 metatomic-core/tests/cmake-project/README.md create mode 100644 metatomic-core/tests/cmake-project/src/main.c create mode 100644 metatomic-core/tests/cmake-project/src/main.cpp rename {metatomic-torch => metatomic-core}/tests/external/.gitattributes (100%) rename {metatomic-torch => metatomic-core}/tests/external/CMakeLists.txt (100%) rename {metatomic-torch => metatomic-core}/tests/external/catch/catch.cpp (100%) rename {metatomic-torch => metatomic-core}/tests/external/catch/catch.hpp (100%) create mode 100644 metatomic-core/tests/misc.cpp create mode 100644 metatomic-core/tests/run-cxx-tests.rs create mode 100644 metatomic-core/tests/utils/mod.rs mode change 100644 => 120000 metatomic-torch/tests/utils/mod.rs diff --git a/.github/workflows/build-wheels.yml b/.github/workflows/build-wheels.yml index db3921f18..16ed75626 100644 --- a/.github/workflows/build-wheels.yml +++ b/.github/workflows/build-wheels.yml @@ -101,8 +101,17 @@ jobs: CIBW_BUILD_VERBOSITY: 1 CIBW_MANYLINUX_X86_64_IMAGE: gcc11-manylinux_2_28_x86_64 CIBW_MANYLINUX_AARCH64_IMAGE: gcc11-manylinux_2_28_aarch64 + # METATOMIC_NO_LOCAL_DEPS is set to 1 when building a tag of + # metatomic-torch, which will force to use the version of + # metatomic-core already released on PyPI. Otherwise, this will use + # the version of metatomic-core from git checkout (in case there are + # unreleased breaking changes). + # + # This means that when releasing a breaking change in metatomic-core, + # the full release should be available on PyPI before pushing the new + # metatomic-torch tag. CIBW_ENVIRONMENT: > - METATOMIC_NO_LOCAL_DEPS=1 + METATOMIC_NO_LOCAL_DEPS=${{ startsWith(github.ref, 'refs/tags/metatomic-torch-v') && '1' || '0' }} METATOMIC_TORCH_BUILD_WITH_TORCH_VERSION=${{ matrix.torch-version }}.* PIP_EXTRA_INDEX_URL=https://download.pytorch.org/whl/cpu MACOSX_DEPLOYMENT_TARGET=11 diff --git a/.github/workflows/torch-tests.yml b/.github/workflows/torch-tests.yml index 0a3819575..93661740a 100644 --- a/.github/workflows/torch-tests.yml +++ b/.github/workflows/torch-tests.yml @@ -55,6 +55,12 @@ jobs: with: fetch-depth: 0 + - name: setup Python + uses: actions/setup-python@v6 + if: matrix.container == null + with: + python-version: ${{ matrix.python-version }} + - name: Configure git safe directory if: matrix.container == 'ubuntu:22.04' run: git config --global --add safe.directory /__w/metatomic/metatomic @@ -95,3 +101,4 @@ jobs: PIP_EXTRA_INDEX_URL: https://download.pytorch.org/whl/cpu METATOMIC_TESTS_TORCH_VERSION: ${{ matrix.torch-version }} CXXFLAGS: ${{ matrix.cxx-flags }} + RUST_BACKTRACE: full diff --git a/CONTRIBUTING.rst b/CONTRIBUTING.rst index ff0e53af3..e13c2c6f1 100644 --- a/CONTRIBUTING.rst +++ b/CONTRIBUTING.rst @@ -109,6 +109,12 @@ workflows. You can also run only a subset of tests with one of these commands: - ``cargo test`` runs everything +- ``cargo test --package=metatomic-core`` to run the C++ tests only; + + - ``cargo test --test=run-cxx-tests`` will run the unit tests C and C++ API; + - ``cargo test --test=check-cxx-install`` will try to build a basic project + depending on metatomic-core with cmake; + - ``cargo test --package=metatomic-torch`` to run the C++ TorchScript tests only; - ``cargo test --test=run-torch-tests`` will run the unit tests for the diff --git a/Cargo.toml b/Cargo.toml index 5256b9601..1a233774c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -2,6 +2,7 @@ resolver = "2" members = [ + "metatomic-core", "metatomic-torch", "python", ] diff --git a/metatomic-core/CHANGELOG.md b/metatomic-core/CHANGELOG.md new file mode 100644 index 000000000..160995db2 --- /dev/null +++ b/metatomic-core/CHANGELOG.md @@ -0,0 +1,18 @@ +# Changelog + +All notable changes to metatomic-core are documented here, following the [keep +a changelog](https://keepachangelog.com/en/1.1.0/) format. This project follows +[Semantic Versioning](https://semver.org/spec/v2.0.0.html). + +## [Unreleased](https://github.com/metatensor/metatensor/) + + diff --git a/metatomic-core/CMakeLists.txt b/metatomic-core/CMakeLists.txt new file mode 100644 index 000000000..717e52e81 --- /dev/null +++ b/metatomic-core/CMakeLists.txt @@ -0,0 +1,506 @@ +# This file defines the CMake build system for the C and C++ API of metatomic. +# +# This API is implemented in Rust, in the metatomic-core crate, but Rust users +# of the API should use the metatomic crate instead, wrapping metatomic-core in +# an easier to use, idiomatic Rust API. +cmake_minimum_required(VERSION 3.22) + +# Is metatomic the main project configured by the user? Or is this being used +# as a submodule/subdirectory? +if (${CMAKE_CURRENT_SOURCE_DIR} STREQUAL ${CMAKE_SOURCE_DIR}) + set(METATOMIC_MAIN_PROJECT ON) +else() + set(METATOMIC_MAIN_PROJECT OFF) +endif() + +if(${METATOMIC_MAIN_PROJECT} AND NOT "${CACHED_LAST_CMAKE_VERSION}" VERSION_EQUAL ${CMAKE_VERSION}) + # We use CACHED_LAST_CMAKE_VERSION to only print the cmake version + # once in the configuration log + set(CACHED_LAST_CMAKE_VERSION ${CMAKE_VERSION} CACHE INTERNAL "Last version of cmake used to configure") + message(STATUS "Running CMake version ${CMAKE_VERSION}") +endif() + +if (POLICY CMP0077) + # use variables to set OPTIONS + cmake_policy(SET CMP0077 NEW) +endif() + +file(STRINGS "Cargo.toml" CARGO_TOML_CONTENT) +foreach(line ${CARGO_TOML_CONTENT}) + string(REGEX REPLACE "^version = \"(.*)\"" "\\1" METATOMIC_VERSION ${line}) + if (NOT ${CMAKE_MATCH_COUNT} EQUAL 0) + # stop on the first regex match, this should be the right version + break() + endif() +endforeach() + +include(cmake/dev-versions.cmake) +create_development_version("${METATOMIC_VERSION}" METATOMIC_FULL_VERSION "metatomic-core-v") +message(STATUS "Building metatomic-core v${METATOMIC_FULL_VERSION}") + +# strip any -dev/-rc suffix on the version since project(VERSION) does not support it +string(REGEX REPLACE "([0-9]*)\\.([0-9]*)\\.([0-9]*).*" "\\1.\\2.\\3" METATOMIC_VERSION ${METATOMIC_FULL_VERSION}) +project(metatomic + VERSION ${METATOMIC_VERSION} + LANGUAGES C CXX # we need to declare a language to access CMAKE_SIZEOF_VOID_P later +) +set(PROJECT_VERSION ${METATOMIC_FULL_VERSION}) + + +# We follow the standard CMake convention of using BUILD_SHARED_LIBS to provide +# either a shared or static library as a default target. But since cargo always +# builds both versions by default, we also install both versions by default. +# `METATOMIC_INSTALL_BOTH_STATIC_SHARED=OFF` allow to disable this behavior, and +# only install the file corresponding to `BUILD_SHARED_LIBS=ON/OFF`. +# +# BUILD_SHARED_LIBS controls the `metatomic` cmake target, making it an alias of +# either `metatomic::static` or `metatomic::shared`. This is mainly relevant +# when using metatomic from another cmake project, either as a submodule or from +# an installed library (see cmake/metatomic-config.cmake) +option(BUILD_SHARED_LIBS "Use a shared library by default instead of a static one" ON) +option(METATOMIC_INSTALL_BOTH_STATIC_SHARED "Install both shared and static libraries" ON) + +set(RUST_BUILD_TARGET "${RUST_BUILD_TARGET}" CACHE STRING "Cross-compilation target for rust code. Leave empty to build for the host") +set(EXTRA_RUST_FLAGS "${EXTRA_RUST_FLAGS}" CACHE STRING "Flags used to build rust code") + +include(GNUInstallDirs) + +if("${CMAKE_BUILD_TYPE}" STREQUAL "" AND "${CMAKE_CONFIGURATION_TYPES}" STREQUAL "") + message(STATUS "Setting build type to 'release' as none was specified.") + set(CMAKE_BUILD_TYPE "release" + CACHE STRING + "Choose the type of build, options are: debug or release" + FORCE) + set_property(CACHE CMAKE_BUILD_TYPE PROPERTY STRINGS release debug) +endif() + +if(${METATOMIC_MAIN_PROJECT} AND NOT "${CACHED_LAST_CMAKE_BUILD_TYPE}" STREQUAL "${CMAKE_BUILD_TYPE}") + set(CACHED_LAST_CMAKE_BUILD_TYPE ${CMAKE_BUILD_TYPE} CACHE INTERNAL "Last build type used in configuration") + message(STATUS "Building metatomic in ${CMAKE_BUILD_TYPE} mode") +endif() + + +function(check_compatible_versions _actual_ _requested_) + if(${_actual_} MATCHES "^([0-9]+)\\.([0-9]+)") + set(_actual_major_ "${CMAKE_MATCH_1}") + set(_actual_minor_ "${CMAKE_MATCH_2}") + else() + message(FATAL_ERROR "Failed to parse actual version: ${_actual_}") + endif() + + if(${_requested_} MATCHES "^([0-9]+)\\.([0-9]+)") + set(_requested_major_ "${CMAKE_MATCH_1}") + set(_requested_minor_ "${CMAKE_MATCH_2}") + else() + message(FATAL_ERROR "Failed to parse requested version: ${_requested_}") + endif() + + if (${_requested_major_} EQUAL 0 AND ${_actual_minor_} EQUAL ${_requested_minor_}) + # major version is 0 and same minor version, everything is fine + elseif (${_actual_major_} EQUAL ${_requested_major_}) + # same major version, everything is fine + else() + # not compatible + message(FATAL_ERROR "Incompatible versions: we need ${_requested_}, but we got ${_actual_}") + endif() +endfunction() + + +set(REQUIRED_METATENSOR_VERSION "0.2.0") +# Either metatensor is built as part of the same CMake project, or we try to +# find the corresponding CMake package +if (TARGET metatensor) + get_target_property(METATENSOR_BUILD_VERSION metatensor BUILD_VERSION) + check_compatible_versions(${METATENSOR_BUILD_VERSION} ${REQUIRED_METATENSOR_VERSION}) +else() + find_package(metatensor ${REQUIRED_METATENSOR_VERSION} CONFIG REQUIRED) +endif() + + +find_program(CARGO_EXE "cargo" DOC "path to cargo (Rust build system)") +if (NOT CARGO_EXE) + message(FATAL_ERROR + "could not find cargo, please make sure the Rust compiler is installed \ + (see https://www.rust-lang.org/tools/install) or set CARGO_EXE" + ) +endif() + +execute_process( + COMMAND ${CARGO_EXE} "--version" "--verbose" + RESULT_VARIABLE CARGO_STATUS + OUTPUT_VARIABLE CARGO_VERSION_RAW +) + +if(CARGO_STATUS AND NOT CARGO_STATUS EQUAL 0) + message(FATAL_ERROR + "could not run cargo, please make sure the Rust compiler is installed \ + (see https://www.rust-lang.org/tools/install)" + ) +endif() + +set(REQUIRED_RUST_VERSION "1.74.0") +if (CARGO_VERSION_RAW MATCHES "cargo ([0-9]+\\.[0-9]+\\.[0-9]+).*") + set(CARGO_VERSION "${CMAKE_MATCH_1}") +else() + message(FATAL_ERROR "failed to determine cargo version, output was: ${CARGO_VERSION_RAW}") +endif() + +if (${CARGO_VERSION} VERSION_LESS ${REQUIRED_RUST_VERSION}) + message(FATAL_ERROR + "your Rust installation is too old (you have version ${CARGO_VERSION}), \ + at least ${REQUIRED_RUST_VERSION} is required" + ) +else() + if(NOT "${CACHED_LAST_CARGO_VERSION}" STREQUAL ${CARGO_VERSION}) + set(CACHED_LAST_CARGO_VERSION ${CARGO_VERSION} CACHE INTERNAL "Last version of cargo used in configuration") + message(STATUS "Using cargo version ${CARGO_VERSION} at ${CARGO_EXE}") + set(CARGO_VERSION_CHANGED TRUE) + endif() +endif() + +# ============================================================================ # +# determine Cargo flags + +set(CARGO_BUILD_ARG "") + +if (EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/Cargo.lock) + set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--locked") +endif() + +# TODO: support multiple configuration generators (MSVC, ...) +string(TOLOWER ${CMAKE_BUILD_TYPE} BUILD_TYPE) +if ("${BUILD_TYPE}" STREQUAL "debug") + set(CARGO_BUILD_TYPE "debug") +elseif("${BUILD_TYPE}" STREQUAL "release") + set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--release") + set(CARGO_BUILD_TYPE "release") +elseif("${BUILD_TYPE}" STREQUAL "relwithdebinfo") + set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--release") + set(CARGO_BUILD_TYPE "release") +else() + message(FATAL_ERROR "unsuported build type: ${CMAKE_BUILD_TYPE}") +endif() + +set(CARGO_TARGET_DIR ${CMAKE_CURRENT_BINARY_DIR}/target) +set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--target-dir=${CARGO_TARGET_DIR}") + +if (CARGO_VERSION_RAW MATCHES "host: ([a-zA-Z0-9_\\-]*)\n") + set(RUST_HOST_TARGET "${CMAKE_MATCH_1}") + if (RUST_HOST_TARGET MATCHES "([a-zA-Z0-9_]*)\\-") + set(RUST_HOST_ARCH "${CMAKE_MATCH_1}") + else() + message(FATAL_ERROR "failed to determine host CPU arch, target was: ${RUST_HOST_TARGET}") + endif() +else() + message(FATAL_ERROR "failed to determine host target, output was: ${CARGO_VERSION_RAW}") +endif() + +if (WIN32) + # on Windows, we need to use the same ABI in both CMake and cargo. If the + # user did not explicitly request a target, we can try to set it ourself, + # otherwise we just check that it matches what we expect. + if (MSVC) + if ("${RUST_BUILD_TARGET}" STREQUAL "") + set(RUST_BUILD_TARGET "${RUST_HOST_ARCH}-pc-windows-msvc") + message(STATUS "Setting rust target to ${RUST_BUILD_TARGET}") + elseif(NOT "${RUST_BUILD_TARGET}" MATCHES "-pc-windows-msvc") + message(FATAL_ERROR "CMake is building with MSVC but the Rust target is ${RUST_BUILD_TARGET}") + endif() + endif() + + if (MINGW) + if ("${RUST_BUILD_TARGET}" STREQUAL "") + set(RUST_BUILD_TARGET "${RUST_HOST_ARCH}-pc-windows-gnu") + message(STATUS "Setting rust target to ${RUST_BUILD_TARGET}") + elseif(NOT "${RUST_BUILD_TARGET}" MATCHES "-pc-windows-gnu") + message(FATAL_ERROR "CMake is building with MinGW but the Rust target is ${RUST_BUILD_TARGET}") + endif() + endif() +endif() + +# Handle cross compilation with RUST_BUILD_TARGET +if ("${RUST_BUILD_TARGET}" STREQUAL "") + if (${METATOMIC_MAIN_PROJECT}) + message(STATUS "Compiling to host (${RUST_HOST_TARGET})") + endif() + + set(CARGO_OUTPUT_DIR "${CARGO_TARGET_DIR}/${CARGO_BUILD_TYPE}") + set(RUST_BUILD_TARGET ${RUST_HOST_TARGET}) +else() + if (${METATOMIC_MAIN_PROJECT}) + message(STATUS "Cross-compiling to ${RUST_BUILD_TARGET}") + endif() + + set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--target=${RUST_BUILD_TARGET}") + set(CARGO_OUTPUT_DIR "${CARGO_TARGET_DIR}/${RUST_BUILD_TARGET}/${CARGO_BUILD_TYPE}") +endif() + +# Get the list of libraries linked by default by cargo/rustc to add when linking +# to metatomic::static +if (CARGO_VERSION_CHANGED) + include(cmake/tempdir.cmake) + get_tempdir(TMPDIR) + + # Adapted from https://github.com/corrosion-rs/corrosion/blob/dc1e4e5/cmake/FindRust.cmake + execute_process( + COMMAND "${CARGO_EXE}" new --lib _cargo_required_libs + WORKING_DIRECTORY "${TMPDIR}" + RESULT_VARIABLE cargo_new_result + ERROR_QUIET + ) + + if (cargo_new_result) + message(FATAL_ERROR "could not create empty project to find default static libs: ${cargo_new_result}") + endif() + + file(APPEND "${TMPDIR}/_cargo_required_libs/Cargo.toml" "[lib]\ncrate-type=[\"staticlib\"]") + + execute_process( + COMMAND ${CARGO_EXE} rustc --color never --target=${RUST_BUILD_TARGET} -- --print=native-static-libs + WORKING_DIRECTORY "${TMPDIR}/_cargo_required_libs" + RESULT_VARIABLE cargo_static_libs_result + ERROR_VARIABLE cargo_static_libs_stderr + ) + + # clean up the files + file(REMOVE_RECURSE "${TMPDIR}") + + if (cargo_static_libs_result) + message(FATAL_ERROR + "could not extract default static libs (status ${cargo_static_libs_result}), stderr:\n${cargo_static_libs_stderr}" + ) + endif() + + # The pattern starts with `native-static-libs:` and goes to the end of the line. + if (cargo_static_libs_stderr MATCHES "native-static-libs: ([^\r\n]+)\r?\n") + string(REPLACE " " ";" "libs_list" "${CMAKE_MATCH_1}") + set(stripped_lib_list "") + foreach(lib ${libs_list}) + # Strip leading `-l` (unix) and potential .lib suffix (windows) + string(REGEX REPLACE "^-l" "" "stripped_lib" "${lib}") + string(REGEX REPLACE "\.lib$" "" "stripped_lib" "${stripped_lib}") + list(APPEND stripped_lib_list "${stripped_lib}") + endforeach() + + # Special case `msvcrt` to link with the debug version in Debug mode. + list(TRANSFORM stripped_lib_list REPLACE "^msvcrt$" "\$<\$:msvcrtd>") + # Don't try to pass a linker *flag* where CMake expects libraries + list(REMOVE_ITEM stripped_lib_list "/defaultlib:msvcrt") + + if (APPLE) + # Prevent warnings about duplicated `System` in linked libraries + # from Apple's `ld` + list(REMOVE_ITEM stripped_lib_list "System") + endif() + + list(REMOVE_DUPLICATES stripped_lib_list) + set(CARGO_DEFAULT_LIBRARIES "${stripped_lib_list}" CACHE INTERNAL "list of implicitly linked libraries") + + if (${METATOMIC_MAIN_PROJECT}) + message(STATUS "Cargo default link libraries are: ${CARGO_DEFAULT_LIBRARIES}") + endif() + else() + message(FATAL_ERROR "could not find default static libs: `native-static-libs` not found in: `${cargo_static_libs_stderr}`") + endif() +endif() + +file(GLOB_RECURSE ALL_RUST_SOURCES + ${PROJECT_SOURCE_DIR}/Cargo.toml + ${PROJECT_SOURCE_DIR}/src/**.rs +) + +add_library(metatomic::shared SHARED IMPORTED GLOBAL) +set(METATOMIC_SHARED_LOCATION "${CARGO_OUTPUT_DIR}/${CMAKE_SHARED_LIBRARY_PREFIX}metatomic${CMAKE_SHARED_LIBRARY_SUFFIX}") +set(METATOMIC_IMPLIB_LOCATION "${METATOMIC_SHARED_LOCATION}.lib") + +if (MINGW) + # `rustc` does not follow the usual naming scheme for DLL with mingw (it + # would typically be 'libmetatomic.dll') + set(METATOMIC_SHARED_LOCATION "${CARGO_OUTPUT_DIR}/metatomic.dll") + set(METATOMIC_IMPLIB_LOCATION "${CARGO_OUTPUT_DIR}/libmetatomic.dll.a") +endif() + +add_library(metatomic::static STATIC IMPORTED GLOBAL) +set(METATOMIC_STATIC_LOCATION "${CARGO_OUTPUT_DIR}/${CMAKE_STATIC_LIBRARY_PREFIX}metatomic${CMAKE_STATIC_LIBRARY_SUFFIX}") + +get_filename_component(METATOMIC_SHARED_LIB_NAME ${METATOMIC_SHARED_LOCATION} NAME) +get_filename_component(METATOMIC_IMPLIB_NAME ${METATOMIC_IMPLIB_LOCATION} NAME) +get_filename_component(METATOMIC_STATIC_LIB_NAME ${METATOMIC_STATIC_LOCATION} NAME) + +# We need to add some metadata to the shared library to enable linking to it +# without using an absolute path. +if (UNIX) + if (APPLE) + # set the install name to `@rpath/libmetatomic.dylib` + set(CARGO_RUSTC_ARGS "-Clink-arg=-Wl,-install_name,@rpath/${METATOMIC_SHARED_LIB_NAME}") + set_target_properties(metatomic::shared PROPERTIES + IMPORTED_SONAME @rpath/${METATOMIC_SHARED_LIB_NAME} + ) + else() # LINUX + # set the SONAME to libmetatomic.so + set(CARGO_RUSTC_ARGS "-Clink-arg=-Wl,-soname,${METATOMIC_SHARED_LIB_NAME}") + set_target_properties(metatomic::shared PROPERTIES + IMPORTED_SONAME ${METATOMIC_SHARED_LIB_NAME} + ) + endif() +else() + set(CARGO_RUSTC_ARGS "") +endif() + +if (NOT "${EXTRA_RUST_FLAGS}" STREQUAL "") + set(CARGO_RUSTC_ARGS "${CARGO_RUSTC_ARGS};${EXTRA_RUST_FLAGS}") +endif() + +# Set environment variables for cargo build +set(CARGO_ENV "METATOMIC_FULL_VERSION=${METATOMIC_FULL_VERSION}") +if (NOT "${CMAKE_OSX_DEPLOYMENT_TARGET}" STREQUAL "") + list(APPEND CARGO_ENV "MACOSX_DEPLOYMENT_TARGET=${CMAKE_OSX_DEPLOYMENT_TARGET}") +endif() +if (NOT "$ENV{RUSTC_WRAPPER}" STREQUAL "") + list(APPEND CARGO_ENV "RUSTC_WRAPPER=$ENV{RUSTC_WRAPPER}") +endif() + +if (METATOMIC_INSTALL_BOTH_STATIC_SHARED) + set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--crate-type=cdylib;--crate-type=staticlib") + set(CARGO_OUTPUTS ${METATOMIC_SHARED_LOCATION} ${METATOMIC_STATIC_LOCATION}) + if (WIN32) + list(APPEND CARGO_OUTPUTS ${METATOMIC_IMPLIB_LOCATION}) + set(FILE_CREATED_MESSAGE "${METATOMIC_SHARED_LIB_NAME}, ${METATOMIC_STATIC_LIB_NAME}, and ${METATOMIC_IMPLIB_NAME}") + else() + set(FILE_CREATED_MESSAGE "${METATOMIC_SHARED_LIB_NAME} and ${METATOMIC_STATIC_LIB_NAME}") + endif() +else() + if (BUILD_SHARED_LIBS) + set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--crate-type=cdylib") + set(CARGO_OUTPUTS ${METATOMIC_SHARED_LOCATION}) + if (WIN32) + list(APPEND CARGO_OUTPUTS ${METATOMIC_IMPLIB_LOCATION}) + set(FILE_CREATED_MESSAGE "${METATOMIC_SHARED_LIB_NAME} and ${METATOMIC_IMPLIB_NAME}") + else() + set(FILE_CREATED_MESSAGE "${METATOMIC_SHARED_LIB_NAME}") + endif() + else() + set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--crate-type=staticlib") + set(CARGO_OUTPUTS ${METATOMIC_STATIC_LOCATION}) + set(FILE_CREATED_MESSAGE "${METATOMIC_STATIC_LIB_NAME}") + endif() +endif() + +add_custom_command( + OUTPUT ${CARGO_OUTPUTS} + COMMAND ${CMAKE_COMMAND} -E env ${CARGO_ENV} + cargo rustc ${CARGO_BUILD_ARG} -- ${CARGO_RUSTC_ARGS} + WORKING_DIRECTORY ${PROJECT_SOURCE_DIR} + DEPENDS ${ALL_RUST_SOURCES} + COMMENT "Building ${FILE_CREATED_MESSAGE} with cargo" + VERBATIM +) +add_custom_target(cargo-build-metatomic ALL DEPENDS ${CARGO_OUTPUTS}) + +# Auto-generate a header containing the version number as #define +set(_path_ "${CMAKE_CURRENT_BINARY_DIR}/generated-version.h") +file(WRITE ${_path_} "#pragma once\n\n") +file(APPEND ${_path_} "/** Full version of metatomic as a string */\n") +file(APPEND ${_path_} "#define METATOMIC_VERSION \"${METATOMIC_FULL_VERSION}\"\n\n") +file(APPEND ${_path_} "/** Major version number of metatomic as an integer */\n") +file(APPEND ${_path_} "#define METATOMIC_VERSION_MAJOR ${PROJECT_VERSION_MAJOR}\n\n") +file(APPEND ${_path_} "/** Minor version number of metatomic as an integer */\n") +file(APPEND ${_path_} "#define METATOMIC_VERSION_MINOR ${PROJECT_VERSION_MINOR}\n\n") +file(APPEND ${_path_} "/** Patch version number of metatomic as an integer */\n") +file(APPEND ${_path_} "#define METATOMIC_VERSION_PATCH ${PROJECT_VERSION_PATCH}\n") + +file(MAKE_DIRECTORY ${PROJECT_BINARY_DIR}/include/metatomic) +set(_destination_ "${CMAKE_CURRENT_BINARY_DIR}/include/metatomic/version.h") +file(COPY_FILE ${_path_} ${_destination_} ONLY_IF_DIFFERENT) + +add_dependencies(metatomic::shared cargo-build-metatomic) +add_dependencies(metatomic::static cargo-build-metatomic) + +set_target_properties(metatomic::shared PROPERTIES + IMPORTED_LOCATION ${METATOMIC_SHARED_LOCATION} + INTERFACE_INCLUDE_DIRECTORIES "${CMAKE_CURRENT_SOURCE_DIR}/include;${CMAKE_CURRENT_BINARY_DIR}/include" + BUILD_VERSION "${METATOMIC_FULL_VERSION}" +) +target_compile_features(metatomic::shared INTERFACE cxx_std_17) + +if (WIN32) + set_target_properties(metatomic::shared PROPERTIES + IMPORTED_IMPLIB ${METATOMIC_IMPLIB_LOCATION} + ) +endif() + +set_target_properties(metatomic::static PROPERTIES + IMPORTED_LOCATION ${METATOMIC_STATIC_LOCATION} + INTERFACE_INCLUDE_DIRECTORIES "${CMAKE_CURRENT_SOURCE_DIR}/include;${CMAKE_CURRENT_BINARY_DIR}/include" + INTERFACE_LINK_LIBRARIES "${CARGO_DEFAULT_LIBRARIES}" + BUILD_VERSION "${METATOMIC_FULL_VERSION}" +) +target_compile_features(metatomic::static INTERFACE cxx_std_17) + +if (TARGET metatensor::static) + target_link_libraries(metatomic::static INTERFACE metatensor::static) +else() + target_link_libraries(metatomic::static INTERFACE metatensor) +endif() + +if (TARGET metatensor::shared) + target_link_libraries(metatomic::shared INTERFACE metatensor::shared) +else() + target_link_libraries(metatomic::shared INTERFACE metatensor) +endif() + + +if (BUILD_SHARED_LIBS) + add_library(metatomic ALIAS metatomic::shared) +else() + add_library(metatomic ALIAS metatomic::static) +endif() + +#------------------------------------------------------------------------------# +# Installation configuration +#------------------------------------------------------------------------------# +include(CMakePackageConfigHelpers) +configure_package_config_file( + ${PROJECT_SOURCE_DIR}/cmake/metatomic-config.in.cmake + ${PROJECT_BINARY_DIR}/metatomic-config.cmake + INSTALL_DESTINATION ${CMAKE_INSTALL_LIBDIR}/cmake/metatomic +) +write_basic_package_version_file( + metatomic-config-version.cmake + VERSION ${METATOMIC_FULL_VERSION} + COMPATIBILITY SameMinorVersion +) + +install(FILES "include/metatomic.h" DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}) +install(FILES "include/metatomic.hpp" DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}) +install(DIRECTORY "include/metatomic" DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}) +install(FILES "${CMAKE_CURRENT_BINARY_DIR}/include/metatomic/version.h" DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}/metatomic) + +if (METATOMIC_INSTALL_BOTH_STATIC_SHARED OR BUILD_SHARED_LIBS) + if (WIN32) + # DLL files should go in /bin + install( + FILES ${METATOMIC_SHARED_LOCATION} + DESTINATION ${CMAKE_INSTALL_BINDIR} + PERMISSIONS OWNER_EXECUTE OWNER_WRITE OWNER_READ GROUP_EXECUTE GROUP_READ WORLD_READ WORLD_EXECUTE + ) + # .lib files should go in /lib + install(FILES ${METATOMIC_IMPLIB_LOCATION} DESTINATION ${CMAKE_INSTALL_LIBDIR}) + else() + install( + FILES ${METATOMIC_SHARED_LOCATION} + DESTINATION ${CMAKE_INSTALL_LIBDIR} + PERMISSIONS OWNER_EXECUTE OWNER_WRITE OWNER_READ GROUP_EXECUTE GROUP_READ WORLD_READ WORLD_EXECUTE + ) + endif() +endif() + +if (METATOMIC_INSTALL_BOTH_STATIC_SHARED OR NOT BUILD_SHARED_LIBS) + install(FILES ${METATOMIC_STATIC_LOCATION} DESTINATION ${CMAKE_INSTALL_LIBDIR}) +endif() + +install(FILES + ${PROJECT_BINARY_DIR}/metatomic-config-version.cmake + ${PROJECT_BINARY_DIR}/metatomic-config.cmake + DESTINATION ${CMAKE_INSTALL_LIBDIR}/cmake/metatomic +) diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml new file mode 100644 index 000000000..2a32c1c09 --- /dev/null +++ b/metatomic-core/Cargo.toml @@ -0,0 +1,26 @@ +[package] +name = "metatomic-core" +version = "0.1.0" +edition = "2021" +publish = false +rust-version = "1.74" +exclude = [ + "tests" +] + +[lib] +crate-type = ["cdylib", "staticlib"] +name = "metatomic" +bench = false + +[dependencies] +once_cell = "1" + + +[build-dependencies] +cbindgen = { version = "0.29", default-features = false } + + +[dev-dependencies] +lazy_static = "1" +which = "8" diff --git a/metatomic-core/Clippy.toml b/metatomic-core/Clippy.toml new file mode 100644 index 000000000..49c5aa7b9 --- /dev/null +++ b/metatomic-core/Clippy.toml @@ -0,0 +1 @@ +doc-valid-idents = ["DLPack", "ROCm", ".."] diff --git a/metatomic-core/build.rs b/metatomic-core/build.rs new file mode 100644 index 000000000..edec71e60 --- /dev/null +++ b/metatomic-core/build.rs @@ -0,0 +1,48 @@ +#![allow(clippy::field_reassign_with_default)] + +use std::path::PathBuf; + +fn main() { + let crate_dir = std::env::var("CARGO_MANIFEST_DIR").unwrap(); + + let generated_comment = "\ +/* ============ Automatically generated file, DO NOT EDIT. ============== * + * * + * This file is automatically generated from the metatomic sources, * + * using cbindgen. If you want to change this file (including documentation), * + * make the corresponding changes in the rust sources and regenerate it. * + * ============================================================================= */"; + + let mut config: cbindgen::Config = Default::default(); + config.language = cbindgen::Language::C; + config.cpp_compat = true; + config.include_guard = Some("METATOMIC_H".into()); + config.include_version = false; + config.documentation = true; + config.documentation_style = cbindgen::DocumentationStyle::Doxy; + config.line_endings = cbindgen::LineEndingStyle::LF; + config.autogen_warning = Some(generated_comment.into()); + config.includes.push("metatomic/version.h".into()); + + let result = cbindgen::Builder::new() + .with_crate(crate_dir) + .with_config(config) + .generate() + .map(|data| { + let mut path = PathBuf::from("include"); + path.push("metatomic.h"); + data.write_to_file(&path); + }); + + // if not ok, rerun the build script unconditionally + if result.is_ok() { + println!("cargo:rerun-if-changed=src"); + println!("cargo:rerun-if-changed=build.rs"); + } + + if std::env::var("METATOMIC_FULL_VERSION").is_err() { + let version = std::env::var("CARGO_PKG_VERSION").expect("missing CARGO_PKG_VERSION"); + println!("cargo:rustc-env=METATOMIC_FULL_VERSION={}+rust", version); + } + println!("cargo:rerun-if-env-changed=METATOMIC_FULL_VERSION"); +} diff --git a/metatomic-core/cmake/dev-versions.cmake b/metatomic-core/cmake/dev-versions.cmake new file mode 100644 index 000000000..543296493 --- /dev/null +++ b/metatomic-core/cmake/dev-versions.cmake @@ -0,0 +1,91 @@ +# Parse a `_version_` number, and store its components in `_major_` `_minor_` +# `_patch_` and `_rc_` +function(parse_version _version_ _major_ _minor_ _patch_ _rc_) + string(REGEX MATCH "([0-9]+)\\.([0-9]+)\\.([0-9]+)(-rc)?([0-9]+)?" _ "${_version_}") + + if(${CMAKE_MATCH_COUNT} EQUAL 3) + set(${_rc_} "" PARENT_SCOPE) + elseif(${CMAKE_MATCH_COUNT} EQUAL 5) + set(${_rc_} ${CMAKE_MATCH_5} PARENT_SCOPE) + else() + message(FATAL_ERROR "invalid version string ${_version_}") + endif() + + set(${_major_} ${CMAKE_MATCH_1} PARENT_SCOPE) + set(${_minor_} ${CMAKE_MATCH_2} PARENT_SCOPE) + set(${_patch_} ${CMAKE_MATCH_3} PARENT_SCOPE) +endfunction() + +# Get the time of the last modification since the last tag/release, and a hash +# of the latest commit/full state of a dirty repository +function(git_version_info _tag_prefix_ _output_n_commits_ _output_git_hash_) + set(_script_ "${CMAKE_CURRENT_FUNCTION_LIST_DIR}/../../scripts/git-version-info.py") + + if (EXISTS "${CMAKE_CURRENT_FUNCTION_LIST_DIR}/git_version_info") + # When building from a tarball, the script is executed and the result + # put in this file + file(STRINGS "${CMAKE_CURRENT_FUNCTION_LIST_DIR}/git_version_info" _file_content_) + list(GET _file_content_ 0 _n_commits_) + list(GET _file_content_ 1 _git_hash_) + + elseif (EXISTS "${_script_}") + # When building from a checkout, we'll need to run the script + find_package(Python COMPONENTS Interpreter REQUIRED) + execute_process( + COMMAND "${Python_EXECUTABLE}" "${_script_}" "${_tag_prefix_}" + RESULT_VARIABLE _status_ + OUTPUT_VARIABLE _stdout_ + ERROR_VARIABLE _stderr_ + WORKING_DIRECTORY ${CMAKE_CURRENT_FUNCTION_LIST_DIR} + ) + + if (NOT ${_status_} EQUAL 0) + message(WARNING + "git-version-info.py failed, version number might be wrong:\nstdout: ${_stdout_}\nstderr: ${_stderr_}") + set(${_output_} 0 PARENT_SCOPE) + return() + endif() + + if (NOT "${_stderr_}" STREQUAL "") + message(WARNING "git-version-info.py gave some errors, version number might be wrong:\nstdout: ${_stdout_}\nstderr: ${_stderr_}") + endif() + + string(REPLACE "\n" ";" _lines_ ${_stdout_}) + list(GET _lines_ 0 _n_commits_) + list(GET _lines_ 1 _git_hash_) + else() + message(FATAL_ERROR "could not update git version information") + endif() + + string(STRIP ${_n_commits_} _n_commits_) + set(${_output_n_commits_} ${_n_commits_} PARENT_SCOPE) + + string(STRIP ${_git_hash_} _git_hash_) + set(${_output_git_hash_} ${_git_hash_} PARENT_SCOPE) +endfunction() + + +# Take the version declared in the package, and increase the right number if we +# are actually installing a developement version from after the latest git tag +function(create_development_version _version_ _output_ _tag_prefix_) + git_version_info("${_tag_prefix_}" _n_commits_ _git_hash_) + + parse_version(${_version_} _major_ _minor_ _patch_ _rc_) + if(${_n_commits_} STREQUAL "0") + # we are building a release, leave the version number as-is + if("${_rc_}" STREQUAL "") + set(${_output_} "${_major_}.${_minor_}.${_patch_}" PARENT_SCOPE) + else() + set(${_output_} "${_major_}.${_minor_}.${_patch_}-rc${_rc_}" PARENT_SCOPE) + endif() + else() + # we are building a development version, increase the right part of the version + if("${_rc_}" STREQUAL "") + math(EXPR _minor_ "${_minor_} + 1") + set(${_output_} "${_major_}.${_minor_}.0-dev${_n_commits_}+${_git_hash_}" PARENT_SCOPE) + else() + math(EXPR _rc_ "${_rc_} + 1") + set(${_output_} "${_major_}.${_minor_}.${_patch_}-rc${_rc_}-dev${_n_commits_}+${_git_hash_}" PARENT_SCOPE) + endif() + endif() +endfunction() diff --git a/metatomic-core/cmake/metatomic-config.in.cmake b/metatomic-core/cmake/metatomic-config.in.cmake new file mode 100644 index 000000000..310f54364 --- /dev/null +++ b/metatomic-core/cmake/metatomic-config.in.cmake @@ -0,0 +1,91 @@ +@PACKAGE_INIT@ + +cmake_minimum_required(VERSION 3.22) + +include(CMakeFindDependencyMacro) +include(FindPackageHandleStandardArgs) + +if(metatomic_FOUND) + return() +endif() + +enable_language(CXX) + +# use the same version for metatensor-core as the main CMakeLists.txt +set(REQUIRED_METATENSOR_VERSION @REQUIRED_METATENSOR_VERSION@) +find_package(metatensor ${REQUIRED_METATENSOR_VERSION} CONFIG REQUIRED) + +get_filename_component(METATOMIC_PREFIX_DIR "${CMAKE_CURRENT_LIST_DIR}/@PACKAGE_RELATIVE_PATH@" ABSOLUTE) + +if (WIN32) + set(METATOMIC_SHARED_LOCATION ${METATOMIC_PREFIX_DIR}/@CMAKE_INSTALL_BINDIR@/@METATOMIC_SHARED_LIB_NAME@) + set(METATOMIC_IMPLIB_LOCATION ${METATOMIC_PREFIX_DIR}/@CMAKE_INSTALL_LIBDIR@/@METATOMIC_IMPLIB_NAME@) +else() + set(METATOMIC_SHARED_LOCATION ${METATOMIC_PREFIX_DIR}/@CMAKE_INSTALL_LIBDIR@/@METATOMIC_SHARED_LIB_NAME@) +endif() + +set(METATOMIC_STATIC_LOCATION ${METATOMIC_PREFIX_DIR}/@CMAKE_INSTALL_LIBDIR@/@METATOMIC_STATIC_LIB_NAME@) +set(METATOMIC_INCLUDE ${METATOMIC_PREFIX_DIR}/@CMAKE_INSTALL_INCLUDEDIR@/) + +if (NOT EXISTS ${METATOMIC_INCLUDE}/metatomic.h OR NOT EXISTS ${METATOMIC_INCLUDE}/metatomic.hpp) + message(FATAL_ERROR "could not find metatomic headers in '${METATOMIC_INCLUDE}', please re-install metatomic") +endif() + + +# Shared library target +if (@METATOMIC_INSTALL_BOTH_STATIC_SHARED@ OR @BUILD_SHARED_LIBS@) + if (NOT EXISTS ${METATOMIC_SHARED_LOCATION}) + message(FATAL_ERROR "could not find metatomic library at '${METATOMIC_SHARED_LOCATION}', please re-install metatomic") + endif() + + add_library(metatomic::shared SHARED IMPORTED) + set_target_properties(metatomic::shared PROPERTIES + IMPORTED_LOCATION ${METATOMIC_SHARED_LOCATION} + INTERFACE_INCLUDE_DIRECTORIES ${METATOMIC_INCLUDE} + BUILD_VERSION "@METATOMIC_FULL_VERSION@" + ) + + target_compile_features(metatomic::shared INTERFACE cxx_std_17) + + if (WIN32) + if (NOT EXISTS ${METATOMIC_IMPLIB_LOCATION}) + message(FATAL_ERROR "could not find metatomic library at '${METATOMIC_IMPLIB_LOCATION}', please re-install metatomic") + endif() + + set_target_properties(metatomic::shared PROPERTIES + IMPORTED_IMPLIB ${METATOMIC_IMPLIB_LOCATION} + ) + endif() +endif() + + +# Static library target +if (@METATOMIC_INSTALL_BOTH_STATIC_SHARED@ OR NOT @BUILD_SHARED_LIBS@) + if (NOT EXISTS ${METATOMIC_STATIC_LOCATION}) + message(FATAL_ERROR "could not find metatomic library at '${METATOMIC_STATIC_LOCATION}', please re-install metatomic") + endif() + + add_library(metatomic::static STATIC IMPORTED) + set_target_properties(metatomic::static PROPERTIES + IMPORTED_LOCATION ${METATOMIC_STATIC_LOCATION} + INTERFACE_INCLUDE_DIRECTORIES ${METATOMIC_INCLUDE} + INTERFACE_LINK_LIBRARIES "@CARGO_DEFAULT_LIBRARIES@" + BUILD_VERSION "@METATOMIC_FULL_VERSION@" + ) + + target_compile_features(metatomic::static INTERFACE cxx_std_17) +endif() + +# Export either the shared or static library as the metatomic target +if (@BUILD_SHARED_LIBS@) + add_library(metatomic ALIAS metatomic::shared) +else() + add_library(metatomic ALIAS metatomic::static) +endif() + + +if (@BUILD_SHARED_LIBS@) + find_package_handle_standard_args(metatomic DEFAULT_MSG METATOMIC_SHARED_LOCATION METATOMIC_INCLUDE) +else() + find_package_handle_standard_args(metatomic DEFAULT_MSG METATOMIC_STATIC_LOCATION METATOMIC_INCLUDE) +endif() diff --git a/metatomic-core/cmake/tempdir.cmake b/metatomic-core/cmake/tempdir.cmake new file mode 100644 index 000000000..52e4805fc --- /dev/null +++ b/metatomic-core/cmake/tempdir.cmake @@ -0,0 +1,51 @@ +# Create a temporary directory using mktemp on *nix and powershell on windows +function(get_tempdir _outvar_) + # special case for github actions, where $TEMP might + # exist but point to nowhere/a non writable location + # https://docs.github.com/en/actions/learn-github-actions/variables + if (DEFINED ENV{RUNNER_TEMP}) + string(RANDOM LENGTH 12 _dirname_) + set(_output_ $ENV{RUNNER_TEMP}/${_dirname_}) + file(TO_NATIVE_PATH "${_output_}" _output_) + file(MAKE_DIRECTORY ${_output_}) + set(${_outvar_} ${_output_} PARENT_SCOPE) + return() + endif() + + find_program(MKTEMP_EXE NAMES mktemp) + if(MKTEMP_EXE) + execute_process( + COMMAND ${MKTEMP_EXE} -d + OUTPUT_VARIABLE _output_ + OUTPUT_STRIP_TRAILING_WHITESPACE + RESULT_VARIABLE _status_ + ) + + if(_status_ EQUAL 0) + file(MAKE_DIRECTORY ${_output_}) + set(${_outvar_} ${_output_} PARENT_SCOPE) + return() + endif() + endif() + + + find_program(POWERSHELL_EXE NAMES pwsh powershell) + if(POWERSHELL_EXE) + execute_process( + COMMAND ${POWERSHELL_EXE} -c "[System.IO.Path]::GetTempPath()" + OUTPUT_VARIABLE _output_ + OUTPUT_STRIP_TRAILING_WHITESPACE + RESULT_VARIABLE _status_ + ) + + if(_status_ EQUAL 0) + string(RANDOM LENGTH 12 _dirname_) + set(_output_ ${_output_}${_dirname_}) + file(MAKE_DIRECTORY ${_output_}) + set(${_outvar_} ${_output_} PARENT_SCOPE) + return() + endif() + endif() + + message(FATAL_ERROR "Could not find mktemp or PowerShell to make temporary directory") +endfunction() diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h new file mode 100644 index 000000000..83b6cb1fc --- /dev/null +++ b/metatomic-core/include/metatomic.h @@ -0,0 +1,32 @@ +#ifndef METATOMIC_H +#define METATOMIC_H + +/* ============ Automatically generated file, DO NOT EDIT. ============== * + * * + * This file is automatically generated from the metatomic sources, * + * using cbindgen. If you want to change this file (including documentation), * + * make the corresponding changes in the rust sources and regenerate it. * + * ============================================================================= */ + +#include +#include +#include +#include +#include "metatomic/version.h" + +#ifdef __cplusplus +extern "C" { +#endif // __cplusplus + +/** + * Get the runtime version of the metatomic library as a string. + * + * This version follows the `..[-]` format. + */ +const char *mta_version(void); + +#ifdef __cplusplus +} // extern "C" +#endif // __cplusplus + +#endif /* METATOMIC_H */ diff --git a/metatomic-core/include/metatomic.hpp b/metatomic-core/include/metatomic.hpp new file mode 100644 index 000000000..016f26bc5 --- /dev/null +++ b/metatomic-core/include/metatomic.hpp @@ -0,0 +1,2 @@ +#include "metatomic/system.hpp" // IWYU pragma: export +#include "metatomic/model.hpp" // IWYU pragma: export diff --git a/metatomic-core/include/metatomic/model.hpp b/metatomic-core/include/metatomic/model.hpp new file mode 100644 index 000000000..1cae91bdf --- /dev/null +++ b/metatomic-core/include/metatomic/model.hpp @@ -0,0 +1,7 @@ +#pragma once + +#include + +namespace metatomic { + +} // namespace metatomic diff --git a/metatomic-core/include/metatomic/system.hpp b/metatomic-core/include/metatomic/system.hpp new file mode 100644 index 000000000..1cae91bdf --- /dev/null +++ b/metatomic-core/include/metatomic/system.hpp @@ -0,0 +1,7 @@ +#pragma once + +#include + +namespace metatomic { + +} // namespace metatomic diff --git a/metatomic-core/src/c_api/mod.rs b/metatomic-core/src/c_api/mod.rs new file mode 100644 index 000000000..33e0786dc --- /dev/null +++ b/metatomic-core/src/c_api/mod.rs @@ -0,0 +1,18 @@ +use std::ffi::CString; +use std::os::raw::c_char; + +use once_cell::sync::Lazy; + + +static VERSION: Lazy = Lazy::new(|| { + CString::new(env!("METATOMIC_FULL_VERSION")).expect("version contains NULL byte") +}); + + +/// Get the runtime version of the metatomic library as a string. +/// +/// This version follows the `..[-]` format. +#[no_mangle] +pub extern "C" fn mta_version() -> *const c_char { + return VERSION.as_ptr(); +} diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs new file mode 100644 index 000000000..bc47948b4 --- /dev/null +++ b/metatomic-core/src/lib.rs @@ -0,0 +1,13 @@ +#![warn(clippy::all, clippy::pedantic)] + +// disable some style lints +#![allow(clippy::needless_return, clippy::must_use_candidate, clippy::comparison_chain)] +#![allow(clippy::redundant_field_names, clippy::redundant_closure_for_method_calls, clippy::redundant_else)] +#![allow(clippy::unreadable_literal, clippy::option_if_let_else, clippy::module_name_repetitions)] +#![allow(clippy::missing_errors_doc, clippy::missing_panics_doc, clippy::missing_safety_doc)] +#![allow(clippy::similar_names, clippy::borrow_as_ptr, clippy::uninlined_format_args)] +#![allow(clippy::let_underscore_untyped, clippy::manual_let_else, clippy::empty_line_after_doc_comments)] + + +#[doc(hidden)] +mod c_api; diff --git a/metatomic-core/tests/CMakeLists.txt b/metatomic-core/tests/CMakeLists.txt new file mode 100644 index 000000000..1a96108d5 --- /dev/null +++ b/metatomic-core/tests/CMakeLists.txt @@ -0,0 +1,86 @@ +cmake_minimum_required(VERSION 3.22) +project(metatomic-tests) + +if (${CMAKE_CURRENT_SOURCE_DIR} STREQUAL ${CMAKE_SOURCE_DIR}) + if("${CMAKE_BUILD_TYPE}" STREQUAL "" AND "${CMAKE_CONFIGURATION_TYPES}" STREQUAL "") + message(STATUS "Setting build type to 'release' as none was specified.") + set(CMAKE_BUILD_TYPE "release" + CACHE STRING + "Choose the type of build, options are: debug or release" + FORCE) + set_property(CACHE CMAKE_BUILD_TYPE PROPERTY STRINGS release debug) + endif() +endif() + +if (MINGW) + # CI can't find libsdc++, so we statically link it + set(CMAKE_EXE_LINKER_FLAGS "-static-libstdc++") +endif() + +add_subdirectory(../ metatomic) +get_target_property(METATOMIC_IMPORTED_LOCATION metatomic::shared IMPORTED_LOCATION) +get_filename_component(METATOMIC_DIR ${METATOMIC_IMPORTED_LOCATION} DIRECTORY) + +add_subdirectory(external) + +find_program(VALGRIND valgrind) +if (VALGRIND) + if (NOT "$ENV{METATOMIC_DISABLE_VALGRIND}" EQUAL "1") + message(STATUS "Running tests using valgrind") + set(TEST_COMMAND + "${VALGRIND}" "--tool=memcheck" "--dsymutil=yes" "--error-exitcode=125" + "--leak-check=full" "--show-leak-kinds=definite,indirect,possible" "--track-origins=yes" + "--gen-suppressions=all" + ) + endif() +else() + set(TEST_COMMAND "") +endif() + +if (CMAKE_CXX_COMPILER_ID MATCHES "Clang") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Weverything") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-c++98-compat") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-c++98-compat-pedantic") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-weak-vtables") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-float-equal") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-missing-prototypes") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-shadow-uncaptured-local") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-padded") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-unsafe-buffer-usage") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-poison-system-directories") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-allocator-wrappers") +endif() + + +enable_testing() + +file(GLOB ALL_TESTS *.cpp) +foreach(_file_ ${ALL_TESTS}) + get_filename_component(_name_ ${_file_} NAME_WE) + add_executable(${_name_} ${_file_}) + target_link_libraries(${_name_} metatomic catch) + + set_target_properties(${_name_} PROPERTIES + # Ensure that the binaries find the right shared library. + # + # Without this, when configuring with cmake before the library is built, + # cmake does not find the library on the filesystem and does not add the + # RPATH to executables linking to it + BUILD_RPATH ${METATOMIC_DIR} + NO_SYSTEM_FROM_IMPORTED ON + ) + + add_test( + NAME ${_name_} + COMMAND ${TEST_COMMAND} $ + ) + + if(WIN32) + # We need to set the path to allow access to metatomic.dll + # this does a similar job to the BUILD_RPATH above + STRING(REPLACE ";" "\\;" PATH_STRING "$ENV{PATH}") + set_tests_properties(${_name_} PROPERTIES + ENVIRONMENT "PATH=${PATH_STRING}\;$" + ) + endif() +endforeach() diff --git a/metatomic-core/tests/check-cxx-install.rs b/metatomic-core/tests/check-cxx-install.rs new file mode 100644 index 000000000..d66f4883b --- /dev/null +++ b/metatomic-core/tests/check-cxx-install.rs @@ -0,0 +1,64 @@ +use std::path::PathBuf; +use std::sync::Mutex; + +mod utils; + +lazy_static::lazy_static! { + // Make sure only one of the tests below run at the time, since they both + // try to modify the same files + static ref LOCK: Mutex<()> = Mutex::new(()); +} + + +/// Check that metatomic can be built and installed with cmake, and that the +/// installed version can be used from another cmake project with `find_package` +#[test] +fn check_cxx_install() { + let _guard = match LOCK.lock() { + Ok(guard) => guard, + Err(_) => { + panic!("another test failed, stopping") + } + }; + + const CARGO_TARGET_TMPDIR: &str = env!("CARGO_TARGET_TMPDIR"); + + // ====================================================================== // + // build and install metatensor with cmake + let mut build_dir = PathBuf::from(CARGO_TARGET_TMPDIR); + build_dir.push("cxx-install"); + build_dir.push("cmake-find-package"); + std::fs::create_dir_all(&build_dir).expect("failed to create build dir"); + + let deps_dir = build_dir.join("deps"); + let virtualenv_dir = deps_dir.join("virtualenv"); + std::fs::create_dir_all(&virtualenv_dir).expect("failed to create virtualenv dir"); + let python_exe = utils::create_python_venv(virtualenv_dir); + let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python_exe); + + let metatomic_dep = deps_dir.join("metatomic-core"); + let source_dir = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); + + let cmake_args = vec![ + format!("-DCMAKE_PREFIX_PATH={}", metatensor_cmake_prefix.display()), + ]; + let metatomic_cmake_prefix = utils::setup_metatomic_cmake(&source_dir, &metatomic_dep, cmake_args); + + // ====================================================================== // + // try to use the installed metatomic from cmake + let mut tests_source_dir = source_dir; + tests_source_dir.extend(["tests", "cmake-project"]); + + // configure cmake for the test cmake project + let mut cmake_config = utils::cmake_config(&tests_source_dir, &build_dir); + cmake_config.arg(format!("-DCMAKE_PREFIX_PATH={};{}", metatensor_cmake_prefix.display(), metatomic_cmake_prefix.display())); + utils::run_command(cmake_config, "cmake configuration"); + + // build the code, linking to metatensor + let cmake_build = utils::cmake_build(&build_dir); + utils::run_command(cmake_build, "cmake build"); + + // run the executables + let ctest = utils::ctest(&build_dir); + utils::run_command(ctest, "ctest"); +} diff --git a/metatomic-core/tests/cmake-project/CMakeLists.txt b/metatomic-core/tests/cmake-project/CMakeLists.txt new file mode 100644 index 000000000..2b04acfa4 --- /dev/null +++ b/metatomic-core/tests/cmake-project/CMakeLists.txt @@ -0,0 +1,84 @@ +cmake_minimum_required(VERSION 3.22) + +message(STATUS "Running with CMake version ${CMAKE_VERSION}") + +project(metatomic-test-cmake-project C CXX) + +option(USE_CMAKE_SUBDIRECTORY OFF) + +if (MINGW) + # CI can't find libsdc++, so we statically link it + set(CMAKE_EXE_LINKER_FLAGS "-static-libstdc++") +endif() + + +if (USE_CMAKE_SUBDIRECTORY) + message(STATUS "Using metatomic with add_subdirectory") + # build metatomic as part of this project + add_subdirectory(../../ metatomic) + + # load metatomic from the build path + set(CMAKE_BUILD_RPATH "$") +else() + message(STATUS "Using metatomic with find_package") + # If building a dev version, we also need to update the REQUIRED_METATOMIC_VERSION + # in the same way we update the metatomic-torch version + include(../../cmake/dev-versions.cmake) + set(REQUIRED_METATOMIC_VERSION "0.1.0") + create_development_version("${REQUIRED_METATOMIC_VERSION}" METATOMIC_CORE_FULL_VERSION "metatomic-core-v") + string(REGEX REPLACE "([0-9]*)\\.([0-9]*).*" "\\1.\\2" REQUIRED_METATOMIC_VERSION ${METATOMIC_CORE_FULL_VERSION}) + + find_package(metatomic ${REQUIRED_METATOMIC_VERSION} REQUIRED) + + if(TARGET metatomic::shared) + get_target_property(mta_build_version metatomic::shared BUILD_VERSION) + if (NOT ${mta_build_version} STREQUAL ${METATOMIC_CORE_FULL_VERSION}) + message(FATAL_ERROR "Invalid BUILD_VERSION for metatomic::shared, expected ${METATOMIC_CORE_FULL_VERSION} but got ${mta_build_version}") + endif() + endif() + + if(TARGET metatomic::static) + get_target_property(mta_build_version metatomic::static BUILD_VERSION) + if (NOT ${mta_build_version} STREQUAL ${METATOMIC_CORE_FULL_VERSION}) + message(FATAL_ERROR "Invalid BUILD_VERSION for metatomic::static, expected ${METATOMIC_CORE_FULL_VERSION} but got ${mta_build_version}") + endif() + endif() +endif() + +enable_testing() + + +if(TARGET metatomic::shared) + add_executable(c-main src/main.c) + target_link_libraries(c-main metatomic::shared) + + add_executable(cxx-main src/main.cpp) + target_link_libraries(cxx-main metatomic::shared) + + add_test(NAME c-main COMMAND c-main) + add_test(NAME cxx-main COMMAND cxx-main) + + if(WIN32) + # We need to set the path to allow access to metatomic.dll + STRING(REPLACE ";" "\\;" PATH_STRING "$ENV{PATH}") + set_tests_properties(c-main PROPERTIES + ENVIRONMENT "PATH=${PATH_STRING}\;$" + ) + + set_tests_properties(cxx-main PROPERTIES + ENVIRONMENT "PATH=${PATH_STRING}\;$" + ) + endif() +endif() + + +if(TARGET metatomic::static) + add_executable(c-main-static src/main.c) + target_link_libraries(c-main-static metatomic::static) + + add_executable(cxx-main-static src/main.cpp) + target_link_libraries(cxx-main-static metatomic::static) + + add_test(NAME c-main-static COMMAND c-main-static) + add_test(NAME cxx-main-static COMMAND cxx-main-static) +endif() diff --git a/metatomic-core/tests/cmake-project/README.md b/metatomic-core/tests/cmake-project/README.md new file mode 100644 index 000000000..70a687bf0 --- /dev/null +++ b/metatomic-core/tests/cmake-project/README.md @@ -0,0 +1,3 @@ +# Sample CMake project using metatomic + +This is a basic cmake project linking to metatomic from C and C++ code. diff --git a/metatomic-core/tests/cmake-project/src/main.c b/metatomic-core/tests/cmake-project/src/main.c new file mode 100644 index 000000000..dcad0f764 --- /dev/null +++ b/metatomic-core/tests/cmake-project/src/main.c @@ -0,0 +1,8 @@ +#include + +#include + +int main(void) { + printf("Metatomic version: %s\n", mta_version()); + return 0; +} diff --git a/metatomic-core/tests/cmake-project/src/main.cpp b/metatomic-core/tests/cmake-project/src/main.cpp new file mode 100644 index 000000000..04ec152b6 --- /dev/null +++ b/metatomic-core/tests/cmake-project/src/main.cpp @@ -0,0 +1,9 @@ +#include + +#include + + +int main() { + std::cout << "Metatomic version: " << mta_version() << std::endl; + return 0; +} diff --git a/metatomic-torch/tests/external/.gitattributes b/metatomic-core/tests/external/.gitattributes similarity index 100% rename from metatomic-torch/tests/external/.gitattributes rename to metatomic-core/tests/external/.gitattributes diff --git a/metatomic-torch/tests/external/CMakeLists.txt b/metatomic-core/tests/external/CMakeLists.txt similarity index 100% rename from metatomic-torch/tests/external/CMakeLists.txt rename to metatomic-core/tests/external/CMakeLists.txt diff --git a/metatomic-torch/tests/external/catch/catch.cpp b/metatomic-core/tests/external/catch/catch.cpp similarity index 100% rename from metatomic-torch/tests/external/catch/catch.cpp rename to metatomic-core/tests/external/catch/catch.cpp diff --git a/metatomic-torch/tests/external/catch/catch.hpp b/metatomic-core/tests/external/catch/catch.hpp similarity index 100% rename from metatomic-torch/tests/external/catch/catch.hpp rename to metatomic-core/tests/external/catch/catch.hpp diff --git a/metatomic-core/tests/misc.cpp b/metatomic-core/tests/misc.cpp new file mode 100644 index 000000000..bf0ce275f --- /dev/null +++ b/metatomic-core/tests/misc.cpp @@ -0,0 +1,15 @@ +#include + +#include "metatomic.h" + + +TEST_CASE("Version macros") { + CHECK(std::string(METATOMIC_VERSION) == mta_version()); + + auto version = std::to_string(METATOMIC_VERSION_MAJOR) + "." + + std::to_string(METATOMIC_VERSION_MINOR) + "." + + std::to_string(METATOMIC_VERSION_PATCH); + + // METATOMIC_VERSION should start with `x.y.z` + CHECK(std::string(METATOMIC_VERSION).find(version) == 0); +} diff --git a/metatomic-core/tests/run-cxx-tests.rs b/metatomic-core/tests/run-cxx-tests.rs new file mode 100644 index 000000000..0d3b48d9d --- /dev/null +++ b/metatomic-core/tests/run-cxx-tests.rs @@ -0,0 +1,40 @@ +use std::path::PathBuf; + +mod utils; + +#[test] +fn run_cxx_tests() { + const CARGO_TARGET_TMPDIR: &str = env!("CARGO_TARGET_TMPDIR"); + + let mut build_dir = PathBuf::from(CARGO_TARGET_TMPDIR); + build_dir.push("cxx-tests"); + std::fs::create_dir_all(&build_dir).expect("failed to create build dir"); + + // ====================================================================== // + // setup dependencies for the torch tests + let deps_dir = build_dir.join("deps"); + let virtualenv_dir = deps_dir.join("virtualenv"); + std::fs::create_dir_all(&virtualenv_dir).expect("failed to create virtualenv dir"); + let python_exe = utils::create_python_venv(virtualenv_dir); + let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python_exe); + + // ====================================================================== // + // build the metatomic C++ tests and run them + + let mut source_dir = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); + source_dir.push("tests"); + + // configure cmake for the tests + let mut cmake_config = utils::cmake_config(&source_dir, &build_dir); + cmake_config.arg("-DCMAKE_EXPORT_COMPILE_COMMANDS=ON"); + cmake_config.arg(format!("-DCMAKE_PREFIX_PATH={}", metatensor_cmake_prefix.display())); + utils::run_command(cmake_config, "cmake configuration"); + + // build the tests + let cmake_build = utils::cmake_build(&build_dir); + utils::run_command(cmake_build, "cmake build"); + + // run the tests + let ctest = utils::ctest(&build_dir); + utils::run_command(ctest, "ctest"); +} diff --git a/metatomic-core/tests/utils/mod.rs b/metatomic-core/tests/utils/mod.rs new file mode 100644 index 000000000..a04e8a194 --- /dev/null +++ b/metatomic-core/tests/utils/mod.rs @@ -0,0 +1,473 @@ +#![allow(dead_code)] +#![allow(clippy::needless_return)] + +use std::io::{Read, Write}; +use std::path::{Path, PathBuf}; +use std::process::{Command, Stdio}; + +fn build_type() -> &'static str { + // assume that debug assertion means that we are building the code in + // debug mode, even if that could be not true in some cases + if cfg!(debug_assertions) { + "debug" + } else { + "release" + } +} + +fn append_flags(existing: Option, extra: &str) -> String { + match existing { + Some(flags) if !flags.trim().is_empty() => format!("{flags} {extra}"), + _ => extra.into(), + } +} + +pub fn cmake_config(source_dir: &Path, build_dir: &Path) -> Command { + let cmake = which::which("cmake").expect("could not find cmake"); + + let mut cmake_config = Command::new(cmake); + cmake_config.current_dir(build_dir); + cmake_config.arg(source_dir); + cmake_config.arg("--no-warn-unused-cli"); + cmake_config.arg(format!("-DCMAKE_BUILD_TYPE={}", build_type())); + + // the cargo executable currently running + let cargo_exe = std::env::var("CARGO").expect("CARGO env var is not set"); + cmake_config.arg(format!("-DCARGO_EXE={}", cargo_exe)); + + if std::env::var_os("CARGO_LLVM_COV").is_some() { + let coverage_compile_flags = "-fprofile-instr-generate -fcoverage-mapping"; + let coverage_link_flags = "-fprofile-instr-generate"; + + let c_flags = append_flags(std::env::var("CFLAGS").ok(), coverage_compile_flags); + let cxx_flags = append_flags(std::env::var("CXXFLAGS").ok(), coverage_compile_flags); + let exe_linker_flags = + append_flags(std::env::var("LDFLAGS").ok(), coverage_link_flags); + + cmake_config.arg(format!("-DCMAKE_C_FLAGS={c_flags}")); + cmake_config.arg(format!("-DCMAKE_CXX_FLAGS={cxx_flags}")); + cmake_config.arg(format!("-DCMAKE_EXE_LINKER_FLAGS={exe_linker_flags}")); + cmake_config.arg(format!("-DCMAKE_SHARED_LINKER_FLAGS={exe_linker_flags}")); + } + + return cmake_config; +} + +pub fn cmake_build(build_dir: &Path) -> Command { + let cmake = which::which("cmake").expect("could not find cmake"); + + let mut cmake_build = Command::new(cmake); + cmake_build.current_dir(build_dir); + cmake_build.arg("--build"); + cmake_build.arg("."); + cmake_build.arg("--parallel"); + cmake_build.arg("--config"); + cmake_build.arg(build_type()); + + return cmake_build; +} + + +pub fn ctest(build_dir: &Path) -> Command { + let ctest = which::which("ctest").expect("could not find ctest"); + + let mut ctest = Command::new(ctest); + ctest.current_dir(build_dir); + ctest.arg("--output-on-failure"); + ctest.arg("--build-config"); + ctest.arg(build_type()); + + return ctest +} + +/// Find the path to the uv binary, or None if not present +fn find_uv() -> Option { + which::which("uv").ok() +} + +/// Find the path to the `python`or `python3` binary on the user system +fn find_python() -> PathBuf { + if let Ok(python) = which::which("python") { + let output = Command::new(&python) + .arg("-c") + .arg("import sys; print(sys.version_info.major)") + .output() + .expect("could not run python"); + + if output.status.success() { + let stdout = String::from_utf8_lossy(&output.stdout); + + if stdout.trim() == "3" { + // we found Python 3 + return python; + } + } + } + + // try python3 + let python = which::which("python3").expect("failed to run `which python3`"); + let output = Command::new(&python) + .arg("-c") + .arg("import sys; print(sys.version_info.major)") + .output() + .expect("could not run python"); + + if output.status.success() { + let stdout = String::from_utf8_lossy(&output.stdout); + if stdout.trim() == "3" { + // we found Python 3 + return python; + } + } + + panic!("could not find Python 3") +} + +/// Helper: get python executable path inside a venv +fn python_in_venv(venv_dir: &Path) -> PathBuf { + let mut python = venv_dir.to_path_buf(); + if cfg!(target_os = "windows") { + python.extend(["Scripts", "python.exe"]); + } else { + python.extend(["bin", "python"]); + } + python +} + +/// Create a fresh Python virtualenv using uv if available, else fallback to +/// `python -m venv`, and return the path to the python executable in the venv +pub fn create_python_venv(build_dir: PathBuf) -> PathBuf { + if let Some(uv_bin) = find_uv() { + let mut cmd = Command::new(&uv_bin); + cmd.arg("venv"); + cmd.arg("--clear"); + cmd.arg(&build_dir); + + run_command(cmd, "uv venv creation"); + } else { + let mut cmd = Command::new(find_python()); + cmd.arg("-m"); + cmd.arg("venv"); + cmd.arg(&build_dir); + + run_command(cmd, "python to create virtualenv with `venv`"); + + // update pip in case the system uses a very old one + let python = python_in_venv(&build_dir); + let mut cmd = Command::new(&python); + cmd.arg("-m"); + cmd.arg("pip"); + cmd.arg("install"); + cmd.arg("--upgrade"); + cmd.arg("pip"); + + run_command(cmd, "pip upgrade in virtualenv"); + } + + python_in_venv(&build_dir) +} + +#[derive(Default)] +pub struct PipInstallOptions { + pub upgrade: bool, + pub no_deps: bool, + pub no_build_isolation: bool, +} + +/// Install a package with pip (uses uv if present, else falls back to python) +fn pip_install( + python: &Path, + packages: &[&str], + options: PipInstallOptions, +) { + if let Some(uv_bin) = find_uv() { + let mut cmd = Command::new(&uv_bin); + cmd.arg("pip").arg("install").arg("--python").arg(python); + + // follow the same behavior as pip when there are multiple indexes + cmd.arg("--index-strategy"); + cmd.arg("unsafe-best-match"); + + if options.upgrade { + cmd.arg("--upgrade"); + } + if options.no_deps { + cmd.arg("--no-deps"); + } + if options.no_build_isolation { + cmd.arg("--no-build-isolation"); + // uv doesn't support --check-build-dependencies + } + + for package in packages { + cmd.arg(package); + } + + run_command(cmd, "uv pip install"); + } else { + let mut cmd = Command::new(python); + cmd.arg("-m").arg("pip").arg("install"); + if options.upgrade { + cmd.arg("--upgrade"); + } + if options.no_deps { + cmd.arg("--no-deps"); + } + if options.no_build_isolation { + // If pip, add both supported options + cmd.arg("--no-build-isolation"); + cmd.arg("--check-build-dependencies"); + } + + for package in packages { + cmd.arg(package); + } + + run_command(cmd, "pip install"); + } +} + +/// Download PyTorch in a Python virtualenv, and return the +/// CMAKE_PREFIX_PATH for the corresponding libtorch +pub fn setup_torch_pip(python: &Path) -> PathBuf { + let torch_version = std::env::var("METATOMIC_TESTS_TORCH_VERSION").unwrap_or("2.13".into()); + pip_install( + python, + &[&format!("torch=={}.*", torch_version)], + PipInstallOptions { upgrade: true, no_deps: false, no_build_isolation: false } + ); + + let mut cmd = Command::new(python); + cmd.arg("-c"); + cmd.arg("import torch; print(torch.utils.cmake_prefix_path)"); + + let output = run_command(cmd, "python to get torch cmake prefix"); + + let stdout = String::from_utf8_lossy(&output.stdout); + let prefix = PathBuf::from(stdout.trim()); + if !prefix.exists() { + panic!("'torch.utils.cmake_prefix' at '{}' does not exist", prefix.display()); + } + + return prefix; +} + +/// Install metatensor in a Python virtualenv with pip, and return the +/// CMAKE_PREFIX_PATH for the installed libmetatensor. +pub fn setup_metatensor_pip(python: &Path) -> PathBuf { + pip_install(python, &["metatensor-core >=0.2.0,<0.3"], PipInstallOptions::default()); + + let mut cmd = Command::new(python); + cmd.arg("-c"); + cmd.arg("import metatensor; print(metatensor.utils.cmake_prefix_path)"); + + let output = run_command(cmd, "python to get metatensor cmake prefix"); + + let stdout = String::from_utf8_lossy(&output.stdout); + let prefix = PathBuf::from(stdout.trim()); + if !prefix.exists() { + panic!("'metatensor.utils.cmake_prefix' at '{}' does not exist", prefix.display()); + } + + return prefix; +} + +/// Install metatensor-torch in a Python virtualenv with pip, and return the +/// CMAKE_PREFIX_PATH for the installed libmetatensor_torch. +pub fn setup_metatensor_torch_pip(python: &Path) -> PathBuf { + pip_install(python, &["metatensor-torch >=0.9.0,<0.10"], PipInstallOptions::default()); + + let mut cmd = Command::new(python); + cmd.arg("-c"); + cmd.arg("import metatensor.torch; print(metatensor.torch.utils.cmake_prefix_path)"); + + let output = run_command(cmd, "python to get metatensor_torch cmake prefix"); + + let stdout = String::from_utf8_lossy(&output.stdout); + let prefix = PathBuf::from(stdout.trim()); + if !prefix.exists() { + panic!("'metatensor.torch.utils.cmake_prefix' at '{}' does not exist", prefix.display()); + } + + return prefix; +} + +/// Build metatomic-torch located in `source_dir` inside `build_dir`, and return +/// the installation prefix. +pub fn setup_metatomic_torch_cmake(source_dir: &Path, build_dir: &Path, cmake_args: Vec) -> PathBuf { + std::fs::create_dir_all(build_dir).expect("failed to create metatomic build dir"); + + // configure cmake for metatomic-torch + let mut cmake_config = cmake_config(source_dir, build_dir); + + let install_prefix = build_dir.join("usr"); + cmake_config.arg(format!("-DCMAKE_INSTALL_PREFIX={}", install_prefix.display())); + + // Add any additional cmake arguments + for arg in cmake_args { + cmake_config.arg(arg); + } + + run_command(cmake_config, "cmake configuration for metatomic_torch"); + + // build and install metatomic-torch + let mut cmake_build = cmake_build(build_dir); + cmake_build.arg("--target"); + cmake_build.arg("install"); + + run_command(cmake_build, "cmake build for metatomic_torch"); + + install_prefix +} + +/// Build metatomic-core located in `source_dir` inside `build_dir`, and return +/// the installation prefix +pub fn setup_metatomic_cmake(source_dir: &Path, build_dir: &Path, cmake_args: Vec) -> PathBuf { + std::fs::create_dir_all(build_dir).expect("failed to create metatomic build dir"); + + // configure cmake for metatomic + let mut cmake_config = cmake_config(source_dir, build_dir); + + let install_prefix = build_dir.join("usr"); + cmake_config.arg(format!("-DCMAKE_INSTALL_PREFIX={}", install_prefix.display())); + + // Add any additional cmake arguments + for arg in cmake_args { + cmake_config.arg(arg); + } + + run_command(cmake_config, "cmake configuration for metatomic"); + + // build and install metatomic + let mut cmake_build = cmake_build(build_dir); + cmake_build.arg("--target"); + cmake_build.arg("install"); + + run_command(cmake_build, "cmake build for metatomic"); + + install_prefix +} + +/// Install metatomic-core in a Python virtualenv with pip, and return the +/// CMAKE_PREFIX_PATH for the installed libmetatomic. +pub fn setup_metatomic_core_pip(python: &Path, source_dir: &Path) -> PathBuf { + // build dependencies + pip_install( + python, + &["cmake", "packaging >=26", "setuptools >=77"], + PipInstallOptions::default() + ); + // runtime dependencies which are not just metatensor and metatensor-torch + pip_install(python, &["wigners"], PipInstallOptions::default()); + + pip_install( + python, + &[&source_dir.display().to_string()], + PipInstallOptions { + upgrade: true, + no_deps: true, + no_build_isolation: true + } + ); + + // let mut cmd = Command::new(python); + // cmd.arg("-c"); + // cmd.arg("import metatomic; print(metatomic.utils.cmake_prefix_path)"); + + // let output = run_command(cmd, "python to get metatomic cmake prefix"); + + // let stdout = String::from_utf8_lossy(&output.stdout); + // let prefix = PathBuf::from(stdout.trim()); + // if !prefix.exists() { + // panic!("'metatomic.utils.cmake_prefix' at '{}' does not exist", prefix.display()); + // } + + // return prefix; + return PathBuf::new(); +} + + +/// Install metatomic-torch in a Python virtualenv with pip, and return the +/// CMAKE_PREFIX_PATH for the installed libmetatomic_torch. +pub fn setup_metatomic_torch_pip(python: &Path, source_dir: &Path) -> PathBuf { + pip_install( + python, + &[&source_dir.display().to_string()], + PipInstallOptions { + upgrade: true, + no_deps: true, + no_build_isolation: true + } + ); + + let mut cmd = Command::new(python); + cmd.arg("-c"); + cmd.arg("import metatomic.torch; print(metatomic.torch.utils.cmake_prefix_path)"); + + let output = run_command(cmd, "python to get metatomic_torch cmake prefix"); + + let stdout = String::from_utf8_lossy(&output.stdout); + let prefix = PathBuf::from(stdout.trim()); + if !prefix.exists() { + panic!("'metatomic.torch.utils.cmake_prefix' at '{}' does not exist", prefix.display()); + } + + return prefix; +} + +pub fn run_command(mut command: Command, context: &str) -> std::process::Output { + write!(std::io::stdout().lock(), "\n\n[Running] {:?}\n\n", command).unwrap(); + + let mut child = command + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn().unwrap_or_else(|_| panic!("failed to spawn {}", context)); + + let mut child_stdout = child.stdout.take().expect("missing stdout"); + let mut child_stderr = child.stderr.take().expect("missing stderr"); + + let out_handle = std::thread::spawn(move || -> std::io::Result> { + let mut buf = [0u8; 8192]; + let mut captured = Vec::new(); + let mut sink = std::io::stdout().lock(); + loop { + let n = child_stdout.read(&mut buf)?; + if n == 0 { + break; + } + sink.write_all(&buf[..n])?; + sink.flush()?; + captured.extend_from_slice(&buf[..n]); + } + Ok(captured) + }); + + let err_handle = std::thread::spawn(move || -> std::io::Result> { + let mut buf = [0u8; 8192]; + let mut captured = Vec::new(); + let mut sink = std::io::stderr().lock(); + loop { + let n = child_stderr.read(&mut buf)?; + if n == 0 { + break; + } + sink.write_all(&buf[..n])?; + sink.flush()?; + captured.extend_from_slice(&buf[..n]); + } + Ok(captured) + }); + + let status = child.wait().unwrap_or_else(|_| panic!("failed to run {}", context)); + let stdout = String::from_utf8_lossy(&out_handle.join().unwrap().unwrap()).into_owned(); + let stderr = String::from_utf8_lossy(&err_handle.join().unwrap().unwrap()).into_owned(); + + if !status.success() { + panic!( + "{} failed, status: {}\nstderr:\n\n{}\nstdout:\n\n{}\n", + context, status, stderr, stdout + ); + } + + return std::process::Output { status, stdout: stdout.into_bytes(), stderr: stderr.into_bytes() }; +} diff --git a/metatomic-torch/tests/CMakeLists.txt b/metatomic-torch/tests/CMakeLists.txt index 8a64a4f33..7d6257a0d 100644 --- a/metatomic-torch/tests/CMakeLists.txt +++ b/metatomic-torch/tests/CMakeLists.txt @@ -1,4 +1,5 @@ -add_subdirectory(external) +# re-use catch from metatomic-core C++ tests +add_subdirectory(../../metatomic-core/tests/external external) # make sure we compile catch with the flags that torch requires. In particular, # torch sets -D_GLIBCXX_USE_CXX11_ABI=0 on Linux, which changes some of the diff --git a/metatomic-torch/tests/check-torch-install.rs b/metatomic-torch/tests/check-torch-install.rs index 8883d916e..ad8cfb604 100644 --- a/metatomic-torch/tests/check-torch-install.rs +++ b/metatomic-torch/tests/check-torch-install.rs @@ -123,8 +123,11 @@ fn check_python_install() { let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python_exe); let metatensor_torch_cmake_prefix = utils::setup_metatensor_torch_pip(&python_exe); - let python_source_dir = cargo_manifest_dir.parent().unwrap().join("python").join("metatomic_torch"); - let metatomic_torch_cmake_prefix = utils::setup_metatomic_torch_pip(&python_exe, &python_source_dir); + let mta_core_source_dir = cargo_manifest_dir.parent().unwrap().join("python").join("metatomic_core"); + let metatomic_core_cmake_prefix = utils::setup_metatomic_core_pip(&python_exe, &mta_core_source_dir); + + let mta_torch_source_dir = cargo_manifest_dir.parent().unwrap().join("python").join("metatomic_torch"); + let metatomic_torch_cmake_prefix = utils::setup_metatomic_torch_pip(&python_exe, &mta_torch_source_dir); // ====================================================================== // // try to use the installed metatensor-torch from cmake @@ -134,10 +137,11 @@ fn check_python_install() { // configure cmake for the test cmake project let mut cmake_config = utils::cmake_config(&source_dir, &build_dir); cmake_config.arg(format!( - "-DCMAKE_PREFIX_PATH={};{};{};{}", + "-DCMAKE_PREFIX_PATH={};{};{};{};{}", pytorch_cmake_prefix.display(), metatensor_cmake_prefix.display(), metatensor_torch_cmake_prefix.display(), + metatomic_core_cmake_prefix.display(), metatomic_torch_cmake_prefix.display(), )); diff --git a/metatomic-torch/tests/utils/mod.rs b/metatomic-torch/tests/utils/mod.rs deleted file mode 100644 index 66890380e..000000000 --- a/metatomic-torch/tests/utils/mod.rs +++ /dev/null @@ -1,413 +0,0 @@ -#![allow(dead_code)] -#![allow(clippy::needless_return)] - -use std::io::{Read, Write}; -use std::path::{Path, PathBuf}; -use std::process::{Command, Stdio}; - -fn build_type() -> &'static str { - // assume that debug assertion means that we are building the code in - // debug mode, even if that could be not true in some cases - if cfg!(debug_assertions) { - "debug" - } else { - "release" - } -} - -fn append_flags(existing: Option, extra: &str) -> String { - match existing { - Some(flags) if !flags.trim().is_empty() => format!("{flags} {extra}"), - _ => extra.into(), - } -} - -pub fn cmake_config(source_dir: &Path, build_dir: &Path) -> Command { - let cmake = which::which("cmake").expect("could not find cmake"); - - let mut cmake_config = Command::new(cmake); - cmake_config.current_dir(build_dir); - cmake_config.arg(source_dir); - cmake_config.arg("--no-warn-unused-cli"); - cmake_config.arg(format!("-DCMAKE_BUILD_TYPE={}", build_type())); - - // the cargo executable currently running - let cargo_exe = std::env::var("CARGO").expect("CARGO env var is not set"); - cmake_config.arg(format!("-DCARGO_EXE={}", cargo_exe)); - - if std::env::var_os("CARGO_LLVM_COV").is_some() { - let coverage_compile_flags = "-fprofile-instr-generate -fcoverage-mapping"; - let coverage_link_flags = "-fprofile-instr-generate"; - - let c_flags = append_flags(std::env::var("CFLAGS").ok(), coverage_compile_flags); - let cxx_flags = append_flags(std::env::var("CXXFLAGS").ok(), coverage_compile_flags); - let exe_linker_flags = - append_flags(std::env::var("LDFLAGS").ok(), coverage_link_flags); - - cmake_config.arg(format!("-DCMAKE_C_FLAGS={c_flags}")); - cmake_config.arg(format!("-DCMAKE_CXX_FLAGS={cxx_flags}")); - cmake_config.arg(format!("-DCMAKE_EXE_LINKER_FLAGS={exe_linker_flags}")); - cmake_config.arg(format!("-DCMAKE_SHARED_LINKER_FLAGS={exe_linker_flags}")); - } - - return cmake_config; -} - -pub fn cmake_build(build_dir: &Path) -> Command { - let cmake = which::which("cmake").expect("could not find cmake"); - - let mut cmake_build = Command::new(cmake); - cmake_build.current_dir(build_dir); - cmake_build.arg("--build"); - cmake_build.arg("."); - cmake_build.arg("--parallel"); - cmake_build.arg("--config"); - cmake_build.arg(build_type()); - - return cmake_build; -} - - -pub fn ctest(build_dir: &Path) -> Command { - let ctest = which::which("ctest").expect("could not find ctest"); - - let mut ctest = Command::new(ctest); - ctest.current_dir(build_dir); - ctest.arg("--output-on-failure"); - ctest.arg("--build-config"); - ctest.arg(build_type()); - - return ctest -} - -/// Find the path to the uv binary, or None if not present -fn find_uv() -> Option { - which::which("uv").ok() -} - -/// Find the path to the `python`or `python3` binary on the user system -fn find_python() -> PathBuf { - if let Ok(python) = which::which("python") { - let output = Command::new(&python) - .arg("-c") - .arg("import sys; print(sys.version_info.major)") - .output() - .expect("could not run python"); - - if output.status.success() { - let stdout = String::from_utf8_lossy(&output.stdout); - - if stdout.trim() == "3" { - // we found Python 3 - return python; - } - } - } - - // try python3 - let python = which::which("python3").expect("failed to run `which python3`"); - let output = Command::new(&python) - .arg("-c") - .arg("import sys; print(sys.version_info.major)") - .output() - .expect("could not run python"); - - if output.status.success() { - let stdout = String::from_utf8_lossy(&output.stdout); - if stdout.trim() == "3" { - // we found Python 3 - return python; - } - } - - panic!("could not find Python 3") -} - -/// Helper: get python executable path inside a venv -fn python_in_venv(venv_dir: &Path) -> PathBuf { - let mut python = venv_dir.to_path_buf(); - if cfg!(target_os = "windows") { - python.extend(["Scripts", "python.exe"]); - } else { - python.extend(["bin", "python"]); - } - python -} - -/// Create a fresh Python virtualenv using uv if available, else fallback to -/// `python -m venv`, and return the path to the python executable in the venv -pub fn create_python_venv(build_dir: PathBuf) -> PathBuf { - if let Some(uv_bin) = find_uv() { - let mut cmd = Command::new(&uv_bin); - cmd.arg("venv"); - cmd.arg("--clear"); - cmd.arg(&build_dir); - - run_command(cmd, "uv venv creation"); - } else { - let mut cmd = Command::new(find_python()); - cmd.arg("-m"); - cmd.arg("venv"); - cmd.arg(&build_dir); - - run_command(cmd, "python to create virtualenv with `venv`"); - - // update pip in case the system uses a very old one - let python = python_in_venv(&build_dir); - let mut cmd = Command::new(&python); - cmd.arg("-m"); - cmd.arg("pip"); - cmd.arg("install"); - cmd.arg("--upgrade"); - cmd.arg("pip"); - - run_command(cmd, "pip upgrade in virtualenv"); - } - - python_in_venv(&build_dir) -} - -#[derive(Default)] -pub struct PipInstallOptions { - pub upgrade: bool, - pub no_deps: bool, - pub no_build_isolation: bool, -} - -/// Install a package with pip (uses uv if present, else falls back to python) -fn pip_install( - python: &Path, - packages: &[&str], - options: PipInstallOptions, -) { - if let Some(uv_bin) = find_uv() { - let mut cmd = Command::new(&uv_bin); - cmd.arg("pip").arg("install").arg("--python").arg(python); - - // follow the same behavior as pip when there are multiple indexes - cmd.arg("--index-strategy"); - cmd.arg("unsafe-best-match"); - - if options.upgrade { - cmd.arg("--upgrade"); - } - if options.no_deps { - cmd.arg("--no-deps"); - } - if options.no_build_isolation { - cmd.arg("--no-build-isolation"); - // uv doesn't support --check-build-dependencies - } - - for package in packages { - cmd.arg(package); - } - - run_command(cmd, "uv pip install"); - } else { - let mut cmd = Command::new(python); - cmd.arg("-m").arg("pip").arg("install"); - if options.upgrade { - cmd.arg("--upgrade"); - } - if options.no_deps { - cmd.arg("--no-deps"); - } - if options.no_build_isolation { - // If pip, add both supported options - cmd.arg("--no-build-isolation"); - cmd.arg("--check-build-dependencies"); - } - - for package in packages { - cmd.arg(package); - } - - run_command(cmd, "pip install"); - } -} - -/// Download PyTorch in a Python virtualenv, and return the -/// CMAKE_PREFIX_PATH for the corresponding libtorch -pub fn setup_torch_pip(python: &Path) -> PathBuf { - let torch_version = std::env::var("METATOMIC_TESTS_TORCH_VERSION").unwrap_or("2.13".into()); - pip_install( - python, - &[&format!("torch=={}.*", torch_version)], - PipInstallOptions { upgrade: true, no_deps: false, no_build_isolation: false } - ); - - let mut cmd = Command::new(python); - cmd.arg("-c"); - cmd.arg("import torch; print(torch.utils.cmake_prefix_path)"); - - let output = run_command(cmd, "python to get torch cmake prefix"); - - let stdout = String::from_utf8_lossy(&output.stdout); - let prefix = PathBuf::from(stdout.trim()); - if !prefix.exists() { - panic!("'torch.utils.cmake_prefix' at '{}' does not exist", prefix.display()); - } - - return prefix; -} - -/// Install metatensor in a Python virtualenv with pip, and return the -/// CMAKE_PREFIX_PATH for the installed libmetatensor. -pub fn setup_metatensor_pip(python: &Path) -> PathBuf { - pip_install(python, &["metatensor-core >=0.2.0,<0.3"], PipInstallOptions::default()); - - let mut cmd = Command::new(python); - cmd.arg("-c"); - cmd.arg("import metatensor; print(metatensor.utils.cmake_prefix_path)"); - - let output = run_command(cmd, "python to get metatensor cmake prefix"); - - let stdout = String::from_utf8_lossy(&output.stdout); - let prefix = PathBuf::from(stdout.trim()); - if !prefix.exists() { - panic!("'metatensor.utils.cmake_prefix' at '{}' does not exist", prefix.display()); - } - - return prefix; -} - -/// Install metatensor-torch in a Python virtualenv with pip, and return the -/// CMAKE_PREFIX_PATH for the installed libmetatensor_torch. -pub fn setup_metatensor_torch_pip(python: &Path) -> PathBuf { - pip_install(python, &["metatensor-torch >=0.9.0,<0.10"], PipInstallOptions::default()); - - let mut cmd = Command::new(python); - cmd.arg("-c"); - cmd.arg("import metatensor.torch; print(metatensor.torch.utils.cmake_prefix_path)"); - - let output = run_command(cmd, "python to get metatensor_torch cmake prefix"); - - let stdout = String::from_utf8_lossy(&output.stdout); - let prefix = PathBuf::from(stdout.trim()); - if !prefix.exists() { - panic!("'metatensor.torch.utils.cmake_prefix' at '{}' does not exist", prefix.display()); - } - - return prefix; -} - -/// Build metatomic-torch located in `source_dir` inside `build_dir`, and return -/// the installation prefix. -pub fn setup_metatomic_torch_cmake(source_dir: &Path, build_dir: &Path, cmake_args: Vec) -> PathBuf { - std::fs::create_dir_all(build_dir).expect("failed to create metatomic build dir"); - - // configure cmake for metatomic-torch - let mut cmake_config = cmake_config(source_dir, build_dir); - - let install_prefix = build_dir.join("usr"); - cmake_config.arg(format!("-DCMAKE_INSTALL_PREFIX={}", install_prefix.display())); - - // Add any additional cmake arguments - for arg in cmake_args { - cmake_config.arg(arg); - } - - run_command(cmake_config, "cmake configuration for metatomic_torch"); - - // build and install metatomic-torch - let mut cmake_build = cmake_build(build_dir); - cmake_build.arg("--target"); - cmake_build.arg("install"); - - run_command(cmake_build, "cmake build for metatomic_torch"); - - install_prefix -} - - -/// Install metatomic-torch in a Python virtualenv with pip, and return the -/// CMAKE_PREFIX_PATH for the installed libmetatomic_torch. -pub fn setup_metatomic_torch_pip(python: &Path, source_dir: &Path) -> PathBuf { - // build dependencies - pip_install(python, &["setuptools>=77", "packaging>=23", "cmake"], PipInstallOptions::default()); - // runtime dependencies which are not just metatensor and metatensor-torch - pip_install(python, &["wigners"], PipInstallOptions::default()); - - pip_install( - python, - &[&source_dir.display().to_string()], - PipInstallOptions { - upgrade: true, - no_deps: false, - no_build_isolation: true - } - ); - - let mut cmd = Command::new(python); - cmd.arg("-c"); - cmd.arg("import metatomic.torch; print(metatomic.torch.utils.cmake_prefix_path)"); - - let output = run_command(cmd, "python to get metatomic_torch cmake prefix"); - - let stdout = String::from_utf8_lossy(&output.stdout); - let prefix = PathBuf::from(stdout.trim()); - if !prefix.exists() { - panic!("'metatomic.torch.utils.cmake_prefix' at '{}' does not exist", prefix.display()); - } - - return prefix; -} - - -pub fn run_command(mut command: Command, context: &str) -> std::process::Output { - write!(std::io::stdout().lock(), "\n\n[Running] {:?}\n\n", command).unwrap(); - - let mut child = command - .stdout(Stdio::piped()) - .stderr(Stdio::piped()) - .spawn().unwrap_or_else(|_| panic!("failed to spawn {}", context)); - - let mut child_stdout = child.stdout.take().expect("missing stdout"); - let mut child_stderr = child.stderr.take().expect("missing stderr"); - - let out_handle = std::thread::spawn(move || -> std::io::Result> { - let mut buf = [0u8; 8192]; - let mut captured = Vec::new(); - let mut sink = std::io::stdout().lock(); - loop { - let n = child_stdout.read(&mut buf)?; - if n == 0 { - break; - } - sink.write_all(&buf[..n])?; - sink.flush()?; - captured.extend_from_slice(&buf[..n]); - } - Ok(captured) - }); - - let err_handle = std::thread::spawn(move || -> std::io::Result> { - let mut buf = [0u8; 8192]; - let mut captured = Vec::new(); - let mut sink = std::io::stderr().lock(); - loop { - let n = child_stderr.read(&mut buf)?; - if n == 0 { - break; - } - sink.write_all(&buf[..n])?; - sink.flush()?; - captured.extend_from_slice(&buf[..n]); - } - Ok(captured) - }); - - let status = child.wait().unwrap_or_else(|_| panic!("failed to run {}", context)); - let stdout = String::from_utf8_lossy(&out_handle.join().unwrap().unwrap()).into_owned(); - let stderr = String::from_utf8_lossy(&err_handle.join().unwrap().unwrap()).into_owned(); - - if !status.success() { - panic!( - "{} failed, status: {}\nstderr:\n\n{}\nstdout:\n\n{}\n", - context, status, stderr, stdout - ); - } - - return std::process::Output { status, stdout: stdout.into_bytes(), stderr: stderr.into_bytes() }; -} diff --git a/metatomic-torch/tests/utils/mod.rs b/metatomic-torch/tests/utils/mod.rs new file mode 120000 index 000000000..20b8b0094 --- /dev/null +++ b/metatomic-torch/tests/utils/mod.rs @@ -0,0 +1 @@ +../../../metatomic-core/tests/utils/mod.rs \ No newline at end of file diff --git a/python/metatomic_torch/build-backend/backend.py b/python/metatomic_torch/build-backend/backend.py index c762d91e6..be0389a2c 100644 --- a/python/metatomic_torch/build-backend/backend.py +++ b/python/metatomic_torch/build-backend/backend.py @@ -1,11 +1,24 @@ # This is a custom Python build backend wrapping setuptool's to only depend on # torch/metatensor-torch when building the wheel and not the sdist import os +import pathlib from setuptools import build_meta -ROOT = os.path.realpath(os.path.dirname(__file__)) +ROOT = pathlib.Path(__file__).parent.resolve() + +METATOMIC_CORE = (ROOT / ".." / ".." / "metatomic_core").resolve() +METATOMIC_NO_LOCAL_DEPS = os.environ.get("METATOMIC_NO_LOCAL_DEPS", "0") == "1" + + +if not METATOMIC_NO_LOCAL_DEPS and METATOMIC_CORE.exists(): + # we are building from a git checkout + METATOMIC_CORE_DEP = f"metatomic-core @ {METATOMIC_CORE.as_uri()}" +else: + # we are building from a sdist + METATOMIC_CORE_DEP = "metatomic-core >=0.1.0,<0.2" + FORCED_TORCH_VERSION = os.environ.get("METATOMIC_TORCH_BUILD_WITH_TORCH_VERSION") if FORCED_TORCH_VERSION is not None: @@ -27,7 +40,7 @@ # Special dependencies to build the wheels def get_requires_for_build_wheel(config_settings=None): defaults = build_meta.get_requires_for_build_wheel(config_settings) - return defaults + [TORCH_DEP] + return defaults + [TORCH_DEP, METATOMIC_CORE_DEP] def build_editable(wheel_directory, config_settings=None, metadata_directory=None): From 4ef28cc1c1b2ec19e363af8a0cc9967d4f053f1d Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Wed, 27 May 2026 16:25:27 +0200 Subject: [PATCH 05/43] Draft the C API for metatomic-core --- .github/workflows/rust-tests.yml | 186 ++++++++++++++ .github/workflows/torch-tests.yml | 9 +- docs/Doxyfile | 4 +- docs/src/core/CHANGELOG.md | 1 + docs/src/core/index.rst | 17 ++ docs/src/core/reference/c/index.rst | 17 ++ docs/src/core/reference/c/misc.rst | 56 +++++ docs/src/core/reference/c/model.rst | 16 ++ docs/src/core/reference/c/plugin.rst | 16 ++ docs/src/core/reference/c/system.rst | 42 ++++ docs/src/index.rst | 1 + metatomic-core/Cargo.toml | 11 + metatomic-core/build.rs | 1 + .../cmake/metatomic-config.in.cmake | 2 + metatomic-core/include/metatomic.h | 237 ++++++++++++++++++ metatomic-core/include/metatomic.hpp | 4 +- metatomic-core/include/metatomic/plugin.hpp | 7 + metatomic-core/include/metatomic/utils.hpp | 7 + metatomic-core/src/c_api/mod.rs | 25 +- metatomic-core/src/c_api/model.rs | 76 ++++++ metatomic-core/src/c_api/plugin.rs | 41 +++ metatomic-core/src/c_api/status.rs | 42 ++++ metatomic-core/src/c_api/system.rs | 131 ++++++++++ metatomic-core/src/c_api/utils.rs | 102 ++++++++ metatomic-core/src/lib.rs | 43 +++- metatomic-core/src/metadata.rs | 132 ++++++++++ metatomic-core/src/model.rs | 20 ++ metatomic-core/src/plugin.rs | 37 +++ metatomic-core/src/system.rs | 53 ++++ metatomic-core/src/units.rs | 7 + metatomic-core/tests/check-cxx-install.rs | 8 +- metatomic-torch/tests/check-torch-install.rs | 31 ++- rustfmt.toml | 1 + scripts/check-c-api-docs.py | 101 ++++++++ scripts/include/README | 4 + scripts/include/metatensor.h | 8 + scripts/include/metatomic/version.h | 0 scripts/include/stdarg.h | 0 scripts/include/stdbool.h | 1 + scripts/include/stddef.h | 6 + scripts/include/stdint.h | 7 + scripts/include/stdlib.h | 1 + 42 files changed, 1475 insertions(+), 36 deletions(-) create mode 100644 .github/workflows/rust-tests.yml create mode 120000 docs/src/core/CHANGELOG.md create mode 100644 docs/src/core/index.rst create mode 100644 docs/src/core/reference/c/index.rst create mode 100644 docs/src/core/reference/c/misc.rst create mode 100644 docs/src/core/reference/c/model.rst create mode 100644 docs/src/core/reference/c/plugin.rst create mode 100644 docs/src/core/reference/c/system.rst create mode 100644 metatomic-core/include/metatomic/plugin.hpp create mode 100644 metatomic-core/include/metatomic/utils.hpp create mode 100644 metatomic-core/src/c_api/model.rs create mode 100644 metatomic-core/src/c_api/plugin.rs create mode 100644 metatomic-core/src/c_api/status.rs create mode 100644 metatomic-core/src/c_api/system.rs create mode 100644 metatomic-core/src/c_api/utils.rs create mode 100644 metatomic-core/src/metadata.rs create mode 100644 metatomic-core/src/model.rs create mode 100644 metatomic-core/src/plugin.rs create mode 100644 metatomic-core/src/system.rs create mode 100644 metatomic-core/src/units.rs create mode 100644 rustfmt.toml create mode 100755 scripts/check-c-api-docs.py create mode 100644 scripts/include/README create mode 100644 scripts/include/metatensor.h create mode 100644 scripts/include/metatomic/version.h create mode 100644 scripts/include/stdarg.h create mode 100644 scripts/include/stdbool.h create mode 100644 scripts/include/stddef.h create mode 100644 scripts/include/stdint.h create mode 100644 scripts/include/stdlib.h diff --git a/.github/workflows/rust-tests.yml b/.github/workflows/rust-tests.yml new file mode 100644 index 000000000..77ac78387 --- /dev/null +++ b/.github/workflows/rust-tests.yml @@ -0,0 +1,186 @@ +name: Rust tests + +on: + push: + branches: [main] + pull_request: + # Check all PR + +concurrency: + group: rust-tests-${{ github.ref }} + cancel-in-progress: ${{ github.ref != 'refs/heads/main' }} + +jobs: + rust-tests: + name: ${{ matrix.os }} / Rust ${{ matrix.rust-version }}${{ matrix.extra-name }} + runs-on: ${{ matrix.os }} + container: ${{ matrix.container }} + defaults: + run: + shell: "bash" + env: + CMAKE_CXX_COMPILER: ${{ matrix.cxx }} + CMAKE_C_COMPILER: ${{ matrix.cc }} + CMAKE_GENERATOR: ${{ matrix.cmake-generator }} + strategy: + matrix: + include: + - os: ubuntu-24.04 + rust-version: stable + rust-target: x86_64-unknown-linux-gnu + cxx: g++ + cc: gcc + cmake-generator: Unix Makefiles + + # check the build on a stock Ubuntu 22.04, which uses cmake 3.22, and + # with our minimal supported rust version + - os: ubuntu-24.04 + rust-version: 1.74 + container: ubuntu:22.04 + rust-target: x86_64-unknown-linux-gnu + extra-name: ", cmake 3.22" + cxx: g++ + cc: gcc + cmake-generator: Unix Makefiles + + - os: macos-15 + rust-version: stable + rust-target: aarch64-apple-darwin + extra-name: "" + cxx: clang++ + cc: clang + cmake-generator: Unix Makefiles + + # - os: windows-2022 + # rust-version: stable + # rust-target: x86_64-pc-windows-msvc + # extra-name: " / MSVC" + # cxx: cl.exe + # cc: cl.exe + # cmake-generator: Visual Studio 17 2022 + + # - os: windows-2022 + # rust-version: stable + # rust-target: x86_64-pc-windows-gnu + # extra-name: " / MinGW" + # cxx: g++.exe + # cc: gcc.exe + # cmake-generator: MinGW Makefiles + steps: + - name: install dependencies in container + if: matrix.container == 'ubuntu:22.04' + run: | + apt update + apt install -y software-properties-common + apt install -y cmake make gcc g++ git curl python3-venv + + - uses: actions/checkout@v6 + with: + fetch-depth: 0 + + - name: Configure git safe directory + if: matrix.container == 'ubuntu:22.04' + run: git config --global --add safe.directory /__w/metatomic/metatomic + + - name: setup rust + uses: dtolnay/rust-toolchain@master + with: + toolchain: ${{ matrix.rust-version }} + target: ${{ matrix.rust-target }} + + - name: setup Python + uses: actions/setup-python@v6 + if: matrix.container == null + with: + # Python 3.14.5 fails with "No module named pip.__main__; 'pip' is a + # package and cannot be directly executed" when using a venv, so we + # use 3.14.4 for now + python-version: "3.14.4" + + - name: Cache Rust dependencies + uses: Leafwing-Studios/cargo-cache@v2.6.1 + with: + sweep-cache: true + + - name: install valgrind + if: matrix.do-valgrind + run: | + sudo apt-get install -y valgrind + + - name: Setup sccache + if: ${{ !env.ACT }} + uses: mozilla-actions/sccache-action@v0.0.10 + with: + version: "v0.15.0" + + - name: Setup sccache environnement variables + if: ${{ !env.ACT }} + run: | + echo "SCCACHE_GHA_ENABLED=true" >> $GITHUB_ENV + echo "RUSTC_WRAPPER=sccache" >> $GITHUB_ENV + echo "CMAKE_C_COMPILER_LAUNCHER=sccache" >> $GITHUB_ENV + echo "CMAKE_CXX_COMPILER_LAUNCHER=sccache" >> $GITHUB_ENV + + - name: run tests + run: | + cargo test --package metatomic-core --target ${{ matrix.rust-target }} + env: + RUST_BACKTRACE: full + + - name: check that the header was already up to date + run: | + git diff --exit-code + + # check that the C API declarations are correctly documented and used + prevent-bitrot: + runs-on: ubuntu-24.04 + name: check C API declarations + steps: + - uses: actions/checkout@v6 + + - name: setup Python + uses: actions/setup-python@v6 + with: + python-version: "3.14" + + - name: install python dependencies + run: | + pip install pycparser + + - name: check that C API functions are all documented + run: | + python scripts/check-c-api-docs.py + + # make sure no debug print stays in the code + check-debug-prints: + runs-on: ubuntu-24.04 + name: check leftover debug print + + steps: + - uses: actions/checkout@v6 + + - name: install ripgrep + run: | + wget https://github.com/BurntSushi/ripgrep/releases/download/13.0.0/ripgrep-13.0.0-x86_64-unknown-linux-musl.tar.gz + tar xf ripgrep-13.0.0-x86_64-unknown-linux-musl.tar.gz + echo "$(pwd)/ripgrep-13.0.0-x86_64-unknown-linux-musl" >> $GITHUB_PATH + + - name: check for leftover dbg! + run: | + # use ripgrep (rg) to check for instances of `dbg!` in rust files. + # rg will return 1 if it fails to find a match, so we invert it again + # with the `!` builtin to get the error/success in CI + + ! rg "dbg!" --type=rust --quiet + + - name: check for leftover \#include + run: | + ! rg "" --iglob "\!metatomic-core/tests/cpp/external/catch/catch.hpp" --quiet + + - name: check for leftover std::cout + run: | + ! rg "cout" --iglob "\!metatomic-core/tests/cpp/external/catch/catch.hpp" --quiet + + - name: check for leftover std::cerr + run: | + ! rg "cerr" --iglob "\!metatomic-core/tests/cpp/external/catch/catch.hpp" --quiet diff --git a/.github/workflows/torch-tests.yml b/.github/workflows/torch-tests.yml index 93661740a..af066c42b 100644 --- a/.github/workflows/torch-tests.yml +++ b/.github/workflows/torch-tests.yml @@ -20,7 +20,10 @@ jobs: include: - os: ubuntu-24.04 torch-version: "2.13" - python-version: "3.14" + # Python 3.14.5 fails with "No module named pip.__main__; 'pip' is a + # package and cannot be directly executed" when using a venv, so we + # use 3.14.4 for now + python-version: "3.14.4" cargo-test-flags: --release do-valgrind: true @@ -33,12 +36,12 @@ jobs: - os: macos-15 torch-version: "2.13" - python-version: "3.14" + python-version: "3.14.4" cargo-test-flags: --release - os: windows-2022 torch-version: "2.13" - python-version: "3.14" + python-version: "3.14.4" cargo-test-flags: --release steps: - name: install dependencies in container diff --git a/docs/Doxyfile b/docs/Doxyfile index f48f15ed9..5cf71fe6f 100644 --- a/docs/Doxyfile +++ b/docs/Doxyfile @@ -991,7 +991,9 @@ WARN_LOGFILE = # spaces. See also FILE_PATTERNS and EXTENSION_MAPPING # Note: If this tag is empty the current directory is searched. -INPUT = ../metatomic-torch/include/metatomic \ +INPUT = ../metatomic-core/include/ \ + ../metatomic-core/include/metatomic \ + ../metatomic-torch/include/metatomic \ ../metatomic-torch/include/metatomic/torch # This tag can be used to specify the character encoding of the source files diff --git a/docs/src/core/CHANGELOG.md b/docs/src/core/CHANGELOG.md new file mode 120000 index 000000000..a344bc46b --- /dev/null +++ b/docs/src/core/CHANGELOG.md @@ -0,0 +1 @@ +../../../metatomic-core/CHANGELOG.md \ No newline at end of file diff --git a/docs/src/core/index.rst b/docs/src/core/index.rst new file mode 100644 index 000000000..60512b353 --- /dev/null +++ b/docs/src/core/index.rst @@ -0,0 +1,17 @@ +Core Classes +============ + +WIP + + +.. toctree:: + :maxdepth: 2 + + reference/c/index + + +.. toctree:: + :maxdepth: 1 + :hidden: + + CHANGELOG.md diff --git a/docs/src/core/reference/c/index.rst b/docs/src/core/reference/c/index.rst new file mode 100644 index 000000000..f190a5e74 --- /dev/null +++ b/docs/src/core/reference/c/index.rst @@ -0,0 +1,17 @@ +.. _c-api-core: + +C API reference +=============== + +WIP + +The functions and types provided in ``metatomic.h`` can be grouped in four +main groups: + +.. toctree:: + :maxdepth: 1 + + system + model + plugin + misc diff --git a/docs/src/core/reference/c/misc.rst b/docs/src/core/reference/c/misc.rst new file mode 100644 index 000000000..6aec886bc --- /dev/null +++ b/docs/src/core/reference/c/misc.rst @@ -0,0 +1,56 @@ +Miscellaneous +============= + +Version number +^^^^^^^^^^^^^^ + +.. doxygenfunction:: mta_version + +.. c:macro:: METATOMIC_VERSION + + Macro containing the compile-time version of metatomic, as a string + +.. c:macro:: METATOMIC_VERSION_MAJOR + + Macro containing the compile-time **major** version number of metatomic, as + an integer + +.. c:macro:: METATOMIC_VERSION_MINOR + + Macro containing the compile-time **minor** version number of metatomic, as + an integer + +.. c:macro:: METATOMIC_VERSION_PATCH + + Macro containing the compile-time **patch** version number of metatomic, as + an integer + + +Error handling +^^^^^^^^^^^^^^ + +.. doxygenfunction:: mta_last_error + +.. doxygenfunction:: mta_set_last_error + +.. doxygenenum:: mta_status_t + + +String manipulation +^^^^^^^^^^^^^^^^^^^ + +.. doxygentypedef:: mta_string_t + +.. doxygenfunction:: mta_string_create + +.. doxygenfunction:: mta_string_free + +.. doxygenfunction:: mta_string_view + +.. doxygenfunction:: mta_format_metadata + + +Unit conversion +^^^^^^^^^^^^^^^ + +.. doxygenfunction:: mta_unit_conversion_factor diff --git a/docs/src/core/reference/c/model.rst b/docs/src/core/reference/c/model.rst new file mode 100644 index 000000000..6a3d9ee38 --- /dev/null +++ b/docs/src/core/reference/c/model.rst @@ -0,0 +1,16 @@ +Model +===== + +.. doxygenstruct:: mta_model_t + :members: + +The following functions operate on :c:type:`mta_model_t`: + +- :c:func:`mta_load_model`: TODO summary +- :c:func:`mta_execute_model`: TODO summary + +-------------------------------------------------------------------------------- + +.. doxygenfunction:: mta_load_model + +.. doxygenfunction:: mta_execute_model diff --git a/docs/src/core/reference/c/plugin.rst b/docs/src/core/reference/c/plugin.rst new file mode 100644 index 000000000..952650f4c --- /dev/null +++ b/docs/src/core/reference/c/plugin.rst @@ -0,0 +1,16 @@ +Plugin system +============= + +.. doxygenstruct:: mta_plugin_t + :members: + +The following functions operate on :c:type:`mta_plugin_t`: + +- :c:func:`mta_register_plugin`: TODO summary +- :c:func:`mta_load_plugin`: TODO summary + +-------------------------------------------------------------------------------- + +.. doxygenfunction:: mta_register_plugin + +.. doxygenfunction:: mta_load_plugin diff --git a/docs/src/core/reference/c/system.rst b/docs/src/core/reference/c/system.rst new file mode 100644 index 000000000..155245253 --- /dev/null +++ b/docs/src/core/reference/c/system.rst @@ -0,0 +1,42 @@ +System +====== + +.. doxygentypedef:: mta_system_t + +The following functions operate on :c:type:`mta_system_t`: + +- :c:func:`mta_system_create`: TODO summary +- :c:func:`mta_system_free`: TODO summary +- :c:func:`mta_system_size`: TODO summary +- :c:func:`mta_system_get_data`: TODO summary +- :c:func:`mta_system_get_length_unit`: TODO summary +- :c:func:`mta_system_add_pairs`: TODO summary +- :c:func:`mta_system_get_pairs`: TODO summary +- :c:func:`mta_system_known_pairs`: TODO summary +- :c:func:`mta_system_add_custom_data`: TODO summary +- :c:func:`mta_system_get_custom_data`: TODO summary +- :c:func:`mta_system_known_custom_data`: TODO summary + +-------------------------------------------------------------------------------- + +.. doxygenfunction:: mta_system_create + +.. doxygenfunction:: mta_system_free + +.. doxygenfunction:: mta_system_size + +.. doxygenfunction:: mta_system_get_data + +.. doxygenfunction:: mta_system_get_length_unit + +.. doxygenfunction:: mta_system_add_pairs + +.. doxygenfunction:: mta_system_get_pairs + +.. doxygenfunction:: mta_system_known_pairs + +.. doxygenfunction:: mta_system_add_custom_data + +.. doxygenfunction:: mta_system_get_custom_data + +.. doxygenfunction:: mta_system_known_custom_data diff --git a/docs/src/index.rst b/docs/src/index.rst index 170c25c19..441356c29 100644 --- a/docs/src/index.rst +++ b/docs/src/index.rst @@ -92,6 +92,7 @@ existing trained models, look into the metatrain_ project instead. overview installation + core/index torch/index quantities/index engines/index diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index 2a32c1c09..2335505a3 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -14,12 +14,23 @@ name = "metatomic" bench = false [dependencies] +metatensor = { version = "0.3.0" } once_cell = "1" +dlpk = "0.3" +json = "0.12" [build-dependencies] cbindgen = { version = "0.29", default-features = false } +# the last versions that supports Rust 1.74 +serde_spanned = "=1.0.1" +toml = "=0.9.6" +toml_datetime = "=0.7.1" +toml_parser = "=1.0.2" +toml_writer = "=1.0.2" +tempfile = "=3.24.0" +indexmap = "=2.11.4" [dev-dependencies] lazy_static = "1" diff --git a/metatomic-core/build.rs b/metatomic-core/build.rs index edec71e60..b92cc2925 100644 --- a/metatomic-core/build.rs +++ b/metatomic-core/build.rs @@ -22,6 +22,7 @@ fn main() { config.documentation_style = cbindgen::DocumentationStyle::Doxy; config.line_endings = cbindgen::LineEndingStyle::LF; config.autogen_warning = Some(generated_comment.into()); + config.includes.push("metatensor.h".into()); config.includes.push("metatomic/version.h".into()); let result = cbindgen::Builder::new() diff --git a/metatomic-core/cmake/metatomic-config.in.cmake b/metatomic-core/cmake/metatomic-config.in.cmake index 310f54364..90fca167a 100644 --- a/metatomic-core/cmake/metatomic-config.in.cmake +++ b/metatomic-core/cmake/metatomic-config.in.cmake @@ -46,6 +46,7 @@ if (@METATOMIC_INSTALL_BOTH_STATIC_SHARED@ OR @BUILD_SHARED_LIBS@) ) target_compile_features(metatomic::shared INTERFACE cxx_std_17) + target_link_libraries(metatomic::shared INTERFACE metatensor) if (WIN32) if (NOT EXISTS ${METATOMIC_IMPLIB_LOCATION}) @@ -74,6 +75,7 @@ if (@METATOMIC_INSTALL_BOTH_STATIC_SHARED@ OR NOT @BUILD_SHARED_LIBS@) ) target_compile_features(metatomic::static INTERFACE cxx_std_17) + target_link_libraries(metatomic::static INTERFACE metatensor) endif() # Export either the shared or static library as the metatomic target diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index 83b6cb1fc..1e69263e1 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -12,12 +12,118 @@ #include #include #include +#include "metatensor.h" #include "metatomic/version.h" +/** + * TODO + */ +#define MTA_ABI_VERSION 1 + +typedef enum mta_status_t { + MTA_SUCCESS = 0, + MTA_ERROR_OTHER = 255, +} mta_status_t; + +/** + * TODO + */ +typedef enum mta_system_data_kind { + MTA_SYSTEM_DATA_TYPES = 0, + MTA_SYSTEM_DATA_POSITIONS = 1, + MTA_SYSTEM_DATA_CELL = 2, + MTA_SYSTEM_DATA_PBC = 3, +} mta_system_data_kind; + +/** + * TODO + */ +typedef struct mta_opaque_string_t mta_opaque_string_t; + +/** + * TODO + */ +typedef struct mta_system_t mta_system_t; + +/** + * TODO + */ +typedef struct mta_opaque_string_t *mta_string_t; + +/** + * TODO + */ +typedef struct mta_model_t { + /** + * TODO + */ + void *data; + /** + * TODO + */ + enum mta_status_t (*unload)(void *model_data); + /** + * TODO + */ + enum mta_status_t (*metadata)(const void *model_data, mta_string_t *metadata_json); + /** + * TODO + */ + enum mta_status_t (*supported_outputs)(const void *model_data, mta_string_t *outputs_json); + /** + * TODO + */ + enum mta_status_t (*requested_pair_lists)(const void *model_data, mta_string_t *pair_options_json); + /** + * TODO + */ + enum mta_status_t (*requested_inputs)(const void *model_data, mta_string_t *inputs_json); + /** + * TODO + */ + enum mta_status_t (*execute_inner)(void *model_data, + const struct mta_system_t *const *systems, + uintptr_t systems_count, + const mts_labels_t *selected_atoms, + const char *const *requested_outputs_json, + uintptr_t requested_outputs_count, + mts_tensormap_t **outputs, + uintptr_t outputs_count); +} mta_model_t; + +/** + * TODO + */ +typedef struct mta_plugin_t { + /** + * TODO + */ + const char *name; + /** + * TODO + */ + enum mta_status_t (*load_model)(const char *load_from, + const char *options_json, + struct mta_model_t *model); +} mta_plugin_t; + #ifdef __cplusplus extern "C" { #endif // __cplusplus +/** + * TODO + */ +enum mta_status_t mta_last_error(const char **message, const char **origin, void **data); + +/** + * TODO + */ +enum mta_status_t mta_set_last_error(const char *message, + const char *origin, + void *data, + void (*data_deleter)(void*)); + /** * Get the runtime version of the metatomic library as a string. * @@ -25,6 +131,137 @@ extern "C" { */ const char *mta_version(void); +/** + * TODO + */ +mta_string_t mta_string_create(const char *raw); + +/** + * TODO + */ +void mta_string_free(mta_string_t string); + +/** + * TODO + */ +const char *mta_string_view(mta_string_t string); + +/** + * TODO + */ +enum mta_status_t mta_unit_conversion_factor(const char *from_unit, + const char *to_unit, + double *conversion); + +/** + * TODO + */ +enum mta_status_t mta_system_create(const char *length_unit, + DLManagedTensorVersioned *types, + DLManagedTensorVersioned *positions, + DLManagedTensorVersioned *cell, + DLManagedTensorVersioned *pbc, + struct mta_system_t **system); + +/** + * TODO + */ +enum mta_status_t mta_system_free(struct mta_system_t *system); + +/** + * TODO + */ +enum mta_status_t mta_system_size(const struct mta_system_t *system, uintptr_t *size); + +/** + * TODO + */ +enum mta_status_t mta_system_get_data(const struct mta_system_t *system, + enum mta_system_data_kind request, + DLManagedTensorVersioned **data); + +/** + * TODO + */ +enum mta_status_t mta_system_get_length_unit(const struct mta_system_t *system, + mta_string_t *length_unit); + +/** + * TODO + */ +enum mta_status_t mta_system_add_pairs(struct mta_system_t *system, + const char *options, + mts_block_t *pairs); + +/** + * TODO + */ +enum mta_status_t mta_system_get_pairs(const struct mta_system_t *system, + const char *options, + const mts_block_t **pairs); + +/** + * TODO + */ +enum mta_status_t mta_system_known_pairs(const struct mta_system_t *system, + mta_string_t *pairs_options); + +/** + * TODO + */ +enum mta_status_t mta_system_add_custom_data(struct mta_system_t *system, + const char *name, + mts_tensormap_t *data); + +/** + * TODO + */ +enum mta_status_t mta_system_get_custom_data(const struct mta_system_t *system, + const char *name, + const mts_tensormap_t **data); + +/** + * TODO + */ +enum mta_status_t mta_system_known_custom_data(const struct mta_system_t *system, + mta_string_t *names); + +/** + * TODO + */ +enum mta_status_t mta_execute_model(struct mta_model_t model, + const struct mta_system_t *const *systems, + uintptr_t systems_count, + const mts_labels_t *selected_atoms, + const char *const *requested_outputs_json, + uintptr_t requested_outputs_count, + bool check_consistency, + mts_tensormap_t **outputs, + uintptr_t outputs_count); + +/** + * TODO + */ +enum mta_status_t mta_format_metadata(const char *metadata, mta_string_t *printed); + +/** + * TODO + */ +void mta_register_plugin(struct mta_plugin_t plugin); + +/** + * TODO + */ +enum mta_status_t mta_load_plugin(const char *path); + +/** + * TODO + */ +enum mta_status_t mta_load_model(const char *plugin_name, + const char *load_from, + const char *options_json, + struct mta_model_t *model); + #ifdef __cplusplus } // extern "C" #endif // __cplusplus diff --git a/metatomic-core/include/metatomic.hpp b/metatomic-core/include/metatomic.hpp index 016f26bc5..3b5c8ac2a 100644 --- a/metatomic-core/include/metatomic.hpp +++ b/metatomic-core/include/metatomic.hpp @@ -1,2 +1,4 @@ +#include "metatomic/utils.hpp" // IWYU pragma: export #include "metatomic/system.hpp" // IWYU pragma: export -#include "metatomic/model.hpp" // IWYU pragma: export +#include "metatomic/model.hpp" // IWYU pragma: export +#include "metatomic/plugin.hpp" // IWYU pragma: export diff --git a/metatomic-core/include/metatomic/plugin.hpp b/metatomic-core/include/metatomic/plugin.hpp new file mode 100644 index 000000000..1cae91bdf --- /dev/null +++ b/metatomic-core/include/metatomic/plugin.hpp @@ -0,0 +1,7 @@ +#pragma once + +#include + +namespace metatomic { + +} // namespace metatomic diff --git a/metatomic-core/include/metatomic/utils.hpp b/metatomic-core/include/metatomic/utils.hpp new file mode 100644 index 000000000..1cae91bdf --- /dev/null +++ b/metatomic-core/include/metatomic/utils.hpp @@ -0,0 +1,7 @@ +#pragma once + +#include + +namespace metatomic { + +} // namespace metatomic diff --git a/metatomic-core/src/c_api/mod.rs b/metatomic-core/src/c_api/mod.rs index 33e0786dc..cf6c6176d 100644 --- a/metatomic-core/src/c_api/mod.rs +++ b/metatomic-core/src/c_api/mod.rs @@ -1,18 +1,15 @@ -use std::ffi::CString; -use std::os::raw::c_char; +mod status; +pub use self::status::mta_status_t; -use once_cell::sync::Lazy; +mod utils; +pub use self::utils::mta_string_t; +pub use self::utils::{mta_string_create, mta_string_free, mta_string_view}; +mod system; +pub use self::system::mta_system_t; -static VERSION: Lazy = Lazy::new(|| { - CString::new(env!("METATOMIC_FULL_VERSION")).expect("version contains NULL byte") -}); +mod model; +pub use self::model::mta_model_t; - -/// Get the runtime version of the metatomic library as a string. -/// -/// This version follows the `..[-]` format. -#[no_mangle] -pub extern "C" fn mta_version() -> *const c_char { - return VERSION.as_ptr(); -} +mod plugin; +pub use self::plugin::{mta_plugin_t, mta_register_plugin, mta_load_model}; diff --git a/metatomic-core/src/c_api/model.rs b/metatomic-core/src/c_api/model.rs new file mode 100644 index 000000000..b586f73c8 --- /dev/null +++ b/metatomic-core/src/c_api/model.rs @@ -0,0 +1,76 @@ +use std::ffi::{c_void, c_char}; +use metatensor::c_api::{mts_labels_t, mts_tensormap_t}; + +use super::{mta_status_t, mta_string_t, mta_system_t}; + +/// TODO +#[repr(C)] +#[allow(non_camel_case_types)] +pub struct mta_model_t { + /// TODO + pub data: *mut c_void, + + /// TODO + pub unload: Option mta_status_t>, + + /// TODO + pub metadata: Option mta_status_t>, + + /// TODO + pub supported_outputs: Option mta_status_t>, + + /// TODO + pub requested_pair_lists: Option mta_status_t>, + + /// TODO + pub requested_inputs: Option mta_status_t>, + + /// TODO + pub execute_inner: Option mta_status_t>, +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_execute_model( + model: mta_model_t, + systems: *const *const mta_system_t, + systems_count: usize, + selected_atoms: *const mts_labels_t, + requested_outputs_json: *const *const c_char, + requested_outputs_count: usize, + check_consistency: bool, + outputs: *mut *mut mts_tensormap_t, + outputs_count: usize, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_format_metadata( + metadata: *const c_char, + printed: *mut mta_string_t, +) -> mta_status_t { + todo!() +} diff --git a/metatomic-core/src/c_api/plugin.rs b/metatomic-core/src/c_api/plugin.rs new file mode 100644 index 000000000..6dbfc4add --- /dev/null +++ b/metatomic-core/src/c_api/plugin.rs @@ -0,0 +1,41 @@ +use std::ffi::c_char; + +use super::{mta_model_t, mta_status_t}; + +/// TODO +#[allow(non_camel_case_types)] +#[repr(C)] +pub struct mta_plugin_t { + /// TODO + pub name: *const c_char, + + /// TODO + pub load_model: Option mta_status_t>, +} + +/// TODO +#[no_mangle] +pub extern "C" fn mta_register_plugin(plugin: mta_plugin_t) { + todo!() +} + +/// TODO +#[no_mangle] +pub extern "C" fn mta_load_plugin(path: *const c_char) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub extern "C" fn mta_load_model( + plugin_name: *const c_char, + load_from: *const c_char, + options_json: *const c_char, + model: *mut mta_model_t, +) -> mta_status_t { + todo!() +} diff --git a/metatomic-core/src/c_api/status.rs b/metatomic-core/src/c_api/status.rs new file mode 100644 index 000000000..0c48707cc --- /dev/null +++ b/metatomic-core/src/c_api/status.rs @@ -0,0 +1,42 @@ +use std::ffi::{c_char, c_void}; + +use crate::Error; + + +// TODO +#[allow(non_camel_case_types)] +#[repr(C)] +#[derive(PartialEq, Eq, Debug)] +pub enum mta_status_t { + MTA_SUCCESS = 0, + // ... + MTA_ERROR_OTHER = 255, +} + + +impl From for mta_status_t { + fn from(err: Error) -> Self { + todo!() + } +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_last_error( + message: *mut *const c_char, + origin: *mut *const c_char, + data: *mut *mut c_void, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_set_last_error( + message: *const c_char, + origin: *const c_char, + data: *mut c_void, + data_deleter: Option, +) -> mta_status_t { + todo!() +} diff --git a/metatomic-core/src/c_api/system.rs b/metatomic-core/src/c_api/system.rs new file mode 100644 index 000000000..1697e5155 --- /dev/null +++ b/metatomic-core/src/c_api/system.rs @@ -0,0 +1,131 @@ +use std::ffi::c_char; + +use dlpk::sys::DLManagedTensorVersioned; +use metatensor::c_api::{mts_block_t, mts_tensormap_t}; + +use crate::System; +use super::{mta_status_t, mta_string_t}; + +/// TODO +#[allow(non_camel_case_types)] +pub struct mta_system_t(pub(crate) System); + + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_create( + length_unit: *const c_char, + types: *mut DLManagedTensorVersioned, + positions: *mut DLManagedTensorVersioned, + cell: *mut DLManagedTensorVersioned, + pbc: *mut DLManagedTensorVersioned, + system: *mut *mut mta_system_t, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_free(system: *mut mta_system_t) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_size( + system: *const mta_system_t, + size: *mut usize, +) -> mta_status_t { + todo!() +} + +/// TODO +#[allow(non_camel_case_types)] +#[repr(C)] +#[non_exhaustive] +pub enum mta_system_data_kind { + MTA_SYSTEM_DATA_TYPES = 0, + MTA_SYSTEM_DATA_POSITIONS = 1, + MTA_SYSTEM_DATA_CELL = 2, + MTA_SYSTEM_DATA_PBC = 3, +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_get_data( + system: *const mta_system_t, + request: mta_system_data_kind, + data: *mut *mut DLManagedTensorVersioned, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_get_length_unit( + system: *const mta_system_t, + length_unit: *mut mta_string_t, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_add_pairs( + system: *mut mta_system_t, + options: *const c_char, + pairs: *mut mts_block_t, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_get_pairs( + system: *const mta_system_t, + options: *const c_char, + pairs: *mut *const mts_block_t, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_known_pairs( + system: *const mta_system_t, + pairs_options: *mut mta_string_t, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_add_custom_data( + system: *mut mta_system_t, + name: *const c_char, + data: *mut mts_tensormap_t, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_get_custom_data( + system: *const mta_system_t, + name: *const c_char, + data: *mut *const mts_tensormap_t, +) -> mta_status_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_system_known_custom_data( + system: *const mta_system_t, + names: *mut mta_string_t, +) -> mta_status_t { + todo!() +} + + +// TODO: mta_system_to(device, dtype) diff --git a/metatomic-core/src/c_api/utils.rs b/metatomic-core/src/c_api/utils.rs new file mode 100644 index 000000000..350a9552d --- /dev/null +++ b/metatomic-core/src/c_api/utils.rs @@ -0,0 +1,102 @@ +use std::ffi::{CString, c_char}; + +use once_cell::sync::Lazy; + +use super::mta_status_t; + + +static VERSION: Lazy = Lazy::new(|| { + CString::new(env!("METATOMIC_FULL_VERSION")).expect("version contains NULL byte") +}); + + +/// Get the runtime version of the metatomic library as a string. +/// +/// This version follows the `..[-]` format. +#[no_mangle] +pub extern "C" fn mta_version() -> *const c_char { + return VERSION.as_ptr(); +} + +/// TODO +#[allow(non_camel_case_types)] +pub struct mta_opaque_string_t(CString); + +/// TODO +#[allow(non_camel_case_types)] +#[repr(transparent)] +pub struct mta_string_t(*mut mta_opaque_string_t); + +impl std::fmt::Debug for mta_string_t { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let mut builder = f.debug_tuple("mta_string_t"); + + if self.0.is_null() { + builder.field(&"NULL"); + } else { + builder.field(&self.as_str()); + } + builder.finish() + } +} + +impl mta_string_t { + /// TODO + pub fn new(value: impl Into) -> Self { + let cstring = CString::new(value.into()).unwrap(); + let boxed = Box::new(mta_opaque_string_t(cstring)); + mta_string_t(Box::into_raw(boxed)) + } + + /// TODO + pub fn null() -> Self { + mta_string_t(std::ptr::null_mut()) + } + + /// TODO + pub fn as_str(&self) -> &str { + if self.0.is_null() { + return ""; + } + unsafe { + return (*(self.0)).0.to_str().expect("mta_string_t is not valid UTF8") + } + } +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_string_create( + raw: *const c_char, +) -> mta_string_t { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_string_free(string: mta_string_t) { + todo!() +} + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_string_view( + string: mta_string_t, +) -> *const c_char { + todo!() +} + + +/// TODO +#[no_mangle] +pub unsafe extern "C" fn mta_unit_conversion_factor( + from_unit: *const c_char, + to_unit: *const c_char, + conversion: *mut f64, +) -> mta_status_t { + todo!() +} + + + +// TODO: logging & warnings? diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index bc47948b4..8e4c828b7 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -9,5 +9,46 @@ #![allow(clippy::let_underscore_untyped, clippy::manual_let_else, clippy::empty_line_after_doc_comments)] +// To be removed lated +#![allow(unused_variables, dead_code, clippy::needless_pass_by_value)] + + #[doc(hidden)] -mod c_api; +pub mod c_api; + +mod metadata; +pub use self::metadata::{ModelMetadata, Quantity, PairListOptions}; + +mod system; +pub use self::system::System; + +mod model; +pub use self::model::Model; + +mod plugin; +pub use self::plugin::{Plugin, load_plugin, load_model}; + +mod units; +pub use self::units::unit_conversion_factor; + +/// TODO +#[derive(Debug)] +pub enum Error { + // TODO +} + +impl std::fmt::Display for Error { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + todo!() + } +} + +impl std::error::Error for Error { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + todo!() + } + + fn cause(&self) -> Option<&dyn std::error::Error> { + self.source() + } +} diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs new file mode 100644 index 000000000..d56a4674a --- /dev/null +++ b/metatomic-core/src/metadata.rs @@ -0,0 +1,132 @@ +use json::JsonValue; + +use crate::Error; + +/// TODO +pub struct PairListOptions { + /// TODO + cutoff: f64, + /// TODO + full_list: bool, + /// TODO + strict: bool, + /// TODO + requestors: Vec, +} + +impl std::cmp::PartialEq for PairListOptions { + fn eq(&self, other: &Self) -> bool { + self.cutoff == other.cutoff + && self.full_list == other.full_list + && self.strict == other.strict + } +} + +impl std::cmp::Eq for PairListOptions {} + +impl std::cmp::PartialOrd for PairListOptions { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl std::cmp::Ord for PairListOptions { + fn cmp(&self, other: &Self) -> std::cmp::Ordering { + self.cutoff.partial_cmp(&other.cutoff).expect("cutoff is NaN") + .then_with(|| self.full_list.cmp(&other.full_list)) + .then_with(|| self.strict.cmp(&other.strict)) + } +} + +// TODO +// { +// "type": "metatomic_pair_options", +// "cutoff": "0xaeabf23", <== hex of the int corresponding to the f64 bits to keep full precision +// "full_list": false, +// "strict": false, +// "requestors": ["..."] +// } +impl From for JsonValue { + fn from(value: PairListOptions) -> Self { + todo!() + } +} + +impl TryFrom for PairListOptions { + type Error = Error; + + fn try_from(value: JsonValue) -> Result { + todo!() + } +} + +// ========================================================================== // +// ========================================================================== // +// ========================================================================== // + +/// TODO +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct ModelMetadata { + pub name: String, + // TODO +} + +// { +// "type": "metatomic_model_metadata", +// "name": "...", +// "authors": ["..."], +// "references": { +// "implementation": ["..."], +// "architecture": ["..."], +// "model": ["..."] +// }, +// "extra": { +// "key...": "value..." +// } +// }, +impl From for JsonValue { + fn from(value: ModelMetadata) -> Self { + todo!() + } +} + +impl TryFrom for ModelMetadata { + type Error = Error; + + fn try_from(value: JsonValue) -> Result { + todo!() + } +} + +// ========================================================================== // +// ========================================================================== // +// ========================================================================== // + +/// TODO, previously `ModelOutput` +#[derive(Debug)] +pub struct Quantity { + pub name: String, + // TODO +} + +// TODO: +// { +// "type": "metatomic_quantity", +// "name": "...", +// "unit": "...", +// "gradients": ["...", "..."], +// "sample_kind": "atom" | "system" | "atom-pair", +// }, +impl From for JsonValue { + fn from(value: Quantity) -> Self { + todo!() + } +} + +impl TryFrom for Quantity { + type Error = Error; + + fn try_from(value: JsonValue) -> Result { + todo!() + } +} diff --git a/metatomic-core/src/model.rs b/metatomic-core/src/model.rs new file mode 100644 index 000000000..16b208ac1 --- /dev/null +++ b/metatomic-core/src/model.rs @@ -0,0 +1,20 @@ +use metatensor::{Labels, TensorMap}; + +use crate::{Error, Quantity, System}; + +use crate::c_api::mta_model_t; + +/// TODO +pub struct Model(pub(crate) mta_model_t); + + +/// TODO +pub fn execute_model( + model: &Model, + systems: &[System], + selected_atoms: Option, + requested_outputs: &[Quantity], + check_consistency: bool, +) -> Result, Error> { + todo!() +} diff --git a/metatomic-core/src/plugin.rs b/metatomic-core/src/plugin.rs new file mode 100644 index 000000000..60d145803 --- /dev/null +++ b/metatomic-core/src/plugin.rs @@ -0,0 +1,37 @@ +use std::collections::BTreeMap; + +use crate::c_api::mta_plugin_t; +use crate::{Error, Model}; + +/// TODO +pub const MTA_ABI_VERSION: i32 = 1; + +/// TODO +pub struct Plugin(mta_plugin_t); + +impl Plugin { + /// TODO + pub fn new(c_plugin: mta_plugin_t) -> Self { + Self(c_plugin) + } + + /// TODO + pub fn name(&self) -> &str { + todo!() + } + + /// TODO + pub fn load_model(&self, load_from: &str, options: BTreeMap) -> Result { + todo!() + } +} + +/// TODO +pub fn load_plugin(path: &str) -> Result<(), Error> { + todo!() +} + +/// TODO +pub fn load_model(plugin: Option<&str>, load_from: &str, options: BTreeMap) -> Result { + todo!() +} diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs new file mode 100644 index 000000000..30677f5f9 --- /dev/null +++ b/metatomic-core/src/system.rs @@ -0,0 +1,53 @@ +use std::collections::{BTreeMap, HashMap}; + +use dlpk::DLPackTensor; +use metatensor::{TensorBlock, TensorMap}; + +use crate::PairListOptions; + + +/// TODO +pub struct System { + length_unit: String, + types: DLPackTensor, + positions: DLPackTensor, + cell: DLPackTensor, + pbc: DLPackTensor, + + pairs: BTreeMap, + custom_data: HashMap, +} + + +impl System { + /// TODO + pub fn new( + length_unit: String, + types: DLPackTensor, + positions: DLPackTensor, + cell: DLPackTensor, + pbc: DLPackTensor + ) -> Self { + todo!() + } + + /// TODO + pub fn add_pairs(&mut self, options: PairListOptions, pairs: TensorBlock, check_consistency: bool) { + todo!() + } + + /// TODO + pub fn get_pairs(&mut self, options: PairListOptions) -> Option<&TensorBlock> { + todo!() + } + + /// TODO + pub fn set_custom_data(&mut self, name: String, data: TensorMap) { + todo!() + } + + /// TODO + pub fn get_custom_data(&self, name: &str) -> Option<&TensorMap> { + todo!() + } +} diff --git a/metatomic-core/src/units.rs b/metatomic-core/src/units.rs new file mode 100644 index 000000000..d06eab413 --- /dev/null +++ b/metatomic-core/src/units.rs @@ -0,0 +1,7 @@ +use crate::Error; + + +/// TODO +pub fn unit_conversion_factor(from_unit: &str, to_unit: &str) -> Result { + todo!() +} diff --git a/metatomic-core/tests/check-cxx-install.rs b/metatomic-core/tests/check-cxx-install.rs index d66f4883b..6baa5b4e1 100644 --- a/metatomic-core/tests/check-cxx-install.rs +++ b/metatomic-core/tests/check-cxx-install.rs @@ -23,19 +23,21 @@ fn check_cxx_install() { const CARGO_TARGET_TMPDIR: &str = env!("CARGO_TARGET_TMPDIR"); - // ====================================================================== // - // build and install metatensor with cmake let mut build_dir = PathBuf::from(CARGO_TARGET_TMPDIR); build_dir.push("cxx-install"); build_dir.push("cmake-find-package"); std::fs::create_dir_all(&build_dir).expect("failed to create build dir"); + // ====================================================================== // + // install dependencies with pip let deps_dir = build_dir.join("deps"); let virtualenv_dir = deps_dir.join("virtualenv"); std::fs::create_dir_all(&virtualenv_dir).expect("failed to create virtualenv dir"); let python_exe = utils::create_python_venv(virtualenv_dir); let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python_exe); + // ====================================================================== // + // build and install metatomic with cmake let metatomic_dep = deps_dir.join("metatomic-core"); let source_dir = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); @@ -54,7 +56,7 @@ fn check_cxx_install() { cmake_config.arg(format!("-DCMAKE_PREFIX_PATH={};{}", metatensor_cmake_prefix.display(), metatomic_cmake_prefix.display())); utils::run_command(cmake_config, "cmake configuration"); - // build the code, linking to metatensor + // build the code, linking to metatomic let cmake_build = utils::cmake_build(&build_dir); utils::run_command(cmake_build, "cmake build"); diff --git a/metatomic-torch/tests/check-torch-install.rs b/metatomic-torch/tests/check-torch-install.rs index ad8cfb604..14e85628a 100644 --- a/metatomic-torch/tests/check-torch-install.rs +++ b/metatomic-torch/tests/check-torch-install.rs @@ -24,14 +24,13 @@ fn check_torch_install() { const CARGO_TARGET_TMPDIR: &str = env!("CARGO_TARGET_TMPDIR"); let cargo_manifest_dir = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); - // ====================================================================== // - // build and install metatensor-torch with cmake let mut build_dir = PathBuf::from(CARGO_TARGET_TMPDIR); build_dir.push("torch-install"); build_dir.push("cmake-find-package"); std::fs::create_dir_all(&build_dir).expect("failed to create build dir"); - + // ====================================================================== // + // install dependencies with pip let deps_dir = build_dir.join("deps"); let torch_dep = deps_dir.join("virtualenv"); @@ -41,7 +40,8 @@ fn check_torch_install() { let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python); let metatensor_torch_cmake_prefix = utils::setup_metatensor_torch_pip(&python); - // configure cmake for metatomic-torch + // ====================================================================== // + // build and install metatomic-torch with cmake let metatomic_torch_dep = deps_dir.join("metatomic-torch"); let cmake_options = vec![ @@ -68,7 +68,7 @@ fn check_torch_install() { ); // ====================================================================== // - // // try to use the installed metatomic-torch from cmake + // try to use the installed metatomic-torch from cmake let mut source_dir = PathBuf::from(&cargo_manifest_dir); source_dir.extend(["tests", "cmake-project"]); @@ -93,7 +93,7 @@ fn check_torch_install() { utils::run_command(ctest, "ctest"); } -/// Same as above, but using pre-built metatensor-torch from the Python wheel, +/// Same as above, but using metatomic-torch from the Python wheel, /// instead of building it from source with cmake. #[test] fn check_python_install() { @@ -106,13 +106,13 @@ fn check_python_install() { const CARGO_TARGET_TMPDIR: &str = env!("CARGO_TARGET_TMPDIR"); - // ====================================================================== // - // build and install metatensor and metatensor-torch with pip let mut build_dir = PathBuf::from(CARGO_TARGET_TMPDIR); build_dir.push("torch-install"); build_dir.push("python-wheels"); std::fs::create_dir_all(&build_dir).expect("failed to create build dir"); + // ====================================================================== // + // install dependencies with pip let mut venv_dir = build_dir.clone(); venv_dir.push("virtualenv"); @@ -123,6 +123,8 @@ fn check_python_install() { let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python_exe); let metatensor_torch_cmake_prefix = utils::setup_metatensor_torch_pip(&python_exe); + // ====================================================================== // + // build and install metatomic and metatomic-torch with pip let mta_core_source_dir = cargo_manifest_dir.parent().unwrap().join("python").join("metatomic_core"); let metatomic_core_cmake_prefix = utils::setup_metatomic_core_pip(&python_exe, &mta_core_source_dir); @@ -130,7 +132,7 @@ fn check_python_install() { let metatomic_torch_cmake_prefix = utils::setup_metatomic_torch_pip(&python_exe, &mta_torch_source_dir); // ====================================================================== // - // try to use the installed metatensor-torch from cmake + // try to use the installed metatomic-torch from cmake let mut source_dir = PathBuf::from(&cargo_manifest_dir); source_dir.extend(["tests", "cmake-project"]); @@ -147,7 +149,7 @@ fn check_python_install() { utils::run_command(cmake_config, "cmake configuration"); - // build the code, linking to metatensor-torch + // build the code, linking to metatomic-torch let cmake_build = utils::cmake_build(&build_dir); utils::run_command(cmake_build, "cmake build"); @@ -175,16 +177,19 @@ fn check_cmake_subdirectory() { build_dir.push("cmake-subdirectory"); std::fs::create_dir_all(&build_dir).expect("failed to create build dir"); + // ====================================================================== // + // install dependencies with pip let deps_dir = build_dir.join("deps"); - let torch_dep = deps_dir.join("virtualenv"); - std::fs::create_dir_all(&torch_dep).expect("failed to create virtualenv dir"); - let python = utils::create_python_venv(torch_dep); + let virtualenv_dir = deps_dir.join("virtualenv"); + std::fs::create_dir_all(&virtualenv_dir).expect("failed to create virtualenv dir"); + let python = utils::create_python_venv(virtualenv_dir); let pytorch_cmake_prefix = utils::setup_torch_pip(&python); let metatensor_cmake_prefix = utils::setup_metatensor_pip(&python); let metatensor_torch_cmake_prefix = utils::setup_metatensor_torch_pip(&python); // ====================================================================== // + // build metatomic-torch with cmake, using add_subdirectory let cargo_manifest_dir = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()); let mut source_dir = PathBuf::from(&cargo_manifest_dir); source_dir.extend(["tests", "cmake-project"]); diff --git a/rustfmt.toml b/rustfmt.toml new file mode 100644 index 000000000..c7ad93baf --- /dev/null +++ b/rustfmt.toml @@ -0,0 +1 @@ +disable_all_formatting = true diff --git a/scripts/check-c-api-docs.py b/scripts/check-c-api-docs.py new file mode 100755 index 000000000..73ee7d921 --- /dev/null +++ b/scripts/check-c-api-docs.py @@ -0,0 +1,101 @@ +#!/usr/bin/env python +""" +A small script checking that all the C API functions are documented +""" + +import os +import sys + +from pycparser import c_ast, parse_file + + +ROOT = os.path.realpath(os.path.join(os.path.dirname(__file__), "..")) +C_API_DOCS = os.path.join(ROOT, "docs", "src", "core", "reference", "c") +FAKE_INCLUDES = [os.path.join(ROOT, "scripts", "include")] +METATOMIC_HEADER = os.path.relpath( + os.path.join(ROOT, "metatomic-core", "include", "metatomic.h") +) + + +ERRORS = 0 + + +def error(message): + global ERRORS + ERRORS += 1 + print(message) + + +def documented_functions(): + functions = [] + + for root, _, paths in os.walk(C_API_DOCS): + for path in paths: + with open(os.path.join(root, path), encoding="utf8") as fd: + for line in fd: + if line.startswith(".. doxygenfunction::"): + name = line.split()[2] + functions.append(name) + + return functions + + +def functions_in_outline(): + # function from the "miscellaneous" section of the docs don't require an outline + # (since they are not related to a specific struct type) + functions = [ + "mta_version", + "mta_last_error", + "mta_set_last_error", + "mta_string_create", + "mta_string_free", + "mta_string_view", + "mta_format_metadata", + "mta_unit_conversion_factor", + ] + + for root, _, paths in os.walk(C_API_DOCS): + for path in paths: + with open(os.path.join(root, path), encoding="utf8") as fd: + for line in fd: + if ":c:func:" in line: + name = line.split("`")[1] + functions.append(name) + return functions + + +def all_functions(): + cpp_args = ["-E"] + for path in FAKE_INCLUDES: + cpp_args += ["-I", path] + ast = parse_file(METATOMIC_HEADER, use_cpp=True, cpp_path="gcc", cpp_args=cpp_args) + + functions = [] + + class AstVisitor(c_ast.NodeVisitor): + def visit_Decl(self, node): + if not isinstance(node.type, c_ast.FuncDecl): + return + + if not node.name.startswith("mta_"): + return + + functions.append(node.name) + + visitor = AstVisitor() + visitor.visit(ast) + + return functions + + +if __name__ == "__main__": + docs = documented_functions() + outline = functions_in_outline() + for function in all_functions(): + if function not in docs: + error("Missing documentation for {}".format(function)) + if function not in outline: + error("Missing outline for {}".format(function)) + + if ERRORS != 0: + sys.exit(1) diff --git a/scripts/include/README b/scripts/include/README new file mode 100644 index 000000000..d56dd0788 --- /dev/null +++ b/scripts/include/README @@ -0,0 +1,4 @@ +This directory contains fake headers used to allow pycparser to parse the code +without having to deal with all the complexity of actual stdlib implementations + +See https://eli.thegreenplace.net/2015/on-parsing-c-type-declarations-and-fake-headers for more information diff --git a/scripts/include/metatensor.h b/scripts/include/metatensor.h new file mode 100644 index 000000000..fb8e88f0d --- /dev/null +++ b/scripts/include/metatensor.h @@ -0,0 +1,8 @@ +// empty header with minimal content, to be used to parse metatomic.h + +typedef struct mts_labels_t mts_labels_t; +typedef struct mts_block_t mts_block_t; +typedef struct mts_tensormap_t mts_tensormap_t; + + +typedef struct DLManagedTensorVersioned DLManagedTensorVersioned; diff --git a/scripts/include/metatomic/version.h b/scripts/include/metatomic/version.h new file mode 100644 index 000000000..e69de29bb diff --git a/scripts/include/stdarg.h b/scripts/include/stdarg.h new file mode 100644 index 000000000..e69de29bb diff --git a/scripts/include/stdbool.h b/scripts/include/stdbool.h new file mode 100644 index 000000000..3bd41ef29 --- /dev/null +++ b/scripts/include/stdbool.h @@ -0,0 +1 @@ +typedef _Bool bool; \ No newline at end of file diff --git a/scripts/include/stddef.h b/scripts/include/stddef.h new file mode 100644 index 000000000..48b3db663 --- /dev/null +++ b/scripts/include/stddef.h @@ -0,0 +1,6 @@ +#ifndef FAKE_STDDEF_H +#define FAKE_STDDEF_H + +typedef void nullptr_t; + +#endif /* FAKE_STDDEF_H */ diff --git a/scripts/include/stdint.h b/scripts/include/stdint.h new file mode 100644 index 000000000..43ccc01dd --- /dev/null +++ b/scripts/include/stdint.h @@ -0,0 +1,7 @@ +typedef int uint64_t; +typedef int int64_t; +typedef int int32_t; +typedef int uint32_t; +typedef int uint16_t; +typedef int uint8_t; +typedef int uintptr_t; diff --git a/scripts/include/stdlib.h b/scripts/include/stdlib.h new file mode 100644 index 000000000..8b1378917 --- /dev/null +++ b/scripts/include/stdlib.h @@ -0,0 +1 @@ + From 5bd57ea9a8d7ef4a0562dbbc3ea15d900ee6db30 Mon Sep 17 00:00:00 2001 From: Sofiia Chorna Date: Thu, 28 May 2026 16:58:17 +0200 Subject: [PATCH 06/43] Implement PairListOptions json serialization --- docs/src/core/index.rst | 1 + docs/src/core/reference/json-formats.rst | 52 ++++++ metatomic-core/src/lib.rs | 11 +- metatomic-core/src/metadata.rs | 204 +++++++++++++++++++++-- 4 files changed, 250 insertions(+), 18 deletions(-) create mode 100644 docs/src/core/reference/json-formats.rst diff --git a/docs/src/core/index.rst b/docs/src/core/index.rst index 60512b353..4d0cf24a4 100644 --- a/docs/src/core/index.rst +++ b/docs/src/core/index.rst @@ -8,6 +8,7 @@ WIP :maxdepth: 2 reference/c/index + reference/json-formats .. toctree:: diff --git a/docs/src/core/reference/json-formats.rst b/docs/src/core/reference/json-formats.rst new file mode 100644 index 000000000..d1589fa2d --- /dev/null +++ b/docs/src/core/reference/json-formats.rst @@ -0,0 +1,52 @@ +.. _core-json-formats: + +JSON data formats +================= + +Some metatomic data structures are exchanged across the C API as JSON-encoded +strings rather than dedicated C types. This page documents the exact JSON +representation of each such structure, so that engines and models written in any +language can produce and consume them. + +Pair list options +----------------- + +Options describing a requested pair list (also known as a neighbor list). This +is the JSON representation of ``PairListOptions``, used for example by +:c:func:`mta_system_set_pairs`, :c:func:`mta_system_get_pairs` and +:c:func:`mta_system_pairs_options`. + +.. code-block:: json + + { + "type": "metatomic_pair_options", + "cutoff": "0x400c000000000000", + "full_list": false, + "strict": false, + "requestors": ["my-model"] + } + +``type`` + Must be the string ``"metatomic_pair_options"``. + +``cutoff`` + Cutoff radius for the pair list in the length unit of the model. Must be a + positive finite number. + + It is stored as a string containing the hexadecimal representation of the + 64-bit integer with the same bit pattern as the ``cutoff`` floating-point + value (i.e. reinterpreting the ``double`` as a ``uint64_t``). + +``full_list`` + Boolean. If ``true``, the list is a full list containing both ``i -> j`` + and ``j -> i`` for each pair, if ``false``, it is a half list containing + only ``i -> j``. + +``strict`` + Boolean. If ``true``, the list is guaranteed to contain only atoms within + the cutoff, if ``false``, it may also include some pairs slightly beyond the + cutoff. + +``requestors`` + Optional array of strings identifying who requested this pair list. May be + omitted, in which case it is treated as an empty list. diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index 8e4c828b7..894e70e3a 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -34,18 +34,23 @@ pub use self::units::unit_conversion_factor; /// TODO #[derive(Debug)] pub enum Error { - // TODO + /// Error while serializing data to or deserializing data from JSON + Serialization(String), } impl std::fmt::Display for Error { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - todo!() + match self { + Error::Serialization(message) => write!(f, "{}", message), + } } } impl std::error::Error for Error { fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { - todo!() + match self { + Error::Serialization(_) => None, + } } fn cause(&self) -> Option<&dyn std::error::Error> { diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs index d56a4674a..913edcdfc 100644 --- a/metatomic-core/src/metadata.rs +++ b/metatomic-core/src/metadata.rs @@ -2,15 +2,19 @@ use json::JsonValue; use crate::Error; -/// TODO +/// Options for the calculation of a pair list (neighbor list) +#[derive(Debug, Clone)] pub struct PairListOptions { - /// TODO + /// Cutoff radius for this pair list in the length unit of the model cutoff: f64, - /// TODO + /// Whether the list is a full list (contains both the pair `i -> j` and `j -> i`) + /// or a half list (contains only `i -> j`) full_list: bool, - /// TODO + /// Whether the list guarantees that only atoms within the cutoff are + /// included (strict) or may also include pairs slightly beyond the cutoff + /// (non-strict) strict: bool, - /// TODO + /// List of strings describing who requested this pair list requestors: Vec, } @@ -38,17 +42,16 @@ impl std::cmp::Ord for PairListOptions { } } -// TODO -// { -// "type": "metatomic_pair_options", -// "cutoff": "0xaeabf23", <== hex of the int corresponding to the f64 bits to keep full precision -// "full_list": false, -// "strict": false, -// "requestors": ["..."] -// } impl From for JsonValue { fn from(value: PairListOptions) -> Self { - todo!() + let mut result = JsonValue::new_object(); + result["type"] = "metatomic_pair_options".into(); + // store the bit pattern so the float round-trips exactly + result["cutoff"] = format!("{:#x}", value.cutoff.to_bits()).into(); + result["full_list"] = value.full_list.into(); + result["strict"] = value.strict.into(); + result["requestors"] = value.requestors.into(); + return result; } } @@ -56,7 +59,61 @@ impl TryFrom for PairListOptions { type Error = Error; fn try_from(value: JsonValue) -> Result { - todo!() + if !value.is_object() { + return Err(Error::Serialization( + "invalid JSON data for PairListOptions, expected an object".into() + )); + } + + if value["type"].as_str() != Some("metatomic_pair_options") { + return Err(Error::Serialization( + "'type' in JSON for PairListOptions must be 'metatomic_pair_options'".into() + )); + } + + let cutoff = value["cutoff"].as_str().ok_or_else(|| Error::Serialization( + "'cutoff' in JSON for PairListOptions must be a hex-encoded string".into() + ))?; + let bits = u64::from_str_radix(cutoff.strip_prefix("0x").unwrap_or(cutoff), 16) + .map_err(|_| Error::Serialization( + "'cutoff' in JSON for PairListOptions must be a hex-encoded string".into() + ))?; + let cutoff = f64::from_bits(bits); + + if !cutoff.is_finite() || cutoff <= 0.0 { + return Err(Error::Serialization( + "'cutoff' in JSON for PairListOptions must be a finite positive number".into() + )); + } + + let full_list = value["full_list"].as_bool().ok_or_else(|| Error::Serialization( + "'full_list' in JSON for PairListOptions must be a boolean".into() + ))?; + + let strict = value["strict"].as_bool().ok_or_else(|| Error::Serialization( + "'strict' in JSON for PairListOptions must be a boolean".into() + ))?; + + let mut requestors = Vec::new(); + if value.has_key("requestors") { + if !value["requestors"].is_array() { + return Err(Error::Serialization( + "'requestors' in JSON for PairListOptions must be an array".into() + )); + } + + for requestor in value["requestors"].members() { + let requestor = requestor.as_str().ok_or_else(|| Error::Serialization( + "'requestors' in JSON for PairListOptions must be an array of strings".into() + ))?; + // ignore empty strings and duplicates, keeping first-seen order + if !requestor.is_empty() && !requestors.iter().any(|r| r == requestor) { + requestors.push(requestor.to_string()); + } + } + } + + return Ok(PairListOptions { cutoff, full_list, strict, requestors }); } } @@ -130,3 +187,120 @@ impl TryFrom for Quantity { todo!() } } + + +#[cfg(test)] +mod tests { + mod pair_list_options { + use super::super::*; + + fn example() -> PairListOptions { + PairListOptions { + cutoff: 3.5, + full_list: true, + strict: false, + requestors: vec!["nl-1".to_string(), "nl-2".to_string()], + } + } + + #[test] + fn roundtrip() { + let options = example(); + let json: JsonValue = options.clone().into(); + + assert_eq!(json["type"].as_str(), Some("metatomic_pair_options")); + assert_eq!(json["cutoff"].as_str(), Some(format!("{:#x}", 3.5_f64.to_bits()).as_str())); + assert_eq!(json["full_list"].as_bool(), Some(true)); + assert_eq!(json["strict"].as_bool(), Some(false)); + + let parsed = PairListOptions::try_from(json).unwrap(); + assert_eq!(parsed.cutoff.to_bits(), options.cutoff.to_bits()); + assert_eq!(parsed.full_list, options.full_list); + assert_eq!(parsed.strict, options.strict); + assert_eq!(parsed.requestors, options.requestors); + } + + #[test] + fn cutoff_keeps_full_precision() { + let mut options = example(); + options.cutoff = 1.0 / 3.0; + let parsed = PairListOptions::try_from(JsonValue::from(options.clone())).unwrap(); + assert_eq!(parsed.cutoff.to_bits(), options.cutoff.to_bits()); + } + + #[test] + fn requestors_are_optional() { + let mut json: JsonValue = example().into(); + json.remove("requestors"); + let parsed = PairListOptions::try_from(json).unwrap(); + assert!(parsed.requestors.is_empty()); + } + + #[test] + fn rejects_invalid_json() { + // each case corrupts exactly one field of an otherwise valid object + let with_cutoff = |value: f64| { + let mut json = JsonValue::from(example()); + json["cutoff"] = format!("{:#x}", value.to_bits()).into(); + json + }; + + let mut wrong_type = JsonValue::from(example()); + wrong_type["type"] = "something-else".into(); + + let mut missing_cutoff = JsonValue::from(example()); + missing_cutoff.remove("cutoff"); + + let mut non_hex_cutoff = JsonValue::from(example()); + non_hex_cutoff["cutoff"] = "not-hex".into(); + + let mut non_boolean_flag = JsonValue::from(example()); + non_boolean_flag["full_list"] = "yes".into(); + + let mut non_array_requestors = JsonValue::from(example()); + non_array_requestors["requestors"] = "nl-1".into(); + + let mut non_string_requestor = JsonValue::from(example()); + non_string_requestor["requestors"] = json::array![ "nl-1", 42 ]; + + let cases = [ + (JsonValue::from("not an object"), + "invalid JSON data for PairListOptions, expected an object"), + (wrong_type, + "'type' in JSON for PairListOptions must be 'metatomic_pair_options'"), + (missing_cutoff, + "'cutoff' in JSON for PairListOptions must be a hex-encoded string"), + (non_hex_cutoff, + "'cutoff' in JSON for PairListOptions must be a hex-encoded string"), + (with_cutoff(f64::NAN), + "'cutoff' in JSON for PairListOptions must be a finite positive number"), + (with_cutoff(f64::INFINITY), + "'cutoff' in JSON for PairListOptions must be a finite positive number"), + (with_cutoff(-1.0), + "'cutoff' in JSON for PairListOptions must be a finite positive number"), + (with_cutoff(0.0), + "'cutoff' in JSON for PairListOptions must be a finite positive number"), + (non_boolean_flag, + "'full_list' in JSON for PairListOptions must be a boolean"), + (non_array_requestors, + "'requestors' in JSON for PairListOptions must be an array"), + (non_string_requestor, + "'requestors' in JSON for PairListOptions must be an array of strings"), + ]; + + for (json, expected) in cases { + let error = PairListOptions::try_from(json).expect_err("expected an error"); + assert_eq!(error.to_string(), expected); + } + } + + #[test] + fn requestors_skip_empty_and_duplicates() { + let mut json: JsonValue = example().into(); + json["requestors"] = json::array![ "a", "", "b", "a" ]; + + let parsed = PairListOptions::try_from(json).unwrap(); + assert_eq!(parsed.requestors, vec!["a".to_string(), "b".to_string()]); + } + } +} From f6d4c8156f6827981f4e468fc51bcc93900f6fab Mon Sep 17 00:00:00 2001 From: GardevoirX Date: Thu, 28 May 2026 22:50:39 +0200 Subject: [PATCH 07/43] Implement JSON serialization for `Quantity` Co-Authored-By: Guillaume Fraux --- docs/src/core/reference/json-formats.rst | 47 +++- metatomic-core/src/lib.rs | 7 +- metatomic-core/src/metadata.rs | 33 --- metatomic-core/src/quantities.rs | 281 +++++++++++++++++++++++ 4 files changed, 329 insertions(+), 39 deletions(-) create mode 100644 metatomic-core/src/quantities.rs diff --git a/docs/src/core/reference/json-formats.rst b/docs/src/core/reference/json-formats.rst index d1589fa2d..d12ff6da5 100644 --- a/docs/src/core/reference/json-formats.rst +++ b/docs/src/core/reference/json-formats.rst @@ -11,10 +11,9 @@ language can produce and consume them. Pair list options ----------------- -Options describing a requested pair list (also known as a neighbor list). This -is the JSON representation of ``PairListOptions``, used for example by -:c:func:`mta_system_set_pairs`, :c:func:`mta_system_get_pairs` and -:c:func:`mta_system_pairs_options`. +The JSON representation of a requested pair list (also known as a neighbor +list). This is used for example by :c:func:`mta_system_add_pairs`, +:c:func:`mta_system_get_pairs` and :c:func:`mta_system_known_pairs`. .. code-block:: json @@ -50,3 +49,43 @@ is the JSON representation of ``PairListOptions``, used for example by ``requestors`` Optional array of strings identifying who requested this pair list. May be omitted, in which case it is treated as an empty list. + + +Quantities +---------- + +The JSON representation of a physical quantity, used to represent custom models +inputs and outputs. This is used for example in +:c:member:`mta_model_t.requested_inputs` and +:c:member:`mta_model_t.supported_outputs`. + +.. code-block:: json + + { + "type": "metatomic_quantity", + "name": "energy", + "unit": "eV", + "sample_kind": "system" + "gradients": ["positions"] + "description": "Potential energy of the system", + } + +``type`` + Must be the string ``"metatomic_quantity"``. + +``name`` + Name of the quantity, this this can be a standard name from the list of + :ref:`standard-quantities`, or a custom name of the form + ``::[/]`` + +``unit`` + Unit of the quantity. + +``gradients`` + Array of strings identifying the gradients for this quantity. This can be an + empty array if the quantity has no gradients. Valid values for the gradients + are ``"positions"``, and ``"strain"``. + +``sample_kind`` + Kind of sample for which this quantity is defined. This can be one of the + following: ``"atom"``, ``"system"`` or ``"atom_pair"``. diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index 894e70e3a..c778867bb 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -17,7 +17,10 @@ pub mod c_api; mod metadata; -pub use self::metadata::{ModelMetadata, Quantity, PairListOptions}; +pub use self::metadata::{ModelMetadata, PairListOptions}; + +mod quantities; +pub use self::quantities::Quantity; mod system; pub use self::system::System; @@ -31,7 +34,7 @@ pub use self::plugin::{Plugin, load_plugin, load_model}; mod units; pub use self::units::unit_conversion_factor; -/// TODO +/// Error type used throughout `metatomic-core`. #[derive(Debug)] pub enum Error { /// Error while serializing data to or deserializing data from JSON diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs index 913edcdfc..f28732919 100644 --- a/metatomic-core/src/metadata.rs +++ b/metatomic-core/src/metadata.rs @@ -155,39 +155,6 @@ impl TryFrom for ModelMetadata { } } -// ========================================================================== // -// ========================================================================== // -// ========================================================================== // - -/// TODO, previously `ModelOutput` -#[derive(Debug)] -pub struct Quantity { - pub name: String, - // TODO -} - -// TODO: -// { -// "type": "metatomic_quantity", -// "name": "...", -// "unit": "...", -// "gradients": ["...", "..."], -// "sample_kind": "atom" | "system" | "atom-pair", -// }, -impl From for JsonValue { - fn from(value: Quantity) -> Self { - todo!() - } -} - -impl TryFrom for Quantity { - type Error = Error; - - fn try_from(value: JsonValue) -> Result { - todo!() - } -} - #[cfg(test)] mod tests { diff --git a/metatomic-core/src/quantities.rs b/metatomic-core/src/quantities.rs new file mode 100644 index 000000000..072cdc41b --- /dev/null +++ b/metatomic-core/src/quantities.rs @@ -0,0 +1,281 @@ +use json::JsonValue; + +use crate::Error; + + +/// Different kind of samples a quantity can be associated with +#[derive(Debug, Clone, PartialEq)] +pub enum SampleKind { + /// The quantity is defined for each atom (e.g. atomic energy, charge, ...) + Atom, + /// The quantity is defined for the whole system (e.g. total energy, ...) + System, + /// The quantity is defined for each pair of atoms (e.g. hamiltonian elements, ...) + AtomPair, +} + +impl From for JsonValue { + fn from(value: SampleKind) -> Self { + let s = match value { + SampleKind::Atom => "atom", + SampleKind::System => "system", + SampleKind::AtomPair => "atom_pair", + }; + JsonValue::from(s) + } +} + +impl<'a> TryFrom<&'a JsonValue> for SampleKind { + type Error = Error; + + fn try_from(value: &'a JsonValue) -> Result { + let s = value.as_str().ok_or_else(|| Error::Serialization( + "'sample_kind' in JSON for Quantity must be a string".into() + ))?; + match s { + "atom" => Ok(SampleKind::Atom), + "system" => Ok(SampleKind::System), + "atom_pair" => Ok(SampleKind::AtomPair), + _ => Err(Error::Serialization(format!( + "'sample_kind' in JSON for Quantity must be 'atom', 'system' or 'atom_pair', got '{}'", s + ))), + } + } +} + +/// Different gradients that a quantity can have +#[derive(Debug, Clone, PartialEq)] +pub enum Gradients { + /// Gradients with respect to atomic positions + Positions, + /// Gradients with respect to the strain (typically used for stress) + Strain, +} + +impl From for JsonValue { + fn from(value: Gradients) -> Self { + let s = match value { + Gradients::Positions => "positions", + Gradients::Strain => "strain", + }; + JsonValue::from(s) + } +} + +impl<'a> TryFrom<&'a JsonValue> for Gradients { + type Error = Error; + + fn try_from(value: &'a JsonValue) -> Result { + let s = value.as_str().ok_or_else(|| Error::Serialization( + "'gradients' in JSON for Quantity must be a string".into() + ))?; + match s { + "positions" => Ok(Gradients::Positions), + "strain" => Ok(Gradients::Strain), + _ => Err(Error::Serialization(format!( + "'gradients' in JSON for Quantity must be 'positions' or 'strain', got '{}'", s + ))), + } + } +} + +/// A quantity that a model can use as input or output +#[derive(Debug, Clone)] +pub struct Quantity { + /// Name of the quantity, this can be a standard name from + /// , or + /// a custom name of the form `::[/]` + pub name: String, + /// Unit of the quantity + pub unit: String, + /// Description of the quantity, used to provide more details about the + /// quantity, especially when a model defines multiple variants of the same + /// quantity. + pub description: Option, + /// List of explicit gradients for this quantity, stored in the + /// corresponding `TensorMap` + pub gradients: Vec, + /// The kind of samples this quantity is associated with (e.g. per-atom, + /// per-system, ...) + pub sample_kind: SampleKind, +} + +impl From for JsonValue { + fn from(value: Quantity) -> Self { + let mut result = JsonValue::new_object(); + result["type"] = "metatomic_quantity".into(); + result["name"] = value.name.into(); + result["unit"] = value.unit.into(); + if let Some(description) = value.description { + result["description"] = description.into(); + } + result["gradients"] = value.gradients.into(); + result["sample_kind"] = value.sample_kind.into(); + return result; + } +} + + +impl TryFrom for Quantity { + type Error = Error; + + fn try_from(value: JsonValue) -> Result { + if !value.is_object() { + return Err(Error::Serialization( + "invalid JSON data for Quantity, expected an object".into() + )); + } + + if value["type"].as_str() != Some("metatomic_quantity") { + return Err(Error::Serialization( + "'type' in JSON for Quantity must be 'metatomic_quantity'".into() + )); + } + + let name = value["name"].as_str().ok_or_else(|| Error::Serialization( + "'name' in JSON for Quantity must be a string".into() + ))?; + + let unit = value["unit"].as_str().ok_or_else(|| Error::Serialization( + "'unit' in JSON for Quantity must be a string".into() + ))?; + + let mut description = value["description"].as_str().map(|s| s.to_string()); + if description == Some(String::new()) { + // Treat empty description as None + description = None; + } + + let gradients = &value["gradients"]; + if !gradients.is_array() { + return Err(Error::Serialization( + "'gradients' in JSON for Quantity must be an array".into() + )); + } + let gradients = gradients.members() + .map(Gradients::try_from) + .collect::, _>>()?; + + let sample_kind = SampleKind::try_from(&value["sample_kind"])?; + + Ok(Quantity { + name: name.to_string(), + unit: unit.to_string(), + description, + gradients, + sample_kind, + }) + } +} + + +#[cfg(test)] +mod tests { + use super::*; + + fn example() -> Quantity { + Quantity { + name: "energy".into(), + unit: "eV".into(), + description: Some("total energy of the system".into()), + gradients: vec![Gradients::Positions], + sample_kind: SampleKind::Atom, + } + } + + #[test] + fn roundtrip() { + let quantity = example(); + let json: JsonValue = quantity.into(); + + assert_eq!(json["type"].as_str(), Some("metatomic_quantity")); + assert_eq!(json["name"].as_str(), Some("energy")); + assert_eq!(json["unit"].as_str(), Some("eV")); + assert_eq!(json["gradients"][0].as_str(), Some("positions")); + assert_eq!(json["sample_kind"].as_str(), Some("atom")); + + let parsed = Quantity::try_from(json).unwrap(); + assert_eq!(parsed.name, "energy"); + assert_eq!(parsed.unit, "eV"); + assert_eq!(parsed.gradients, vec![Gradients::Positions]); + assert!(matches!(parsed.sample_kind, SampleKind::Atom)); + } + + #[test] + fn roundtrip_all_variants() { + for sample in [SampleKind::Atom, SampleKind::System, SampleKind::AtomPair] { + for grads in [ + vec![], + vec![Gradients::Positions], + vec![Gradients::Strain], + vec![Gradients::Positions, Gradients::Strain], + ] { + let quantity = Quantity { + name: "test".into(), + unit: "unit".into(), + description: Some("Hello".to_string()), + gradients: grads.clone(), + sample_kind: sample.clone(), + }; + let parsed = Quantity::try_from(JsonValue::from(quantity.clone())).unwrap(); + assert_eq!(parsed.name, quantity.name); + assert_eq!(parsed.unit, quantity.unit); + assert_eq!(parsed.gradients, grads); + assert_eq!(parsed.sample_kind, sample); + } + } + } + + #[test] + fn rejects_invalid_json() { + let mut wrong_type = JsonValue::from(example()); + wrong_type["type"] = "something-else".into(); + + let mut missing_name = JsonValue::from(example()); + missing_name.remove("name"); + + let mut missing_unit = JsonValue::from(example()); + missing_unit.remove("unit"); + + let mut missing_gradients = JsonValue::from(example()); + missing_gradients.remove("gradients"); + + let mut non_array_gradients = JsonValue::from(example()); + non_array_gradients["gradients"] = "positions".into(); + + let mut invalid_gradient = JsonValue::from(example()); + invalid_gradient["gradients"] = json::array!["positions", "foo"]; + + let mut missing_sample_kind = JsonValue::from(example()); + missing_sample_kind.remove("sample_kind"); + + let mut invalid_sample_kind = JsonValue::from(example()); + invalid_sample_kind["sample_kind"] = "foo".into(); + + let cases: Vec<(JsonValue, &str)> = vec![ + (JsonValue::from("not an object"), + "invalid JSON data for Quantity, expected an object"), + (wrong_type, + "'type' in JSON for Quantity must be 'metatomic_quantity'"), + (missing_name, + "'name' in JSON for Quantity must be a string"), + (missing_unit, + "'unit' in JSON for Quantity must be a string"), + (missing_gradients, + "'gradients' in JSON for Quantity must be an array"), + (non_array_gradients, + "'gradients' in JSON for Quantity must be an array"), + (invalid_gradient, + "'gradients' in JSON for Quantity must be 'positions' or 'strain', got 'foo'"), + (missing_sample_kind, + "'sample_kind' in JSON for Quantity must be a string"), + (invalid_sample_kind, + "'sample_kind' in JSON for Quantity must be 'atom', 'system' or 'atom_pair', got 'foo'"), + ]; + + for (json, expected) in cases { + let error = Quantity::try_from(json).expect_err("expected an error"); + assert_eq!(error.to_string(), expected); + } + } +} From 0d7925856b6e6d752dde870601cae5bc5b5acbef Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Fri, 29 May 2026 11:46:27 +0200 Subject: [PATCH 08/43] Validate quantities names --- metatomic-core/src/lib.rs | 7 +- metatomic-core/src/metadata.rs | 22 ++--- metatomic-core/src/quantities.rs | 164 +++++++++++++++++++++++++++++-- 3 files changed, 171 insertions(+), 22 deletions(-) diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index c778867bb..bf394f738 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -39,12 +39,15 @@ pub use self::units::unit_conversion_factor; pub enum Error { /// Error while serializing data to or deserializing data from JSON Serialization(String), + /// Invalid parameters passed to a function + InvalidParameters(String), } impl std::fmt::Display for Error { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { - Error::Serialization(message) => write!(f, "{}", message), + Error::Serialization(message) => write!(f, "serialization error: {}", message), + Error::InvalidParameters(message) => write!(f, "invalid parameter: {}", message), } } } @@ -52,7 +55,7 @@ impl std::fmt::Display for Error { impl std::error::Error for Error { fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { match self { - Error::Serialization(_) => None, + Error::Serialization(_) | Error::InvalidParameters(_) => None, } } diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs index f28732919..26101c0fa 100644 --- a/metatomic-core/src/metadata.rs +++ b/metatomic-core/src/metadata.rs @@ -232,27 +232,27 @@ mod tests { let cases = [ (JsonValue::from("not an object"), - "invalid JSON data for PairListOptions, expected an object"), + "serialization error: invalid JSON data for PairListOptions, expected an object"), (wrong_type, - "'type' in JSON for PairListOptions must be 'metatomic_pair_options'"), + "serialization error: 'type' in JSON for PairListOptions must be 'metatomic_pair_options'"), (missing_cutoff, - "'cutoff' in JSON for PairListOptions must be a hex-encoded string"), + "serialization error: 'cutoff' in JSON for PairListOptions must be a hex-encoded string"), (non_hex_cutoff, - "'cutoff' in JSON for PairListOptions must be a hex-encoded string"), + "serialization error: 'cutoff' in JSON for PairListOptions must be a hex-encoded string"), (with_cutoff(f64::NAN), - "'cutoff' in JSON for PairListOptions must be a finite positive number"), + "serialization error: 'cutoff' in JSON for PairListOptions must be a finite positive number"), (with_cutoff(f64::INFINITY), - "'cutoff' in JSON for PairListOptions must be a finite positive number"), + "serialization error: 'cutoff' in JSON for PairListOptions must be a finite positive number"), (with_cutoff(-1.0), - "'cutoff' in JSON for PairListOptions must be a finite positive number"), + "serialization error: 'cutoff' in JSON for PairListOptions must be a finite positive number"), (with_cutoff(0.0), - "'cutoff' in JSON for PairListOptions must be a finite positive number"), + "serialization error: 'cutoff' in JSON for PairListOptions must be a finite positive number"), (non_boolean_flag, - "'full_list' in JSON for PairListOptions must be a boolean"), + "serialization error: 'full_list' in JSON for PairListOptions must be a boolean"), (non_array_requestors, - "'requestors' in JSON for PairListOptions must be an array"), + "serialization error: 'requestors' in JSON for PairListOptions must be an array"), (non_string_requestor, - "'requestors' in JSON for PairListOptions must be an array of strings"), + "serialization error: 'requestors' in JSON for PairListOptions must be an array of strings"), ]; for (json, expected) in cases { diff --git a/metatomic-core/src/quantities.rs b/metatomic-core/src/quantities.rs index 072cdc41b..c8d4b45d8 100644 --- a/metatomic-core/src/quantities.rs +++ b/metatomic-core/src/quantities.rs @@ -2,6 +2,83 @@ use json::JsonValue; use crate::Error; +static STANDARD_QUANTITIES: &[&str] = &[ + "charge", + "energy_ensemble", + "energy_uncertainty", + "energy", + "feature", + "heat_flux", + "mass", + "momentum", + "non_conservative_force", + "non_conservative_stress", + "position", + "spin_multiplicity", + "velocity", +]; + +fn is_valid_identifier(s: &str) -> bool { + if s.is_empty() { + return false; + } + let first = s.chars().next().unwrap(); + if !(first.is_ascii_alphabetic() || first == '_') { + return false; + } + s.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') +} + +/// Validate a quantity name. +/// +/// The name can be either a standard name or a custom name with the form +/// `::`, where the namespace can itself contain `::` to define +/// sub-namespaces. +/// +/// Both standard and custom names can also define a variant with the form +/// `/` or `::/`. +/// +/// All components (namespace, name, variant) must be non-empty if they are +/// present, and must be valid identifiers (alphanumeric + underscore, not +/// starting with a digit). +fn validate_quantity_name(name: &str) -> Result<(), Error> { + if STANDARD_QUANTITIES.contains(&name) { + return Ok(()); + } + + let (main_part, variant) = if let Some(pos) = name.find('/') { + (&name[..pos], Some(&name[pos + 1..])) + } else { + (name, None) + }; + + if main_part.is_empty() { + return Err(Error::InvalidParameters(format!( + "quantity name cannot be empty in '{}'", name + ))); + } + + if let Some(variant) = variant { + if !is_valid_identifier(variant) { + return Err(Error::InvalidParameters(format!( + "invalid quantity variant '{}' in '{}': must be a valid identifier (alphanumeric or underscore, not starting with a digit)", + variant, name + ))); + } + } + + for component in main_part.split("::") { + if !is_valid_identifier(component) { + return Err(Error::InvalidParameters(format!( + "invalid quantity name component '{}' in '{}': must be a valid identifier (alphanumeric or underscore, not starting with a digit)", + component, name + ))); + } + } + + Ok(()) +} + /// Different kind of samples a quantity can be associated with #[derive(Debug, Clone, PartialEq)] @@ -135,6 +212,7 @@ impl TryFrom for Quantity { let name = value["name"].as_str().ok_or_else(|| Error::Serialization( "'name' in JSON for Quantity must be a string".into() ))?; + validate_quantity_name(name)?; let unit = value["unit"].as_str().ok_or_else(|| Error::Serialization( "'unit' in JSON for Quantity must be a string".into() @@ -254,23 +332,23 @@ mod tests { let cases: Vec<(JsonValue, &str)> = vec![ (JsonValue::from("not an object"), - "invalid JSON data for Quantity, expected an object"), + "serialization error: invalid JSON data for Quantity, expected an object"), (wrong_type, - "'type' in JSON for Quantity must be 'metatomic_quantity'"), + "serialization error: 'type' in JSON for Quantity must be 'metatomic_quantity'"), (missing_name, - "'name' in JSON for Quantity must be a string"), + "serialization error: 'name' in JSON for Quantity must be a string"), (missing_unit, - "'unit' in JSON for Quantity must be a string"), + "serialization error: 'unit' in JSON for Quantity must be a string"), (missing_gradients, - "'gradients' in JSON for Quantity must be an array"), + "serialization error: 'gradients' in JSON for Quantity must be an array"), (non_array_gradients, - "'gradients' in JSON for Quantity must be an array"), + "serialization error: 'gradients' in JSON for Quantity must be an array"), (invalid_gradient, - "'gradients' in JSON for Quantity must be 'positions' or 'strain', got 'foo'"), + "serialization error: 'gradients' in JSON for Quantity must be 'positions' or 'strain', got 'foo'"), (missing_sample_kind, - "'sample_kind' in JSON for Quantity must be a string"), + "serialization error: 'sample_kind' in JSON for Quantity must be a string"), (invalid_sample_kind, - "'sample_kind' in JSON for Quantity must be 'atom', 'system' or 'atom_pair', got 'foo'"), + "serialization error: 'sample_kind' in JSON for Quantity must be 'atom', 'system' or 'atom_pair', got 'foo'"), ]; for (json, expected) in cases { @@ -278,4 +356,72 @@ mod tests { assert_eq!(error.to_string(), expected); } } + + #[test] + fn validate_names() { + for name in STANDARD_QUANTITIES { + assert!(validate_quantity_name(name).is_ok(), "expected '{}' to be valid", name); + } + + let custom = [ + "my_model::energy", + "org::my_model::custom_qty", + "ns1::ns2::ns3::energy", + "custom_name", + "some_ns::name_with_underscores", + "_underscore_start", + "_ns::_name", + ]; + for name in custom { + assert!(validate_quantity_name(name).is_ok(), "expected '{}' to be valid", name); + } + + let variants = [ + "energy/ensemble", + "my_ns::energy/raw", + "ns1::ns2::energy/some_variant", + ]; + for name in variants { + assert!(validate_quantity_name(name).is_ok(), "expected '{}' to be valid", name); + } + + let error = validate_quantity_name("").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: quantity name cannot be empty in ''"); + + let error = validate_quantity_name("/variant").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: quantity name cannot be empty in '/variant'"); + + let error = validate_quantity_name("name/").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity variant '' in 'name/': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("::energy").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '' in '::energy': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("ns::").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '' in 'ns::': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("ns::/variant").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '' in 'ns::/variant': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("::").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '' in '::': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("123name").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '123name' in '123name': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("my_ns::123name").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '123name' in 'my_ns::123name': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("my_ns::name/123variant").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity variant '123variant' in 'my_ns::name/123variant': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("has spaces").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component 'has spaces' in 'has spaces': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("my_ns::name/has spaces").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity variant 'has spaces' in 'my_ns::name/has spaces': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + + let error = validate_quantity_name("has-dash").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component 'has-dash' in 'has-dash': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + } } From ee2ecba4c9edc1ba6d9b49f8a593eeb2f5e7e16b Mon Sep 17 00:00:00 2001 From: Alessandro Forina Date: Thu, 28 May 2026 16:56:37 +0200 Subject: [PATCH 09/43] Implement JSON serialization for ModelMetadata Co-Authored-By: Guillaume Fraux --- docs/src/core/reference/json-formats.rst | 59 +++++ metatomic-core/src/metadata.rs | 265 +++++++++++++++++++++-- 2 files changed, 306 insertions(+), 18 deletions(-) diff --git a/docs/src/core/reference/json-formats.rst b/docs/src/core/reference/json-formats.rst index d12ff6da5..d47cb2715 100644 --- a/docs/src/core/reference/json-formats.rst +++ b/docs/src/core/reference/json-formats.rst @@ -89,3 +89,62 @@ inputs and outputs. This is used for example in ``sample_kind`` Kind of sample for which this quantity is defined. This can be one of the following: ``"atom"``, ``"system"`` or ``"atom_pair"``. + + +Model metadata +-------------- + +The JSON representation of a model's metadata. This is used for example by +:c:member:`mta_model_t.metadata`. + +.. code-block:: json + + { + "type": "metatomic_model_metadata", + "name": "MyCoolModel v1.2", + "authors": ["Alice Smith", "Bob Johnson "], + "description": "A machine learning potential for water", + "references": { + "model": ["doi:10.1234/model-paper"], + "architecture": ["doi:10.1234/arch-paper"], + "implementation": ["https://github.com/example/mycoolmodel"] + }, + "extra": { + "training_set": "QM9", + "cutoff": "4.5" + } + } + +``type`` + Must be the string ``"metatomic_model_metadata"``. + +``name`` + Name of the model, e.g. ``"MyCoolModel v1.2"``. + +``authors`` + Array of strings identifying the authors of the model. Each string can be a + name or a name with an email address, e.g. ``"Alice Smith"`` or + ``"Bob Johnson "``. + +``description`` + A free-text description of the model. + +``references`` + An object with three keys, each containing an array of strings (DOIs, URLs, + or any other format): + + ``model`` + References about the model as a whole, e.g. a paper describing the model + or a website presenting it. + + ``architecture`` + References about the architecture of the model, e.g. papers describing + the mathematical form of the model. + + ``implementation`` + References about the implementation of the model, e.g. a link to the + source code repository or a paper describing the software. + +``extra`` + An object with string values, providing any additional key-value pairs the + model author wishes to include. This can be used for any purpose. diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs index 26101c0fa..f8cb769b0 100644 --- a/metatomic-core/src/metadata.rs +++ b/metatomic-core/src/metadata.rs @@ -1,3 +1,5 @@ +use std::collections::BTreeMap; + use json::JsonValue; use crate::Error; @@ -121,29 +123,97 @@ impl TryFrom for PairListOptions { // ========================================================================== // // ========================================================================== // -/// TODO -#[derive(Debug, Clone, PartialEq, Eq, Hash)] +/// References for a model, divided into three categories: references about the +/// model as a whole, references about the architecture of the model, and +/// references about the implementation of the model. Each category is a list of +/// strings, which can be DOIs, URLs, or any other format the model author finds +/// useful. +#[derive(Debug, Clone)] +pub struct References { + /// The references about the model as a whole, e.g. a paper describing the + /// model or a website presenting it. + model: Vec, + /// The references about the architecture of the model, e.g. papers + /// describing the mathematical form of the model. + architecture: Vec, + /// The references about the implementation of the model, e.g. a link to + /// the source code repository or a paper describing the software. + implementation: Vec, +} + +impl From for JsonValue { + fn from(value: References) -> Self { + let mut result = JsonValue::new_object(); + result["model"] = value.model.into(); + result["architecture"] = value.architecture.into(); + result["implementation"] = value.implementation.into(); + return result; + } +} + + +fn read_references(object: &JsonValue, key: &str) -> Result, Error> { + let mut references = Vec::new(); + if !object[key].is_array() { + return Err(Error::Serialization( + format!("'{}' in references of ModelMetadata must be an array", key) + )); + } + for reference in object[key].members() { + let reference = reference.as_str().ok_or_else(|| Error::Serialization( + format!("'{}' in references of ModelMetadata must be an array of strings", key) + ))?; + references.push(reference.to_string()); + } + Ok(references) +} + +impl TryFrom for References { + type Error = Error; + + fn try_from(value: JsonValue) -> Result { + if !value.is_object() { + return Err(Error::Serialization( + "invalid JSON data for references in ModelMetadata, expected an object".into() + )); + } + + let model = read_references(&value, "model")?; + let architecture = read_references(&value, "architecture")?; + let implementation = read_references(&value, "implementation")?; + + Ok(References { model, architecture, implementation }) + } +} + + +/// Metadata about a model +#[derive(Debug, Clone)] pub struct ModelMetadata { + /// The name of the model, e.g. `"MyCoolModel v1.2"` pub name: String, - // TODO + /// The authors of the model, e.g. `["Alice Smith", "Bob Johnson + /// "]` + pub authors: Vec, + /// A description of the model + pub description: String, + /// References for the model that should be cited when using it + pub references: References, + /// Any other key-value pairs the model author wants to include in the + /// metadata. This can be used for any purpose. + pub extra: BTreeMap, } -// { -// "type": "metatomic_model_metadata", -// "name": "...", -// "authors": ["..."], -// "references": { -// "implementation": ["..."], -// "architecture": ["..."], -// "model": ["..."] -// }, -// "extra": { -// "key...": "value..." -// } -// }, impl From for JsonValue { fn from(value: ModelMetadata) -> Self { - todo!() + let mut result = JsonValue::new_object(); + result["type"] = "metatomic_model_metadata".into(); + result["name"] = value.name.into(); + result["authors"] = value.authors.into(); + result["description"] = value.description.into(); + result["references"] = value.references.into(); + result["extra"] = value.extra.into(); + return result; } } @@ -151,7 +221,61 @@ impl TryFrom for ModelMetadata { type Error = Error; fn try_from(value: JsonValue) -> Result { - todo!() + if !value.is_object() { + return Err(Error::Serialization( + "invalid JSON data for ModelMetadata, expected an object".into() + )); + } + + if value["type"].as_str() != Some("metatomic_model_metadata") { + return Err(Error::Serialization( + "'type' in JSON for ModelMetadata must be 'metatomic_model_metadata'".into() + )); + } + + let name = value["name"].as_str().ok_or_else(|| Error::Serialization( + "'name' in JSON for ModelMetadata must be a string".into() + ))?; + + if !value["authors"].is_array() { + return Err(Error::Serialization( + "'authors' in JSON for ModelMetadata must be an array".into() + )); + } + + let authors = value["authors"].members().map(|author| { + author.as_str().ok_or_else(|| Error::Serialization( + "'authors' in JSON for ModelMetadata must be an array of strings".into() + )).map(|s| s.to_string()) + }).collect::, Error>>()?; + + let description = value["description"].as_str().ok_or_else(|| Error::Serialization( + "'description' in JSON for ModelMetadata must be a string".into() + ))?.to_string(); + + let references = References::try_from(value["references"].clone())?; + + if !value["extra"].is_object() { + return Err(Error::Serialization( + "'extra' in JSON for ModelMetadata must be an object".into() + )); + } + + let mut extra = BTreeMap::new(); + for (key, value) in value["extra"].entries() { + let value = value.as_str().ok_or_else(|| Error::Serialization( + "'extra' in JSON for ModelMetadata must be an object with string values".into() + ))?; + extra.insert(key.to_string(), value.to_string()); + } + + Ok(ModelMetadata { + name: name.to_string(), + authors: authors, + description: description, + references: references, + extra: extra, + }) } } @@ -270,4 +394,109 @@ mod tests { assert_eq!(parsed.requestors, vec!["a".to_string(), "b".to_string()]); } } + + mod model_metadata { + use super::super::*; + + fn example() -> ModelMetadata { + ModelMetadata { + name: "test-model".into(), + authors: vec!["Alice".into(), "Bob ".into()], + description: "A test model".into(), + references: References { + model: vec!["doi:10.1234/test".into()], + architecture: vec!["doi:10.1234/arch".into()], + implementation: vec!["https://github.com/test".into()], + }, + extra: BTreeMap::from([ + ("key1".into(), "value1".into()), + ("key2".into(), "value2".into()), + ]), + } + } + + #[test] + fn roundtrip() { + let metadata = example(); + let json: JsonValue = metadata.clone().into(); + + assert_eq!(json["type"].as_str(), Some("metatomic_model_metadata")); + assert_eq!(json["name"].as_str(), Some("test-model")); + assert_eq!(json["authors"][0].as_str(), Some("Alice")); + assert_eq!(json["authors"][1].as_str(), Some("Bob ")); + assert_eq!(json["description"].as_str(), Some("A test model")); + assert_eq!(json["references"]["model"][0].as_str(), Some("doi:10.1234/test")); + assert_eq!(json["references"]["architecture"][0].as_str(), Some("doi:10.1234/arch")); + assert_eq!(json["references"]["implementation"][0].as_str(), Some("https://github.com/test")); + assert_eq!(json["extra"]["key1"].as_str(), Some("value1")); + assert_eq!(json["extra"]["key2"].as_str(), Some("value2")); + + let parsed = ModelMetadata::try_from(json).unwrap(); + assert_eq!(parsed.name, metadata.name); + assert_eq!(parsed.authors, metadata.authors); + assert_eq!(parsed.description, metadata.description); + assert_eq!(parsed.references.model, metadata.references.model); + assert_eq!(parsed.references.architecture, metadata.references.architecture); + assert_eq!(parsed.references.implementation, metadata.references.implementation); + assert_eq!(parsed.extra, metadata.extra); + } + + #[test] + fn rejects_invalid_json() { + let mut wrong_type = JsonValue::from(example()); + wrong_type["type"] = "something-else".into(); + + let mut missing_name = JsonValue::from(example()); + missing_name.remove("name"); + + let mut non_string_name = JsonValue::from(example()); + non_string_name["name"] = 42.into(); + + let mut non_array_authors = JsonValue::from(example()); + non_array_authors["authors"] = "Alice".into(); + + let mut non_string_author = JsonValue::from(example()); + non_string_author["authors"] = json::array!["Alice", 42]; + + let mut missing_description = JsonValue::from(example()); + missing_description.remove("description"); + + let mut non_object_extra = JsonValue::from(example()); + non_object_extra["extra"] = "not-an-object".into(); + + let mut non_string_extra_value = JsonValue::from(example()); + non_string_extra_value["extra"] = json::object!{ "key" => 42 }; + + let mut non_object_references = JsonValue::from(example()); + non_object_references["references"] = "not-an-object".into(); + + let cases = [ + (JsonValue::from("not an object"), + "serialization error: invalid JSON data for ModelMetadata, expected an object"), + (wrong_type, + "serialization error: 'type' in JSON for ModelMetadata must be 'metatomic_model_metadata'"), + (missing_name, + "serialization error: 'name' in JSON for ModelMetadata must be a string"), + (non_string_name, + "serialization error: 'name' in JSON for ModelMetadata must be a string"), + (non_array_authors, + "serialization error: 'authors' in JSON for ModelMetadata must be an array"), + (non_string_author, + "serialization error: 'authors' in JSON for ModelMetadata must be an array of strings"), + (missing_description, + "serialization error: 'description' in JSON for ModelMetadata must be a string"), + (non_object_extra, + "serialization error: 'extra' in JSON for ModelMetadata must be an object"), + (non_string_extra_value, + "serialization error: 'extra' in JSON for ModelMetadata must be an object with string values"), + (non_object_references, + "serialization error: invalid JSON data for references in ModelMetadata, expected an object"), + ]; + + for (json, expected) in cases { + let error = ModelMetadata::try_from(json).expect_err("expected an error"); + assert_eq!(error.to_string(), expected); + } + } + } } From cab2e98cd2fe2d0ec01ff18bf397c504b7b52689 Mon Sep 17 00:00:00 2001 From: Rocco Meli Date: Thu, 28 May 2026 17:31:09 +0200 Subject: [PATCH 10/43] Add error handling based on metatensor --- metatomic-core/include/metatomic.h | 35 +++++- metatomic-core/src/c_api/mod.rs | 1 + metatomic-core/src/c_api/status.rs | 185 +++++++++++++++++++++++++++-- metatomic-core/src/lib.rs | 44 +++++-- metatomic-core/src/quantities.rs | 6 +- 5 files changed, 247 insertions(+), 24 deletions(-) diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index 1e69263e1..2ffa3a2c1 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -20,9 +20,38 @@ */ #define MTA_ABI_VERSION 1 +/** + * Status type returned by all functions in the C API. + * + * The value 0 (`MTA_SUCCESS`) indicates success, while any non-zero value indicates an error. + */ typedef enum mta_status_t { + /** + * Status code indicating success + */ MTA_SUCCESS = 0, - MTA_ERROR_OTHER = 255, + /** + * Status code indicating invalid function parameters + */ + MTA_INVALID_PARAMETER_ERROR = 1, + /** + * Status code indicating I/O errors + */ + MTA_IO_ERROR = 2, + /** + * Status code indicating serialization/deserialization errors + */ + MTA_SERIALIZATION_ERROR = 3, + /** + * Status code indicating errors that come from callbacks provided by the user. + * The error message and arbitrary data can be stored using `mta_set_last_error`, + * and retrieved using `mta_last_error`. + */ + MTA_CALLBACK_ERROR = 254, + /** + * Status code used when there is an internal error + */ + MTA_INTERNAL_ERROR = 255, } mta_status_t; /** @@ -112,12 +141,12 @@ extern "C" { #endif // __cplusplus /** - * TODO + * Get last error message that was created on the current thread. */ enum mta_status_t mta_last_error(const char **message, const char **origin, void **data); /** - * TODO + * Set last error message for the current thread. */ enum mta_status_t mta_set_last_error(const char *message, const char *origin, diff --git a/metatomic-core/src/c_api/mod.rs b/metatomic-core/src/c_api/mod.rs index cf6c6176d..c1b29c419 100644 --- a/metatomic-core/src/c_api/mod.rs +++ b/metatomic-core/src/c_api/mod.rs @@ -1,3 +1,4 @@ +#[macro_use] mod status; pub use self::status::mta_status_t; diff --git a/metatomic-core/src/c_api/status.rs b/metatomic-core/src/c_api/status.rs index 0c48707cc..8aef16c11 100644 --- a/metatomic-core/src/c_api/status.rs +++ b/metatomic-core/src/c_api/status.rs @@ -1,36 +1,170 @@ -use std::ffi::{c_char, c_void}; +use std::cell::RefCell; +use std::ffi::{c_char, c_void, CStr, CString}; +use std::panic::UnwindSafe; use crate::Error; +#[derive(Debug)] +struct LastError { + message: CString, + origin: CString, + custom_data: *mut c_void, + custom_data_deleter: Option, +} + +// Save the last error message in thread local storage. +thread_local! { + pub static LAST_ERROR: RefCell = RefCell::new(LastError { + message: CString::new("").expect("invalid C string"), + origin: CString::new("").expect("invalid C string"), + custom_data: std::ptr::null_mut(), + custom_data_deleter: None, + }); +} -// TODO +/// Status type returned by all functions in the C API. +/// +/// The value 0 (`MTA_SUCCESS`) indicates success, while any non-zero value indicates an error. #[allow(non_camel_case_types)] #[repr(C)] #[derive(PartialEq, Eq, Debug)] pub enum mta_status_t { + /// Status code indicating success MTA_SUCCESS = 0, - // ... - MTA_ERROR_OTHER = 255, + /// Status code indicating invalid function parameters + MTA_INVALID_PARAMETER_ERROR = 1, + /// Status code indicating I/O errors + MTA_IO_ERROR = 2, + /// Status code indicating serialization/deserialization errors + MTA_SERIALIZATION_ERROR = 3, + /// Status code indicating errors that come from callbacks provided by the user. + /// The error message and arbitrary data can be stored using `mta_set_last_error`, + /// and retrieved using `mta_last_error`. + MTA_CALLBACK_ERROR = 254, + /// Status code used when there is an internal error + MTA_INTERNAL_ERROR = 255, } +/// `std::panic::catch_unwind` that automatically transform +/// the error into `mta_status_t`. +pub fn catch_unwind(function: F) -> mta_status_t +where + F: FnOnce() -> Result<(), Error> + UnwindSafe, +{ + match std::panic::catch_unwind(function) { + Ok(Ok(())) => mta_status_t::MTA_SUCCESS, + Ok(Err(error)) => error.into(), + Err(error) => Error::from(error).into(), + } +} + +/// Check that pointers (used as C API function parameters) are not null. +#[macro_export] +#[doc(hidden)] +macro_rules! check_pointers_non_null { + ($pointer: ident) => { + if $pointer.is_null() { + return Err($crate::Error::InvalidParameter( + format!( + "got invalid NULL pointer for {} at {}:{}", + stringify!($pointer), file!(), line!() + ) + )); + } + }; + ($($pointer: ident),* $(,)?) => { + $(check_pointers_non_null!($pointer);)* + } +} impl From for mta_status_t { - fn from(err: Error) -> Self { - todo!() + fn from(error: Error) -> mta_status_t { + if let Error::CallbackError = error { + // If the error is already a CallbackError, we can directly return the corresponding status code. + return mta_status_t::MTA_CALLBACK_ERROR; + } + + LAST_ERROR.with(|last_error| { + let mut last_error = last_error.borrow_mut(); + + // If there is a custom data deleter, + // use it to free the custom data before overwriting it with the new error. + if let Some(deleter) = last_error.custom_data_deleter { + unsafe { + deleter(last_error.custom_data); + } + } + + *last_error = LastError { + message: CString::new(format!("{}", error)) + .expect("error message contains a null byte"), + origin: CString::new("metatensor-core").expect("invalid C string"), + custom_data: std::ptr::null_mut(), + custom_data_deleter: None, + }; + }); + + match error { + Error::InvalidParameter(_) => mta_status_t::MTA_INVALID_PARAMETER_ERROR, + Error::Io(_) => mta_status_t::MTA_IO_ERROR, + Error::Serialization(_) => mta_status_t::MTA_SERIALIZATION_ERROR, + Error::CallbackError => unreachable!(), + Error::Internal(_) => mta_status_t::MTA_INTERNAL_ERROR, + } } } -/// TODO +/// Get last error message that was created on the current thread. #[no_mangle] pub unsafe extern "C" fn mta_last_error( message: *mut *const c_char, origin: *mut *const c_char, data: *mut *mut c_void, ) -> mta_status_t { - todo!() + let status = std::panic::catch_unwind(|| { + LAST_ERROR.with(|last_error| { + let last_error = last_error.borrow(); + if !message.is_null() { + *message = last_error.message.as_ptr(); + } + if !origin.is_null() { + *origin = last_error.origin.as_ptr(); + } + if !data.is_null() { + *data = last_error.custom_data; + } + }); + }); + + match status { + Ok(()) => mta_status_t::MTA_SUCCESS, + Err(error) => { + let last_error_debug = + LAST_ERROR.with(|last_error| format!("{:?}", last_error.borrow())); + if error.is::() { + eprintln!( + "panic in mta_last_error: {:?}, last_error: {:?}", + error.downcast_ref::(), + last_error_debug + ); + } else if error.is::<&str>() { + eprintln!( + "panic in mta_last_error: {:?}, last_error: {:?}", + error.downcast_ref::<&str>(), + last_error_debug + ); + } else { + eprintln!( + "panic in mta_last_error: unknown panic error type. last_error: {:?}", + last_error_debug + ); + } + mta_status_t::MTA_INTERNAL_ERROR + } + } } -/// TODO +/// Set last error message for the current thread. #[no_mangle] pub unsafe extern "C" fn mta_set_last_error( message: *const c_char, @@ -38,5 +172,36 @@ pub unsafe extern "C" fn mta_set_last_error( data: *mut c_void, data_deleter: Option, ) -> mta_status_t { - todo!() + catch_unwind(move || { + let message = if message.is_null() { + CString::new("").expect("invalid C string") + } else { + CString::from(CStr::from_ptr(message)) + }; + + let origin = if origin.is_null() { + CString::new("").expect("invalid C string") + } else { + CString::from(CStr::from_ptr(origin)) + }; + + LAST_ERROR.with(|last_error| { + let mut last_error = last_error.borrow_mut(); + + // Call custom data deleter before overwriting the custom data with the new one, to avoid memory leaks. + if let Some(deleter) = last_error.custom_data_deleter { + unsafe { + deleter(last_error.custom_data); + } + } + + *last_error = LastError { + message: message, + origin: origin, + custom_data: data, + custom_data_deleter: data_deleter, + }; + }); + Ok(()) + }) } diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index bf394f738..09ec7a962 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -8,11 +8,9 @@ #![allow(clippy::similar_names, clippy::borrow_as_ptr, clippy::uninlined_format_args)] #![allow(clippy::let_underscore_untyped, clippy::manual_let_else, clippy::empty_line_after_doc_comments)] - -// To be removed lated +// To be removed later #![allow(unused_variables, dead_code, clippy::needless_pass_by_value)] - #[doc(hidden)] pub mod c_api; @@ -34,20 +32,31 @@ pub use self::plugin::{Plugin, load_plugin, load_model}; mod units; pub use self::units::unit_conversion_factor; -/// Error type used throughout `metatomic-core`. +/// The possible sources of error in metatomic #[derive(Debug)] pub enum Error { /// Error while serializing data to or deserializing data from JSON Serialization(String), /// Invalid parameters passed to a function - InvalidParameters(String), + InvalidParameter(String), + /// I/O error + Io(std::io::Error), + /// Error coming from an external function used as a callback + CallbackError, + /// Any other internal error, usually these are internal bugs. + Internal(String), } impl std::fmt::Display for Error { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { - Error::Serialization(message) => write!(f, "serialization error: {}", message), - Error::InvalidParameters(message) => write!(f, "invalid parameter: {}", message), + Error::Serialization(e) => write!(f, "serialization error: {}", e), + Error::InvalidParameter(e) => write!(f, "invalid parameter: {}", e), + Error::Io(e) => write!(f, "io error: {}", e), + Error::CallbackError => write!(f, "callback error"), + Error::Internal(e) => write!(f, + "internal metatomic error (this is likely a bug, please report it): {}", e + ), } } } @@ -55,7 +64,11 @@ impl std::fmt::Display for Error { impl std::error::Error for Error { fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { match self { - Error::Serialization(_) | Error::InvalidParameters(_) => None, + Error::InvalidParameter(_) + | Error::Serialization(_) + | Error::Internal(_) + | Error::CallbackError => None, + Error::Io(e) => Some(e), } } @@ -63,3 +76,18 @@ impl std::error::Error for Error { self.source() } } + +// Box is the error type in std::panic::catch_unwind +impl From> for Error { + fn from(error: Box) -> Error { + if error.is::() { + Error::Internal(*error.downcast::().expect("should be a String")) + } else if error.is::<&str>() { + Error::Internal((*error.downcast::<&str>().expect("should be an &str")).to_owned()) + } else if error.is::() { + return *error.downcast::().expect("it should be an Error"); + } else { + panic!("panic message is not a string, something is very wrong") + } + } +} diff --git a/metatomic-core/src/quantities.rs b/metatomic-core/src/quantities.rs index c8d4b45d8..9d1dfebbf 100644 --- a/metatomic-core/src/quantities.rs +++ b/metatomic-core/src/quantities.rs @@ -53,14 +53,14 @@ fn validate_quantity_name(name: &str) -> Result<(), Error> { }; if main_part.is_empty() { - return Err(Error::InvalidParameters(format!( + return Err(Error::InvalidParameter(format!( "quantity name cannot be empty in '{}'", name ))); } if let Some(variant) = variant { if !is_valid_identifier(variant) { - return Err(Error::InvalidParameters(format!( + return Err(Error::InvalidParameter(format!( "invalid quantity variant '{}' in '{}': must be a valid identifier (alphanumeric or underscore, not starting with a digit)", variant, name ))); @@ -69,7 +69,7 @@ fn validate_quantity_name(name: &str) -> Result<(), Error> { for component in main_part.split("::") { if !is_valid_identifier(component) { - return Err(Error::InvalidParameters(format!( + return Err(Error::InvalidParameter(format!( "invalid quantity name component '{}' in '{}': must be a valid identifier (alphanumeric or underscore, not starting with a digit)", component, name ))); From 66eee81413f8d731a2dfb00e2444012a8b0939d1 Mon Sep 17 00:00:00 2001 From: Johannes Spies <13813209+johannes-spies@users.noreply.github.com> Date: Thu, 28 May 2026 16:18:18 +0200 Subject: [PATCH 11/43] Implement mta_string_t in the C API --- metatomic-core/build.rs | 11 ++++ metatomic-core/include/metatomic.h | 41 +++++++++---- metatomic-core/src/c_api/mod.rs | 2 +- metatomic-core/src/c_api/utils.rs | 99 ++++++++++++++++++++++++------ metatomic-core/tests/misc.cpp | 34 ++++++++++ 5 files changed, 155 insertions(+), 32 deletions(-) diff --git a/metatomic-core/build.rs b/metatomic-core/build.rs index b92cc2925..01f95d71c 100644 --- a/metatomic-core/build.rs +++ b/metatomic-core/build.rs @@ -25,6 +25,17 @@ fn main() { config.includes.push("metatensor.h".into()); config.includes.push("metatomic/version.h".into()); + config.export = cbindgen::ExportConfig { + include: vec!["mta_.*".into()], + // This is done manually below + exclude: vec!["mta_opaque_string_t".into()], + ..Default::default() + }; + config.after_includes = Some(" + +/** Heap allocated storage for mta_string_t */ +typedef struct mta_opaque_string_t mta_opaque_string_t;".into()); + let result = cbindgen::Builder::new() .with_crate(crate_dir) .with_config(config) diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index 2ffa3a2c1..25a5f7444 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -15,6 +15,10 @@ #include "metatensor.h" #include "metatomic/version.h" + +/** Heap allocated storage for mta_string_t */ +typedef struct mta_opaque_string_t mta_opaque_string_t; + /** * TODO */ @@ -64,20 +68,21 @@ typedef enum mta_system_data_kind { MTA_SYSTEM_DATA_PBC = 3, } mta_system_data_kind; -/** - * TODO - */ -typedef struct mta_opaque_string_t mta_opaque_string_t; - /** * TODO */ typedef struct mta_system_t mta_system_t; /** - * TODO + * An heap-allocated UTF-8 string passed across the C API boundary. + * + * This is used whenever a C API function or callback needs to return a string. + * + * A null pointer represents an absent or empty string. Use `mta_string_create` + * to allocate, `mta_string_free` to release, and `mta_string_view` to get a + * pointer to the inner C string. */ -typedef struct mta_opaque_string_t *mta_string_t; +typedef mta_opaque_string_t *mta_string_t; /** * TODO @@ -161,17 +166,31 @@ enum mta_status_t mta_set_last_error(const char *message, const char *mta_version(void); /** - * TODO + * Allocate a new `mta_string_t` by copying the null-terminated C string + * `string`. + * + * The returned string must be freed with `mta_string_free`. + * + * @param string A pointer to a null-terminated C string. Must not be null. + * @return A new `mta_string_t` containing a copy of `string`, or null if an + * error occurred. You can check the error with `mta_last_error`. */ -mta_string_t mta_string_create(const char *raw); +mta_string_t mta_string_create(const char *string); /** - * TODO + * Free a `mta_string_t` previously created by `mta_string_create`. + * + * @param string A `mta_string_t` to free. Can be null, in which case this function is a no-op. */ void mta_string_free(mta_string_t string); /** - * TODO + * Return a pointer to the null-terminated string data inside `string`. + * + * The pointer is valid only for the lifetime of `string`. + * + * @param string A `mta_string_t` containing the string to view. Must not be null. + * @return A pointer to the null-terminated C string inside `string` */ const char *mta_string_view(mta_string_t string); diff --git a/metatomic-core/src/c_api/mod.rs b/metatomic-core/src/c_api/mod.rs index c1b29c419..282e230e5 100644 --- a/metatomic-core/src/c_api/mod.rs +++ b/metatomic-core/src/c_api/mod.rs @@ -1,6 +1,6 @@ #[macro_use] mod status; -pub use self::status::mta_status_t; +pub use self::status::{mta_status_t, catch_unwind}; mod utils; pub use self::utils::mta_string_t; diff --git a/metatomic-core/src/c_api/utils.rs b/metatomic-core/src/c_api/utils.rs index 350a9552d..8e22bcc47 100644 --- a/metatomic-core/src/c_api/utils.rs +++ b/metatomic-core/src/c_api/utils.rs @@ -2,7 +2,7 @@ use std::ffi::{CString, c_char}; use once_cell::sync::Lazy; -use super::mta_status_t; +use super::{mta_status_t, catch_unwind}; static VERSION: Lazy = Lazy::new(|| { @@ -18,11 +18,18 @@ pub extern "C" fn mta_version() -> *const c_char { return VERSION.as_ptr(); } -/// TODO +/// Heap-allocated backing storage for `mta_string_t`, opaque to C users. #[allow(non_camel_case_types)] -pub struct mta_opaque_string_t(CString); +#[repr(transparent)] +pub struct mta_opaque_string_t(c_char); -/// TODO +/// An heap-allocated UTF-8 string passed across the C API boundary. +/// +/// This is used whenever a C API function or callback needs to return a string. +/// +/// A null pointer represents an absent or empty string. Use `mta_string_create` +/// to allocate, `mta_string_free` to release, and `mta_string_view` to get a +/// pointer to the inner C string. #[allow(non_camel_case_types)] #[repr(transparent)] pub struct mta_string_t(*mut mta_opaque_string_t); @@ -41,51 +48,104 @@ impl std::fmt::Debug for mta_string_t { } impl mta_string_t { - /// TODO + /// Create a new `mta_string_t` from a Rust string. pub fn new(value: impl Into) -> Self { - let cstring = CString::new(value.into()).unwrap(); - let boxed = Box::new(mta_opaque_string_t(cstring)); - mta_string_t(Box::into_raw(boxed)) + let cstring = CString::new(value.into()).expect("string contains NULL byte"); + let ptr = CString::into_raw(cstring); + return mta_string_t(ptr.cast()); } - /// TODO + /// Create a null `mta_string_t`, representing an absent string. pub fn null() -> Self { mta_string_t(std::ptr::null_mut()) } - /// TODO + /// View the string as a `&str`. Returns `""` for a null string. pub fn as_str(&self) -> &str { if self.0.is_null() { return ""; } unsafe { - return (*(self.0)).0.to_str().expect("mta_string_t is not valid UTF8") + let cstr = std::ffi::CStr::from_ptr(self.0.cast()); + return cstr.to_str().expect("invalid UTF-8 in mta_string_t"); } } } -/// TODO +/// Allocate a new `mta_string_t` by copying the null-terminated C string +/// `string`. +/// +/// The returned string must be freed with `mta_string_free`. +/// +/// @param string A pointer to a null-terminated C string. Must not be null. +/// @return A new `mta_string_t` containing a copy of `string`, or null if an +/// error occurred. You can check the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_string_create( - raw: *const c_char, + string: *const c_char, ) -> mta_string_t { - todo!() + let mut result = mta_string_t::null(); + let unwind_wrapper = std::panic::AssertUnwindSafe(&mut result); + + catch_unwind(move || { + check_pointers_non_null!(string); + + let cstr = std::ffi::CStr::from_ptr(string); + let string = CString::from(cstr); + + let ptr = CString::into_raw(string); + + let _ = &unwind_wrapper; + *unwind_wrapper.0 = mta_string_t(ptr.cast()); + Ok(()) + }); + + return result; } -/// TODO +/// Free a `mta_string_t` previously created by `mta_string_create`. +/// +/// @param string A `mta_string_t` to free. Can be null, in which case this function is a no-op. #[no_mangle] pub unsafe extern "C" fn mta_string_free(string: mta_string_t) { - todo!() + catch_unwind(|| { + if string.0.is_null() { + return Ok(()); + } + + let ptr = string.0.cast::(); + let cstring = CString::from_raw(ptr); + std::mem::drop(cstring); + + Ok(()) + }); } -/// TODO +/// Return a pointer to the null-terminated string data inside `string`. +/// +/// The pointer is valid only for the lifetime of `string`. +/// +/// @param string A `mta_string_t` containing the string to view. Must not be null. +/// @return A pointer to the null-terminated C string inside `string` #[no_mangle] pub unsafe extern "C" fn mta_string_view( string: mta_string_t, ) -> *const c_char { - todo!() -} + let mut result = std::ptr::null(); + let unwind_wrapper = std::panic::AssertUnwindSafe(&mut result); + + catch_unwind(move || { + let string = string.0; + check_pointers_non_null!(string); + + let _ = &unwind_wrapper; + *unwind_wrapper.0 = string.cast(); + Ok(()) + }); + + return result; +} /// TODO #[no_mangle] @@ -98,5 +158,4 @@ pub unsafe extern "C" fn mta_unit_conversion_factor( } - // TODO: logging & warnings? diff --git a/metatomic-core/tests/misc.cpp b/metatomic-core/tests/misc.cpp index bf0ce275f..188eca67c 100644 --- a/metatomic-core/tests/misc.cpp +++ b/metatomic-core/tests/misc.cpp @@ -1,3 +1,5 @@ +#include + #include #include "metatomic.h" @@ -13,3 +15,35 @@ TEST_CASE("Version macros") { // METATOMIC_VERSION should start with `x.y.z` CHECK(std::string(METATOMIC_VERSION).find(version) == 0); } + +TEST_CASE("mta_string_t") { + auto* str = mta_string_create("hello"); + REQUIRE(str != nullptr); + + const char* view = mta_string_view(str); + CHECK(std::strlen(view) == 5); + CHECK(std::string(view) == "hello"); + mta_string_free(str); + + // empty string + str = mta_string_create(""); + REQUIRE(str != nullptr); + CHECK(std::string(mta_string_view(str)) == ""); + mta_string_free(str); + + // special characters + str = mta_string_create("a\nb\tc\xFFºµ"); + REQUIRE(str != nullptr); + CHECK(std::string(mta_string_view(str)) == std::string("a\nb\tc\xFFºµ")); + mta_string_free(str); + + // long string + std::string long_str(10000, 'x'); + str = mta_string_create(long_str.c_str()); + REQUIRE(str != nullptr); + CHECK(std::string(mta_string_view(str)) == long_str); + mta_string_free(str); + + // free on a null pointer should work + mta_string_free(nullptr); +} From e33cd84c7eacff73fe0421c3c7271aaf98298f41 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Fri, 29 May 2026 15:44:44 +0200 Subject: [PATCH 12/43] Port unit parsing from metatomic-torch --- docs/src/core/index.rst | 1 + docs/src/core/units.rst | 97 ++++ docs/src/torch/reference/index.rst | 1 - docs/src/torch/reference/misc.rst | 2 + docs/src/torch/reference/units.rst | 71 --- metatomic-core/include/metatomic.h | 19 +- metatomic-core/src/c_api/mod.rs | 2 + metatomic-core/src/c_api/utils.rs | 39 +- metatomic-core/src/units.rs | 723 ++++++++++++++++++++++++++++- metatomic-core/tests/misc.cpp | 23 + 10 files changed, 900 insertions(+), 78 deletions(-) create mode 100644 docs/src/core/units.rst delete mode 100644 docs/src/torch/reference/units.rst diff --git a/docs/src/core/index.rst b/docs/src/core/index.rst index 4d0cf24a4..2a3691316 100644 --- a/docs/src/core/index.rst +++ b/docs/src/core/index.rst @@ -9,6 +9,7 @@ WIP reference/c/index reference/json-formats + units .. toctree:: diff --git a/docs/src/core/units.rst b/docs/src/core/units.rst new file mode 100644 index 000000000..d5f3dd776 --- /dev/null +++ b/docs/src/core/units.rst @@ -0,0 +1,97 @@ +.. _core-unit-expressions: + +Units +^^^^^ + +Models in metatensor can use arbitrary units for their inputs and outputs. The +unit conversion system allows models to specify the units they expect and +receive data in any compatible unit, with automatic conversion handled by +:c:func:`mta_execute_model`. + +The :c:func:`mta_unit_conversion_factor` function parses two unit expressions, +checks that they have compatible physical dimensions, and returns the +multiplicative conversion factor: + +.. code-block:: c + + // How many eV are in one kJ/mol? + double factor; + mta_unit_conversion_factor("kJ/mol", "eV", &factor); + // factor ≈ 0.01036 + + // How many GPa are in one eV/A^3? + mta_unit_conversion_factor("eV/A^3", "GPa", &factor); + // factor ≈ 160.22 + +If either (or both) unit strings are empty, the conversion returns ``1.0`` +without checking dimensions. This makes it safe to pass optional/unknown units. + +.. _known-base-units: + +Base units +~~~~~~~~~~ + +Unit expressions are built from the following base units. Matching is +case-insensitive, and whitespace is ignored. + +**Temperature**: + ``Kelvin`` (``K``) + +**Length**: + ``angstrom`` (``A``), ``Bohr``, ``meter`` (``m``), ``centimeter`` (``cm``), + ``millimeter`` (``mm``), ``micrometer`` (``um``, ``µm``), ``nanometer`` (``nm``) + +**Energy**: + ``eV``, ``meV``, ``Hartree``, ``kcal``, ``kJ``, ``Joule`` (``J``), ``Rydberg`` (``Ry``) + +**Time**: + ``second`` (``s``), ``millisecond`` (``ms``), ``microsecond`` (``us``, ``µs``), + ``nanosecond`` (``ns``), ``picosecond`` (``ps``), ``femtosecond`` (``fs``) + +**Mass**: + ``Dalton`` (``u``), ``kilogram`` (``kg``), ``gram`` (``g``), ``electron_mass`` (``m_e``) + +**Charge**: + ``e``, ``Coulomb`` (``C``) + +**Pressure**: + ``Pascal`` (``Pa``), ``kiloPascal`` (``kPa``), ``MegaPascal`` (``MPa``), + ``GigaPascal`` (``GPa``), ``bar``, ``atm`` + +**Electric Dipole Moment**: + ``Debye`` (``D``) + +**Dimensionless**: + ``mol`` + +**Derived constants**: + ``hbar`` + +Expression syntax +~~~~~~~~~~~~~~~~~ + +Base units can be combined using the following operators: + +- Multiplication: ``*`` or whitespace (``kJ mol``, ``kJ*mol``) +- Division: ``/`` (``kJ/mol``) +- Exponentiation: ``^`` (``A^3``, ``m^2``) +- Parentheses: ``()`` for grouping (``(eV*u)^(1/2)``) + +Fractional powers + Exponents can be integers (``A^3``) or fractions enclosed in parentheses + (``^(1/2)``, ``^(2/3)``). Fractional powers are supported only when the + result has integer physical dimensions — for example ``(eV*u)^(1/2)`` + computes momentum with dimensions :math:`[L T^{-1} M]`. + +Numeric literals + Bare numbers can be used as dimensionless quantity expressions, e.g. + ``"2"`` evaluates to the conversion factor ``2.0``. This is useful when a + model needs to define a unit that is simply a scalar multiple of another. + +Examples of valid compound expressions: + +- ``kJ/mol`` --- energy per mole +- ``eV/Angstrom^3`` or ``eV/A^3`` --- pressure +- ``(eV*u)^(1/2)`` --- momentum (fractional powers) +- ``Hartree/Bohr`` --- force in atomic units +- ``nm/fs`` --- velocity diff --git a/docs/src/torch/reference/index.rst b/docs/src/torch/reference/index.rst index 7cb577e46..b0419ada3 100644 --- a/docs/src/torch/reference/index.rst +++ b/docs/src/torch/reference/index.rst @@ -8,7 +8,6 @@ API reference systems models/index - units wrappers o3 ase diff --git a/docs/src/torch/reference/misc.rst b/docs/src/torch/reference/misc.rst index 00bf79f06..10d3e636d 100644 --- a/docs/src/torch/reference/misc.rst +++ b/docs/src/torch/reference/misc.rst @@ -7,3 +7,5 @@ simulation engine to use metatomic models. .. autofunction:: metatomic.torch.pick_device .. autofunction:: metatomic.torch.pick_output + +.. autofunction:: metatomic.torch.unit_conversion_factor diff --git a/docs/src/torch/reference/units.rst b/docs/src/torch/reference/units.rst deleted file mode 100644 index cba2a7397..000000000 --- a/docs/src/torch/reference/units.rst +++ /dev/null @@ -1,71 +0,0 @@ -Unit conversions -================ - -.. autofunction:: metatomic.torch.unit_conversion_factor - -The :py:func:`unit_conversion_factor` function accepts any valid unit expression -built from base units combined with operators. There is no need to specify a -physical quantity --- the parser automatically verifies dimensional -compatibility between the source and target units. - -.. _known-base-units: - -Supported base units -~~~~~~~~~~~~~~~~~~~~ - -Unit expressions are built from the following base units. Matching is -case-insensitive, and whitespace is ignored. - - -**Temperature**: - ``Kelvin`` (``K``) - -**Length**: - ``angstrom`` (``A``), ``Bohr``, ``meter`` (``m``), ``centimeter`` (``cm``), - ``millimeter`` (``mm``), ``micrometer`` (``um``, ``µm``), ``nanometer`` (``nm``) - -**Energy**: - ``eV``, ``meV``, ``Hartree``, ``kcal``, ``kJ``, ``Joule`` (``J``), ``Rydberg`` (``Ry``) - -**Time**: - ``second`` (``s``), ``millisecond`` (``ms``), ``microsecond`` (``us``, ``µs``), - ``nanosecond`` (``ns``), ``picosecond`` (``ps``), ``femtosecond`` (``fs``) - -**Mass**: - ``Dalton`` (``u``), ``kilogram`` (``kg``), ``gram`` (``g``), ``electron_mass`` (``m_e``) - -**Charge**: - ``e``, ``Coulomb`` (``C``) - -**Pressure**: - ``Pascal`` (``Pa``), ``kiloPascal`` (``kPa``), ``MegaPascal`` (``MPa``), ``GigaPascal`` (``GPa``), ``bar``, ``atm`` - -**Electric Dipole Moment**: - ``Debye`` (``D``) - -**Dimensionless**: - ``mol`` - -**Derived constants**: - ``hbar`` - -Expression syntax -~~~~~~~~~~~~~~~~~~~ - -Base units can be combined using the following operators: - -- Multiplication: ``*`` or whitespace (``kJ mol``, ``kJ*mol``) -- Division: ``/`` (``kJ/mol``) -- Exponentiation: ``^`` (``A^3``, ``m^2``) -- Parentheses: ``()`` for grouping (``(eV*u)^(1/2)``) - -Examples of valid compound expressions: - -- ``kJ/mol`` --- energy per mole -- ``eV/Angstrom^3`` or ``eV/A^3`` --- pressure -- ``(eV*u)^(1/2)`` --- momentum (fractional powers) -- ``Hartree/Bohr`` --- force in atomic units -- ``nm/fs`` --- velocity - -The parser automatically checks that both unit expressions have matching -physical dimensions before computing the conversion factor. diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index 25a5f7444..1a4fc35f9 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -195,7 +195,24 @@ void mta_string_free(mta_string_t string); const char *mta_string_view(mta_string_t string); /** - * TODO + * Get the multiplicative conversion factor to use to convert from + * `from_unit` to `to_unit`. Both units are parsed as expressions (e.g. + * "kJ/mol/A^2", "(eV*u)^(1/2)") and their dimensions must match. + * + * Unit expressions are built from base units combined with `*`, `/`, `^`, + * and parentheses. Unit lookup is case-insensitive, and whitespace is + * ignored. For example: + * + * - `"kJ/mol"` -- energy per mole + * - `"eV/Angstrom^3"` -- pressure + * - `"(eV*u)^(1/2)"` -- momentum (fractional powers) + * - `"Hartree/Bohr"` -- force in atomic units + * + * @param from_unit A null-terminated C string containing the unit to convert from. + * @param to_unit A null-terminated C string containing the unit to convert to. + * @param conversion A pointer to a `double` where the conversion factor will be stored. + * @return The status code of the operation. If this code is not `MTA_SUCCESS`, + * you can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_unit_conversion_factor(const char *from_unit, const char *to_unit, diff --git a/metatomic-core/src/c_api/mod.rs b/metatomic-core/src/c_api/mod.rs index 282e230e5..bffa5003a 100644 --- a/metatomic-core/src/c_api/mod.rs +++ b/metatomic-core/src/c_api/mod.rs @@ -1,3 +1,5 @@ +#![allow(clippy::doc_markdown)] + #[macro_use] mod status; pub use self::status::{mta_status_t, catch_unwind}; diff --git a/metatomic-core/src/c_api/utils.rs b/metatomic-core/src/c_api/utils.rs index 8e22bcc47..448c2b5d4 100644 --- a/metatomic-core/src/c_api/utils.rs +++ b/metatomic-core/src/c_api/utils.rs @@ -3,7 +3,7 @@ use std::ffi::{CString, c_char}; use once_cell::sync::Lazy; use super::{mta_status_t, catch_unwind}; - +use crate::Error; static VERSION: Lazy = Lazy::new(|| { CString::new(env!("METATOMIC_FULL_VERSION")).expect("version contains NULL byte") @@ -147,14 +147,47 @@ pub unsafe extern "C" fn mta_string_view( return result; } -/// TODO +/// Get the multiplicative conversion factor to use to convert from +/// `from_unit` to `to_unit`. Both units are parsed as expressions (e.g. +/// "kJ/mol/A^2", "(eV*u)^(1/2)") and their dimensions must match. +/// +/// Unit expressions are built from base units combined with `*`, `/`, `^`, +/// and parentheses. Unit lookup is case-insensitive, and whitespace is +/// ignored. For example: +/// +/// - `"kJ/mol"` -- energy per mole +/// - `"eV/Angstrom^3"` -- pressure +/// - `"(eV*u)^(1/2)"` -- momentum (fractional powers) +/// - `"Hartree/Bohr"` -- force in atomic units +/// +/// @param from_unit A null-terminated C string containing the unit to convert from. +/// @param to_unit A null-terminated C string containing the unit to convert to. +/// @param conversion A pointer to a `double` where the conversion factor will be stored. +/// @return The status code of the operation. If this code is not `MTA_SUCCESS`, +/// you can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_unit_conversion_factor( from_unit: *const c_char, to_unit: *const c_char, conversion: *mut f64, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(from_unit, to_unit, conversion); + + let from_cstr = std::ffi::CStr::from_ptr(from_unit); + let to_cstr = std::ffi::CStr::from_ptr(to_unit); + + let from_str = from_cstr.to_str().map_err(|_| { + Error::InvalidParameter("from_unit is not valid UTF-8".into()) + })?; + let to_str = to_cstr.to_str().map_err(|_| { + Error::InvalidParameter("to_unit is not valid UTF-8".into()) + })?; + + *conversion = crate::unit_conversion_factor(from_str, to_str)?; + + Ok(()) + }) } diff --git a/metatomic-core/src/units.rs b/metatomic-core/src/units.rs index d06eab413..4cfff2c1d 100644 --- a/metatomic-core/src/units.rs +++ b/metatomic-core/src/units.rs @@ -1,7 +1,726 @@ use crate::Error; +use once_cell::sync::Lazy; +use std::collections::HashMap; +use std::fmt; +use std::ops::{Add, Sub}; -/// TODO +/// Physical dimension vector with named integer exponents: +/// [Length, Time, Mass, Electric Current, Temperature] +/// +/// Note: quantity of substance (mole) is intentionally not included, since we +/// want `kJ/mol` and `eV` to have the same dimension. +#[derive(Debug, Clone, PartialEq, Eq)] +struct Dimension { + length: i32, + time: i32, + mass: i32, + electric_current: i32, + temperature: i32, +} + +impl Dimension { + /// Dimensionless — all exponents are zero. + const NONE: Dimension = Dimension { length: 0, time: 0, mass: 0, electric_current: 0, temperature: 0 }; + + /// Length dimension + const LENGTH: Dimension = Dimension { length: 1, time: 0, mass: 0, electric_current: 0, temperature: 0 }; + /// Time dimension + const TIME: Dimension = Dimension { length: 0, time: 1, mass: 0, electric_current: 0, temperature: 0 }; + /// Mass dimension + const MASS: Dimension = Dimension { length: 0, time: 0, mass: 1, electric_current: 0, temperature: 0 }; + /// Electric charge dimension (current × time) + const CHARGE: Dimension = Dimension { length: 0, time: 1, mass: 0, electric_current: 1, temperature: 0 }; + /// Temperature dimension + const TEMPERATURE: Dimension = Dimension { length: 0, time: 0, mass: 0, electric_current: 0, temperature: 1 }; + + /// Energy dimension: L² T⁻² M¹ + const ENERGY: Dimension = Dimension { length: 2, time: -2, mass: 1, electric_current: 0, temperature: 0 }; + /// Pressure dimension: L⁻¹ T⁻² M¹ + const PRESSURE: Dimension = Dimension { length: -1, time: -2, mass: 1, electric_current: 0, temperature: 0 }; + /// Electric dipole dimension: L¹ T¹ I¹ + const ELECTRIC_DIPOLE: Dimension = Dimension { length: 1, time: 1, mass: 0, electric_current: 1, temperature: 0 }; + + fn pow(&self, p: f64) -> Dimension { + Dimension { + length: round_if_integer(f64::from(self.length) * p), + time: round_if_integer(f64::from(self.time) * p), + mass: round_if_integer(f64::from(self.mass) * p), + electric_current: round_if_integer(f64::from(self.electric_current) * p), + temperature: round_if_integer(f64::from(self.temperature) * p), + } + } +} + +impl Add<&Dimension> for &Dimension { + type Output = Dimension; + + fn add(self, other: &Dimension) -> Dimension { + Dimension { + length: self.length + other.length, + time: self.time + other.time, + mass: self.mass + other.mass, + electric_current: self.electric_current + other.electric_current, + temperature: self.temperature + other.temperature, + } + } +} + +impl Sub<&Dimension> for &Dimension { + type Output = Dimension; + + fn sub(self, other: &Dimension) -> Dimension { + Dimension { + length: self.length - other.length, + time: self.time - other.time, + mass: self.mass - other.mass, + electric_current: self.electric_current - other.electric_current, + temperature: self.temperature - other.temperature, + } + } +} + +#[allow(clippy::cast_possible_truncation)] +fn round_if_integer(v: f64) -> i32 { + let rounded = v.round(); + assert!((v - rounded).abs() <= 1e-10, "non-integer dimension exponent {} is not supported", v); + return rounded as i32; +} + +impl fmt::Display for Dimension { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + use fmt::Write; + let mut first = true; + f.write_char('[')?; + + for (name, v) in [ + ("L", self.length), + ("T", self.time), + ("M", self.mass), + ("I", self.electric_current), + ("Θ", self.temperature), + ] { + if v == 0 { + continue; + } + + if !first { + f.write_char(' ')?; + } + first = false; + + f.write_str(name)?; + + if v != 1 && v != -1 { + write!(f, "^{}", v)?; + } + + if v == -1 { + f.write_str("^-1")?; + } + } + + if first { + f.write_str("dimensionless")?; + } + f.write_char(']')?; + + Ok(()) + } +} + +/// A parsed unit value: SI conversion factor and physical dimension. +#[derive(Debug, Clone)] +struct UnitValue { + factor: f64, + dim: Dimension, +} + +/// All base units with SI factors and dimensions. +/// Factors are expressed in SI base units (m, s, kg, C, K). +/// Case-insensitive lookup: names are lowercased before searching. +static BASE_UNITS: Lazy> = Lazy::new(|| { + let mut map = HashMap::new(); + + // --- Temperature --- + map.insert("kelvin", UnitValue { factor: 1.0, dim: Dimension::TEMPERATURE }); + map.insert("k", UnitValue { factor: 1.0, dim: Dimension::TEMPERATURE }); + + // --- Length --- + map.insert("angstrom", UnitValue { factor: 1e-10, dim: Dimension::LENGTH }); + map.insert("a", UnitValue { factor: 1e-10, dim: Dimension::LENGTH }); + map.insert("bohr", UnitValue { factor: 5.2917721054482e-11, dim: Dimension::LENGTH }); + map.insert("nm", UnitValue { factor: 1e-9, dim: Dimension::LENGTH }); + map.insert("nanometer", UnitValue { factor: 1e-9, dim: Dimension::LENGTH }); + map.insert("meter", UnitValue { factor: 1.0, dim: Dimension::LENGTH }); + map.insert("m", UnitValue { factor: 1.0, dim: Dimension::LENGTH }); + map.insert("cm", UnitValue { factor: 1e-2, dim: Dimension::LENGTH }); + map.insert("centimeter", UnitValue { factor: 1e-2, dim: Dimension::LENGTH }); + map.insert("mm", UnitValue { factor: 1e-3, dim: Dimension::LENGTH }); + map.insert("millimeter", UnitValue { factor: 1e-3, dim: Dimension::LENGTH }); + map.insert("um", UnitValue { factor: 1e-6, dim: Dimension::LENGTH }); + map.insert("µm", UnitValue { factor: 1e-6, dim: Dimension::LENGTH }); + map.insert("micrometer", UnitValue { factor: 1e-6, dim: Dimension::LENGTH }); + + // --- Energy --- + map.insert("electronvolt", UnitValue { factor: 1.602176634e-19, dim: Dimension::ENERGY }); + map.insert("ev", UnitValue { factor: 1.602176634e-19, dim: Dimension::ENERGY }); + map.insert("mev", UnitValue { factor: 1.602176634e-19 * 1e-3, dim: Dimension::ENERGY }); + map.insert("hartree", UnitValue { factor: 4.359744722206048e-18, dim: Dimension::ENERGY }); + map.insert("ry", UnitValue { factor: 2.179872361103024e-18, dim: Dimension::ENERGY }); + map.insert("rydberg", UnitValue { factor: 2.179872361103024e-18, dim: Dimension::ENERGY }); + map.insert("joule", UnitValue { factor: 1.0, dim: Dimension::ENERGY }); + map.insert("j", UnitValue { factor: 1.0, dim: Dimension::ENERGY }); + map.insert("kcal", UnitValue { factor: 4184.0, dim: Dimension::ENERGY }); + map.insert("kj", UnitValue { factor: 1000.0, dim: Dimension::ENERGY }); + + // --- Time --- + map.insert("s", UnitValue { factor: 1.0, dim: Dimension::TIME }); + map.insert("second", UnitValue { factor: 1.0, dim: Dimension::TIME }); + map.insert("ms", UnitValue { factor: 1e-3, dim: Dimension::TIME }); + map.insert("millisecond", UnitValue { factor: 1e-3, dim: Dimension::TIME }); + map.insert("us", UnitValue { factor: 1e-6, dim: Dimension::TIME }); + map.insert("µs", UnitValue { factor: 1e-6, dim: Dimension::TIME }); + map.insert("microsecond", UnitValue { factor: 1e-6, dim: Dimension::TIME }); + map.insert("ns", UnitValue { factor: 1e-9, dim: Dimension::TIME }); + map.insert("nanosecond", UnitValue { factor: 1e-9, dim: Dimension::TIME }); + map.insert("ps", UnitValue { factor: 1e-12, dim: Dimension::TIME }); + map.insert("picosecond", UnitValue { factor: 1e-12, dim: Dimension::TIME }); + map.insert("fs", UnitValue { factor: 1e-15, dim: Dimension::TIME }); + map.insert("femtosecond", UnitValue { factor: 1e-15, dim: Dimension::TIME }); + + // --- Mass --- + map.insert("u", UnitValue { factor: 1.6605390689252e-27, dim: Dimension::MASS }); + map.insert("dalton", UnitValue { factor: 1.6605390689252e-27, dim: Dimension::MASS }); + map.insert("kg", UnitValue { factor: 1.0, dim: Dimension::MASS }); + map.insert("kilogram", UnitValue { factor: 1.0, dim: Dimension::MASS }); + map.insert("g", UnitValue { factor: 1e-3, dim: Dimension::MASS }); + map.insert("gram", UnitValue { factor: 1e-3, dim: Dimension::MASS }); + map.insert("electron_mass", UnitValue { factor: 9.109383713928e-31, dim: Dimension::MASS }); + map.insert("m_e", UnitValue { factor: 9.109383713928e-31, dim: Dimension::MASS }); + + // --- Charge --- + map.insert("e", UnitValue { factor: 1.602176634e-19, dim: Dimension::CHARGE }); + map.insert("coulomb", UnitValue { factor: 1.0, dim: Dimension::CHARGE }); + map.insert("c", UnitValue { factor: 1.0, dim: Dimension::CHARGE }); + + // --- Pressure --- + map.insert("pa", UnitValue { factor: 1.0, dim: Dimension::PRESSURE }); + map.insert("pascal", UnitValue { factor: 1.0, dim: Dimension::PRESSURE }); + map.insert("kpa", UnitValue { factor: 1e3, dim: Dimension::PRESSURE }); + map.insert("kilopascal", UnitValue { factor: 1e3, dim: Dimension::PRESSURE }); + map.insert("mpa", UnitValue { factor: 1e6, dim: Dimension::PRESSURE }); + map.insert("megapascal", UnitValue { factor: 1e6, dim: Dimension::PRESSURE }); + map.insert("gpa", UnitValue { factor: 1e9, dim: Dimension::PRESSURE }); + map.insert("gigapascal", UnitValue { factor: 1e9, dim: Dimension::PRESSURE }); + map.insert("bar", UnitValue { factor: 100000.0, dim: Dimension::PRESSURE }); + map.insert("atm", UnitValue { factor: 101325.0, dim: Dimension::PRESSURE }); + + // --- Electric dipole moment --- + map.insert("debye", UnitValue { factor: 1.0 / 299792458.0 * 1e-21, dim: Dimension::ELECTRIC_DIPOLE }); + map.insert("d", UnitValue { factor: 1.0 / 299792458.0 * 1e-21, dim: Dimension::ELECTRIC_DIPOLE }); + + // --- Dimensionless --- + map.insert("mol", UnitValue { factor: 6.02214076e23, dim: Dimension::NONE }); + + // --- Derived --- + map.insert("hbar", UnitValue { + factor: 1.0545718176462e-34, + dim: Dimension { length: 2, time: -1, mass: 1, electric_current: 0, temperature: 0 }, + }); + + map +}); + +// ---- Tokenizer ---- + +#[derive(Debug, Clone)] +enum Token { + LParen, + RParen, + Mul, + Div, + Pow, + Value(String), +} + +impl Token { + fn precedence(&self) -> i32 { + match self { + Token::LParen | Token::RParen => 0, + Token::Mul | Token::Div => 10, + Token::Pow => 20, + Token::Value(_) => -1, + } + } +} + +impl fmt::Display for Token { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Token::LParen => write!(f, "("), + Token::RParen => write!(f, ")"), + Token::Mul => write!(f, "*"), + Token::Div => write!(f, "/"), + Token::Pow => write!(f, "^"), + Token::Value(v) => write!(f, "{}", v), + } + } +} + +fn tokenize(unit: &str) -> Vec { + let mut tokens = Vec::new(); + let mut current = String::new(); + + for c in unit.chars() { + if c == '*' || c == '/' || c == '^' || c == '(' || c == ')' { + if !current.is_empty() { + tokens.push(Token::Value(current.clone())); + current.clear(); + } + let t = match c { + '*' => Token::Mul, + '/' => Token::Div, + '^' => Token::Pow, + '(' => Token::LParen, + ')' => Token::RParen, + _ => unreachable!(), + }; + tokens.push(t); + } else if !c.is_whitespace() { + current.push(c); + } + } + + if !current.is_empty() { + tokens.push(Token::Value(current)); + } + + tokens +} + +// ---- Shunting-Yard ---- + +/// Convert infix tokens to [Reverse Polish Notation] (RPN) using the +/// [Shunting-Yard] algorithm. +/// +/// RPN (also called postfix notation) writes operators after their operands, +/// e.g. `kJ / mol` becomes `kJ mol /`. This removes the need for parentheses +/// and precedence rules, making the expression easy to evaluate with a stack. +/// +/// All operators are treated as left-associative. +/// +/// [Reverse Polish Notation]: https://en.wikipedia.org/wiki/Reverse_Polish_notation +/// [Shunting-Yard]: https://en.wikipedia.org/wiki/Shunting-yard_algorithm +fn shunting_yard(tokens: &[Token]) -> Result, Error> { + let mut output: Vec = Vec::new(); + let mut operators: Vec = Vec::new(); + + for token in tokens { + match token { + Token::Value(_) => { + output.push(token.clone()); + } + Token::Mul | Token::Div | Token::Pow => { + while let Some(top) = operators.last() { + if token.precedence() <= top.precedence() { + output.push(operators.pop().unwrap()); + } else { + break; + } + } + operators.push(token.clone()); + } + Token::LParen => { + operators.push(token.clone()); + } + Token::RParen => { + while let Some(top) = operators.last() { + if matches!(top, Token::LParen) { + break; + } + output.push(operators.pop().unwrap()); + } + if operators.is_empty() || !matches!(operators.last(), Some(Token::LParen)) { + return Err(Error::InvalidParameter( + "unit expression has unbalanced parentheses".into(), + )); + } + operators.pop(); // discard LParen + } + } + } + + while let Some(top) = operators.pop() { + if matches!(top, Token::LParen | Token::RParen) { + return Err(Error::InvalidParameter( + "unit expression has unbalanced parentheses".into(), + )); + } + output.push(top); + } + + Ok(output) +} + +// ---- AST evaluator ---- + +struct UnitExpr { + val: UnitExprData, +} + +enum UnitExprData { + Val(UnitValue, String), + Mul(Box, Box), + Div(Box, Box), + Pow(Box, Box), +} + +impl fmt::Display for UnitExpr { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match &self.val { + UnitExprData::Val(_, name) => f.write_str(name), + UnitExprData::Mul(lhs, rhs) => { + write!(f, "({} * {})", lhs, rhs) + } + UnitExprData::Div(lhs, rhs) => { + write!(f, "({} / {})", lhs, rhs) + } + UnitExprData::Pow(base, exponent) => { + write!(f, "({} ^ {})", base, exponent) + } + } + } +} + +impl UnitExpr { + fn eval(&self) -> Result { + match &self.val { + UnitExprData::Val(v, _) => Ok(v.clone()), + UnitExprData::Mul(lhs, rhs) => { + let l = lhs.eval()?; + let r = rhs.eval()?; + let result_factor = l.factor * r.factor; + if !result_factor.is_finite() { + return Err(Error::InvalidParameter(format!( + "unit conversion factor overflows: multiplication result is infinite \ + or NaN for '{}'", + self + ))); + } + Ok(UnitValue { + factor: result_factor, + dim: &l.dim + &r.dim, + }) + } + UnitExprData::Div(lhs, rhs) => { + let l = lhs.eval()?; + let r = rhs.eval()?; + let result_factor = l.factor / r.factor; + if !result_factor.is_finite() { + return Err(Error::InvalidParameter(format!( + "unit conversion factor overflows: division result is infinite \ + or NaN for '{}'", + self + ))); + } + Ok(UnitValue { + factor: result_factor, + dim: &l.dim - &r.dim, + }) + } + UnitExprData::Pow(base, exponent) => { + let b = base.eval()?; + let e = exponent.eval()?; + + if e.dim != Dimension::NONE { + return Err(Error::InvalidParameter(format!( + "exponent in unit expression must be dimensionless, got dimension {} \ + for exponent '{}'", + e.dim, + exponent + ))); + } + let result_factor = b.factor.powf(e.factor); + if !result_factor.is_finite() { + return Err(Error::InvalidParameter(format!( + "unit conversion factor overflows: exponentiation result is infinite \ + or NaN for '{}'", + self + ))); + } + Ok(UnitValue { + factor: result_factor, + dim: b.dim.pow(e.factor), + }) + } + } + } +} + +/// Read one expression from the [RPN] stream (recursive, pops from the back). +/// +/// RPN arranges expressions as `lhs rhs op`, so `rhs` is on top of the stack +/// and must be popped first. For example `kJ mol /` pops `mol` (rhs) then +/// `kJ` (lhs) to build `Div(lhs=kJ, rhs=mol)`. +/// +/// [RPN]: https://en.wikipedia.org/wiki/Reverse_Polish_notation +fn read_expr(stream: &mut Vec) -> Result { + let token = stream.pop().ok_or_else(|| { + Error::InvalidParameter("malformed unit expression: missing a value".into()) + })?; + + match token { + Token::Value(v) => { + let lower = v.to_lowercase(); + if let Some(uv) = BASE_UNITS.get(lower.as_str()) { + return Ok(UnitExpr { + val: UnitExprData::Val(uv.clone(), v), + }); + } + if let Ok(val) = v.parse::() { + return Ok(UnitExpr { + val: UnitExprData::Val(UnitValue { factor: val, dim: Dimension::NONE }, v), + }); + } + Err(Error::InvalidParameter(format!("unknown unit '{}'", v))) + } + // RPN: lhs rhs Mul — pop rhs first, then lhs + Token::Mul => { + let rhs = read_expr(stream)?; + let lhs = read_expr(stream)?; + Ok(UnitExpr { + val: UnitExprData::Mul(Box::new(lhs), Box::new(rhs)), + }) + } + // RPN: lhs rhs Div — pop rhs first, then lhs + Token::Div => { + let rhs = read_expr(stream)?; + let lhs = read_expr(stream)?; + Ok(UnitExpr { + val: UnitExprData::Div(Box::new(lhs), Box::new(rhs)), + }) + } + // RPN: base exponent Pow — pop exponent first, then base + Token::Pow => { + let exponent = read_expr(stream)?; + let base = read_expr(stream)?; + Ok(UnitExpr { + val: UnitExprData::Pow(Box::new(base), Box::new(exponent)), + }) + } + _ => Err(Error::InvalidParameter(format!( + "unexpected symbol in unit expression: '{}'", + token + ))), + } +} + +/// Parse a unit expression string and return the evaluated `UnitValue`. +fn parse_unit_expression(unit: &str) -> Result { + if unit.is_empty() { + return Ok(UnitValue { factor: 1.0, dim: Dimension::NONE }); + } + + let tokens = tokenize(unit); + if tokens.is_empty() { + return Ok(UnitValue { factor: 1.0, dim: Dimension::NONE }); + } + + let mut rpn = shunting_yard(&tokens)?; + let ast = read_expr(&mut rpn)?; + + if !rpn.is_empty() { + let remaining: Vec = rpn.iter().map(|t| t.to_string()).collect(); + return Err(Error::InvalidParameter(format!( + "malformed unit expression: leftover input '{}'", + remaining.join(" ") + ))); + } + + ast.eval() +} + +/// Get the multiplicative conversion factor to use to convert from +/// `from_unit` to `to_unit`. Both units are parsed as expressions (e.g. +/// "kJ/mol/A^2", "(eV*u)^(1/2)") and their dimensions must match. +/// +/// Unit expressions are built from base units combined with `*`, `/`, `^`, +/// and parentheses. Unit lookup is case-insensitive, and whitespace is +/// ignored. For example: +/// +/// - `"kJ/mol"` -- energy per mole +/// - `"eV/Angstrom^3"` -- pressure +/// - `"(eV*u)^(1/2)"` -- momentum (fractional powers) +/// - `"Hartree/Bohr"` -- force in atomic units pub fn unit_conversion_factor(from_unit: &str, to_unit: &str) -> Result { - todo!() + if from_unit.is_empty() || to_unit.is_empty() { + return Ok(1.0); + } + + let from = parse_unit_expression(from_unit)?; + let to = parse_unit_expression(to_unit)?; + + if from.dim != to.dim { + return Err(Error::InvalidParameter(format!( + "dimension mismatch in unit conversion: '{}' has dimension {} but '{}' has dimension {}", + from_unit, + from.dim, + to_unit, + to.dim + ))); + } + + Ok(from.factor / to.factor) +} + +#[cfg(test)] +#[allow(clippy::float_cmp)] +mod tests { + use super::*; + + #[test] + fn test_tokenize_simple() { + let tokens = tokenize("eV"); + assert_eq!(tokens.len(), 1); + assert!(matches!(&tokens[0], Token::Value(v) if v == "eV")); + } + + #[test] + fn test_tokenize_operators() { + let tokens = tokenize("kJ/mol/A^2"); + let types: Vec = tokens.iter().map(|t| t.to_string()).collect(); + assert_eq!(types, vec!["kJ", "/", "mol", "/", "A", "^", "2"]); + } + + #[test] + fn test_tokenize_parens() { + let tokens = tokenize("(eV*u)^(1/2)"); + let types: Vec = tokens.iter().map(|t| t.to_string()).collect(); + assert_eq!(types, vec!["(", "eV", "*", "u", ")", "^", "(", "1", "/", "2", ")"]); + } + + #[test] + fn test_tokenize_whitespace() { + let tokens = tokenize(" kJ / mol "); + let types: Vec = tokens.iter().map(|t| t.to_string()).collect(); + assert_eq!(types, vec!["kJ", "/", "mol"]); + } + + #[test] + fn test_shunting_yard() { + let tokens = tokenize("kJ/mol"); + let rpn = shunting_yard(&tokens).unwrap(); + let types: Vec = rpn.iter().map(|t| t.to_string()).collect(); + assert_eq!(types, vec!["kJ", "mol", "/"]); + + let tokens = tokenize("kJ/mol/A^2"); + let rpn = shunting_yard(&tokens).unwrap(); + let types: Vec = rpn.iter().map(|t| t.to_string()).collect(); + assert_eq!(types, vec!["kJ", "mol", "/", "A", "2", "^", "/"]); + } + + #[test] + fn test_parens_mismatch() { + let tokens = tokenize("("); + let err = shunting_yard(&tokens).expect_err("expected error"); + assert_eq!( + err.to_string(), + "invalid parameter: unit expression has unbalanced parentheses" + ); + + let tokens = tokenize("(eV*u"); + let err = shunting_yard(&tokens).expect_err("expected error"); + assert_eq!( + err.to_string(), + "invalid parameter: unit expression has unbalanced parentheses" + ); + } + + #[test] + fn test_simple_conversion() { + let factor = unit_conversion_factor("eV", "eV").unwrap(); + assert_eq!(factor, 1.0); + + let factor = unit_conversion_factor("m", "A").unwrap(); + assert!((factor - 1e10).abs() < 1e-5); + + let factor = unit_conversion_factor("eV", "kJ").unwrap(); + assert!((factor - 1.602176634e-22).abs() < 1e-30); + } + + #[test] + fn test_dimension_mismatch() { + let err = unit_conversion_factor("eV", "m").expect_err("expected error"); + assert_eq!( + err.to_string(), + "invalid parameter: dimension mismatch in unit conversion: \ + 'eV' has dimension [L^2 T^-2 M] but 'm' has dimension [L]" + ); + } + + #[test] + fn test_empty_units() { + let factor = unit_conversion_factor("", "").unwrap(); + assert_eq!(factor, 1.0); + + let factor = unit_conversion_factor("eV", "").unwrap(); + assert_eq!(factor, 1.0); + } + + #[test] + fn test_compound_units() { + let from = unit_conversion_factor("kJ/mol", "eV").unwrap(); + assert!((from - 0.010364269656262174).abs() < 1e-15); + + let from = unit_conversion_factor("eV/A^3", "GPa").unwrap(); + assert!((from - 160.21766339999996).abs() < 1e-12); + } + + #[test] + fn test_case_insensitive() { + let f1 = unit_conversion_factor("eV", "eV").unwrap(); + let f2 = unit_conversion_factor("EV", "eV").unwrap(); + assert_eq!(f1, f2); + + let factor = unit_conversion_factor("eV", "MeV").unwrap(); + assert!((factor - 1000.0).abs() < 1e-12); + } + + #[test] + fn test_unknown_unit() { + let err = unit_conversion_factor("foo", "eV").expect_err("expected error"); + assert_eq!(err.to_string(), "invalid parameter: unknown unit 'foo'"); + } + + #[test] + fn test_numeric_literal() { + let factor = unit_conversion_factor("2", "1").unwrap(); + assert_eq!(factor, 2.0); + } + + #[test] + fn test_fractional_power() { + let err = unit_conversion_factor("(eV*u)^(1/2)", "eV*u").expect_err("expected error"); + assert_eq!( + err.to_string(), + "invalid parameter: dimension mismatch in unit conversion: \ + '(eV*u)^(1/2)' has dimension [L T^-1 M] but 'eV*u' has dimension [L^2 T^-2 M^2]" + ); + + let factor = unit_conversion_factor("(eV*u)^(1/2)", "(eV*u)^(1/2)").unwrap(); + assert_eq!(factor, 1.0); + } + + #[test] + fn test_dimension_to_string() { + assert_eq!(Dimension::NONE.to_string(), "[dimensionless]"); + assert_eq!(Dimension::LENGTH.to_string(), "[L]"); + assert_eq!(Dimension::ENERGY.to_string(), "[L^2 T^-2 M]"); + assert_eq!(Dimension::PRESSURE.to_string(), "[L^-1 T^-2 M]"); + assert_eq!(Dimension::TEMPERATURE.to_string(), "[Θ]"); + + let velocity = Dimension { length: 1, time: -1, mass: 0, electric_current: 0, temperature: 0 }; + assert_eq!(velocity.to_string(), "[L T^-1]"); + } } diff --git a/metatomic-core/tests/misc.cpp b/metatomic-core/tests/misc.cpp index 188eca67c..8c66d7561 100644 --- a/metatomic-core/tests/misc.cpp +++ b/metatomic-core/tests/misc.cpp @@ -47,3 +47,26 @@ TEST_CASE("mta_string_t") { // free on a null pointer should work mta_string_free(nullptr); } + +TEST_CASE("mta_unit_conversion_factor") { + double factor = 0.0; + + // same unit -> factor = 1.0 + auto status = mta_unit_conversion_factor("m", "m", &factor); + REQUIRE(status == MTA_SUCCESS); + CHECK(factor == 1.0); + + // kJ/mol -> eV + CHECK(mta_unit_conversion_factor("kJ/mol", "eV", &factor) == MTA_SUCCESS); + CHECK(factor == Approx(0.010364269656262174).epsilon(1e-15)); + + // dimension mismatch -> error + status = mta_unit_conversion_factor("m", "kg", &factor); + REQUIRE(status != MTA_SUCCESS); + + const char* error_msg = nullptr; + mta_last_error(&error_msg, nullptr, nullptr); + CHECK(std::string(error_msg) == + "invalid parameter: dimension mismatch in unit conversion: " + "'m' has dimension [L] but 'kg' has dimension [M]"); +} From 000fea4b0c2c4610f163e51372f11f96a910174e Mon Sep 17 00:00:00 2001 From: Sofiia Chorna Date: Sat, 30 May 2026 13:58:51 +0200 Subject: [PATCH 13/43] document mta_model_t and related functions in C API --- docs/src/core/reference/json-formats.rst | 6 + metatomic-core/include/metatomic.h | 149 ++++++++++++++++++++--- metatomic-core/src/c_api/model.rs | 149 ++++++++++++++++++++--- metatomic-core/src/lib.rs | 1 + 4 files changed, 277 insertions(+), 28 deletions(-) diff --git a/docs/src/core/reference/json-formats.rst b/docs/src/core/reference/json-formats.rst index d47cb2715..859f7d985 100644 --- a/docs/src/core/reference/json-formats.rst +++ b/docs/src/core/reference/json-formats.rst @@ -8,6 +8,8 @@ strings rather than dedicated C types. This page documents the exact JSON representation of each such structure, so that engines and models written in any language can produce and consume them. +.. _core-json-pair-options: + Pair list options ----------------- @@ -51,6 +53,8 @@ list). This is used for example by :c:func:`mta_system_add_pairs`, omitted, in which case it is treated as an empty list. +.. _core-json-quantity: + Quantities ---------- @@ -91,6 +95,8 @@ inputs and outputs. This is used for example in following: ``"atom"``, ``"system"`` or ``"atom_pair"``. +.. _core-json-model-metadata: + Model metadata -------------- diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index 1a4fc35f9..a4ba5e95e 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -85,42 +85,136 @@ typedef struct mta_system_t mta_system_t; typedef mta_opaque_string_t *mta_string_t; /** - * TODO + * A model that computes physical properties of atomistic systems. + * + * `mta_model_t` is a small virtual table: `data` holds the model's own state, + * and the function pointers describe what the model can do. A model is usually + * produced by a plugin's `load_model` callback (see `mta_load_model`) and then + * executed with `mta_execute_model`. + * + * Every callback receives `data` as its first argument. metatomic treats + * `data` as opaque and only hands it back to the callbacks. Callbacks should + * report any error by saving it with `mta_set_last_error` and returning a + * non-success `mta_status_t`. */ typedef struct mta_model_t { /** - * TODO + * Opaque pointer to the model's internal state + * + * Its layout and meaning are private to the model implementation. It is + * initialized by whoever creates the model (e.g. a plugin's `load_model`) + * and released by `unload`. */ void *data; /** - * TODO + * Release the resources owned by `model_data` + * + * Called exactly once when the model is no longer needed. May be `NULL` if + * the model owns no resources. + * + * @param model_data the model's `data` pointer + * @return `MTA_SUCCESS` on success, another status code on error */ enum mta_status_t (*unload)(void *model_data); /** - * TODO + * Get metadata describing the model (name, authors, references, ...) as a + * JSON string. + * + * @verbatim embed:rst:leading-asterisk + * The expected JSON structure is documented in :ref:`core-json-model-metadata`. + * @endverbatim + * + * @param model_data the model's `data` pointer + * @param metadata_json output string, set to a JSON-serialized + * `ModelMetadata` object. The + * caller takes ownership and must free it with `mta_string_free`. + * @return `MTA_SUCCESS` on success, another status code on error */ enum mta_status_t (*metadata)(const void *model_data, mta_string_t *metadata_json); /** - * TODO + * List the outputs this model is able to compute as a JSON string. + * + * @verbatim embed:rst:leading-asterisk + * The expected JSON structure for each output is documented in :ref:`core-json-quantity`. + * @endverbatim + * + * @param model_data the model's `data` pointer + * @param outputs_json output string, set to a JSON array of `Quantity` + * objects, one per supported output. The caller takes ownership and + * must free it with `mta_string_free`. + * @return `MTA_SUCCESS` on success, another status code on error */ enum mta_status_t (*supported_outputs)(const void *model_data, mta_string_t *outputs_json); /** - * TODO + * List the pair lists (neighbor lists) the model needs as input as a JSON + * string. + * + * @verbatim embed:rst:leading-asterisk + * + * The engine is expected to compute these and attach them to every system + * with :c:func:`mta_system_add_pairs` before calling + * :c:func:`mta_execute_model`. + * + * The expected JSON structure for each pair list is documented in :ref:`core-json-pair-options`. + * + * @endverbatim + * + * @param model_data the model's `data` pointer + * @param pair_options_json output string, set to a JSON array of + * `PairListOptions` objects. The caller takes ownership and must + * free it with `mta_string_free`. + * @return `MTA_SUCCESS` on success, another status code on error */ enum mta_status_t (*requested_pair_lists)(const void *model_data, mta_string_t *pair_options_json); /** - * TODO + * List the additional per-system inputs the model needs as a JSON string. + * + * @verbatim embed:rst:leading-asterisk + * + * These correspond to custom data the engine should attach to every system + * with :c:func:`mta_system_add_custom_data` before execution. + * + * The expected JSON structure for each input is documented in :ref:`core-json-quantity`. + * + * @endverbatim + * + * @param model_data the model's `data` pointer + * @param inputs_json output string, set to a JSON array of `Quantity` + * objects, one per requested input. The caller takes ownership and + * must free it with `mta_string_free`. + * @return `MTA_SUCCESS` on success, another status code on error */ enum mta_status_t (*requested_inputs)(const void *model_data, mta_string_t *inputs_json); /** - * TODO + * Run the model and compute the requested outputs + * + * @verbatim embed:rst:leading-asterisk + * + * This performs the model's actual computation. This should not be called + * directly, but rather through :c:func:`mta_execute_model`, which handles + * unit conversion and can check inputs and output data for consistency. + * + * @endverbatim + * + * @param model_data the model's `data` pointer + * @param systems array of `systems_count` systems to run the model on + * @param systems_count number of entries in `systems` + * @param selected_atoms optional labels selecting the subset of atoms to + * compute outputs for, or `NULL` to use all atoms. When set, it has the + * dimensions `"system"` and `"atom"` holding 0-based indices. + * @param requested_outputs_json JSON string containing an array of + * `Quantity`, one for each output the model should produce + * @param outputs array of `outputs_count` tensor maps to fill, one per + * requested output and in the same order + * @param outputs_count number of entries in `outputs`, must equal + * `requested_outputs_count` + * @return `MTA_SUCCESS` on success, another status code on error */ enum mta_status_t (*execute_inner)(void *model_data, const struct mta_system_t *const *systems, uintptr_t systems_count, const mts_labels_t *selected_atoms, - const char *const *requested_outputs_json, - uintptr_t requested_outputs_count, + const char *requested_outputs_json, mts_tensormap_t **outputs, uintptr_t outputs_count); } mta_model_t; @@ -292,20 +386,47 @@ enum mta_status_t mta_system_known_custom_data(const struct mta_system_t *system mta_string_t *names); /** - * TODO + * Execute a model to compute the requested outputs for a set of systems + * + * This is the main entry point to run a model loaded through the C API. It + * validates the arguments and delegates the computation to the model's + * `execute_inner` callback. + * + * @param model the model to execute + * @param systems array of `systems_count` systems to run the model on + * @param systems_count number of entries in `systems` + * @param selected_atoms optional labels selecting the subset of atoms to + * compute outputs for, or `NULL` to use all atoms + * @param requested_outputs_json JSON string containing an array of + * `Quantity`, one for each output the model should produce + * @param check_consistency if `true`, run additional checks on the + * inputs and on the data produced by the model + * @param outputs array of `outputs_count` tensor maps to fill, one per + * requested output and in the same order. The caller takes ownership of + * the returned tensor maps. + * @param outputs_count number of entries in `outputs`, must equal + * `requested_outputs_count` + * @return `MTA_SUCCESS` on success, another status code on error (the message + * is available through `mta_last_error`) */ enum mta_status_t mta_execute_model(struct mta_model_t model, const struct mta_system_t *const *systems, uintptr_t systems_count, const mts_labels_t *selected_atoms, - const char *const *requested_outputs_json, - uintptr_t requested_outputs_count, + const char *requested_outputs_json, bool check_consistency, mts_tensormap_t **outputs, uintptr_t outputs_count); /** - * TODO + * Render model metadata as a human-readable string + * + * @param metadata a JSON-serialized `ModelMetadata` object as produced by a + * model's `metadata` callback. Must not be null. + * @param printed output string, set to a human-readable rendering of the + * metadata. The caller takes ownership and must free it with + * `mta_string_free`. + * @return `MTA_SUCCESS` on success, another status code on error */ enum mta_status_t mta_format_metadata(const char *metadata, mta_string_t *printed); diff --git a/metatomic-core/src/c_api/model.rs b/metatomic-core/src/c_api/model.rs index b586f73c8..c9fc83610 100644 --- a/metatomic-core/src/c_api/model.rs +++ b/metatomic-core/src/c_api/model.rs @@ -3,62 +3,176 @@ use metatensor::c_api::{mts_labels_t, mts_tensormap_t}; use super::{mta_status_t, mta_string_t, mta_system_t}; -/// TODO +/// A model that computes physical properties of atomistic systems. +/// +/// `mta_model_t` is a small virtual table: `data` holds the model's own state, +/// and the function pointers describe what the model can do. A model is usually +/// produced by a plugin's `load_model` callback (see `mta_load_model`) and then +/// executed with `mta_execute_model`. +/// +/// Every callback receives `data` as its first argument. metatomic treats +/// `data` as opaque and only hands it back to the callbacks. Callbacks should +/// report any error by saving it with `mta_set_last_error` and returning a +/// non-success `mta_status_t`. #[repr(C)] #[allow(non_camel_case_types)] pub struct mta_model_t { - /// TODO + /// Opaque pointer to the model's internal state + /// + /// Its layout and meaning are private to the model implementation. It is + /// initialized by whoever creates the model (e.g. a plugin's `load_model`) + /// and released by `unload`. pub data: *mut c_void, - /// TODO + /// Release the resources owned by `model_data` + /// + /// Called exactly once when the model is no longer needed. May be `NULL` if + /// the model owns no resources. + /// + /// @param model_data the model's `data` pointer + /// @return `MTA_SUCCESS` on success, another status code on error pub unload: Option mta_status_t>, - /// TODO + /// Get metadata describing the model (name, authors, references, ...) as a + /// JSON string. + /// + /// @verbatim embed:rst:leading-asterisk + /// The expected JSON structure is documented in :ref:`core-json-model-metadata`. + /// @endverbatim + /// + /// @param model_data the model's `data` pointer + /// @param metadata_json output string, set to a JSON-serialized + /// `ModelMetadata` object. The + /// caller takes ownership and must free it with `mta_string_free`. + /// @return `MTA_SUCCESS` on success, another status code on error pub metadata: Option mta_status_t>, - /// TODO + /// List the outputs this model is able to compute as a JSON string. + /// + /// @verbatim embed:rst:leading-asterisk + /// The expected JSON structure for each output is documented in :ref:`core-json-quantity`. + /// @endverbatim + /// + /// @param model_data the model's `data` pointer + /// @param outputs_json output string, set to a JSON array of `Quantity` + /// objects, one per supported output. The caller takes ownership and + /// must free it with `mta_string_free`. + /// @return `MTA_SUCCESS` on success, another status code on error pub supported_outputs: Option mta_status_t>, - /// TODO + /// List the pair lists (neighbor lists) the model needs as input as a JSON + /// string. + /// + /// @verbatim embed:rst:leading-asterisk + /// + /// The engine is expected to compute these and attach them to every system + /// with :c:func:`mta_system_add_pairs` before calling + /// :c:func:`mta_execute_model`. + /// + /// The expected JSON structure for each pair list is documented in :ref:`core-json-pair-options`. + /// + /// @endverbatim + /// + /// @param model_data the model's `data` pointer + /// @param pair_options_json output string, set to a JSON array of + /// `PairListOptions` objects. The caller takes ownership and must + /// free it with `mta_string_free`. + /// @return `MTA_SUCCESS` on success, another status code on error pub requested_pair_lists: Option mta_status_t>, - /// TODO + /// List the additional per-system inputs the model needs as a JSON string. + /// + /// @verbatim embed:rst:leading-asterisk + /// + /// These correspond to custom data the engine should attach to every system + /// with :c:func:`mta_system_add_custom_data` before execution. + /// + /// The expected JSON structure for each input is documented in :ref:`core-json-quantity`. + /// + /// @endverbatim + /// + /// @param model_data the model's `data` pointer + /// @param inputs_json output string, set to a JSON array of `Quantity` + /// objects, one per requested input. The caller takes ownership and + /// must free it with `mta_string_free`. + /// @return `MTA_SUCCESS` on success, another status code on error pub requested_inputs: Option mta_status_t>, - /// TODO + /// Run the model and compute the requested outputs + /// + /// @verbatim embed:rst:leading-asterisk + /// + /// This performs the model's actual computation. This should not be called + /// directly, but rather through :c:func:`mta_execute_model`, which handles + /// unit conversion and can check inputs and output data for consistency. + /// + /// @endverbatim + /// + /// @param model_data the model's `data` pointer + /// @param systems array of `systems_count` systems to run the model on + /// @param systems_count number of entries in `systems` + /// @param selected_atoms optional labels selecting the subset of atoms to + /// compute outputs for, or `NULL` to use all atoms. When set, it has the + /// dimensions `"system"` and `"atom"` holding 0-based indices. + /// @param requested_outputs_json JSON string containing an array of + /// `Quantity`, one for each output the model should produce + /// @param outputs array of `outputs_count` tensor maps to fill, one per + /// requested output and in the same order + /// @param outputs_count number of entries in `outputs`, must equal + /// `requested_outputs_count` + /// @return `MTA_SUCCESS` on success, another status code on error pub execute_inner: Option mta_status_t>, } -/// TODO +/// Execute a model to compute the requested outputs for a set of systems +/// +/// This is the main entry point to run a model loaded through the C API. It +/// validates the arguments and delegates the computation to the model's +/// `execute_inner` callback. +/// +/// @param model the model to execute +/// @param systems array of `systems_count` systems to run the model on +/// @param systems_count number of entries in `systems` +/// @param selected_atoms optional labels selecting the subset of atoms to +/// compute outputs for, or `NULL` to use all atoms +/// @param requested_outputs_json JSON string containing an array of +/// `Quantity`, one for each output the model should produce +/// @param check_consistency if `true`, run additional checks on the +/// inputs and on the data produced by the model +/// @param outputs array of `outputs_count` tensor maps to fill, one per +/// requested output and in the same order. The caller takes ownership of +/// the returned tensor maps. +/// @param outputs_count number of entries in `outputs`, must equal +/// `requested_outputs_count` +/// @return `MTA_SUCCESS` on success, another status code on error (the message +/// is available through `mta_last_error`) #[no_mangle] pub unsafe extern "C" fn mta_execute_model( model: mta_model_t, systems: *const *const mta_system_t, systems_count: usize, selected_atoms: *const mts_labels_t, - requested_outputs_json: *const *const c_char, - requested_outputs_count: usize, + requested_outputs_json: *const c_char, check_consistency: bool, outputs: *mut *mut mts_tensormap_t, outputs_count: usize, @@ -66,7 +180,14 @@ pub unsafe extern "C" fn mta_execute_model( todo!() } -/// TODO +/// Render model metadata as a human-readable string +/// +/// @param metadata a JSON-serialized `ModelMetadata` object as produced by a +/// model's `metadata` callback. Must not be null. +/// @param printed output string, set to a human-readable rendering of the +/// metadata. The caller takes ownership and must free it with +/// `mta_string_free`. +/// @return `MTA_SUCCESS` on success, another status code on error #[no_mangle] pub unsafe extern "C" fn mta_format_metadata( metadata: *const c_char, diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index 09ec7a962..89834e45d 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -6,6 +6,7 @@ #![allow(clippy::unreadable_literal, clippy::option_if_let_else, clippy::module_name_repetitions)] #![allow(clippy::missing_errors_doc, clippy::missing_panics_doc, clippy::missing_safety_doc)] #![allow(clippy::similar_names, clippy::borrow_as_ptr, clippy::uninlined_format_args)] +#![allow(clippy::doc_markdown)] #![allow(clippy::let_underscore_untyped, clippy::manual_let_else, clippy::empty_line_after_doc_comments)] // To be removed later From 36a50e2f109a52d7cada551c930e40cada29deef Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Sun, 31 May 2026 18:41:35 +0200 Subject: [PATCH 14/43] Add ModelCapabilities to the JSON structs --- docs/src/core/reference/json-formats.rst | 70 +++++ docs/src/core/units.rst | 2 +- metatomic-core/include/metatomic.h | 18 +- metatomic-core/src/c_api/model.rs | 20 +- metatomic-core/src/lib.rs | 2 +- metatomic-core/src/metadata.rs | 373 +++++++++++++++++++++-- metatomic-core/src/quantities.rs | 10 +- metatomic-core/src/units.rs | 20 ++ 8 files changed, 486 insertions(+), 29 deletions(-) diff --git a/docs/src/core/reference/json-formats.rst b/docs/src/core/reference/json-formats.rst index 859f7d985..f7e3f468b 100644 --- a/docs/src/core/reference/json-formats.rst +++ b/docs/src/core/reference/json-formats.rst @@ -154,3 +154,73 @@ The JSON representation of a model's metadata. This is used for example by ``extra`` An object with string values, providing any additional key-value pairs the model author wishes to include. This can be used for any purpose. + +.. _core-json-model-capabilities: + +Model capabilities +------------------ + +The JSON representation of a model's capabilities, describing which outputs it +provides, which atomic types it supports, and other constraints. This is used +for example by :c:member:`mta_model_t.capabilities`. + +.. code-block:: json + + { + "type": "metatomic_model_capabilities", + "outputs": [ + { + "type": "metatomic_quantity", + "name": "energy", + "unit": "eV", + "sample_kind": "system", + "gradients": ["positions"], + "description": "Potential energy of the system" + }, + { + "type": "metatomic_quantity", + "name": "energy/pbe0", + "unit": "eV", + "sample_kind": "system", + "gradients": ["positions", "strain"], + "description": "Potential energy of the system" + }, + ], + "atomic_types": [1, 6, 8], + "interaction_range": 5.0, + "length_unit": "angstrom", + "supported_devices": ["cpu", "cuda"], + "dtype": "float32" + } + +``type`` + Must be the string ``"metatomic_model_capabilities"``. + +``outputs`` + Array of :ref:`quantity objects ` describing the + outputs this model can provide. + +``atomic_types`` + Array of integers listing the atomic types this model supports. The meaning + of these integers is up to the model, and is not required to be the atomic + numbers. + +``interaction_range`` + The interaction range of the model in the length unit of the model. This is + the maximum distance between two atoms for which the model's output can + depend on their relative position. Must be a non-negative number. + +``length_unit`` + String identifying the length unit used by the model, e.g. ``"angstrom"`` or + ``"nanometer"``. This must be a valid :ref:`unit expression ` with + dimensions compatible with length. + +``supported_devices`` + Array of strings listing the devices on which the model can run. Valid + values are ``"cpu"``, ``"cuda"``, ``"rocm"``, and ``"metal"``. + +``dtype`` + The data type of the model, used for all inputs and outputs. Must be either + ``"float32"`` or ``"float64"``. The model is free to use different data + types for internal computations, but all inputs and outputs must be in this + data type. diff --git a/docs/src/core/units.rst b/docs/src/core/units.rst index d5f3dd776..c9415ed9f 100644 --- a/docs/src/core/units.rst +++ b/docs/src/core/units.rst @@ -1,4 +1,4 @@ -.. _core-unit-expressions: +.. _units: Units ^^^^^ diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index a4ba5e95e..c8077b5c5 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -116,6 +116,20 @@ typedef struct mta_model_t { * @return `MTA_SUCCESS` on success, another status code on error */ enum mta_status_t (*unload)(void *model_data); + /** + * Get the capabilities of the model as a JSON string. + * + * @verbatim embed:rst:leading-asterisk + * The expected JSON structure is documented in :ref:`core-json-model-capabilities`. + * @endverbatim + * + * @param model_data the model's `data` pointer + * @param capabilities_json output string, set to a JSON-serialized + * `ModelCapabilities` object. The caller takes ownership and must + * free it with `mta_string_free`. + * @return `MTA_SUCCESS` on success, another status code on error + */ + enum mta_status_t (*capabilities)(const void *model_data, mta_string_t *capabilities_json); /** * Get metadata describing the model (name, authors, references, ...) as a * JSON string. @@ -126,8 +140,8 @@ typedef struct mta_model_t { * * @param model_data the model's `data` pointer * @param metadata_json output string, set to a JSON-serialized - * `ModelMetadata` object. The - * caller takes ownership and must free it with `mta_string_free`. + * `ModelMetadata` object. The caller takes ownership and must + * free it with `mta_string_free`. * @return `MTA_SUCCESS` on success, another status code on error */ enum mta_status_t (*metadata)(const void *model_data, mta_string_t *metadata_json); diff --git a/metatomic-core/src/c_api/model.rs b/metatomic-core/src/c_api/model.rs index c9fc83610..bf03e2a46 100644 --- a/metatomic-core/src/c_api/model.rs +++ b/metatomic-core/src/c_api/model.rs @@ -33,6 +33,22 @@ pub struct mta_model_t { /// @return `MTA_SUCCESS` on success, another status code on error pub unload: Option mta_status_t>, + /// Get the capabilities of the model as a JSON string. + /// + /// @verbatim embed:rst:leading-asterisk + /// The expected JSON structure is documented in :ref:`core-json-model-capabilities`. + /// @endverbatim + /// + /// @param model_data the model's `data` pointer + /// @param capabilities_json output string, set to a JSON-serialized + /// `ModelCapabilities` object. The caller takes ownership and must + /// free it with `mta_string_free`. + /// @return `MTA_SUCCESS` on success, another status code on error + pub capabilities: Option mta_status_t>, + /// Get metadata describing the model (name, authors, references, ...) as a /// JSON string. /// @@ -42,8 +58,8 @@ pub struct mta_model_t { /// /// @param model_data the model's `data` pointer /// @param metadata_json output string, set to a JSON-serialized - /// `ModelMetadata` object. The - /// caller takes ownership and must free it with `mta_string_free`. + /// `ModelMetadata` object. The caller takes ownership and must + /// free it with `mta_string_free`. /// @return `MTA_SUCCESS` on success, another status code on error pub metadata: Option for JsonValue { } } -impl TryFrom for PairListOptions { +impl<'a> TryFrom<&'a JsonValue> for PairListOptions { type Error = Error; - fn try_from(value: JsonValue) -> Result { + fn try_from(value: &'a JsonValue) -> Result { if !value.is_object() { return Err(Error::Serialization( "invalid JSON data for PairListOptions, expected an object".into() @@ -168,19 +169,19 @@ fn read_references(object: &JsonValue, key: &str) -> Result, Error> Ok(references) } -impl TryFrom for References { +impl<'a> TryFrom<&'a JsonValue> for References { type Error = Error; - fn try_from(value: JsonValue) -> Result { + fn try_from(value: &'a JsonValue) -> Result { if !value.is_object() { return Err(Error::Serialization( "invalid JSON data for references in ModelMetadata, expected an object".into() )); } - let model = read_references(&value, "model")?; - let architecture = read_references(&value, "architecture")?; - let implementation = read_references(&value, "implementation")?; + let model = read_references(value, "model")?; + let architecture = read_references(value, "architecture")?; + let implementation = read_references(value, "implementation")?; Ok(References { model, architecture, implementation }) } @@ -217,10 +218,10 @@ impl From for JsonValue { } } -impl TryFrom for ModelMetadata { +impl<'a> TryFrom<&'a JsonValue> for ModelMetadata { type Error = Error; - fn try_from(value: JsonValue) -> Result { + fn try_from(value: &'a JsonValue) -> Result { if !value.is_object() { return Err(Error::Serialization( "invalid JSON data for ModelMetadata, expected an object".into() @@ -253,7 +254,7 @@ impl TryFrom for ModelMetadata { "'description' in JSON for ModelMetadata must be a string".into() ))?.to_string(); - let references = References::try_from(value["references"].clone())?; + let references = References::try_from(&value["references"])?; if !value["extra"].is_object() { return Err(Error::Serialization( @@ -279,6 +280,211 @@ impl TryFrom for ModelMetadata { } } +/// The data type of a model, used for all inputs and outputs. The model can +/// still internally use a different data type for its calculations, but it will +/// get inputs in this type and must produce outputs in this type. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum DType { + /// 32-bit floating point, following the IEEE 754 standard + Float32, + /// 64-bit floating point, following the IEEE 754 standard + Float64, +} + +impl From for JsonValue { + fn from(value: DType) -> Self { + match value { + DType::Float32 => "float32".into(), + DType::Float64 => "float64".into(), + } + } +} + +impl<'a> TryFrom<&'a JsonValue> for DType { + type Error = Error; + + fn try_from(value: &'a JsonValue) -> Result { + if let Some(s) = value.as_str() { + match s { + "float32" => Ok(DType::Float32), + "float64" => Ok(DType::Float64), + _ => Err(Error::Serialization( + "invalid string for dtype in JSON for ModelCapabilities, expected 'float32' or 'float64'".into() + )), + } + } else { + Err(Error::Serialization( + "dtype in JSON for ModelCapabilities must be a string".into() + )) + } + } +} + +/// A device on which a model can run. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct Device(dlpk::DLDeviceType); + +impl From for JsonValue { + fn from(value: Device) -> Self { + match value.0 { + dlpk::DLDeviceType::kDLCPU => "cpu".into(), + dlpk::DLDeviceType::kDLCUDA => "cuda".into(), + dlpk::DLDeviceType::kDLROCM => "rocm".into(), + dlpk::DLDeviceType::kDLMetal => "metal".into(), + dlpk::DLDeviceType::kDLCUDAHost | dlpk::DLDeviceType::kDLCUDAManaged => { + // These refer to memory devices more than execution devices + panic!("Do not use kDLCUDAHost or kDLCUDAManaged, use kDLCUDA instead."); + } + dlpk::DLDeviceType::kDLROCMHost => { + // This refers to a memory device more than an execution device + panic!("Do not use kDLROCMHost, use kDLROCM instead."); + } + _ => { + // We don't want to expose other device types until we have a + // use case for them, and we don't want to accidentally leak + // them if they're added in the future + panic!("unsupported device type: {:?}", value.0); + } + } + } +} + +impl<'a> TryFrom<&'a JsonValue> for Device { + type Error = Error; + + fn try_from(value: &'a JsonValue) -> Result { + if let Some(s) = value.as_str() { + match s { + "cpu" => Ok(Device(dlpk::DLDeviceType::kDLCPU)), + "cuda" => Ok(Device(dlpk::DLDeviceType::kDLCUDA)), + "rocm" => Ok(Device(dlpk::DLDeviceType::kDLROCM)), + "metal" => Ok(Device(dlpk::DLDeviceType::kDLMetal)), + _ => Err(Error::Serialization( + "invalid string for device in JSON for ModelCapabilities, expected 'cpu', 'cuda', 'rocm', or 'metal'".into() + )), + } + } else { + Err(Error::Serialization( + "device in JSON for ModelCapabilities must be a string".into() + )) + } + } +} + +/// Capabilities about a model: which outputs it provides, which atoms it +/// supports, etc. +#[derive(Debug, Clone)] +pub struct ModelCapabilities { + /// The outputs this model can provide + pub outputs: Vec, + /// The atomic types this model supports. The meaning of the integers in + /// this list is up to the model, and is not required to be the atomic + /// numbers. + pub atomic_types: Vec, + /// The interaction range of the model (in the length unit of the model), + /// i.e. the maximum distance between two atoms for which the model's output + /// can depend on their relative position. + pub interaction_range: f64, + /// The length unit of the model, e.g. "angstrom" or "nanometer". This is + /// used to interpret the `interaction_range` and convert the inputs. + pub length_unit: String, + /// The devices on which the model can run, e.g. `["cpu", "cuda"]`. + pub supported_devices: Vec, + /// The data type of the model, used for all inputs and outputs. + pub dtype: DType, +} + +impl From for JsonValue { + fn from(value: ModelCapabilities) -> Self { + let mut result = JsonValue::new_object(); + result["type"] = "metatomic_model_capabilities".into(); + result["outputs"] = value.outputs.into(); + result["atomic_types"] = value.atomic_types.into(); + result["interaction_range"] = value.interaction_range.into(); + result["length_unit"] = value.length_unit.into(); + result["supported_devices"] = value.supported_devices.into(); + result["dtype"] = value.dtype.into(); + return result; + } +} + +impl<'a> TryFrom<&'a JsonValue> for ModelCapabilities { + type Error = Error; + + fn try_from(value: &'a JsonValue) -> Result { + if !value.is_object() { + return Err(Error::Serialization( + "invalid JSON data for ModelCapabilities, expected an object".into() + )); + } + + if value["type"].as_str() != Some("metatomic_model_capabilities") { + return Err(Error::Serialization( + "'type' in JSON for ModelCapabilities must be 'metatomic_model_capabilities'".into() + )); + } + + let mut outputs = Vec::new(); + if !value["outputs"].is_array() { + return Err(Error::Serialization( + "'outputs' in JSON for ModelCapabilities must be an array".into() + )); + } + for output in value["outputs"].members() { + outputs.push(Quantity::try_from(output)?); + } + + + let mut atomic_types = Vec::new(); + if !value["atomic_types"].is_array() { + return Err(Error::Serialization( + "'atomic_types' in JSON for ModelCapabilities must be an array".into() + )); + } + + for atomic_type in value["atomic_types"].members() { + let atomic_type = atomic_type.as_i64().ok_or_else(|| Error::Serialization( + "'atomic_types' in JSON for ModelCapabilities must be an array of integers".into() + ))?; + atomic_types.push(atomic_type); + } + + let interaction_range = value["interaction_range"].as_f64().ok_or_else(|| Error::Serialization( + "'interaction_range' in JSON for ModelCapabilities must be a number".into() + ))?; + if interaction_range < 0.0 { + return Err(Error::Serialization( + "'interaction_range' in JSON for ModelCapabilities must be non-negative".into() + )); + } + + let length_unit = value["length_unit"].as_str().ok_or_else(|| Error::Serialization( + "'length_unit' in JSON for ModelCapabilities must be a string".into() + ))?.to_string(); + validate_unit(&length_unit, "m", Some("'length_unit' in JSON for ModelCapabilities"))?; + + let mut supported_devices = Vec::new(); + if !value["supported_devices"].is_array() { + return Err(Error::Serialization( + "'supported_devices' in JSON for ModelCapabilities must be an array".into() + )); + } + for device in value["supported_devices"].members() { + supported_devices.push(Device::try_from(device)?); + } + + let dtype = DType::try_from(&value["dtype"])?; + + Ok(ModelCapabilities { + outputs, + atomic_types, + interaction_range, + length_unit, + supported_devices, + dtype, + }) + } +} #[cfg(test)] mod tests { @@ -304,7 +510,7 @@ mod tests { assert_eq!(json["full_list"].as_bool(), Some(true)); assert_eq!(json["strict"].as_bool(), Some(false)); - let parsed = PairListOptions::try_from(json).unwrap(); + let parsed = PairListOptions::try_from(&json).unwrap(); assert_eq!(parsed.cutoff.to_bits(), options.cutoff.to_bits()); assert_eq!(parsed.full_list, options.full_list); assert_eq!(parsed.strict, options.strict); @@ -315,7 +521,7 @@ mod tests { fn cutoff_keeps_full_precision() { let mut options = example(); options.cutoff = 1.0 / 3.0; - let parsed = PairListOptions::try_from(JsonValue::from(options.clone())).unwrap(); + let parsed = PairListOptions::try_from(&JsonValue::from(options.clone())).unwrap(); assert_eq!(parsed.cutoff.to_bits(), options.cutoff.to_bits()); } @@ -323,7 +529,7 @@ mod tests { fn requestors_are_optional() { let mut json: JsonValue = example().into(); json.remove("requestors"); - let parsed = PairListOptions::try_from(json).unwrap(); + let parsed = PairListOptions::try_from(&json).unwrap(); assert!(parsed.requestors.is_empty()); } @@ -380,7 +586,7 @@ mod tests { ]; for (json, expected) in cases { - let error = PairListOptions::try_from(json).expect_err("expected an error"); + let error = PairListOptions::try_from(&json).expect_err("expected an error"); assert_eq!(error.to_string(), expected); } } @@ -390,7 +596,7 @@ mod tests { let mut json: JsonValue = example().into(); json["requestors"] = json::array![ "a", "", "b", "a" ]; - let parsed = PairListOptions::try_from(json).unwrap(); + let parsed = PairListOptions::try_from(&json).unwrap(); assert_eq!(parsed.requestors, vec!["a".to_string(), "b".to_string()]); } } @@ -431,7 +637,7 @@ mod tests { assert_eq!(json["extra"]["key1"].as_str(), Some("value1")); assert_eq!(json["extra"]["key2"].as_str(), Some("value2")); - let parsed = ModelMetadata::try_from(json).unwrap(); + let parsed = ModelMetadata::try_from(&json).unwrap(); assert_eq!(parsed.name, metadata.name); assert_eq!(parsed.authors, metadata.authors); assert_eq!(parsed.description, metadata.description); @@ -494,7 +700,138 @@ mod tests { ]; for (json, expected) in cases { - let error = ModelMetadata::try_from(json).expect_err("expected an error"); + let error = ModelMetadata::try_from(&json).expect_err("expected an error"); + assert_eq!(error.to_string(), expected); + } + } + } + + mod model_capabilities { + use super::super::*; + + fn example() -> ModelCapabilities { + ModelCapabilities { + outputs: vec![ + Quantity { + name: "energy".into(), + unit: "eV".into(), + description: Some("total energy".into()), + gradients: vec![crate::Gradients::Positions], + sample_kind: crate::SampleKind::System, + }, + Quantity { + name: "charge".into(), + unit: "e".into(), + description: None, + gradients: vec![], + sample_kind: crate::SampleKind::Atom, + }, + ], + atomic_types: vec![1, 6, 8], + interaction_range: 5.0, + length_unit: "Angstrom".into(), + supported_devices: vec![Device(dlpk::DLDeviceType::kDLCPU), Device(dlpk::DLDeviceType::kDLCUDA)], + dtype: DType::Float32, + } + } + + #[test] + fn roundtrip() { + let capabilities = example(); + let json: JsonValue = capabilities.clone().into(); + + assert_eq!(json["type"].as_str(), Some("metatomic_model_capabilities")); + assert_eq!(json["outputs"][0]["name"].as_str(), Some("energy")); + assert_eq!(json["outputs"][1]["name"].as_str(), Some("charge")); + assert_eq!(json["atomic_types"][0].as_i64(), Some(1)); + assert_eq!(json["atomic_types"][1].as_i64(), Some(6)); + assert_eq!(json["atomic_types"][2].as_i64(), Some(8)); + assert_eq!(json["interaction_range"].as_f64(), Some(5.0)); + assert_eq!(json["length_unit"].as_str(), Some("Angstrom")); + assert_eq!(json["supported_devices"][0].as_str(), Some("cpu")); + assert_eq!(json["supported_devices"][1].as_str(), Some("cuda")); + assert_eq!(json["dtype"].as_str(), Some("float32")); + + let parsed = ModelCapabilities::try_from(&json).unwrap(); + assert_eq!(parsed.outputs.len(), 2); + assert_eq!(parsed.outputs[0].name, "energy"); + assert_eq!(parsed.outputs[1].name, "charge"); + assert_eq!(parsed.atomic_types, vec![1, 6, 8]); + assert_eq!(parsed.interaction_range.to_bits(), 5.0_f64.to_bits()); + assert_eq!(parsed.length_unit, "Angstrom"); + assert_eq!(parsed.supported_devices.len(), 2); + assert_eq!(parsed.dtype, DType::Float32); + } + + #[test] + fn rejects_invalid_json() { + let mut wrong_type = JsonValue::from(example()); + wrong_type["type"] = "something-else".into(); + + let mut non_array_outputs = JsonValue::from(example()); + non_array_outputs["outputs"] = "energy".into(); + + let mut non_array_atomic_types = JsonValue::from(example()); + non_array_atomic_types["atomic_types"] = "1".into(); + + let mut non_integer_atomic_type = JsonValue::from(example()); + non_integer_atomic_type["atomic_types"] = json::array![1, "x"]; + + let mut missing_interaction_range = JsonValue::from(example()); + missing_interaction_range.remove("interaction_range"); + + let mut negative_interaction_range = JsonValue::from(example()); + negative_interaction_range["interaction_range"] = (-1.0).into(); + + let mut missing_length_unit = JsonValue::from(example()); + missing_length_unit.remove("length_unit"); + + let mut wrong_dimension_length_unit = JsonValue::from(example()); + wrong_dimension_length_unit["length_unit"] = "eV".into(); + + let mut non_array_supported_devices = JsonValue::from(example()); + non_array_supported_devices["supported_devices"] = "cpu".into(); + + let mut invalid_device = JsonValue::from(example()); + invalid_device["supported_devices"] = json::array!["cpu", "wat"]; + + let mut missing_dtype = JsonValue::from(example()); + missing_dtype.remove("dtype"); + + let mut invalid_dtype = JsonValue::from(example()); + invalid_dtype["dtype"] = "float16".into(); + + let cases: Vec<(JsonValue, &str)> = vec![ + (JsonValue::from("not an object"), + "serialization error: invalid JSON data for ModelCapabilities, expected an object"), + (wrong_type, + "serialization error: 'type' in JSON for ModelCapabilities must be 'metatomic_model_capabilities'"), + (non_array_outputs, + "serialization error: 'outputs' in JSON for ModelCapabilities must be an array"), + (non_array_atomic_types, + "serialization error: 'atomic_types' in JSON for ModelCapabilities must be an array"), + (non_integer_atomic_type, + "serialization error: 'atomic_types' in JSON for ModelCapabilities must be an array of integers"), + (missing_interaction_range, + "serialization error: 'interaction_range' in JSON for ModelCapabilities must be a number"), + (negative_interaction_range, + "serialization error: 'interaction_range' in JSON for ModelCapabilities must be non-negative"), + (missing_length_unit, + "serialization error: 'length_unit' in JSON for ModelCapabilities must be a string"), + (wrong_dimension_length_unit, + "invalid parameter: dimension mismatch in 'length_unit' in JSON for ModelCapabilities: 'eV' has dimension [L^2 T^-2 M] but expected dimension [L]"), + (non_array_supported_devices, + "serialization error: 'supported_devices' in JSON for ModelCapabilities must be an array"), + (invalid_device, + "serialization error: invalid string for device in JSON for ModelCapabilities, expected 'cpu', 'cuda', 'rocm', or 'metal'"), + (missing_dtype, + "serialization error: dtype in JSON for ModelCapabilities must be a string"), + (invalid_dtype, + "serialization error: invalid string for dtype in JSON for ModelCapabilities, expected 'float32' or 'float64'"), + ]; + + for (json, expected) in cases { + let error = ModelCapabilities::try_from(&json).expect_err("expected an error"); assert_eq!(error.to_string(), expected); } } diff --git a/metatomic-core/src/quantities.rs b/metatomic-core/src/quantities.rs index 9d1dfebbf..93727c837 100644 --- a/metatomic-core/src/quantities.rs +++ b/metatomic-core/src/quantities.rs @@ -193,10 +193,10 @@ impl From for JsonValue { } -impl TryFrom for Quantity { +impl<'a> TryFrom<&'a JsonValue> for Quantity { type Error = Error; - fn try_from(value: JsonValue) -> Result { + fn try_from(value: &'a JsonValue) -> Result { if !value.is_object() { return Err(Error::Serialization( "invalid JSON data for Quantity, expected an object".into() @@ -272,7 +272,7 @@ mod tests { assert_eq!(json["gradients"][0].as_str(), Some("positions")); assert_eq!(json["sample_kind"].as_str(), Some("atom")); - let parsed = Quantity::try_from(json).unwrap(); + let parsed = Quantity::try_from(&json).unwrap(); assert_eq!(parsed.name, "energy"); assert_eq!(parsed.unit, "eV"); assert_eq!(parsed.gradients, vec![Gradients::Positions]); @@ -295,7 +295,7 @@ mod tests { gradients: grads.clone(), sample_kind: sample.clone(), }; - let parsed = Quantity::try_from(JsonValue::from(quantity.clone())).unwrap(); + let parsed = Quantity::try_from(&JsonValue::from(quantity.clone())).unwrap(); assert_eq!(parsed.name, quantity.name); assert_eq!(parsed.unit, quantity.unit); assert_eq!(parsed.gradients, grads); @@ -352,7 +352,7 @@ mod tests { ]; for (json, expected) in cases { - let error = Quantity::try_from(json).expect_err("expected an error"); + let error = Quantity::try_from(&json).expect_err("expected an error"); assert_eq!(error.to_string(), expected); } } diff --git a/metatomic-core/src/units.rs b/metatomic-core/src/units.rs index 4cfff2c1d..4239b86be 100644 --- a/metatomic-core/src/units.rs +++ b/metatomic-core/src/units.rs @@ -574,6 +574,26 @@ pub fn unit_conversion_factor(from_unit: &str, to_unit: &str) -> Result) -> Result<(), Error> { + let unit_value = parse_unit_expression(unit)?; + let reference_value = parse_unit_expression(reference_unit)?; + + if unit_value.dim != reference_value.dim { + return Err(Error::InvalidParameter(format!( + "dimension mismatch{}: '{}' has dimension {} but expected dimension {}", + context.map_or_else(String::new, |c| format!(" in {}", c)), + unit, + unit_value.dim, + reference_value.dim + ))); + } + + Ok(()) +} + + #[cfg(test)] #[allow(clippy::float_cmp)] mod tests { From c648ab271ee491290f7a51f6858bb0d9306b80f4 Mon Sep 17 00:00:00 2001 From: frostedoyster Date: Mon, 1 Jun 2026 14:49:24 +0200 Subject: [PATCH 15/43] Implement plugin registration and loading, model loading Co-Authored-By: Guillaume Fraux --- metatomic-core/Cargo.toml | 1 + metatomic-core/build.rs | 42 +++- metatomic-core/include/metatomic.h | 133 +++++++++++-- metatomic-core/src/c_api/mod.rs | 2 +- metatomic-core/src/c_api/model.rs | 15 ++ metatomic-core/src/c_api/plugin.rs | 151 ++++++++++++-- metatomic-core/src/c_api/status.rs | 17 +- metatomic-core/src/lib.rs | 24 ++- metatomic-core/src/model.rs | 11 + metatomic-core/src/plugin.rs | 188 ++++++++++++++++-- metatomic-core/tests/CMakeLists.txt | 3 + metatomic-core/tests/plugins.cpp | 45 +++++ .../tests/test-plugins/CMakeLists.txt | 14 ++ metatomic-core/tests/test-plugins/bad-abi.c | 17 ++ metatomic-core/tests/test-plugins/plugin.c | 17 ++ scripts/include/stdio.h | 0 16 files changed, 619 insertions(+), 61 deletions(-) create mode 100644 metatomic-core/tests/plugins.cpp create mode 100644 metatomic-core/tests/test-plugins/CMakeLists.txt create mode 100644 metatomic-core/tests/test-plugins/bad-abi.c create mode 100644 metatomic-core/tests/test-plugins/plugin.c create mode 100644 scripts/include/stdio.h diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index 2335505a3..dace0f8d5 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -18,6 +18,7 @@ metatensor = { version = "0.3.0" } once_cell = "1" dlpk = "0.3" json = "0.12" +libloading = "0.8" [build-dependencies] diff --git a/metatomic-core/build.rs b/metatomic-core/build.rs index 01f95d71c..1a58845d0 100644 --- a/metatomic-core/build.rs +++ b/metatomic-core/build.rs @@ -22,7 +22,8 @@ fn main() { config.documentation_style = cbindgen::DocumentationStyle::Doxy; config.line_endings = cbindgen::LineEndingStyle::LF; config.autogen_warning = Some(generated_comment.into()); - config.includes.push("metatensor.h".into()); + config.sys_includes.push("stdio.h".into()); + config.sys_includes.push("metatensor.h".into()); config.includes.push("metatomic/version.h".into()); config.export = cbindgen::ExportConfig { @@ -33,6 +34,45 @@ fn main() { }; config.after_includes = Some(" +#ifndef MTA_EXPORT + #if defined(_WIN32) || defined(__CYGWIN__) + #define MTA_EXPORT __declspec(dllexport) + #else + #define MTA_EXPORT __attribute__((visibility(\"default\"))) + #endif +#endif + +#ifndef MTA_EXTERN_C + #ifdef __cplusplus + #define MTA_EXTERN_C extern \"C\" + #else + #define MTA_EXTERN_C + #endif +#endif + +/** + * Define the exported plugin entry points. + * + * This macro should be used once in each plugin shared library with a + * `mta_plugin_t` expression. It exports the plugin ABI version and a + * registration function used by `mta_load_plugin`. + */ +#define MTA_REGISTER_PLUGIN(register_fn_name, ...) \\ + MTA_EXTERN_C MTA_EXPORT mta_status_t mta_plugin_init(int abi, void *data) { \\ + if (abi != MTA_ABI_VERSION) { \\ + char message[256]; \\ + snprintf(message, sizeof(message), \\ + \"Metatomic plugin ABI version mismatch: expected %d, got %d\", \\ + MTA_ABI_VERSION, abi \\ + ); \\ + mta_set_last_error(message, \"MTA_REGISTER_PLUGIN\", NULL, NULL); \\ + return MTA_INVALID_PARAMETER_ERROR; \\ + } \\ + mta_status_t (*register_fn_name)(mta_plugin_t) = (mta_status_t (*)(mta_plugin_t))data; \\ + __VA_ARGS__; \\ + return MTA_SUCCESS; \\ + } + /** Heap allocated storage for mta_string_t */ typedef struct mta_opaque_string_t mta_opaque_string_t;".into()); diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index c8077b5c5..a2a549e78 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -12,15 +12,59 @@ #include #include #include -#include "metatensor.h" +#include +#include #include "metatomic/version.h" +#ifndef MTA_EXPORT + #if defined(_WIN32) || defined(__CYGWIN__) + #define MTA_EXPORT __declspec(dllexport) + #else + #define MTA_EXPORT __attribute__((visibility("default"))) + #endif +#endif + +#ifndef MTA_EXTERN_C + #ifdef __cplusplus + #define MTA_EXTERN_C extern "C" + #else + #define MTA_EXTERN_C + #endif +#endif + +/** + * Define the exported plugin entry points. + * + * This macro should be used once in each plugin shared library with a + * `mta_plugin_t` expression. It exports the plugin ABI version and a + * registration function used by `mta_load_plugin`. + */ +#define MTA_REGISTER_PLUGIN(register_fn_name, ...) \ + MTA_EXTERN_C MTA_EXPORT mta_status_t mta_plugin_init(int abi, void *data) { \ + if (abi != MTA_ABI_VERSION) { \ + char message[256]; \ + snprintf(message, sizeof(message), \ + "Metatomic plugin ABI version mismatch: expected %d, got %d", \ + MTA_ABI_VERSION, abi \ + ); \ + mta_set_last_error(message, "MTA_REGISTER_PLUGIN", NULL, NULL); \ + return MTA_INVALID_PARAMETER_ERROR; \ + } \ + mta_status_t (*register_fn_name)(mta_plugin_t) = (mta_status_t (*)(mta_plugin_t))data; \ + __VA_ARGS__; \ + return MTA_SUCCESS; \ + } + /** Heap allocated storage for mta_string_t */ typedef struct mta_opaque_string_t mta_opaque_string_t; /** - * TODO + * ABI version of the metatomic plugin interface. + * + * This increases anytime the plugin or model C API changes in a non backward + * compatible way. Plugins compiled with an incompatible ABI version will be + * rejected at registration time. */ #define MTA_ABI_VERSION 1 @@ -47,11 +91,10 @@ typedef enum mta_status_t { */ MTA_SERIALIZATION_ERROR = 3, /** - * Status code indicating errors that come from callbacks provided by the user. - * The error message and arbitrary data can be stored using `mta_set_last_error`, - * and retrieved using `mta_last_error`. + * Status code used by plugins when a model is not supported by the + * current plugin */ - MTA_CALLBACK_ERROR = 254, + MTA_MODEL_NOT_SUPPORTED_ERROR = 4, /** * Status code used when there is an internal error */ @@ -234,15 +277,41 @@ typedef struct mta_model_t { } mta_model_t; /** - * TODO + * A metatomic plugin definition. */ typedef struct mta_plugin_t { /** - * TODO + * ABI version this plugin was compiled against, this should be set to + * `MTA_ABI_VERSION` when creating the plugin struct. + */ + int32_t abi_version; + /** + * Name of the plugin, as a null-terminated UTF-8 string. This is the name + * specified in `mta_load_model` when trying to load a model with a + * specific plugin. The name must be unique among all registered plugins. */ const char *name; /** - * TODO + * Callback function to load a model. This function should try to load a + * model from `load_from` (which can be a file path, a model name, etc.) + * and a set of key/values options passed as a JSON string. + * + * If the plugin can load the model, it should fill `model` with a pointer + * to a valid `mta_model_t` struct and return `MTA_SUCCESS`. If the data in + * `load_from` does not correspond to a model supported by the plugin, it + * should return `MTA_MODEL_NOT_SUPPORTED_ERROR`. If an error occurs while + * loading the model, it should return another status code and save an + * error message with `mta_set_last_error`. + * + * @param load_from a null-terminated UTF-8 string describing where to load + * the model from (e.g. a file path, a model name, etc.). The interpretation + * of this string is up to the plugin. + * @param options_json a null-terminated UTF-8 string containing a set of + * string keys and string value options for loading the model. + * @param model output pointer to the loaded model. The caller takes ownership of + * the model and must unload it when the model is no longer needed. + * @return `MTA_SUCCESS` if the model was loaded successfully, `MTA_MODEL_NOT_SUPPORTED_ERROR` + * if the plugin can not load the model, or another status code if an error occurs. */ enum mta_status_t (*load_model)(const char *load_from, const char *options_json, @@ -445,17 +514,55 @@ enum mta_status_t mta_execute_model(struct mta_model_t model, enum mta_status_t mta_format_metadata(const char *metadata, mta_string_t *printed); /** - * TODO + * Register a plugin. This is passed as a callback to the `MTA_REGISTER_PLUGIN` + * macro, and should not be called directly by C or C++ plugin implementations. + * + * @param plugin the plugin to register + * @return `MTA_SUCCESS` if the plugin was registered successfully, or another + * status code if an error occurs. You can get more details about the error + * with `mta_last_error`. */ -void mta_register_plugin(struct mta_plugin_t plugin); +enum mta_status_t mta_register_plugin(struct mta_plugin_t plugin); /** - * TODO + * Load the shared library at `path` and register the plugin contained within. + * + * The library must export the symbols generated by the `MTA_REGISTER_PLUGIN` + * macro. + * + * @param path a null-terminated UTF-8 string containing the path to the plugin + * shared library + * @return `MTA_SUCCESS` if the plugin was loaded successfully, or another + * status code if an error occurs. You can get more details about the + * error with `mta_last_error`. */ enum mta_status_t mta_load_plugin(const char *path); /** - * TODO + * Load a model from `load_from` with the given options. + * + * If `plugin_name` is a NULL pointer, metatomic will try to determine the + * correct plugin to use by checking the `load_from` parameter. If we can not + * determine the correct plugin, we then try to load the model with each + * registered plugin until one succeeds. + * + * If `plugin_name` is given, then we only try to load the model with the + * specified plugin, and return an error if the plugin can not load the model. + * + * @param plugin_name optional null-terminated UTF-8 string containing the name + * of the plugin to use for loading the model, or `NULL` to let metatomic + * search for a correct plugin + * @param load_from a null-terminated UTF-8 string describing where to load the + * model from (e.g. a file path, a model name, etc.). The interpretation + * of this string is up to the plugin. + * @param options_json a null-terminated UTF-8 string containing a set of string + * keys and string value options for loading the model. The interpretation + * of these options is up to the plugin. + * @param model output pointer to the loaded model. The caller takes ownership of + * the model and must unload it when the model is no longer needed. + * @return `MTA_SUCCESS` if the model was loaded successfully, or another + * status code if an error occurs. You can get more details about the + * error with `mta_last_error`. */ enum mta_status_t mta_load_model(const char *plugin_name, const char *load_from, diff --git a/metatomic-core/src/c_api/mod.rs b/metatomic-core/src/c_api/mod.rs index bffa5003a..235b00296 100644 --- a/metatomic-core/src/c_api/mod.rs +++ b/metatomic-core/src/c_api/mod.rs @@ -15,4 +15,4 @@ mod model; pub use self::model::mta_model_t; mod plugin; -pub use self::plugin::{mta_plugin_t, mta_register_plugin, mta_load_model}; +pub use self::plugin::{mta_plugin_t, mta_register_plugin, mta_load_plugin, mta_load_model}; diff --git a/metatomic-core/src/c_api/model.rs b/metatomic-core/src/c_api/model.rs index bf03e2a46..d2050643b 100644 --- a/metatomic-core/src/c_api/model.rs +++ b/metatomic-core/src/c_api/model.rs @@ -160,6 +160,21 @@ pub struct mta_model_t { ) -> mta_status_t>, } +impl mta_model_t { + pub(crate) fn null() -> Self { + return mta_model_t { + data: std::ptr::null_mut(), + unload: None, + capabilities: None, + metadata: None, + supported_outputs: None, + requested_pair_lists: None, + requested_inputs: None, + execute_inner: None, + }; + } +} + /// Execute a model to compute the requested outputs for a set of systems /// /// This is the main entry point to run a model loaded through the C API. It diff --git a/metatomic-core/src/c_api/plugin.rs b/metatomic-core/src/c_api/plugin.rs index 6dbfc4add..76b9989d1 100644 --- a/metatomic-core/src/c_api/plugin.rs +++ b/metatomic-core/src/c_api/plugin.rs @@ -1,15 +1,43 @@ -use std::ffi::c_char; +use std::ffi::{CStr, c_char}; +use super::catch_unwind; use super::{mta_model_t, mta_status_t}; +use crate::Error; +use crate::Plugin; -/// TODO +/// A metatomic plugin definition. #[allow(non_camel_case_types)] #[repr(C)] pub struct mta_plugin_t { - /// TODO + /// ABI version this plugin was compiled against, this should be set to + /// `MTA_ABI_VERSION` when creating the plugin struct. + pub abi_version: i32, + + /// Name of the plugin, as a null-terminated UTF-8 string. This is the name + /// specified in `mta_load_model` when trying to load a model with a + /// specific plugin. The name must be unique among all registered plugins. pub name: *const c_char, - /// TODO + /// Callback function to load a model. This function should try to load a + /// model from `load_from` (which can be a file path, a model name, etc.) + /// and a set of key/values options passed as a JSON string. + /// + /// If the plugin can load the model, it should fill `model` with a pointer + /// to a valid `mta_model_t` struct and return `MTA_SUCCESS`. If the data in + /// `load_from` does not correspond to a model supported by the plugin, it + /// should return `MTA_MODEL_NOT_SUPPORTED_ERROR`. If an error occurs while + /// loading the model, it should return another status code and save an + /// error message with `mta_set_last_error`. + /// + /// @param load_from a null-terminated UTF-8 string describing where to load + /// the model from (e.g. a file path, a model name, etc.). The interpretation + /// of this string is up to the plugin. + /// @param options_json a null-terminated UTF-8 string containing a set of + /// string keys and string value options for loading the model. + /// @param model output pointer to the loaded model. The caller takes ownership of + /// the model and must unload it when the model is no longer needed. + /// @return `MTA_SUCCESS` if the model was loaded successfully, `MTA_MODEL_NOT_SUPPORTED_ERROR` + /// if the plugin can not load the model, or another status code if an error occurs. pub load_model: Option mta_status_t>, } -/// TODO +unsafe impl Send for mta_plugin_t {} + +/// Register a plugin. This is passed as a callback to the `MTA_REGISTER_PLUGIN` +/// macro, and should not be called directly by C or C++ plugin implementations. +/// +/// @param plugin the plugin to register +/// @return `MTA_SUCCESS` if the plugin was registered successfully, or another +/// status code if an error occurs. You can get more details about the error +/// with `mta_last_error`. #[no_mangle] -pub extern "C" fn mta_register_plugin(plugin: mta_plugin_t) { - todo!() +pub unsafe extern "C" fn mta_register_plugin(plugin: mta_plugin_t) -> mta_status_t { + catch_unwind(move || { + let plugin = Plugin::new(plugin)?; + crate::plugin::register_plugin(plugin)?; + Ok(()) + }) } -/// TODO +/// Load the shared library at `path` and register the plugin contained within. +/// +/// The library must export the symbols generated by the `MTA_REGISTER_PLUGIN` +/// macro. +/// +/// @param path a null-terminated UTF-8 string containing the path to the plugin +/// shared library +/// @return `MTA_SUCCESS` if the plugin was loaded successfully, or another +/// status code if an error occurs. You can get more details about the +/// error with `mta_last_error`. #[no_mangle] -pub extern "C" fn mta_load_plugin(path: *const c_char) -> mta_status_t { - todo!() +pub unsafe extern "C" fn mta_load_plugin(path: *const c_char) -> mta_status_t { + catch_unwind(move || { + check_pointers_non_null!(path); + + let path = CStr::from_ptr(path).to_str().map_err(|_| { + Error::InvalidParameter("invalid UTF-8 in plugin path".into()) + })?; + + crate::plugin::load_plugin(path) + }) } -/// TODO +/// Load a model from `load_from` with the given options. +/// +/// If `plugin_name` is a NULL pointer, metatomic will try to determine the +/// correct plugin to use by checking the `load_from` parameter. If we can not +/// determine the correct plugin, we then try to load the model with each +/// registered plugin until one succeeds. +/// +/// If `plugin_name` is given, then we only try to load the model with the +/// specified plugin, and return an error if the plugin can not load the model. +/// +/// @param plugin_name optional null-terminated UTF-8 string containing the name +/// of the plugin to use for loading the model, or `NULL` to let metatomic +/// search for a correct plugin +/// @param load_from a null-terminated UTF-8 string describing where to load the +/// model from (e.g. a file path, a model name, etc.). The interpretation +/// of this string is up to the plugin. +/// @param options_json a null-terminated UTF-8 string containing a set of string +/// keys and string value options for loading the model. The interpretation +/// of these options is up to the plugin. +/// @param model output pointer to the loaded model. The caller takes ownership of +/// the model and must unload it when the model is no longer needed. +/// @return `MTA_SUCCESS` if the model was loaded successfully, or another +/// status code if an error occurs. You can get more details about the +/// error with `mta_last_error`. #[no_mangle] -pub extern "C" fn mta_load_model( +pub unsafe extern "C" fn mta_load_model( plugin_name: *const c_char, load_from: *const c_char, options_json: *const c_char, model: *mut mta_model_t, ) -> mta_status_t { - todo!() + let unwind_wrapper = std::panic::AssertUnwindSafe(model); + + catch_unwind(move || { + check_pointers_non_null!(load_from, model); + + let plugin_name = if plugin_name.is_null() { + None + } else { + Some(CStr::from_ptr(plugin_name).to_str().map_err(|_| { + Error::InvalidParameter("invalid UTF-8 in plugin name".into()) + })?) + }; + + let options_json = if options_json.is_null() { + CStr::from_bytes_with_nul(b"{}\0").expect("invalid CStr") + } else { + CStr::from_ptr(options_json) + }; + + let options_str = options_json.to_str().map_err(|_| { + Error::InvalidParameter("invalid UTF-8 in options JSON".into()) + })?; + + let options = json::parse(options_str).map_err( + |e| Error::Serialization(format!("JSON parsing error: {}", e)) + )?; + if !options.is_object() { + return Err(Error::Serialization("JSON options must be an object in `mta_load_model`".into())) + } + + // just some validation, we pass the raw JSON down to the plugins + for (key, value) in options.entries() { + if !value.is_string() { + return Err(Error::InvalidParameter(format!( + "JSON option '{}' has a non-string value in `mta_load_model`", + key + ))); + } + } + + let loaded = crate::plugin::load_model(plugin_name, CStr::from_ptr(load_from), options_json)?; + + let _ = &unwind_wrapper; + *unwind_wrapper.0 = loaded.into_raw(); + Ok(()) + }) } diff --git a/metatomic-core/src/c_api/status.rs b/metatomic-core/src/c_api/status.rs index 8aef16c11..b9c9a16aa 100644 --- a/metatomic-core/src/c_api/status.rs +++ b/metatomic-core/src/c_api/status.rs @@ -27,7 +27,7 @@ thread_local! { /// The value 0 (`MTA_SUCCESS`) indicates success, while any non-zero value indicates an error. #[allow(non_camel_case_types)] #[repr(C)] -#[derive(PartialEq, Eq, Debug)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum mta_status_t { /// Status code indicating success MTA_SUCCESS = 0, @@ -37,10 +37,9 @@ pub enum mta_status_t { MTA_IO_ERROR = 2, /// Status code indicating serialization/deserialization errors MTA_SERIALIZATION_ERROR = 3, - /// Status code indicating errors that come from callbacks provided by the user. - /// The error message and arbitrary data can be stored using `mta_set_last_error`, - /// and retrieved using `mta_last_error`. - MTA_CALLBACK_ERROR = 254, + /// Status code used by plugins when a model is not supported by the + /// current plugin + MTA_MODEL_NOT_SUPPORTED_ERROR = 4, /// Status code used when there is an internal error MTA_INTERNAL_ERROR = 255, } @@ -79,9 +78,9 @@ macro_rules! check_pointers_non_null { impl From for mta_status_t { fn from(error: Error) -> mta_status_t { - if let Error::CallbackError = error { + if let Error::CallbackError(status) = error { // If the error is already a CallbackError, we can directly return the corresponding status code. - return mta_status_t::MTA_CALLBACK_ERROR; + return status; } LAST_ERROR.with(|last_error| { @@ -98,7 +97,7 @@ impl From for mta_status_t { *last_error = LastError { message: CString::new(format!("{}", error)) .expect("error message contains a null byte"), - origin: CString::new("metatensor-core").expect("invalid C string"), + origin: CString::new("metatomic-core").expect("invalid C string"), custom_data: std::ptr::null_mut(), custom_data_deleter: None, }; @@ -108,7 +107,7 @@ impl From for mta_status_t { Error::InvalidParameter(_) => mta_status_t::MTA_INVALID_PARAMETER_ERROR, Error::Io(_) => mta_status_t::MTA_IO_ERROR, Error::Serialization(_) => mta_status_t::MTA_SERIALIZATION_ERROR, - Error::CallbackError => unreachable!(), + Error::CallbackError(_) => unreachable!("already handled above"), Error::Internal(_) => mta_status_t::MTA_INTERNAL_ERROR, } } diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index b57ab76dd..3450f05df 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -6,16 +6,20 @@ #![allow(clippy::unreadable_literal, clippy::option_if_let_else, clippy::module_name_repetitions)] #![allow(clippy::missing_errors_doc, clippy::missing_panics_doc, clippy::missing_safety_doc)] #![allow(clippy::similar_names, clippy::borrow_as_ptr, clippy::uninlined_format_args)] -#![allow(clippy::doc_markdown)] +#![allow(clippy::doc_markdown, clippy::needless_continue)] #![allow(clippy::let_underscore_untyped, clippy::manual_let_else, clippy::empty_line_after_doc_comments)] // To be removed later #![allow(unused_variables, dead_code, clippy::needless_pass_by_value)] +use std::sync::Arc; + #[doc(hidden)] pub mod c_api; mod metadata; +use crate::c_api::mta_status_t; + pub use self::metadata::{ModelMetadata, PairListOptions}; mod quantities; @@ -28,22 +32,22 @@ mod model; pub use self::model::Model; mod plugin; -pub use self::plugin::{Plugin, load_plugin, load_model}; +pub use self::plugin::Plugin; mod units; pub use self::units::unit_conversion_factor; /// The possible sources of error in metatomic -#[derive(Debug)] +#[derive(Debug, Clone)] pub enum Error { /// Error while serializing data to or deserializing data from JSON Serialization(String), /// Invalid parameters passed to a function InvalidParameter(String), /// I/O error - Io(std::io::Error), + Io(Arc), /// Error coming from an external function used as a callback - CallbackError, + CallbackError(mta_status_t), /// Any other internal error, usually these are internal bugs. Internal(String), } @@ -54,7 +58,7 @@ impl std::fmt::Display for Error { Error::Serialization(e) => write!(f, "serialization error: {}", e), Error::InvalidParameter(e) => write!(f, "invalid parameter: {}", e), Error::Io(e) => write!(f, "io error: {}", e), - Error::CallbackError => write!(f, "callback error"), + Error::CallbackError(e) => write!(f, "callback error, status code: {:?}", e), Error::Internal(e) => write!(f, "internal metatomic error (this is likely a bug, please report it): {}", e ), @@ -68,7 +72,7 @@ impl std::error::Error for Error { Error::InvalidParameter(_) | Error::Serialization(_) | Error::Internal(_) - | Error::CallbackError => None, + | Error::CallbackError(_) => None, Error::Io(e) => Some(e), } } @@ -92,3 +96,9 @@ impl From> for Error { } } } + +impl From for Error { + fn from(error: std::io::Error) -> Self { + Error::Io(Arc::new(error)) + } +} diff --git a/metatomic-core/src/model.rs b/metatomic-core/src/model.rs index 16b208ac1..1b93577d2 100644 --- a/metatomic-core/src/model.rs +++ b/metatomic-core/src/model.rs @@ -7,6 +7,17 @@ use crate::c_api::mta_model_t; /// TODO pub struct Model(pub(crate) mta_model_t); +impl Model { + /// Create a new `Model` from the corresponding C API struct. + pub fn new(model: mta_model_t) -> Self { + return Model(model); + } + + /// Extract the underlying C API struct. + pub fn into_raw(self) -> mta_model_t { + return self.0; + } +} /// TODO pub fn execute_model( diff --git a/metatomic-core/src/plugin.rs b/metatomic-core/src/plugin.rs index 60d145803..4a9da5ab9 100644 --- a/metatomic-core/src/plugin.rs +++ b/metatomic-core/src/plugin.rs @@ -1,37 +1,191 @@ -use std::collections::BTreeMap; +use std::ffi::CStr; +use std::sync::Mutex; -use crate::c_api::mta_plugin_t; +use libloading::Library; +use once_cell::sync::Lazy; + +use crate::c_api::{mta_model_t, mta_plugin_t, mta_register_plugin, mta_status_t}; use crate::{Error, Model}; -/// TODO +/// ABI version of the metatomic plugin interface. +/// +/// This increases anytime the plugin or model C API changes in a non backward +/// compatible way. Plugins compiled with an incompatible ABI version will be +/// rejected at registration time. pub const MTA_ABI_VERSION: i32 = 1; -/// TODO +/// The list of registered plugins in the current process. +static PLUGINS: Lazy>> = Lazy::new(|| Mutex::new(Vec::new())); +/// Keep the loaded libraries alive for the entire process lifetime, to ensure +/// that the plugin code is not unloaded while it's still in use. +static LIBRARIES: Lazy>> = Lazy::new(|| Mutex::new(Vec::new())); + pub struct Plugin(mta_plugin_t); impl Plugin { - /// TODO - pub fn new(c_plugin: mta_plugin_t) -> Self { - Self(c_plugin) + /// Create a new plugin from the C struct + pub fn new(plugin: mta_plugin_t) -> Result { + if plugin.name.is_null() { + return Err(Error::InvalidParameter( + "can not register plugin: plugin `name` is NULL".into(), + )); + } + + let c_str_name = unsafe { CStr::from_ptr(plugin.name) }; + if c_str_name.to_str().is_err() { + return Err(Error::InvalidParameter(format!( + "can not register plugin: plugin `name` is not valid UTF-8: {}", + c_str_name.to_string_lossy() + ))); + } + + if plugin.load_model.is_none() { + return Err(Error::InvalidParameter( + "can not register plugin: plugin `load_model` callback is NULL".into(), + )); + } + + if plugin.abi_version != MTA_ABI_VERSION { + let name = unsafe { + CStr::from_ptr(plugin.name).to_string_lossy() + }; + + return Err(Error::InvalidParameter(format!( + "can not register plugin '{}': plugin ABI version is {}, but metatomic expects {}", + name, + plugin.abi_version, + MTA_ABI_VERSION, + ))); + } + + Ok(Plugin(plugin)) } - /// TODO + /// Get the plugin name. pub fn name(&self) -> &str { - todo!() + unsafe { + return CStr::from_ptr(self.0.name) + .to_str() + .expect("invalid UTF-8 in plugin name"); + } } - /// TODO - pub fn load_model(&self, load_from: &str, options: BTreeMap) -> Result { - todo!() + /// Try to load a model with this plugin. + pub fn load_model( + &self, + load_from: &CStr, + options_json: &CStr, + ) -> Result { + let load_model = self.0.load_model.expect("`load_model` is NULL"); + + let mut model = mta_model_t::null(); + let status = unsafe { + load_model(load_from.as_ptr(), options_json.as_ptr(), &mut model) + }; + + if status != mta_status_t::MTA_SUCCESS { + return Err(Error::CallbackError(status)); + } + + return Ok(Model::new(model)); } } -/// TODO +/// Register a new plugin in the current process. +pub fn register_plugin(plugin: Plugin) -> Result<(), Error> { + let mut plugins = PLUGINS.lock().expect("plugin registry mutex was poisoned"); + if plugins.iter().any(|existing| existing.name() == plugin.name()) { + return Err(Error::InvalidParameter(format!( + "a plugin named '{}' is already registered", + plugin.name() + ))); + } + + plugins.push(plugin); + return Ok(()); +} + +/// Load a plugin from a shared library. +/// +/// The shared library must export the symbols generated by the +/// `MTA_REGISTER_PLUGIN` C macro. pub fn load_plugin(path: &str) -> Result<(), Error> { - todo!() + // this needs to be kept in sync with the definition in `MTA_REGISTER_PLUGIN` in build.rs + type PluginInitFn = unsafe extern "C" fn(abi: i32, data: *mut std::ffi::c_void) -> mta_status_t; + + let library = unsafe { Library::new(path) }; + + let library = library.map_err(|error| { + std::io::Error::other( + format!("failed to load plugin '{}': {}", path, error), + ) + })?; + + let status = unsafe { + let init_plugin = library.get::(b"mta_plugin_init\0") + .map_err(|error| Error::InvalidParameter(format!( + "failed to load plugin registration symbol from '{}': {}", + path, error + )))?; + init_plugin(MTA_ABI_VERSION, mta_register_plugin as *mut std::ffi::c_void) + }; + + if status != mta_status_t::MTA_SUCCESS { + return Err(Error::CallbackError(status)); + } + + LIBRARIES.lock().expect("loaded plugin registry mutex was poisoned").push(library); + + return Ok(()); } -/// TODO -pub fn load_model(plugin: Option<&str>, load_from: &str, options: BTreeMap) -> Result { - todo!() +/// Load a model from `load_from`, using the given options. +pub fn load_model( + plugin_name: Option<&str>, + load_from: &CStr, + options_json: &CStr, +) -> Result { + let plugins = PLUGINS.lock().expect("plugin registry mutex was poisoned"); + + if let Some(plugin_name) = plugin_name { + for plugin in plugins.iter() { + if plugin.name() == plugin_name { + return plugin.load_model(load_from, options_json); + } + } + + return Err(Error::InvalidParameter(format!( + "no plugin named '{}' is registered", + plugin_name + ))); + } + + for plugin in plugins.iter() { + match plugin.load_model(load_from, options_json) { + Ok(model) => return Ok(model), + Err(e) => { + if let Error::CallbackError(mta_status_t::MTA_MODEL_NOT_SUPPORTED_ERROR) = e { + // try the next plugin + continue; + } else { + return Err(e); + } + } + } + } + + let message = if plugins.is_empty() { + "no plugin is registered".into() + } else { + format!( + "tried the following plugins, but none could load the model: {}", + plugins.iter().map(|p| p.name()).collect::>().join(", ") + ) + }; + + return Err(Error::InvalidParameter(format!( + "failed to load model from '{}': {}", + load_from.to_string_lossy(), + message + ))); } diff --git a/metatomic-core/tests/CMakeLists.txt b/metatomic-core/tests/CMakeLists.txt index 1a96108d5..77e251ddc 100644 --- a/metatomic-core/tests/CMakeLists.txt +++ b/metatomic-core/tests/CMakeLists.txt @@ -53,6 +53,7 @@ endif() enable_testing() +add_subdirectory(test-plugins) file(GLOB ALL_TESTS *.cpp) foreach(_file_ ${ALL_TESTS}) @@ -70,6 +71,8 @@ foreach(_file_ ${ALL_TESTS}) NO_SYSTEM_FROM_IMPORTED ON ) + target_compile_definitions(${_name_} PRIVATE PLUGIN_DIR="${CMAKE_CURRENT_BINARY_DIR}/test-plugins") + add_test( NAME ${_name_} COMMAND ${TEST_COMMAND} $ diff --git a/metatomic-core/tests/plugins.cpp b/metatomic-core/tests/plugins.cpp new file mode 100644 index 000000000..79b53947e --- /dev/null +++ b/metatomic-core/tests/plugins.cpp @@ -0,0 +1,45 @@ +#include + +#include "metatomic.h" + + +TEST_CASE("Load plugins") { + auto status = mta_load_plugin(PLUGIN_DIR "/test-c-plugin.so"); + CHECK(status == MTA_SUCCESS); + + // try to load the model with an explicit plugin name + struct mta_model_t model; + status = mta_load_model("test-c-plugin", "some_model", "{}", &model); + CHECK(status == MTA_MODEL_NOT_SUPPORTED_ERROR); + + // load the plugin without specifying the plugin name + status = mta_load_model(nullptr, "some_model", "{}", &model); + CHECK(status == MTA_INVALID_PARAMETER_ERROR); + + const char* error_message; + const char* error_origin; + + status = mta_last_error(&error_message, &error_origin, nullptr); + REQUIRE(status == MTA_SUCCESS); + + CHECK(std::string(error_origin) == "metatomic-core"); + const char* expected_message = ( + "invalid parameter: failed to load model from 'some_model': tried the " + "following plugins, but none could load the model: test-c-plugin" + ); + CHECK(std::string(error_message) == expected_message); + + + status = mta_load_plugin(PLUGIN_DIR "/bad-abi-plugin.so"); + CHECK(status == MTA_INVALID_PARAMETER_ERROR); + + status = mta_last_error(&error_message, &error_origin, nullptr); + REQUIRE(status == MTA_SUCCESS); + + CHECK(std::string(error_origin) == "metatomic-core"); + expected_message = ( + "invalid parameter: can not register plugin 'bad-abi-plugin': " + "plugin ABI version is 2, but metatomic expects 1" + ); + CHECK(std::string(error_message) == expected_message); +} diff --git a/metatomic-core/tests/test-plugins/CMakeLists.txt b/metatomic-core/tests/test-plugins/CMakeLists.txt new file mode 100644 index 000000000..2693ceab7 --- /dev/null +++ b/metatomic-core/tests/test-plugins/CMakeLists.txt @@ -0,0 +1,14 @@ +add_library(test-c-plugin SHARED plugin.c) +target_link_libraries(test-c-plugin metatomic) +# create test plugins with a consistent name across platforms +set_target_properties(test-c-plugin PROPERTIES + PREFIX "" + SUFFIX ".so" +) + +add_library(bad-abi-plugin SHARED bad-abi.c) +target_link_libraries(bad-abi-plugin metatomic) +set_target_properties(bad-abi-plugin PROPERTIES + PREFIX "" + SUFFIX ".so" +) diff --git a/metatomic-core/tests/test-plugins/bad-abi.c b/metatomic-core/tests/test-plugins/bad-abi.c new file mode 100644 index 000000000..35e86bcfe --- /dev/null +++ b/metatomic-core/tests/test-plugins/bad-abi.c @@ -0,0 +1,17 @@ +#include + + +static mta_status_t load_model(const char *load_from, const char *options_json, struct mta_model_t *model) { + // This plugin can not load any model + return MTA_MODEL_NOT_SUPPORTED_ERROR; +} + + +MTA_REGISTER_PLUGIN(register_plugin, { + mta_plugin_t plugin = { + .abi_version = MTA_ABI_VERSION + 1, // incompatible ABI version + .name = "bad-abi-plugin", + .load_model = load_model, + }; + return register_plugin(plugin); +}); diff --git a/metatomic-core/tests/test-plugins/plugin.c b/metatomic-core/tests/test-plugins/plugin.c new file mode 100644 index 000000000..1602dfae7 --- /dev/null +++ b/metatomic-core/tests/test-plugins/plugin.c @@ -0,0 +1,17 @@ +#include + + +static mta_status_t load_model(const char *load_from, const char *options_json, struct mta_model_t *model) { + // This plugin can not load any model + return MTA_MODEL_NOT_SUPPORTED_ERROR; +} + + +MTA_REGISTER_PLUGIN(register_plugin, { + mta_plugin_t plugin = { + .abi_version = MTA_ABI_VERSION, + .name = "test-c-plugin", + .load_model = load_model, + }; + return register_plugin(plugin); +}); diff --git a/scripts/include/stdio.h b/scripts/include/stdio.h new file mode 100644 index 000000000..e69de29bb From 9698edd92896c3526b7e8c3d37f0c9c2a49df86c Mon Sep 17 00:00:00 2001 From: frostedoyster Date: Fri, 29 May 2026 08:32:58 +0200 Subject: [PATCH 16/43] Add a C model registration test --- metatomic-core/tests/c-model.cpp | 197 +++++++++++++++++++++++++++++++ 1 file changed, 197 insertions(+) create mode 100644 metatomic-core/tests/c-model.cpp diff --git a/metatomic-core/tests/c-model.cpp b/metatomic-core/tests/c-model.cpp new file mode 100644 index 000000000..ec1da92f7 --- /dev/null +++ b/metatomic-core/tests/c-model.cpp @@ -0,0 +1,197 @@ +#include + +#include + +#include "metatomic.h" + +#include + + +struct SimpleModelData { + double scale; +}; + +static mta_status_t unload_impl(void* model_data) { + delete static_cast(model_data); + return MTA_SUCCESS; +} + +static mta_status_t metadata_impl(const void* model_data, mta_string_t* metadata_json) { + (void) model_data; + + *metadata_json = mta_string_create(R"({ + "name": "test C model", + "description": "small model used as a C API example", + "authors": [], + "references": { + "model": [], + "implementation": [], + "architecture": [] + } + })"); + return MTA_SUCCESS; +} + +static mta_status_t capabilities_impl(const void* model_data, mta_string_t* capabilities_json) { + (void) model_data; + + *capabilities_json = mta_string_create(R"({ + "outputs": [{ + "quantity": "energy", + "unit": "eV", + "per_atom": false + }], + "atomic_types": [1, 6, 8], + "interaction_range": 4.5, + "length_unit": "nm", + "supported_devices": ["cpu"], + "dtype": "float32" + })"); + return MTA_SUCCESS; +} + +static mta_status_t supported_outputs_impl( + const void* model_data, + mta_string_t* outputs_json +) { + (void) model_data; + *outputs_json = mta_string_create(R"([{ + "quantity": "energy", + "unit": "eV", + "per_atom": false + }])"); + return MTA_SUCCESS; +} + +static mta_status_t requested_pair_lists_impl( + const void* model_data, + mta_string_t* pair_options_json +) { + (void) model_data; + *pair_options_json = mta_string_create("[]"); + return MTA_SUCCESS; +} + +static mta_status_t requested_inputs_impl( + const void* model_data, + mta_string_t* requested_inputs_json +) { + (void) model_data; + *requested_inputs_json = mta_string_create("[]"); + return MTA_SUCCESS; +} + + +mts_tensormap_t* scalar_tensormap(double value) { + auto values = std::make_unique>( + std::vector{1, 1}, + std::vector{value} + ); + + auto array = metatensor::DataArrayBase::to_mts_array(std::move(values)); + + auto samples = metatensor::Labels({"system"}, {{0}}); + auto properties = metatensor::Labels({"energy"}, {{0}}); + + auto* block = mts_block( + std::move(array).release(), + samples.as_mts_labels_t(), + nullptr, + 0, + properties.as_mts_labels_t() + ); + if (block == nullptr) { + return nullptr; + } + + auto keys = metatensor::Labels({"_"}, {{0}}); + auto blocks = std::vector{block}; + return mts_tensormap(keys.as_mts_labels_t(), blocks.data(), blocks.size()); +} + +static mta_status_t execute_inner_impl( + void* model_data, + const mta_system_t* const* systems, + uintptr_t systems_count, + const mts_labels_t* selected_atoms, + const char* requested_outputs_json, + mts_tensormap_t** outputs, + uintptr_t outputs_count +) { + (void)model_data; + (void)systems; + (void)systems_count; + (void)selected_atoms; + (void)requested_outputs_json; + (void)outputs; + (void)outputs_count; + + return MTA_INTERNAL_ERROR; +} + +static mta_status_t load_model_impl( + const char* load_from, + const char* options_json, + mta_model_t* model +) { + (void)options_json; + assert(model != nullptr); + + if (std::strcmp(load_from, "test-c-model") != 0) { + return MTA_MODEL_NOT_SUPPORTED_ERROR; + } + + model->data = new SimpleModelData{2.0}; + model->unload = unload_impl; + model->metadata = metadata_impl; + model->capabilities = capabilities_impl; + model->supported_outputs = supported_outputs_impl; + model->requested_pair_lists = requested_pair_lists_impl; + model->requested_inputs = requested_inputs_impl; + model->execute_inner = execute_inner_impl; + + return MTA_SUCCESS; +} + +TEST_CASE("simple C model can be registered and loaded through the C API") { + static auto PLUGIN = mta_plugin_t { + MTA_ABI_VERSION, + "test-c-plugin", + load_model_impl, + }; + mta_register_plugin(PLUGIN); + + auto model = mta_model_t{}; + auto status = mta_load_model("test-c-plugin", "test-c-model", nullptr, &model); + REQUIRE(status == MTA_SUCCESS); + + CHECK(model.data != nullptr); + CHECK(model.unload != nullptr); + CHECK(model.metadata != nullptr); + CHECK(model.capabilities != nullptr); + CHECK(model.supported_outputs != nullptr); + CHECK(model.requested_pair_lists != nullptr); + CHECK(model.requested_inputs != nullptr); + CHECK(model.execute_inner != nullptr); + + mta_string_t metadata = nullptr; + status = model.metadata(model.data, &metadata); + REQUIRE(status == MTA_SUCCESS); + + CHECK(metadata != nullptr); + auto metadata_str = std::string(mta_string_view(metadata)); + mta_string_free(metadata); + + CHECK(metadata_str.find("\"name\": \"test C model\"") != std::string::npos); + + + mta_string_t pair_lists = nullptr; + status = model.requested_pair_lists(model.data, &pair_lists); + REQUIRE(status == MTA_SUCCESS); + + CHECK(pair_lists != nullptr); + CHECK(std::strcmp(mta_string_view(pair_lists), "[]") == 0); + mta_string_free(pair_lists); + + REQUIRE(model.unload(model.data) == MTA_SUCCESS); +} From 84cbc36fc7efb0bde74f495325c51568c3940c8f Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Mon, 1 Jun 2026 16:23:16 +0200 Subject: [PATCH 17/43] Implement System in metatomic-core Co-Authored-By: frostedoyster --- metatomic-core/Cargo.toml | 3 +- metatomic-core/include/metatomic.h | 10 +- metatomic-core/src/c_api/status.rs | 9 +- metatomic-core/src/lib.rs | 20 + metatomic-core/src/metadata.rs | 8 +- metatomic-core/src/quantities.rs | 23 +- metatomic-core/src/system.rs | 739 ++++++++++++++++++++++++++++- 7 files changed, 779 insertions(+), 33 deletions(-) diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index dace0f8d5..b88a91cda 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -16,9 +16,10 @@ bench = false [dependencies] metatensor = { version = "0.3.0" } once_cell = "1" -dlpk = "0.3" +dlpk = { version = "0.3", features = ["ndarray"]} json = "0.12" libloading = "0.8" +ndarray = "0.17" [build-dependencies] diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index a2a549e78..9273007d0 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -90,11 +90,19 @@ typedef enum mta_status_t { * Status code indicating serialization/deserialization errors */ MTA_SERIALIZATION_ERROR = 3, + /** + * Status code indicating dlpack errors + */ + MTA_DLPACK_ERROR = 4, + /** + * Status code indicating metatensor errors + */ + MTA_METATENSOR_ERROR = 5, /** * Status code used by plugins when a model is not supported by the * current plugin */ - MTA_MODEL_NOT_SUPPORTED_ERROR = 4, + MTA_MODEL_NOT_SUPPORTED_ERROR = 6, /** * Status code used when there is an internal error */ diff --git a/metatomic-core/src/c_api/status.rs b/metatomic-core/src/c_api/status.rs index b9c9a16aa..48d658777 100644 --- a/metatomic-core/src/c_api/status.rs +++ b/metatomic-core/src/c_api/status.rs @@ -37,9 +37,13 @@ pub enum mta_status_t { MTA_IO_ERROR = 2, /// Status code indicating serialization/deserialization errors MTA_SERIALIZATION_ERROR = 3, + /// Status code indicating dlpack errors + MTA_DLPACK_ERROR = 4, + /// Status code indicating metatensor errors + MTA_METATENSOR_ERROR = 5, /// Status code used by plugins when a model is not supported by the /// current plugin - MTA_MODEL_NOT_SUPPORTED_ERROR = 4, + MTA_MODEL_NOT_SUPPORTED_ERROR = 6, /// Status code used when there is an internal error MTA_INTERNAL_ERROR = 255, } @@ -107,8 +111,11 @@ impl From for mta_status_t { Error::InvalidParameter(_) => mta_status_t::MTA_INVALID_PARAMETER_ERROR, Error::Io(_) => mta_status_t::MTA_IO_ERROR, Error::Serialization(_) => mta_status_t::MTA_SERIALIZATION_ERROR, + Error::Dlpack(_) => mta_status_t::MTA_DLPACK_ERROR, + Error::Metatensor(_) => mta_status_t::MTA_METATENSOR_ERROR, Error::CallbackError(_) => unreachable!("already handled above"), Error::Internal(_) => mta_status_t::MTA_INTERNAL_ERROR, + } } } diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index 3450f05df..cc20d6961 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -46,6 +46,10 @@ pub enum Error { InvalidParameter(String), /// I/O error Io(Arc), + /// Error related to dlpack tensors, such as invalid tensor shapes or types + Dlpack(Arc), + /// Error coming from metatensor + Metatensor(metatensor::Error), /// Error coming from an external function used as a callback CallbackError(mta_status_t), /// Any other internal error, usually these are internal bugs. @@ -58,6 +62,8 @@ impl std::fmt::Display for Error { Error::Serialization(e) => write!(f, "serialization error: {}", e), Error::InvalidParameter(e) => write!(f, "invalid parameter: {}", e), Error::Io(e) => write!(f, "io error: {}", e), + Error::Dlpack(e) => write!(f, "dlpack error: {}", e), + Error::Metatensor(e) => write!(f, "metatensor error: {}", e), Error::CallbackError(e) => write!(f, "callback error, status code: {:?}", e), Error::Internal(e) => write!(f, "internal metatomic error (this is likely a bug, please report it): {}", e @@ -74,6 +80,8 @@ impl std::error::Error for Error { | Error::Internal(_) | Error::CallbackError(_) => None, Error::Io(e) => Some(e), + Error::Dlpack(e) => Some(e), + Error::Metatensor(e) => Some(e), } } @@ -102,3 +110,15 @@ impl From for Error { Error::Io(Arc::new(error)) } } + +impl From for Error { + fn from(error: dlpk::ndarray::DLPackNDarrayError) -> Self { + Error::Dlpack(Arc::new(error)) + } +} + +impl From for Error { + fn from(error: metatensor::Error) -> Self { + Error::Metatensor(error) + } +} diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs index 48bfeaad5..a0f0c3875 100644 --- a/metatomic-core/src/metadata.rs +++ b/metatomic-core/src/metadata.rs @@ -9,16 +9,16 @@ use crate::units::validate_unit; #[derive(Debug, Clone)] pub struct PairListOptions { /// Cutoff radius for this pair list in the length unit of the model - cutoff: f64, + pub cutoff: f64, /// Whether the list is a full list (contains both the pair `i -> j` and `j -> i`) /// or a half list (contains only `i -> j`) - full_list: bool, + pub full_list: bool, /// Whether the list guarantees that only atoms within the cutoff are /// included (strict) or may also include pairs slightly beyond the cutoff /// (non-strict) - strict: bool, + pub strict: bool, /// List of strings describing who requested this pair list - requestors: Vec, + pub requestors: Vec, } impl std::cmp::PartialEq for PairListOptions { diff --git a/metatomic-core/src/quantities.rs b/metatomic-core/src/quantities.rs index 93727c837..cc313d555 100644 --- a/metatomic-core/src/quantities.rs +++ b/metatomic-core/src/quantities.rs @@ -41,7 +41,7 @@ fn is_valid_identifier(s: &str) -> bool { /// All components (namespace, name, variant) must be non-empty if they are /// present, and must be valid identifiers (alphanumeric + underscore, not /// starting with a digit). -fn validate_quantity_name(name: &str) -> Result<(), Error> { +pub(crate) fn validate_quantity_name(name: &str) -> Result<(), Error> { if STANDARD_QUANTITIES.contains(&name) { return Ok(()); } @@ -67,7 +67,12 @@ fn validate_quantity_name(name: &str) -> Result<(), Error> { } } - for component in main_part.split("::") { + if STANDARD_QUANTITIES.contains(&main_part) { + return Ok(()); + } + + let components: Vec<_> = main_part.split("::").collect(); + for component in &components { if !is_valid_identifier(component) { return Err(Error::InvalidParameter(format!( "invalid quantity name component '{}' in '{}': must be a valid identifier (alphanumeric or underscore, not starting with a digit)", @@ -76,6 +81,13 @@ fn validate_quantity_name(name: &str) -> Result<(), Error> { } } + if components.len() == 1 { + return Err(Error::InvalidParameter(format!( + "'{}' is not a standard quantity name; custom quantity names must use '::'", + name + ))); + } + Ok(()) } @@ -289,7 +301,7 @@ mod tests { vec![Gradients::Positions, Gradients::Strain], ] { let quantity = Quantity { - name: "test".into(), + name: "test_ns::test".into(), unit: "unit".into(), description: Some("Hello".to_string()), gradients: grads.clone(), @@ -367,9 +379,7 @@ mod tests { "my_model::energy", "org::my_model::custom_qty", "ns1::ns2::ns3::energy", - "custom_name", "some_ns::name_with_underscores", - "_underscore_start", "_ns::_name", ]; for name in custom { @@ -388,6 +398,9 @@ mod tests { let error = validate_quantity_name("").expect_err("expected an error"); assert_eq!(error.to_string(), "invalid parameter: quantity name cannot be empty in ''"); + let error = validate_quantity_name("not_a_standard_name").expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: 'not_a_standard_name' is not a standard quantity name; custom quantity names must use '::'"); + let error = validate_quantity_name("/variant").expect_err("expected an error"); assert_eq!(error.to_string(), "invalid parameter: quantity name cannot be empty in '/variant'"); diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs index 30677f5f9..d30dccce4 100644 --- a/metatomic-core/src/system.rs +++ b/metatomic-core/src/system.rs @@ -1,12 +1,21 @@ -use std::collections::{BTreeMap, HashMap}; +use std::collections::{BTreeMap, HashMap, HashSet}; +use once_cell::sync::Lazy; -use dlpk::DLPackTensor; +use dlpk::sys::{DLDataType, DLDevice, DLDeviceType}; +use dlpk::{DLPackTensor, DLPackTensorRef}; use metatensor::{TensorBlock, TensorMap}; -use crate::PairListOptions; +use crate::{Error, PairListOptions}; +/// Names that can never be used as custom data in a system +static INVALID_DATA_NAMES: Lazy> = Lazy::new(|| { + HashSet::from(["types", "type", "positions", "position", "cell", "neighbors", "neighbor", "pair", "pairs"]) +}); -/// TODO +/// Storage for an atomistic system. +/// +/// This owns the raw DLPack tensors and metatensor objects used at FFI +/// boundaries. pub struct System { length_unit: String, types: DLPackTensor, @@ -18,36 +27,724 @@ pub struct System { custom_data: HashMap, } - impl System { - /// TODO + /// Create a `System` from raw DLPack tensors pub fn new( length_unit: String, types: DLPackTensor, positions: DLPackTensor, cell: DLPackTensor, - pbc: DLPackTensor - ) -> Self { - todo!() + pbc: DLPackTensor, + ) -> Result { + validate_system_tensors(&types, &positions, &cell, &pbc)?; + + let system = System { + length_unit, + types, + positions, + cell, + pbc, + pairs: BTreeMap::new(), + custom_data: HashMap::new(), + }; + + if system.device().device_type == DLDeviceType::kDLCPU { + validate_cpu_system_data(&system)?; + } + + return Ok(system); + } + + /// Get the length unit used by this system + pub fn length_unit(&self) -> &str { + &self.length_unit + } + + /// Get the number of atoms/particles in this system + #[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)] + pub fn size(&self) -> usize { + let size = self.types.shape()[0]; + debug_assert!(usize::try_from(size).is_ok()); + return size as usize; + } + + /// Get the particle types + pub fn types(&self) -> DLPackTensorRef<'_> { + self.types.as_ref() + } + + /// Get the particle positions + pub fn positions(&self) -> DLPackTensorRef<'_> { + self.positions.as_ref() + } + + /// Get the unit cell + pub fn cell(&self) -> DLPackTensorRef<'_> { + self.cell.as_ref() + } + + /// Get the periodic boundary condition flags + pub fn pbc(&self) -> DLPackTensorRef<'_> { + self.pbc.as_ref() + } + + /// Add a pair list to this system + pub fn add_pairs( + &mut self, + options: PairListOptions, + pairs: TensorBlock, + ) -> Result<(), Error> { + if self.pairs.contains_key(&options) { + return Err(Error::InvalidParameter( + "the pair list for these options already exists in this system".into(), + )); + } + + let samples = pairs.samples(); + let samples_names = samples.names(); + if samples_names != ["first_atom", "second_atom", "cell_shift_a", "cell_shift_b", "cell_shift_c"] { + return Err(Error::InvalidParameter( + "invalid samples for `pairs`: the samples names must be \ + 'first_atom', 'second_atom', 'cell_shift_a', 'cell_shift_b', \ + 'cell_shift_c'".into(), + )); + } + + let components = pairs.components(); + if components.len() != 1 || components[0].names() != ["xyz"] || components[0].count() != 3 { + return Err(Error::InvalidParameter( + "invalid components for `pairs`: there should be a \ + single 'xyz'=[0, 1, 2] component".into() + )); + } + + #[allow(clippy::collapsible_if)] + if components[0].device().device_type == DLDeviceType::kDLCPU { + if components[0][0] != [0] || components[0][1] != [1] || components[0][2] != [2] { + return Err(Error::InvalidParameter( + "invalid components for `pairs`: the 'xyz' \ + component should contain [0, 1, 2]".into() + )); + } + } + + let properties = pairs.properties(); + if properties.names() != ["distance"] || properties.count() != 1 { + return Err(Error::InvalidParameter( + "invalid properties for `pairs`: there should be a single \ + 'distance'=0 property".into() + )); + } + + #[allow(clippy::collapsible_if)] + if properties.device().device_type == DLDeviceType::kDLCPU { + if properties[0] != [0] { + return Err(Error::InvalidParameter( + "invalid properties for `pairs`: the 'distance' property \ + should contain [0]".into() + )); + } + } + + if !pairs.as_ref().gradient_list().is_empty() { + return Err(Error::InvalidParameter( + "`pairs` should not have any gradients".into() + )); + } + + // TODO: add TensorBlock::device/dtype and use them here + let values = pairs.values(); + let values_device = values.device()?; + if values_device != self.device() { + return Err(Error::InvalidParameter(format!( + "`pairs` device ({}) does not match this system's device ({})", + values_device, self.device(), + ))); + } + + let values_dtype = values.dtype()?; + if values_dtype != self.dtype() { + return Err(Error::InvalidParameter(format!( + "`pairs` dtype ({}) does not match this system's dtype ({})", + values_dtype, self.dtype(), + ))); + } + + self.pairs.insert(options, pairs); + return Ok(()); + } + + /// Get a pair list from this system + pub fn get_pairs(&self, options: &PairListOptions) -> Option<&TensorBlock> { + return self.pairs.get(options); + } + + /// Get all pair list options known by this system + pub fn known_pairs(&self) -> Vec<&PairListOptions> { + return self.pairs.keys().collect(); + } + + /// Add custom data to this system + /// + /// If `override_` is `true`, existing data with the same name will be + /// replaced. + pub fn add_custom_data(&mut self, name: String, data: TensorMap, override_: bool) -> Result<(), Error> { + if INVALID_DATA_NAMES.contains(name.to_lowercase().as_str()) { + return Err(Error::InvalidParameter(format!( + "custom data can not be named '{}'", name + ))); + } + + crate::quantities::validate_quantity_name(&name)?; + + if !override_ && self.custom_data.contains_key(&name) { + return Err(Error::InvalidParameter(format!( + "custom data '{}' is already present in this system", + name + ))); + } + + if data.keys().count() == 0 { + return Err(Error::InvalidParameter(format!( + "custom data '{}' has no blocks", name + ))); + } + + // TODO: add TensorMap::device/dtype and use them here + let block = data.block_by_id(0); + let values = block.values(); + let data_device = values.device()?; + if data_device != self.device() { + return Err(Error::InvalidParameter(format!( + "device ({}:{}) of the custom data '{}' does not match this system device ({}:{})", + data_device.device_type, data_device.device_id, name, + self.device().device_type, self.device().device_id, + ))); + } + + let values_dtype = values.dtype()?; + if values_dtype != self.dtype() { + return Err(Error::InvalidParameter(format!( + "dtype of custom data '{}' does not match this system dtype", + name, + ))); + } + + self.custom_data.insert(name, data); + return Ok(()); + } + + /// Get custom data from this system. + pub fn get_custom_data(&self, name: &str) -> Result<&TensorMap, Error> { + let lower = name.to_lowercase(); + if INVALID_DATA_NAMES.contains(lower.as_str()) { + return Err(Error::InvalidParameter(format!( + "custom data can not be named '{}'", name + ))); + } + + return self.custom_data.get(name).ok_or_else(|| Error::InvalidParameter(format!( + "no data for '{}' found in this system", name + ))); + } + + /// Get all custom data names known by this system. + pub fn known_custom_data(&self) -> Vec<&str> { + return self.custom_data.keys().map(String::as_str).collect(); + } + + /// The device used for all tensors in this system + fn device(&self) -> DLDevice { + self.types.device() + } + + /// The data type used for the `positions` and `cell` tensors in this + /// system, as well as any pair lists and custom data added to this system. + fn dtype(&self) -> DLDataType { + self.positions.dtype() + } +} + +fn validate_system_tensors( + types: &DLPackTensor, + positions: &DLPackTensor, + cell: &DLPackTensor, + pbc: &DLPackTensor, +) -> Result<(), Error> { + let device = types.device(); + if positions.device() != device || cell.device() != device || pbc.device() != device { + return Err(Error::InvalidParameter( + "`types`, `positions`, `cell`, and `pbc` must be on the same device".into() + )); + } + + let dtype_i32 = ::get_dlpack_data_type(); + let dtype_f32 = ::get_dlpack_data_type(); + let dtype_f64 = ::get_dlpack_data_type(); + let dtype_bool = ::get_dlpack_data_type(); + + if types.dtype() != dtype_i32 { + return Err(Error::InvalidParameter( + "`types` must be a tensor of 32-bit integers".into() + )); + } + + let types_shape = types.shape(); + if types_shape.len() != 1 || types_shape[0] < 0 { + return Err(Error::InvalidParameter(format!( + "`types` must be a (n_atoms,) tensor, got a tensor with shape [{}]", + types_shape.iter().map(|dim| dim.to_string()).collect::>().join(", ") + ))); + } + + let n_atoms = types_shape[0]; + + let positions_shape = positions.shape(); + if positions_shape.len() != 2 || positions_shape[0] != n_atoms || positions_shape[1] != 3 { + return Err(Error::InvalidParameter(format!( + "`positions` must be a (n_atoms x 3) tensor, got a tensor with shape [{}]", + positions_shape.iter().map(|dim| dim.to_string()).collect::>().join(", ") + ))); + } + + if positions.dtype() != dtype_f32 && positions.dtype() != dtype_f64 { + return Err(Error::InvalidParameter( + "`positions` must be a tensor of 32 or 64-bit floating point data".into() + )); + } + + let cell_shape = cell.shape(); + if cell_shape.len() != 2 || cell_shape[0] != 3 || cell_shape[1] != 3 { + return Err(Error::InvalidParameter(format!( + "`cell` must be a (3 x 3) tensor, got a tensor with shape [{}]", + cell_shape.iter().map(|dim| dim.to_string()).collect::>().join(", ") + ))); } - /// TODO - pub fn add_pairs(&mut self, options: PairListOptions, pairs: TensorBlock, check_consistency: bool) { - todo!() + if cell.dtype() != positions.dtype() { + return Err(Error::InvalidParameter( + "`cell` must have the same dtype as `positions`".into() + )); } - /// TODO - pub fn get_pairs(&mut self, options: PairListOptions) -> Option<&TensorBlock> { - todo!() + let pbc_shape = pbc.shape(); + if pbc_shape.len() != 1 || pbc_shape[0] != 3 { + return Err(Error::InvalidParameter(format!( + "`pbc` must contain 3 entries, got a tensor with shape [{}]", + pbc_shape.iter().map(|dim| dim.to_string()).collect::>().join(", ") + ))); } - /// TODO - pub fn set_custom_data(&mut self, name: String, data: TensorMap) { - todo!() + if pbc.dtype() != dtype_bool { + return Err(Error::InvalidParameter( + "`pbc` must be a tensor of booleans".into() + )); } - /// TODO - pub fn get_custom_data(&self, name: &str) -> Option<&TensorMap> { - todo!() + return Ok(()); +} + +fn validate_cpu_system_data(system: &System) -> Result<(), Error> { + let pbc_array: ndarray::ArrayView1 = system.pbc().try_into()?; + + if system.dtype().bits == 32 { + let cell_array: ndarray::ArrayView2 = system.cell().try_into()?; + for i in 0..3 { + if !pbc_array[i] && !cell_array.row(i).iter().all(|&x| x == 0.0) { + return Err(Error::InvalidParameter(format!( + "invalid cell: for non-periodic dimensions, the corresponding \ + cell vector must be zero, but cell[{}] contains non-zero values", + i + ))); + } + } + } else { + let cell_array: ndarray::ArrayView2 = system.cell().try_into()?; + for i in 0..3 { + if !pbc_array[i] && !cell_array.row(i).iter().all(|&x| x == 0.0) { + return Err(Error::InvalidParameter(format!( + "invalid cell: for non-periodic dimensions, the corresponding \ + cell vector must be zero, but cell[{}] contains non-zero values", + i + ))); + } + } + } + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use super::*; + use metatensor::Labels; +use ndarray::{Array1, Array2}; + + // ----------------------------------------------------------------------- + // helpers to create DLPack tensors + // ----------------------------------------------------------------------- + fn type_tensor(data: &[i32]) -> DLPackTensor { + Array1::from_vec(data.to_vec()).try_into().unwrap() + } + + #[allow(clippy::cast_precision_loss)] + fn positions_tensor(n_atoms: usize, dtype: &str) -> DLPackTensor { + match dtype { + "f32" => { + let mut data = Vec::with_capacity(3 * n_atoms); + for i in 0..n_atoms { + data.extend_from_slice(&[i as f32, 0.0, 0.0]); + } + Array2::from_shape_vec((n_atoms, 3), data).unwrap().try_into().unwrap() + } + "f64" => { + let mut data = Vec::with_capacity(3 * n_atoms); + for i in 0..n_atoms { + data.extend_from_slice(&[i as f64, 0.0, 0.0]); + } + Array2::from_shape_vec((n_atoms, 3), data).unwrap().try_into().unwrap() + } + _ => panic!("unsupported dtype '{}'", dtype), + } + } + + #[allow(clippy::cast_possible_truncation)] + fn cell_tensor(size: f64, dtype: &str) -> DLPackTensor { + match dtype { + "f32" => { + Array2::::from_shape_vec( + (3, 3), + vec![ + size as f32, 0.0, 0.0, + 0.0, size as f32, 0.0, + 0.0, 0.0, size as f32, + ], + ).unwrap().try_into().unwrap() + } + "f64" => Array2::::from_shape_vec( + (3, 3), + vec![ + size, 0.0, 0.0, + 0.0, size, 0.0, + 0.0, 0.0, size, + ], + ).unwrap().try_into().unwrap(), + _ => panic!("unsupported dtype '{}'", dtype), + } + } + + fn pbc_tensor(data: &[bool]) -> DLPackTensor { + Array1::from_vec(data.to_vec()).try_into().unwrap() + } + + fn valid_pair_block(dtype: &str) -> TensorBlock { + let samples = Labels::new( + ["first_atom", "second_atom", "cell_shift_a", "cell_shift_b", "cell_shift_c"], + &[[0i32, 1, 0, 0, 0]], + ); + let components = vec![Labels::new(["xyz"], &[[0i32], [1], [2]])]; + let properties = Labels::new(["distance"], &[[0i32]]); + + match dtype { + "f32" => { + let values = ndarray::ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.5, 2.5, 3.5]).unwrap(); + TensorBlock::new(values, &samples, &components, &properties).unwrap() + } + "f64" => { + let values = ndarray::ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.5, 2.5, 3.5]).unwrap(); + TensorBlock::new(values, &samples, &components, &properties).unwrap() + } + _ => panic!("unsupported dtype '{}'", dtype), + } + } + + fn valid_custom_data(dtype: &str) -> TensorMap { + let keys = Labels::new(["key"], &[[0i32]]); + let samples = Labels::new(["sample"], &[[0i32]]); + let properties = Labels::new(["property"], &[[0i32]]); + + let block = match dtype { + "f32" => { + let values = ndarray::ArrayD::::from_shape_vec(vec![1, 1], vec![42.0]).unwrap(); + TensorBlock::new(values, &samples, &[], &properties).unwrap() + } + "f64" => { + let values = ndarray::ArrayD::::from_shape_vec(vec![1, 1], vec![42.0]).unwrap(); + TensorBlock::new(values, &samples, &[], &properties).unwrap() + } + _ => panic!("unsupported dtype '{}'", dtype), + }; + + TensorMap::new(keys, vec![block]).unwrap() + } + + fn assert_error(result: Result, expected: &str) { + let error = match result { + Ok(_) => panic!("expected error"), + Err(error) => error, + }; + assert_eq!(error.to_string(), expected); + } + + #[test] + fn system() { + let system = System::new( + "Angstrom".into(), + type_tensor(&[1, 6, 8]), + positions_tensor(3, "f32"), + cell_tensor(10.0, "f32"), + pbc_tensor(&[true, true, true]), + ).unwrap(); + + assert_eq!(system.length_unit(), "Angstrom"); + assert_eq!(system.size(), 3); + assert_eq!(system.device(), DLDevice::cpu()); + assert_eq!(system.dtype().bits, 32); + + let system = System::new( + "Angstrom".into(), + type_tensor(&[1, 6, 8]), + positions_tensor(3, "f64"), + cell_tensor(10.0, "f64"), + pbc_tensor(&[true, true, true]), + ).unwrap(); + assert_eq!(system.length_unit(), "Angstrom"); + assert_eq!(system.size(), 3); + assert_eq!(system.device(), DLDevice::cpu()); + assert_eq!(system.dtype().bits, 64); + } + + #[test] + fn system_invalid_tensors() { + let length_unit = "Angstrom".to_string(); + + let bad_types: DLPackTensor = Array1::::from_vec(vec![1.0, 2.0]).try_into().unwrap(); + let positions = positions_tensor(2, "f32"); + let cell = cell_tensor(0.0, "f32"); + let pbc = pbc_tensor(&[true, true, true]); + + assert_error( + System::new(length_unit.clone(), bad_types, positions, cell, pbc), + "invalid parameter: `types` must be a tensor of 32-bit integers", + ); + + let bad_types: DLPackTensor = Array2::::from_shape_vec((2, 2), vec![1, 2, 3, 4]).unwrap().try_into().unwrap(); + let positions = positions_tensor(2, "f32"); + let cell = cell_tensor(0.0, "f32"); + let pbc = pbc_tensor(&[true, true, true]); + assert_error( + System::new(length_unit.clone(), bad_types, positions, cell, pbc), + "invalid parameter: `types` must be a (n_atoms,) tensor, got a tensor with shape [2, 2]", + ); + + let types = type_tensor(&[1]); + let bad_positions: DLPackTensor = Array2::::from_shape_vec((1, 3), vec![1, 2, 3]).unwrap().try_into().unwrap(); + let cell = cell_tensor(0.0, "f32"); + let pbc = pbc_tensor(&[true, true, true]); + assert_error( + System::new(length_unit.clone(), types, bad_positions, cell, pbc), + "invalid parameter: `positions` must be a tensor of 32 or 64-bit floating point data", + ); + + let types = type_tensor(&[1, 6]); + let bad_positions = Array2::::from_shape_vec((2, 2), vec![0.0; 4]).unwrap().try_into().unwrap(); + let cell = cell_tensor(0.0, "f32"); + let pbc = pbc_tensor(&[true, true, true]); + assert_error( + System::new("Angstrom".into(), types, bad_positions, cell, pbc), + "invalid parameter: `positions` must be a (n_atoms x 3) tensor, got a tensor with shape [2, 2]", + ); + + let types = type_tensor(&[1, 6]); + let positions = positions_tensor(2, "f32"); + let bad_cell = Array2::::from_shape_vec((2, 3), vec![0.0; 6]).unwrap().try_into().unwrap(); + let pbc = pbc_tensor(&[true, true, true]); + assert_error( + System::new(length_unit.clone(), types, positions, bad_cell, pbc), + "invalid parameter: `cell` must be a (3 x 3) tensor, got a tensor with shape [2, 3]", + ); + + let types = type_tensor(&[1, 6]); + let positions = positions_tensor(2, "f32"); + let cell = cell_tensor(0.0, "f64"); + let pbc = pbc_tensor(&[true, true, true]); + assert_error( + System::new(length_unit.clone(), types, positions, cell, pbc), + "invalid parameter: `cell` must have the same dtype as `positions`", + ); + + let bad_pbc_dtype: DLPackTensor = Array1::::from_vec(vec![1, 0, 1]).try_into().unwrap(); + let types = type_tensor(&[1, 6]); + let positions = positions_tensor(2, "f32"); + let cell = cell_tensor(0.0, "f32"); + assert_error( + System::new(length_unit.clone(), types, positions, cell, bad_pbc_dtype), + "invalid parameter: `pbc` must be a tensor of booleans", + ); + + let types = type_tensor(&[1, 6]); + let positions = positions_tensor(2, "f32"); + let cell = cell_tensor(0.0, "f32"); + let bad_pbc = pbc_tensor(&[true, true]); + assert_error( + System::new(length_unit, types, positions, cell, bad_pbc), + "invalid parameter: `pbc` must contain 3 entries, got a tensor with shape [2]", + ); + } + + #[test] + fn system_periodic() { + let length_unit = "Angstrom".to_string(); + + // valid periodicity combinations: (1) fully periodic + let types = type_tensor(&[1]); + let positions = positions_tensor(1, "f32"); + let cell = cell_tensor(10.0, "f32"); + let pbc = pbc_tensor(&[true, true, true]); + System::new(length_unit.clone(), types, positions, cell, pbc).unwrap(); + + // (2) fully non-periodic with zero cell + let types = type_tensor(&[1]); + let positions = positions_tensor(1, "f32"); + let cell = cell_tensor(0.0, "f32"); + let pbc = pbc_tensor(&[false, false, false]); + System::new(length_unit.clone(), types, positions, cell, pbc).unwrap(); + + // (3) mixed periodic/non-periodic + let types = type_tensor(&[1]); + let positions = positions_tensor(1, "f32"); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + let pbc = pbc_tensor(&[true, false, true]); + System::new(length_unit.clone(), types, positions, cell, pbc).unwrap(); + + // invalid periodicity/cell + let types = type_tensor(&[1]); + let positions = positions_tensor(1, "f32"); + let cell = cell_tensor(10.0, "f32"); + let pbc = pbc_tensor(&[true, false, true]); + assert_error( + System::new(length_unit.clone(), types, positions, cell, pbc), + "invalid parameter: invalid cell: for non-periodic dimensions, the corresponding cell vector must be zero, but cell[1] contains non-zero values", + ); + } + + #[test] + fn add_pairs() { + let mut system = System::new( + "Angstrom".into(), + type_tensor(&[1, 6, 8]), + positions_tensor(3, "f32"), + cell_tensor(10.0, "f32"), + pbc_tensor(&[true, true, true]), + ).unwrap(); + + let options = PairListOptions { cutoff: 3.5, full_list: true, strict: false, requestors: vec![] }; + let pairs = valid_pair_block("f32"); + system.add_pairs(options.clone(), pairs).unwrap(); + assert_eq!(system.known_pairs().len(), 1); + assert_eq!(system.get_pairs(&options).unwrap().properties().names(), ["distance"]); + + let options_with_requestor = PairListOptions { + cutoff: 3.5, + full_list: true, + strict: false, + requestors: vec!["test-requestor".into()], + }; + // TODO: check that this is the exact same block once we can get the + // pointer to check for id. + assert!(system.get_pairs(&options_with_requestor).is_some()); + + system.add_pairs( + PairListOptions { cutoff: 5.0, full_list: false, strict: true, requestors: vec![] }, + valid_pair_block("f32"), + ).unwrap(); + assert_eq!(system.known_pairs().len(), 2); + } + + + #[test] + fn custom_data() { + let mut system = System::new( + "Angstrom".into(), + type_tensor(&[1, 6, 8]), + positions_tensor(3, "f32"), + cell_tensor(10.0, "f32"), + pbc_tensor(&[true, true, true]), + ).unwrap(); + + let data = valid_custom_data("f32"); + system.add_custom_data("test::my_data".into(), data, false).unwrap(); + assert_eq!(system.known_custom_data(), vec!["test::my_data"]); + assert_eq!(system.get_custom_data("test::my_data").unwrap().keys().names(), ["key"]); + + assert_error( + system.add_custom_data("test::my_data".into(), valid_custom_data("f32"), false), + "invalid parameter: custom data 'test::my_data' is already present in this system", + ); + + let replacement = valid_custom_data("f32"); + system.add_custom_data("test::my_data".into(), replacement, true).unwrap(); + assert_eq!(system.known_custom_data(), vec!["test::my_data"]); + + let mut system = System::new( + "Angstrom".into(), + type_tensor(&[1, 6, 8]), + positions_tensor(3, "f32"), + cell_tensor(10.0, "f32"), + pbc_tensor(&[true, true, true]), + ).unwrap(); + system.add_custom_data("test::a".into(), valid_custom_data("f32"), false).unwrap(); + system.add_custom_data("test::b".into(), valid_custom_data("f32"), false).unwrap(); + let mut names = system.known_custom_data(); + names.sort_unstable(); + assert_eq!(names, vec!["test::a", "test::b"]); + + // TODO: check we get back the same pointer + assert!(system.get_custom_data("test::a").is_ok()); + assert!(system.get_custom_data("test::b").is_ok()); + + assert_error( + system.get_custom_data("no_such_data"), + "invalid parameter: no data for 'no_such_data' found in this system", + ); + } + + #[test] + fn custom_data_validation() { + let mut system = System::new( + "Angstrom".into(), + type_tensor(&[1, 6, 8]), + positions_tensor(3, "f32"), + cell_tensor(10.0, "f32"), + pbc_tensor(&[true, true, true]), + ).unwrap(); + for name in ["types", "type", "Positions", "position", "CELL", "neighbors", "neighbor", "pair", "pairs", "Types", "POSITIONS", "Cell", "Neighbors"] { + let data = valid_custom_data("f32"); + assert_error( + system.add_custom_data(name.to_string(), data, false), + &format!("invalid parameter: custom data can not be named '{}'", name), + ); + } + + assert_error( + system.add_custom_data("my_data".into(), valid_custom_data("f32"), false), + "invalid parameter: 'my_data' is not a standard quantity name; custom quantity names must use '::'", + ); + + let keys = Labels::empty(vec!["key"]); + let empty = TensorMap::new(keys, vec![]).unwrap(); + assert_error( + system.add_custom_data("test::empty".into(), empty, false), + "invalid parameter: custom data 'test::empty' has no blocks", + ); + + let dtype_mismatch = valid_custom_data("f64"); + assert_error( + system.add_custom_data("test::dtype".into(), dtype_mismatch, false), + "invalid parameter: dtype of custom data 'test::dtype' does not match this system dtype", + ); } } From ade7e37a0dfe91b0ada053cb51137e1d6617fd66 Mon Sep 17 00:00:00 2001 From: Qianjun Xu <92628709+GardevoirX@users.noreply.github.com> Date: Wed, 3 Jun 2026 13:53:45 +0200 Subject: [PATCH 18/43] Implement `mta_format_metadata` --- metatomic-core/src/c_api/model.rs | 21 +++- metatomic-core/src/metadata.rs | 191 +++++++++++++++++++++++++++++- metatomic-core/tests/misc.cpp | 46 +++++++ 3 files changed, 254 insertions(+), 4 deletions(-) diff --git a/metatomic-core/src/c_api/model.rs b/metatomic-core/src/c_api/model.rs index d2050643b..869db8ede 100644 --- a/metatomic-core/src/c_api/model.rs +++ b/metatomic-core/src/c_api/model.rs @@ -1,6 +1,9 @@ use std::ffi::{c_void, c_char}; use metatensor::c_api::{mts_labels_t, mts_tensormap_t}; +use super::catch_unwind; +use crate::{Error, ModelMetadata}; + use super::{mta_status_t, mta_string_t, mta_system_t}; /// A model that computes physical properties of atomistic systems. @@ -224,5 +227,21 @@ pub unsafe extern "C" fn mta_format_metadata( metadata: *const c_char, printed: *mut mta_string_t, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(metadata, printed); + + let metadata = std::ffi::CStr::from_ptr(metadata); + let metadata = metadata.to_str().map_err(|_| { + Error::InvalidParameter("metadata is not valid UTF-8".into()) + })?; + + let metadata = json::parse(metadata).map_err(|e| { + Error::Serialization(format!("invalid JSON for ModelMetadata: {e}")) + })?; + + let metadata = ModelMetadata::try_from(&metadata)?; + + *printed = mta_string_t::new(metadata.print()); + Ok(()) + }) } diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs index a0f0c3875..3bf1ce00d 100644 --- a/metatomic-core/src/metadata.rs +++ b/metatomic-core/src/metadata.rs @@ -1,4 +1,5 @@ use std::collections::BTreeMap; +use std::fmt::Write; use json::JsonValue; @@ -188,6 +189,66 @@ impl<'a> TryFrom<&'a JsonValue> for References { } +fn normalize_whitespace(data: &str) -> String { + let mut normalized_string = String::new(); + for c in data.chars() { + if c == '\n' || c == '\r' || c == '\t' { + normalized_string.push(' '); + } else { + normalized_string.push(c); + } + } + normalized_string +} + + +fn wrap_80_chars(output: &mut String, data: &str, indent: usize) { + let string = normalize_whitespace(data); + assert!(indent < 30); + let line_length = 80 - indent; + assert!(line_length > 50); + let mut first_line = true; + let mut start = 0; + + loop { + let remaining = &string[start..]; + + if remaining.len() <= line_length { + if !first_line { + output.push_str(&" ".repeat(indent)); + } + output.push_str(remaining); + break; + } + + // byte offset of the character just past the first `line_length` chars + let end = remaining.char_indices().nth(line_length).map_or(remaining.len(), |(i, _)| i); + + if let Some(space_pos) = remaining[..end].rfind(' ') { + if !first_line { + output.push_str(&" ".repeat(indent)); + } + output.push_str(&remaining[..space_pos]); + output.push('\n'); + start += space_pos + 1; + first_line = false; + } else { + let word_end = remaining.find(' ').unwrap_or(remaining.len()); + if !first_line { + output.push_str(&" ".repeat(indent)); + } + output.push_str(&remaining[..word_end]); + output.push('\n'); + first_line = false; + if word_end < remaining.len() { + start += word_end + 1; + } else { + break; + } + } + } +} + /// Metadata about a model #[derive(Debug, Clone)] pub struct ModelMetadata { @@ -270,13 +331,103 @@ impl<'a> TryFrom<&'a JsonValue> for ModelMetadata { extra.insert(key.to_string(), value.to_string()); } - Ok(ModelMetadata { + // Validate the contents of `authors` and `references` + for author in &authors { + if author.is_empty() { + return Err(Error::InvalidParameter("author can not be empty string in ModelMetadata".into())); + } + } + + for model_ref in &references.model { + if model_ref.is_empty() { + return Err(Error::InvalidParameter("reference can not be empty string (in 'model' section)".into())); + } + } + + for architecture_ref in &references.architecture { + if architecture_ref.is_empty() { + return Err(Error::InvalidParameter("reference can not be empty string (in 'architecture' section)".into())); + } + } + + for implementation_ref in &references.implementation { + if implementation_ref.is_empty() { + return Err(Error::InvalidParameter("reference can not be empty string (in 'implementation' section)".into())); + } + } + + let metadata = ModelMetadata { name: name.to_string(), authors: authors, description: description, references: references, extra: extra, - }) + }; + Ok(metadata) + } +} + +impl ModelMetadata{ + pub fn print(&self) -> String { + let mut output = String::new(); + if self.name.is_empty() { + let _ = writeln!(output, "This is an unnamed model"); + let _ = writeln!(output, "========================"); + } else { + let _ = writeln!(output, "This is the {} model", &self.name); + let _ = writeln!(output, "============{}======", "=".repeat(self.name.len())); + } + + if !self.description.is_empty() { + let _ = writeln!(output); + wrap_80_chars(&mut output, &(self.description), 0); + let _ = writeln!(output); + } + + if !self.authors.is_empty() { + let _ = writeln!(output, "\nModel authors\n-------------\n"); + for author in &self.authors { + let _ = write!(output, "- "); + wrap_80_chars(&mut output, author, 2); + output.push('\n'); + } + } + + let mut references_output = String::new(); + if !self.references.model.is_empty() { + references_output.push_str("- about this specific model:\n"); + for reference in &self.references.model { + references_output.push_str(" * "); + wrap_80_chars(&mut references_output, reference, 4); + references_output.push('\n'); + } + } + + if !self.references.architecture.is_empty() { + references_output.push_str("- about the architecture of this model:\n"); + for reference in &self.references.architecture { + references_output.push_str(" * "); + wrap_80_chars(&mut references_output, reference, 4); + references_output.push('\n'); + } + } + + if !self.references.implementation.is_empty() { + references_output.push_str("- about the implementation of this model:\n"); + for reference in &self.references.implementation { + references_output.push_str(" * "); + wrap_80_chars(&mut references_output, reference, 4); + references_output.push('\n'); + } + } + + if !references_output.is_empty() { + output.push_str("\nModel references\n----------------\n\n"); + output.push_str("Please cite the following references when using this model:\n"); + output.push_str(&references_output); + } + + return output; } } @@ -486,6 +637,7 @@ impl<'a> TryFrom<&'a JsonValue> for ModelCapabilities { } } + #[cfg(test)] mod tests { mod pair_list_options { @@ -602,7 +754,8 @@ mod tests { } mod model_metadata { - use super::super::*; + +use super::super::*; fn example() -> ModelMetadata { ModelMetadata { @@ -704,6 +857,38 @@ mod tests { assert_eq!(error.to_string(), expected); } } + + #[test] + fn printing() { + let metadata = example(); + let output = metadata.print(); + let expected = String::from( + "This is the test-model model +============================ + +A test model + +Model authors +------------- + +- Alice +- Bob + +Model references +---------------- + +Please cite the following references when using this model: +- about this specific model: + * doi:10.1234/test +- about the architecture of this model: + * doi:10.1234/arch +- about the implementation of this model: + * https://github.com/test +" +); + + assert_eq!(output, expected); + } } mod model_capabilities { diff --git a/metatomic-core/tests/misc.cpp b/metatomic-core/tests/misc.cpp index 8c66d7561..93baecb67 100644 --- a/metatomic-core/tests/misc.cpp +++ b/metatomic-core/tests/misc.cpp @@ -70,3 +70,49 @@ TEST_CASE("mta_unit_conversion_factor") { "invalid parameter: dimension mismatch in unit conversion: " "'m' has dimension [L] but 'kg' has dimension [M]"); } + +TEST_CASE("mta_format_metadata") { + std::string json =R"({ + "type": "metatomic_model_metadata", + "name": "name", + "description": "Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua. Ut enim ad minim veniam, quis nostrud exercitation.", + "authors": ["Short author", "Some extremely long author that will take more than one line in the printed output"], + "references": { + "architecture": ["ref-2", "ref-3"], + "model": ["a very long reference that will take more than one line in the printed output"], + "implementation": [] + }, + "extra": {} +})"; + auto* mta_string = mta_string_create(""); + REQUIRE(mta_string != nullptr); + auto status = mta_format_metadata(json.c_str(), &mta_string); + REQUIRE(status == MTA_SUCCESS); + const auto expected = R"(This is the name model +====================== + +Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor +incididunt ut labore et dolore magna aliqua. Ut enim ad minim veniam, quis +nostrud exercitation. + +Model authors +------------- + +- Short author +- Some extremely long author that will take more than one line in the printed + output + +Model references +---------------- + +Please cite the following references when using this model: +- about this specific model: + * a very long reference that will take more than one line in the printed + output +- about the architecture of this model: + * ref-2 + * ref-3 +)"; + CHECK(std::string(mta_string_view(mta_string)) == expected); + mta_string_free(mta_string); +} From 63185466289d16551ddb05d790e0deff88617a89 Mon Sep 17 00:00:00 2001 From: Rocco Meli Date: Tue, 2 Jun 2026 10:56:37 +0200 Subject: [PATCH 19/43] Add unit conversion and error handling to C++ --- metatomic-core/include/metatomic.hpp | 1 + metatomic-core/include/metatomic/errors.hpp | 108 ++++++++++++++++++++ metatomic-core/include/metatomic/utils.hpp | 15 +++ metatomic-core/tests/misc.cpp | 69 ++++++++----- 4 files changed, 170 insertions(+), 23 deletions(-) create mode 100644 metatomic-core/include/metatomic/errors.hpp diff --git a/metatomic-core/include/metatomic.hpp b/metatomic-core/include/metatomic.hpp index 3b5c8ac2a..e41f09542 100644 --- a/metatomic-core/include/metatomic.hpp +++ b/metatomic-core/include/metatomic.hpp @@ -2,3 +2,4 @@ #include "metatomic/system.hpp" // IWYU pragma: export #include "metatomic/model.hpp" // IWYU pragma: export #include "metatomic/plugin.hpp" // IWYU pragma: export +#include "metatomic/errors.hpp" // IWYU pragma: export diff --git a/metatomic-core/include/metatomic/errors.hpp b/metatomic-core/include/metatomic/errors.hpp new file mode 100644 index 000000000..a926b388b --- /dev/null +++ b/metatomic-core/include/metatomic/errors.hpp @@ -0,0 +1,108 @@ +#pragma once + +#include +#include +#include +#include +#include + +#include + +namespace metatomic { + + /// Exception class used for all errors in metatomic + class Error: public std::runtime_error { + public: + /// Create a new MetatomicError with the given `message` + Error(const std::string& message): std::runtime_error(message) {} + }; + + namespace details { + /// Check if a return status from the C API indicates an error, and if it is + /// the case, throw an exception of type `metatomic::Error` with the last + /// error message from the library. + inline void check_status(mta_status_t status) { + if (status == MTA_SUCCESS) { + return; + } else if (status == MTA_MODEL_NOT_SUPPORTED_ERROR) { + const char* message = nullptr; + const char* origin = nullptr; + void* data = nullptr; + mta_last_error(&message, &origin, &data); + if (origin != nullptr &&std::strcmp(origin, "C++ exception") == 0 && data != nullptr) { + std::rethrow_exception(*static_cast(data)); + } else { + throw Error(message == nullptr ? "unknown error" : message); + } + } else { + const char* message = nullptr; + mta_last_error(&message, nullptr, nullptr); + throw Error(message == nullptr ? "unknown error" : message); + } + } + + /// Call the given `function` with the given `args` (the function should + /// return an `mta_status_t`), catching any C++ exception, and translating + /// them to native metatomic error code. + /// + /// This is required to prevent callbacks unwinding through the C API. + template + inline mta_status_t catch_exceptions(Function function, Args ...args) { + try { + function(std::move(args)...); + return MTA_SUCCESS; + } catch (...) { + auto* exception_ptr = new std::exception_ptr(std::current_exception()); + + const char* message = nullptr; + try { + std::rethrow_exception(*exception_ptr); + } catch (const std::exception& e) { + message = e.what(); + } catch (...) { + message = "C++ code threw an exception that was not a std::exception"; + } + + auto status = mta_set_last_error( + message, + "C++ exception", + exception_ptr, + [](void *ptr) { delete static_cast(ptr); } + ); + + if (status != MTA_SUCCESS) { + // If we failed to set the error, we are in a very bad state, + // but we should still try to report the original error + // message if possible. + std::fprintf(stderr, "INTERNAL ERROR: unable to set last error after C++ callback failure (status: %d). ", status); + if (message != nullptr) { + fprintf(stderr, "C++ error was: %s\n", message); + } else { + fprintf(stderr, "Unknown C++ error\n"); + } + delete exception_ptr; + } + + return MTA_MODEL_NOT_SUPPORTED_ERROR; + } + } + + /// Check if a pointer allocated by the C API is null, and if it is the + /// case, throw an exception of type `metatomic::Error` with the last + /// error message from the library. + inline void check_pointer(const void* pointer) { + if (pointer == nullptr) { + const char* message = nullptr; + const char* origin = nullptr; + void* data = nullptr; + mta_last_error(&message, &origin, &data); + if (std::strcmp(origin, "C++ exception") == 0 && data != nullptr) { + std::rethrow_exception(*static_cast(data)); + } else { + throw Error(message); + } + } + } + } // namespace details + +} // namespace metatomic diff --git a/metatomic-core/include/metatomic/utils.hpp b/metatomic-core/include/metatomic/utils.hpp index 1cae91bdf..402fee1a7 100644 --- a/metatomic-core/include/metatomic/utils.hpp +++ b/metatomic-core/include/metatomic/utils.hpp @@ -1,7 +1,22 @@ #pragma once +#include + #include +#include namespace metatomic { + inline double unit_conversion_factor( + const std::string& from_unit, + const std::string& to_unit + ) { + double conversion = 0.0; + + auto status = mta_unit_conversion_factor(from_unit.c_str(), to_unit.c_str(), &conversion); + details::check_status(status); + + return conversion; + } + } // namespace metatomic diff --git a/metatomic-core/tests/misc.cpp b/metatomic-core/tests/misc.cpp index 93baecb67..0028f1b88 100644 --- a/metatomic-core/tests/misc.cpp +++ b/metatomic-core/tests/misc.cpp @@ -3,6 +3,7 @@ #include #include "metatomic.h" +#include "metatomic.hpp" TEST_CASE("Version macros") { @@ -48,30 +49,52 @@ TEST_CASE("mta_string_t") { mta_string_free(nullptr); } -TEST_CASE("mta_unit_conversion_factor") { - double factor = 0.0; - - // same unit -> factor = 1.0 - auto status = mta_unit_conversion_factor("m", "m", &factor); - REQUIRE(status == MTA_SUCCESS); - CHECK(factor == 1.0); - - // kJ/mol -> eV - CHECK(mta_unit_conversion_factor("kJ/mol", "eV", &factor) == MTA_SUCCESS); - CHECK(factor == Approx(0.010364269656262174).epsilon(1e-15)); - - // dimension mismatch -> error - status = mta_unit_conversion_factor("m", "kg", &factor); - REQUIRE(status != MTA_SUCCESS); - - const char* error_msg = nullptr; - mta_last_error(&error_msg, nullptr, nullptr); - CHECK(std::string(error_msg) == - "invalid parameter: dimension mismatch in unit conversion: " - "'m' has dimension [L] but 'kg' has dimension [M]"); +TEST_CASE("unit conversion factor") { + SECTION("C API") { + double factor = 0.0; + + // same unit -> factor = 1.0 + auto status = mta_unit_conversion_factor("m", "m", &factor); + REQUIRE(status == MTA_SUCCESS); + CHECK(factor == 1.0); + + // kJ/mol -> eV + CHECK(mta_unit_conversion_factor("kJ/mol", "eV", &factor) == MTA_SUCCESS); + CHECK(factor == Approx(0.010364269656262174).epsilon(1e-15)); + + // dimension mismatch -> error + status = mta_unit_conversion_factor("m", "kg", &factor); + REQUIRE(status != MTA_SUCCESS); + + const char* error_msg = nullptr; + mta_last_error(&error_msg, nullptr, nullptr); + CHECK(std::string(error_msg) == + "invalid parameter: dimension mismatch in unit conversion: " + "'m' has dimension [L] but 'kg' has dimension [M]" + ); + } + + SECTION("C++ API") { + // same unit -> factor = 1.0 + auto factor = metatomic::unit_conversion_factor("m", "m"); + CHECK(factor == 1.0); + + // kJ/mol -> eV + factor = metatomic::unit_conversion_factor("kJ/mol", "eV"); + CHECK(factor == Approx(0.010364269656262174).epsilon(1e-15)); + + // dimension mismatch -> error + try{ + factor = metatomic::unit_conversion_factor("m", "kg"); + } + catch(metatomic::Error& e){ + CHECK(std::string(e.what()) == "invalid parameter: dimension mismatch in unit conversion: 'm' has dimension [L] but 'kg' has dimension [M]"); + } + } } -TEST_CASE("mta_format_metadata") { + +TEST_CASE("metatdata formatting") { std::string json =R"({ "type": "metatomic_model_metadata", "name": "name", @@ -88,7 +111,7 @@ TEST_CASE("mta_format_metadata") { REQUIRE(mta_string != nullptr); auto status = mta_format_metadata(json.c_str(), &mta_string); REQUIRE(status == MTA_SUCCESS); - const auto expected = R"(This is the name model + const auto* expected = R"(This is the name model ====================== Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor From d7a4053fb5ea34d0b43017890a529cda6b495202 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Wed, 3 Jun 2026 13:51:04 +0200 Subject: [PATCH 20/43] Add C++ documentation --- docs/src/core/index.rst | 1 + docs/src/core/reference/cxx/index.rst | 17 +++++++++++++++++ docs/src/core/reference/cxx/misc.rst | 14 ++++++++++++++ docs/src/core/reference/cxx/model.rst | 2 ++ docs/src/core/reference/cxx/plugin.rst | 2 ++ docs/src/core/reference/cxx/system.rst | 2 ++ docs/src/core/units.rst | 20 ++++++++++++-------- metatomic-core/include/metatomic.h | 20 ++++++++++---------- metatomic-core/include/metatomic/utils.hpp | 18 ++++++++++++++++-- metatomic-core/src/c_api/utils.rs | 20 ++++++++++---------- 10 files changed, 86 insertions(+), 30 deletions(-) create mode 100644 docs/src/core/reference/cxx/index.rst create mode 100644 docs/src/core/reference/cxx/misc.rst create mode 100644 docs/src/core/reference/cxx/model.rst create mode 100644 docs/src/core/reference/cxx/plugin.rst create mode 100644 docs/src/core/reference/cxx/system.rst diff --git a/docs/src/core/index.rst b/docs/src/core/index.rst index 2a3691316..66bbcffb7 100644 --- a/docs/src/core/index.rst +++ b/docs/src/core/index.rst @@ -8,6 +8,7 @@ WIP :maxdepth: 2 reference/c/index + reference/cxx/index reference/json-formats units diff --git a/docs/src/core/reference/cxx/index.rst b/docs/src/core/reference/cxx/index.rst new file mode 100644 index 000000000..9a4a7add3 --- /dev/null +++ b/docs/src/core/reference/cxx/index.rst @@ -0,0 +1,17 @@ +.. _cxx-api-core: + +C++ API reference +================= + +WIP + +The functions and types provided in ``metatomic.hpp`` can be grouped in four +main groups: + +.. toctree:: + :maxdepth: 1 + + system + model + plugin + misc diff --git a/docs/src/core/reference/cxx/misc.rst b/docs/src/core/reference/cxx/misc.rst new file mode 100644 index 000000000..26ba29607 --- /dev/null +++ b/docs/src/core/reference/cxx/misc.rst @@ -0,0 +1,14 @@ +Miscellaneous +============= + + +Error handling +^^^^^^^^^^^^^^ + +.. doxygenclass:: metatomic::Error + + +Unit conversion +^^^^^^^^^^^^^^^ + +.. doxygenfunction:: metatomic::unit_conversion_factor diff --git a/docs/src/core/reference/cxx/model.rst b/docs/src/core/reference/cxx/model.rst new file mode 100644 index 000000000..75338c89d --- /dev/null +++ b/docs/src/core/reference/cxx/model.rst @@ -0,0 +1,2 @@ +Model +===== diff --git a/docs/src/core/reference/cxx/plugin.rst b/docs/src/core/reference/cxx/plugin.rst new file mode 100644 index 000000000..67cd50b04 --- /dev/null +++ b/docs/src/core/reference/cxx/plugin.rst @@ -0,0 +1,2 @@ +Plugin system +============= diff --git a/docs/src/core/reference/cxx/system.rst b/docs/src/core/reference/cxx/system.rst new file mode 100644 index 000000000..3dcbaeea1 --- /dev/null +++ b/docs/src/core/reference/cxx/system.rst @@ -0,0 +1,2 @@ +System +====== diff --git a/docs/src/core/units.rst b/docs/src/core/units.rst index c9415ed9f..6c50603ca 100644 --- a/docs/src/core/units.rst +++ b/docs/src/core/units.rst @@ -6,21 +6,25 @@ Units Models in metatensor can use arbitrary units for their inputs and outputs. The unit conversion system allows models to specify the units they expect and receive data in any compatible unit, with automatic conversion handled by -:c:func:`mta_execute_model`. +during model execution. -The :c:func:`mta_unit_conversion_factor` function parses two unit expressions, -checks that they have compatible physical dimensions, and returns the -multiplicative conversion factor: +Unit parsing is handled by one of the following functions: -.. code-block:: c +- :c:func:`mta_unit_conversion_factor` in C +- :cpp:func:`metatomic::unit_conversion_factor` in C++ + +These functions parses two unit expressions, checks that they have compatible +physical dimensions, and returns the multiplicative conversion factor. For +example, in C++: + +.. code-block:: C++ // How many eV are in one kJ/mol? - double factor; - mta_unit_conversion_factor("kJ/mol", "eV", &factor); + double factor = metatomic::unit_conversion_factor("kJ/mol", "eV"); // factor ≈ 0.01036 // How many GPa are in one eV/A^3? - mta_unit_conversion_factor("eV/A^3", "GPa", &factor); + factor = metatomic::unit_conversion_factor("eV/A^3", "GPa"); // factor ≈ 160.22 If either (or both) unit strings are empty, the conversion returns ``1.0`` diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index 9273007d0..e7de4118b 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -380,18 +380,18 @@ void mta_string_free(mta_string_t string); const char *mta_string_view(mta_string_t string); /** - * Get the multiplicative conversion factor to use to convert from - * `from_unit` to `to_unit`. Both units are parsed as expressions (e.g. - * "kJ/mol/A^2", "(eV*u)^(1/2)") and their dimensions must match. + * Get the multiplicative conversion factor to use to convert from `from_unit` + * to `to_unit`. Both units are parsed as expressions (e.g. `kJ / mol / A^2`, + * `(eV * u)^(1/2)`) and their dimensions must match. * - * Unit expressions are built from base units combined with `*`, `/`, `^`, - * and parentheses. Unit lookup is case-insensitive, and whitespace is - * ignored. For example: + * @verbatim embed:rst:leading-asterisk * - * - `"kJ/mol"` -- energy per mole - * - `"eV/Angstrom^3"` -- pressure - * - `"(eV*u)^(1/2)"` -- momentum (fractional powers) - * - `"Hartree/Bohr"` -- force in atomic units + * .. seealso:: + * + * The general documentation for :ref:`units`, with the expression + * syntax and list of supported base units. + * + * @endverbatim * * @param from_unit A null-terminated C string containing the unit to convert from. * @param to_unit A null-terminated C string containing the unit to convert to. diff --git a/metatomic-core/include/metatomic/utils.hpp b/metatomic-core/include/metatomic/utils.hpp index 402fee1a7..38f6aaf79 100644 --- a/metatomic-core/include/metatomic/utils.hpp +++ b/metatomic-core/include/metatomic/utils.hpp @@ -6,7 +6,22 @@ #include namespace metatomic { - + /// Get the multiplicative conversion factor to use to convert from + /// `from_unit` to `to_unit`. Both units are parsed as expressions + /// (e.g. `kJ / mol / A^2`, `(eV * u)^(1/2)`) and their dimensions must + /// match. + /// + /// @verbatim embed:rst:leading-slashes + /// + /// .. seealso:: + /// + /// The general documentation for :ref:`units`, with the expression + /// syntax and list of supported base units. + /// + /// @endverbatim + /// + /// @param from_unit the unit to convert from + /// @param to_unit the unit to convert to inline double unit_conversion_factor( const std::string& from_unit, const std::string& to_unit @@ -18,5 +33,4 @@ namespace metatomic { return conversion; } - } // namespace metatomic diff --git a/metatomic-core/src/c_api/utils.rs b/metatomic-core/src/c_api/utils.rs index 448c2b5d4..38e5222c4 100644 --- a/metatomic-core/src/c_api/utils.rs +++ b/metatomic-core/src/c_api/utils.rs @@ -147,18 +147,18 @@ pub unsafe extern "C" fn mta_string_view( return result; } -/// Get the multiplicative conversion factor to use to convert from -/// `from_unit` to `to_unit`. Both units are parsed as expressions (e.g. -/// "kJ/mol/A^2", "(eV*u)^(1/2)") and their dimensions must match. +/// Get the multiplicative conversion factor to use to convert from `from_unit` +/// to `to_unit`. Both units are parsed as expressions (e.g. `kJ / mol / A^2`, +/// `(eV * u)^(1/2)`) and their dimensions must match. /// -/// Unit expressions are built from base units combined with `*`, `/`, `^`, -/// and parentheses. Unit lookup is case-insensitive, and whitespace is -/// ignored. For example: +/// @verbatim embed:rst:leading-asterisk /// -/// - `"kJ/mol"` -- energy per mole -/// - `"eV/Angstrom^3"` -- pressure -/// - `"(eV*u)^(1/2)"` -- momentum (fractional powers) -/// - `"Hartree/Bohr"` -- force in atomic units +/// .. seealso:: +/// +/// The general documentation for :ref:`units`, with the expression +/// syntax and list of supported base units. +/// +/// @endverbatim /// /// @param from_unit A null-terminated C string containing the unit to convert from. /// @param to_unit A null-terminated C string containing the unit to convert to. From 52b55f8a4a9b6ecbf3201c5a379856326b85e9d7 Mon Sep 17 00:00:00 2001 From: alessandroforina Date: Thu, 11 Jun 2026 11:39:20 +0200 Subject: [PATCH 21/43] Expose mta_model_t function in rust's Model struct --- metatomic-core/src/lib.rs | 2 +- metatomic-core/src/model.rs | 293 +++++++++++++++++++++++++++++++++++- 2 files changed, 288 insertions(+), 7 deletions(-) diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index cc20d6961..9e1acb29b 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -20,7 +20,7 @@ pub mod c_api; mod metadata; use crate::c_api::mta_status_t; -pub use self::metadata::{ModelMetadata, PairListOptions}; +pub use self::metadata::{Device, DType, ModelCapabilities, ModelMetadata, PairListOptions}; mod quantities; pub use self::quantities::{Quantity, SampleKind, Gradients}; diff --git a/metatomic-core/src/model.rs b/metatomic-core/src/model.rs index 1b93577d2..a9ab74b81 100644 --- a/metatomic-core/src/model.rs +++ b/metatomic-core/src/model.rs @@ -1,21 +1,153 @@ -use metatensor::{Labels, TensorMap}; +use std::ffi::c_void; -use crate::{Error, Quantity, System}; +use metatensor::{Labels, TensorMap}; -use crate::c_api::mta_model_t; +use crate::{Error, ModelCapabilities, ModelMetadata, PairListOptions, Quantity, System}; +use crate::c_api::{mta_model_t, mta_status_t, mta_string_t, mta_string_free}; -/// TODO +/// A loaded atomistic model, ready to be executed on a set of systems. +/// +/// `Model` wraps a [`mta_model_t`] vtable provided by a plugin. It gives +/// access to the model's metadata and capabilities, and can be run with +/// [`execute_model`]. pub struct Model(pub(crate) mta_model_t); +impl Drop for Model { + fn drop(&mut self) { + if let Some(unload) = self.0.unload { + unsafe { unload(self.0.data) }; + } + } +} + +fn call_string_callback( + callback: unsafe extern "C" fn(*const c_void, *mut mta_string_t) -> mta_status_t, + data: *const c_void, +) -> Result { + let mut output = mta_string_t::null(); + let status = unsafe { callback(data, &mut output) }; + if status != mta_status_t::MTA_SUCCESS { + unsafe { mta_string_free(output) }; + return Err(Error::CallbackError(status)); + } + let json_str = output.as_str().to_owned(); + unsafe { mta_string_free(output) }; + return Ok(json_str); +} + impl Model { /// Create a new `Model` from the corresponding C API struct. + /// + /// The `Model` takes ownership of `model` and will call its `unload` + /// callback when dropped. pub fn new(model: mta_model_t) -> Self { return Model(model); } - /// Extract the underlying C API struct. + /// Extract the underlying C API struct, transferring ownership to the caller. + /// + /// The caller is responsible for eventually calling the `unload` callback + /// on the returned [`mta_model_t`] to free its resources. The `Model`'s + /// own `Drop` implementation is skipped. pub fn into_raw(self) -> mta_model_t { - return self.0; + let model = std::mem::ManuallyDrop::new(self); + return unsafe { std::ptr::read(&model.0) }; + } + + /// Get the metadata describing this model (name, authors, description, + /// references, ...). + pub fn metadata(&self) -> Result { + let callback = self.0.metadata.ok_or_else(|| { + Error::Internal("model is missing a 'metadata' callback".into()) + })?; + let json_str = call_string_callback(callback, self.0.data)?; + let json = json::parse(&json_str).map_err(|e| { + Error::Serialization(format!("model returned invalid JSON for metadata: {}", e)) + })?; + return ModelMetadata::try_from(&json); + } + + /// Get the capabilities of this model: which outputs it can compute, which + /// atomic types it supports, its interaction range, length unit, supported + /// devices, and data type. + pub fn capabilities(&self) -> Result { + let callback = self.0.capabilities.ok_or_else(|| { + Error::Internal("model is missing a 'capabilities' callback".into()) + })?; + let json_str = call_string_callback(callback, self.0.data)?; + let json = json::parse(&json_str).map_err(|e| { + Error::Serialization(format!("model returned invalid JSON for capabilities: {}", e)) + })?; + return ModelCapabilities::try_from(&json); + } + + /// Get the pair lists (neighbor lists) this model needs as input. + /// + /// The engine must compute these and attach them to every system with + /// `mta_system_add_pairs` before calling [`execute_model`]. + pub fn requested_pair_lists(&self) -> Result, Error> { + let callback = self.0.requested_pair_lists.ok_or_else(|| { + Error::Internal("model is missing a 'requested_pair_lists' callback".into()) + })?; + let json_str = call_string_callback(callback, self.0.data)?; + let json = json::parse(&json_str).map_err(|e| { + Error::Serialization(format!("model returned invalid JSON for requested_pair_lists: {}", e)) + })?; + if !json.is_array() { + return Err(Error::Serialization( + "model returned invalid JSON for requested_pair_lists, expected an array".into() + )); + } + let mut result = Vec::new(); + for item in json.members() { + result.push(PairListOptions::try_from(item)?); + } + return Ok(result); + } + + /// Get the additional per-system inputs this model needs. + /// + /// The engine must attach these to every system with + /// `mta_system_add_custom_data` before calling [`execute_model`]. + pub fn requested_inputs(&self) -> Result, Error> { + let callback = self.0.requested_inputs.ok_or_else(|| { + Error::Internal("model is missing a 'requested_inputs' callback".into()) + })?; + let json_str = call_string_callback(callback, self.0.data)?; + let json = json::parse(&json_str).map_err(|e| { + Error::Serialization(format!("model returned invalid JSON for requested_inputs: {}", e)) + })?; + if !json.is_array() { + return Err(Error::Serialization( + "model returned invalid JSON for requested_inputs, expected an array".into() + )); + } + let mut result = Vec::new(); + for item in json.members() { + result.push(Quantity::try_from(item)?); + } + return Ok(result); + } + + /// Get the outputs this model can compute. + pub fn supported_outputs(&self) -> Result, Error> { + let callback = self.0.supported_outputs.ok_or_else(|| { + Error::Internal("model is missing a 'supported_outputs' callback".into()) + })?; + let json_str = call_string_callback(callback, self.0.data)?; + let json = json::parse(&json_str).map_err(|e| { + Error::Serialization(format!("model returned invalid JSON for supported_outputs: {}", e)) + })?; + if !json.is_array() { + return Err(Error::Serialization( + "model returned invalid JSON for supported_outputs, expected an array".into() + )); + } + let mut result = Vec::new(); + for item in json.members() { + result.push(Quantity::try_from(item)?); + } + return Ok(result); } } @@ -29,3 +161,152 @@ pub fn execute_model( ) -> Result, Error> { todo!() } + +#[cfg(test)] +mod tests { + use super::*; + use crate::c_api::{mta_model_t, mta_status_t, mta_string_t}; + + + // Each function below is a stand-in for what a real plugin would implement. + // They simply write a hard-coded JSON string into the output mta_string_t + // and return MTA_SUCCESS. + unsafe extern "C" fn metadata_impl( + _data: *const c_void, + out: *mut mta_string_t, + ) -> mta_status_t { + *out = mta_string_t::new(r#"{ + "type": "metatomic_model_metadata", + "name": "test-model", + "authors": ["Alice"], + "description": "A test model", + "references": {"model": [], "architecture": [], "implementation": []}, + "extra": {} + }"#); + return mta_status_t::MTA_SUCCESS; + } + + unsafe extern "C" fn capabilities_impl( + _data: *const c_void, + out: *mut mta_string_t, + ) -> mta_status_t { + *out = mta_string_t::new(r#"{ + "type": "metatomic_model_capabilities", + "outputs": [{"type": "metatomic_quantity", "name": "energy", "unit": "eV", "gradients": [], "sample_kind": "system"}], + "atomic_types": [1, 6], + "interaction_range": 5.0, + "length_unit": "Angstrom", + "supported_devices": ["cpu"], + "dtype": "float32" + }"#); + return mta_status_t::MTA_SUCCESS; + } + + unsafe extern "C" fn requested_pair_lists_impl( + _data: *const c_void, + out: *mut mta_string_t, + ) -> mta_status_t { + *out = mta_string_t::new(format!(r#"[{{ + "type": "metatomic_pair_options", + "cutoff": "{:#x}", + "full_list": true, + "strict": true + }}]"#, 3.5_f64.to_bits())); + return mta_status_t::MTA_SUCCESS; + } + + unsafe extern "C" fn requested_inputs_impl( + _data: *const c_void, + out: *mut mta_string_t, + ) -> mta_status_t { + *out = mta_string_t::new(r#"[{ + "type": "metatomic_quantity", + "name": "charge", + "unit": "e", + "gradients": [], + "sample_kind": "atom" + }]"#); + return mta_status_t::MTA_SUCCESS; + } + + unsafe extern "C" fn supported_outputs_impl( + _data: *const c_void, + out: *mut mta_string_t, + ) -> mta_status_t { + *out = mta_string_t::new(r#"[ + { + "type": "metatomic_quantity", + "name": "energy", + "unit": "eV", + "gradients": ["positions"], + "sample_kind": "system" + }, + { + "type": "metatomic_quantity", + "name": "custom::output", + "unit": "", + "gradients": [], + "sample_kind": "atom_pair" + }]"#); + return mta_status_t::MTA_SUCCESS; + } + + + fn test_model() -> Model { + Model(mta_model_t { + metadata: Some(metadata_impl), + capabilities: Some(capabilities_impl), + requested_pair_lists: Some(requested_pair_lists_impl), + requested_inputs: Some(requested_inputs_impl), + supported_outputs:Some(supported_outputs_impl), + ..mta_model_t::null() + }) + } + + #[test] + fn metadata() { + let metadata = test_model().metadata().unwrap(); + assert_eq!(metadata.name, "test-model"); + assert_eq!(metadata.authors, vec!["Alice"]); + assert_eq!(metadata.description, "A test model"); + } + + + #[test] + fn capabilities() { + let capabilities = test_model().capabilities().unwrap(); + assert_eq!(capabilities.outputs.len(), 1); + assert_eq!(capabilities.outputs[0].name, "energy"); + assert_eq!(capabilities.atomic_types, vec![1, 6]); + assert_eq!(capabilities.interaction_range.to_bits(), 5.0_f64.to_bits()); + assert_eq!(capabilities.length_unit, "Angstrom"); + } + + #[test] + fn requested_pair_lists() { + let options = test_model().requested_pair_lists().unwrap(); + assert_eq!(options.len(), 1); + assert_eq!(options[0].cutoff.to_bits(), 3.5_f64.to_bits()); + assert!(options[0].full_list); + assert!(options[0].strict); + } + + #[test] + fn requested_inputs() { + let inputs = test_model().requested_inputs().unwrap(); + assert_eq!(inputs.len(), 1); + assert_eq!(inputs[0].name, "charge"); + assert_eq!(inputs[0].unit, "e"); + } + + #[test] + fn supported_outputs() { + let outputs = test_model().supported_outputs().unwrap(); + assert_eq!(outputs.len(), 2); + assert_eq!(outputs[0].name, "energy"); + assert_eq!(outputs[0].unit, "eV"); + + assert_eq!(outputs[1].name, "custom::output"); + assert_eq!(outputs[1].unit, ""); + } +} From 9c228095a450c03b3d4fd0e7eb56c69dfcf4c288 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Tue, 16 Jun 2026 13:56:37 +0200 Subject: [PATCH 22/43] Cleanup CMake code calling cargo --- metatomic-core/CMakeLists.txt | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/metatomic-core/CMakeLists.txt b/metatomic-core/CMakeLists.txt index 717e52e81..6a2169173 100644 --- a/metatomic-core/CMakeLists.txt +++ b/metatomic-core/CMakeLists.txt @@ -356,9 +356,6 @@ set(CARGO_ENV "METATOMIC_FULL_VERSION=${METATOMIC_FULL_VERSION}") if (NOT "${CMAKE_OSX_DEPLOYMENT_TARGET}" STREQUAL "") list(APPEND CARGO_ENV "MACOSX_DEPLOYMENT_TARGET=${CMAKE_OSX_DEPLOYMENT_TARGET}") endif() -if (NOT "$ENV{RUSTC_WRAPPER}" STREQUAL "") - list(APPEND CARGO_ENV "RUSTC_WRAPPER=$ENV{RUSTC_WRAPPER}") -endif() if (METATOMIC_INSTALL_BOTH_STATIC_SHARED) set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--crate-type=cdylib;--crate-type=staticlib") @@ -389,7 +386,7 @@ endif() add_custom_command( OUTPUT ${CARGO_OUTPUTS} COMMAND ${CMAKE_COMMAND} -E env ${CARGO_ENV} - cargo rustc ${CARGO_BUILD_ARG} -- ${CARGO_RUSTC_ARGS} + ${CARGO_EXE} rustc ${CARGO_BUILD_ARG} -- ${CARGO_RUSTC_ARGS} WORKING_DIRECTORY ${PROJECT_SOURCE_DIR} DEPENDS ${ALL_RUST_SOURCES} COMMENT "Building ${FILE_CREATED_MESSAGE} with cargo" From f2ae54870f257df745c262431bed40381a2415e9 Mon Sep 17 00:00:00 2001 From: Rocco Meli Date: Tue, 23 Jun 2026 13:22:29 +0200 Subject: [PATCH 23/43] C++ API to format model metadata (#266) --- metatomic-core/include/metatomic/model.hpp | 17 +++++++++++++++++ metatomic-core/tests/misc.cpp | 20 ++++++++++++++------ 2 files changed, 31 insertions(+), 6 deletions(-) diff --git a/metatomic-core/include/metatomic/model.hpp b/metatomic-core/include/metatomic/model.hpp index 1cae91bdf..fc4df3911 100644 --- a/metatomic-core/include/metatomic/model.hpp +++ b/metatomic-core/include/metatomic/model.hpp @@ -1,7 +1,24 @@ #pragma once +#include + #include +#include namespace metatomic { + /// Render model metadata as a human-readable string. + /// + /// @param metadata a JSON-serialized `ModelMetadata` object as produced by a + /// model's `metadata` callback + /// @return a human-readable rendering of the metadata + inline std::string format_metadata(const std::string& metadata) { + mta_string_t printed = nullptr; + auto status = mta_format_metadata(metadata.c_str(), &printed); + details::check_status(status); + + auto result = std::string(mta_string_view(printed)); + mta_string_free(printed); + return result; + } } // namespace metatomic diff --git a/metatomic-core/tests/misc.cpp b/metatomic-core/tests/misc.cpp index 0028f1b88..41bbe5b23 100644 --- a/metatomic-core/tests/misc.cpp +++ b/metatomic-core/tests/misc.cpp @@ -107,10 +107,6 @@ TEST_CASE("metatdata formatting") { }, "extra": {} })"; - auto* mta_string = mta_string_create(""); - REQUIRE(mta_string != nullptr); - auto status = mta_format_metadata(json.c_str(), &mta_string); - REQUIRE(status == MTA_SUCCESS); const auto* expected = R"(This is the name model ====================== @@ -136,6 +132,18 @@ Please cite the following references when using this model: * ref-2 * ref-3 )"; - CHECK(std::string(mta_string_view(mta_string)) == expected); - mta_string_free(mta_string); + + SECTION("C API") { + auto* mta_string = mta_string_create(""); + REQUIRE(mta_string != nullptr); + auto status = mta_format_metadata(json.c_str(), &mta_string); + REQUIRE(status == MTA_SUCCESS); + CHECK(std::string(mta_string_view(mta_string)) == expected); + mta_string_free(mta_string); + } + + SECTION("C++ API") { + auto result = metatomic::format_metadata(json); + CHECK(result == expected); + } } From 9eb7f7b464edecb4ada7aee5b632e0151969e633 Mon Sep 17 00:00:00 2001 From: Rocco Meli Date: Wed, 24 Jun 2026 16:24:13 +0200 Subject: [PATCH 24/43] Add C++ API for loading plugins --- metatomic-core/include/metatomic.h | 4 +- metatomic-core/include/metatomic/plugin.hpp | 44 +++++++++++ metatomic-core/src/c_api/plugin.rs | 4 +- metatomic-core/src/plugin.rs | 14 +++- metatomic-core/tests/c-model.cpp | 2 +- metatomic-core/tests/misc.cpp | 12 ++- metatomic-core/tests/plugins.cpp | 84 ++++++++++++++------- 7 files changed, 121 insertions(+), 43 deletions(-) diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index e7de4118b..b1899af0d 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -572,9 +572,9 @@ enum mta_status_t mta_load_plugin(const char *path); * status code if an error occurs. You can get more details about the * error with `mta_last_error`. */ -enum mta_status_t mta_load_model(const char *plugin_name, - const char *load_from, +enum mta_status_t mta_load_model(const char *load_from, const char *options_json, + const char *plugin_name, struct mta_model_t *model); #ifdef __cplusplus diff --git a/metatomic-core/include/metatomic/plugin.hpp b/metatomic-core/include/metatomic/plugin.hpp index 1cae91bdf..42300984f 100644 --- a/metatomic-core/include/metatomic/plugin.hpp +++ b/metatomic-core/include/metatomic/plugin.hpp @@ -1,7 +1,51 @@ #pragma once +#include + #include +#include namespace metatomic { + /// Load the shared library at `path` and register the plugin contained + /// within. The library must export the symbols generated by the + /// `MTA_REGISTER_PLUGIN` macro. + /// + /// @param path path to the plugin shared library + inline void load_plugin(const std::string& path) { + auto status = mta_load_plugin(path.c_str()); + details::check_status(status); + } + + /// Load a model from `load_from` with the given options. + /// + /// If `plugin_name` is empty, metatomic will try to determine the correct + /// plugin to use by checking the `load_from` parameter. If we can not + /// determine the correct plugin, we then try to load the model with each + /// registered plugin until one succeeds. + /// + /// If `plugin_name` is given, then we only try to load the model with the + /// specified plugin, and return an error if the plugin can not load the + /// model. + /// + /// @param load_from where to load the model from (e.g. a file path, a + /// model name, etc.) + /// @param plugin_name optional name of the plugin to use for loading the + /// model, or empty to let metatomic search + /// @param options_json optional JSON object containing string keys and + /// string values for loading the model + /// @return the loaded model + inline mta_model_t load_model( + const std::string& load_from, + const std::string& options_json = "", + const std::string& plugin_name = "" + ) { + mta_model_t model; + const char* plugin_name_ptr = plugin_name.empty() ? nullptr : plugin_name.c_str(); + const char* options_json_ptr = options_json.empty() ? nullptr : options_json.c_str(); + + auto status = mta_load_model(load_from.c_str(), options_json_ptr, plugin_name_ptr, &model); + details::check_status(status); + return model; + } } // namespace metatomic diff --git a/metatomic-core/src/c_api/plugin.rs b/metatomic-core/src/c_api/plugin.rs index 76b9989d1..42ee45dee 100644 --- a/metatomic-core/src/c_api/plugin.rs +++ b/metatomic-core/src/c_api/plugin.rs @@ -112,9 +112,9 @@ pub unsafe extern "C" fn mta_load_plugin(path: *const c_char) -> mta_status_t { /// error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_load_model( - plugin_name: *const c_char, load_from: *const c_char, options_json: *const c_char, + plugin_name: *const c_char, model: *mut mta_model_t, ) -> mta_status_t { let unwind_wrapper = std::panic::AssertUnwindSafe(model); @@ -157,7 +157,7 @@ pub unsafe extern "C" fn mta_load_model( } } - let loaded = crate::plugin::load_model(plugin_name, CStr::from_ptr(load_from), options_json)?; + let loaded = crate::plugin::load_model(CStr::from_ptr(load_from), options_json, plugin_name)?; let _ = &unwind_wrapper; *unwind_wrapper.0 = loaded.into_raw(); diff --git a/metatomic-core/src/plugin.rs b/metatomic-core/src/plugin.rs index 4a9da5ab9..cd087965c 100644 --- a/metatomic-core/src/plugin.rs +++ b/metatomic-core/src/plugin.rs @@ -141,16 +141,26 @@ pub fn load_plugin(path: &str) -> Result<(), Error> { /// Load a model from `load_from`, using the given options. pub fn load_model( - plugin_name: Option<&str>, load_from: &CStr, options_json: &CStr, + plugin_name: Option<&str>, ) -> Result { let plugins = PLUGINS.lock().expect("plugin registry mutex was poisoned"); if let Some(plugin_name) = plugin_name { for plugin in plugins.iter() { if plugin.name() == plugin_name { - return plugin.load_model(load_from, options_json); + return plugin.load_model(load_from, options_json).map_err(|e| { + if let Error::CallbackError(mta_status_t::MTA_MODEL_NOT_SUPPORTED_ERROR) = e { + Error::InvalidParameter(format!( + "failed to load model from '{}': plugin '{}' could not load the model", + load_from.to_string_lossy(), + plugin_name + )) + } else { + e + } + }); } } diff --git a/metatomic-core/tests/c-model.cpp b/metatomic-core/tests/c-model.cpp index ec1da92f7..500e2a956 100644 --- a/metatomic-core/tests/c-model.cpp +++ b/metatomic-core/tests/c-model.cpp @@ -162,7 +162,7 @@ TEST_CASE("simple C model can be registered and loaded through the C API") { mta_register_plugin(PLUGIN); auto model = mta_model_t{}; - auto status = mta_load_model("test-c-plugin", "test-c-model", nullptr, &model); + auto status = mta_load_model("test-c-model", "{}", "test-c-plugin", &model); REQUIRE(status == MTA_SUCCESS); CHECK(model.data != nullptr); diff --git a/metatomic-core/tests/misc.cpp b/metatomic-core/tests/misc.cpp index 41bbe5b23..db9f21d40 100644 --- a/metatomic-core/tests/misc.cpp +++ b/metatomic-core/tests/misc.cpp @@ -83,13 +83,11 @@ TEST_CASE("unit conversion factor") { factor = metatomic::unit_conversion_factor("kJ/mol", "eV"); CHECK(factor == Approx(0.010364269656262174).epsilon(1e-15)); - // dimension mismatch -> error - try{ - factor = metatomic::unit_conversion_factor("m", "kg"); - } - catch(metatomic::Error& e){ - CHECK(std::string(e.what()) == "invalid parameter: dimension mismatch in unit conversion: 'm' has dimension [L] but 'kg' has dimension [M]"); - } + REQUIRE_THROWS_WITH( + metatomic::unit_conversion_factor("m", "kg"), + "invalid parameter: dimension mismatch in unit conversion: " + "'m' has dimension [L] but 'kg' has dimension [M]" + ); } } diff --git a/metatomic-core/tests/plugins.cpp b/metatomic-core/tests/plugins.cpp index 79b53947e..2652d534c 100644 --- a/metatomic-core/tests/plugins.cpp +++ b/metatomic-core/tests/plugins.cpp @@ -1,45 +1,71 @@ #include #include "metatomic.h" +#include "metatomic.hpp" TEST_CASE("Load plugins") { - auto status = mta_load_plugin(PLUGIN_DIR "/test-c-plugin.so"); - CHECK(status == MTA_SUCCESS); + SECTION("C API") { + auto status = mta_load_plugin(PLUGIN_DIR "/test-c-plugin.so"); + CHECK(status == MTA_SUCCESS); - // try to load the model with an explicit plugin name - struct mta_model_t model; - status = mta_load_model("test-c-plugin", "some_model", "{}", &model); - CHECK(status == MTA_MODEL_NOT_SUPPORTED_ERROR); + const char* error_message; + const char* error_origin; - // load the plugin without specifying the plugin name - status = mta_load_model(nullptr, "some_model", "{}", &model); - CHECK(status == MTA_INVALID_PARAMETER_ERROR); + struct mta_model_t model; + status = mta_load_model("some_model", "{}", "test-c-plugin", &model); + CHECK(status == MTA_INVALID_PARAMETER_ERROR); - const char* error_message; - const char* error_origin; + status = mta_last_error(&error_message, &error_origin, nullptr); + REQUIRE(status == MTA_SUCCESS); - status = mta_last_error(&error_message, &error_origin, nullptr); - REQUIRE(status == MTA_SUCCESS); + CHECK(std::string(error_origin) == "metatomic-core"); + CHECK(std::string(error_message) == ( + "invalid parameter: failed to load model from 'some_model': plugin 'test-c-plugin' could not load the model" + )); - CHECK(std::string(error_origin) == "metatomic-core"); - const char* expected_message = ( - "invalid parameter: failed to load model from 'some_model': tried the " - "following plugins, but none could load the model: test-c-plugin" - ); - CHECK(std::string(error_message) == expected_message); + status = mta_load_model("some_model", "{}", nullptr, &model); + CHECK(status == MTA_INVALID_PARAMETER_ERROR); + status = mta_last_error(&error_message, &error_origin, nullptr); + REQUIRE(status == MTA_SUCCESS); - status = mta_load_plugin(PLUGIN_DIR "/bad-abi-plugin.so"); - CHECK(status == MTA_INVALID_PARAMETER_ERROR); + CHECK(std::string(error_origin) == "metatomic-core"); + CHECK(std::string(error_message) == ( + "invalid parameter: failed to load model from 'some_model': tried the " + "following plugins, but none could load the model: test-c-plugin" + )); - status = mta_last_error(&error_message, &error_origin, nullptr); - REQUIRE(status == MTA_SUCCESS); - CHECK(std::string(error_origin) == "metatomic-core"); - expected_message = ( - "invalid parameter: can not register plugin 'bad-abi-plugin': " - "plugin ABI version is 2, but metatomic expects 1" - ); - CHECK(std::string(error_message) == expected_message); + status = mta_load_plugin(PLUGIN_DIR "/bad-abi-plugin.so"); + CHECK(status == MTA_INVALID_PARAMETER_ERROR); + + status = mta_last_error(&error_message, &error_origin, nullptr); + REQUIRE(status == MTA_SUCCESS); + + CHECK(std::string(error_origin) == "metatomic-core"); + CHECK(std::string(error_message) == ( + "invalid parameter: can not register plugin 'bad-abi-plugin': " + "plugin ABI version is 2, but metatomic expects 1" + )); + } + + SECTION("C++ API") { + REQUIRE_THROWS_WITH( + metatomic::load_model("some_model", "{}", "test-c-plugin"), + "invalid parameter: failed to load model from 'some_model': plugin 'test-c-plugin' could not load the model" + ); + + REQUIRE_THROWS_WITH( + metatomic::load_model("some_model"), + "invalid parameter: failed to load model from 'some_model': tried the " + "following plugins, but none could load the model: test-c-plugin" + ); + + REQUIRE_THROWS_WITH( + metatomic::load_plugin(PLUGIN_DIR "/bad-abi-plugin.so"), + "invalid parameter: can not register plugin 'bad-abi-plugin': " + "plugin ABI version is 2, but metatomic expects 1" + ); + } } From 62e2476bcaef552a0544433ef89c81ef9e59a7bf Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Thu, 28 May 2026 13:22:48 +0200 Subject: [PATCH 25/43] Re-enable Windows Rust tests --- .github/workflows/rust-tests.yml | 30 ++++++++++++++-------------- metatomic-core/Cargo.toml | 2 +- metatomic-core/src/system.rs | 12 +++++------ metatomic-core/tests/CMakeLists.txt | 18 +++++++++++++---- metatomic-core/tests/utils/mod.rs | 4 ++-- python/metatomic_core/pyproject.toml | 2 +- python/metatomic_core/setup.py | 2 +- 7 files changed, 40 insertions(+), 30 deletions(-) diff --git a/.github/workflows/rust-tests.yml b/.github/workflows/rust-tests.yml index 77ac78387..9a29bf23b 100644 --- a/.github/workflows/rust-tests.yml +++ b/.github/workflows/rust-tests.yml @@ -51,21 +51,21 @@ jobs: cc: clang cmake-generator: Unix Makefiles - # - os: windows-2022 - # rust-version: stable - # rust-target: x86_64-pc-windows-msvc - # extra-name: " / MSVC" - # cxx: cl.exe - # cc: cl.exe - # cmake-generator: Visual Studio 17 2022 - - # - os: windows-2022 - # rust-version: stable - # rust-target: x86_64-pc-windows-gnu - # extra-name: " / MinGW" - # cxx: g++.exe - # cc: gcc.exe - # cmake-generator: MinGW Makefiles + - os: windows-2022 + rust-version: stable + rust-target: x86_64-pc-windows-msvc + extra-name: " / MSVC" + cxx: cl.exe + cc: cl.exe + cmake-generator: Visual Studio 17 2022 + + - os: windows-2022 + rust-version: stable + rust-target: x86_64-pc-windows-gnu + extra-name: " / MinGW" + cxx: g++.exe + cc: gcc.exe + cmake-generator: MinGW Makefiles steps: - name: install dependencies in container if: matrix.container == 'ubuntu:22.04' diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index b88a91cda..4693e30dd 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -14,7 +14,7 @@ name = "metatomic" bench = false [dependencies] -metatensor = { version = "0.3.0" } +metatensor = { version = "0.4.1" } once_cell = "1" dlpk = { version = "0.3", features = ["ndarray"]} json = "0.12" diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs index d30dccce4..d014fcc80 100644 --- a/metatomic-core/src/system.rs +++ b/metatomic-core/src/system.rs @@ -440,10 +440,10 @@ use ndarray::{Array1, Array2}; fn valid_pair_block(dtype: &str) -> TensorBlock { let samples = Labels::new( ["first_atom", "second_atom", "cell_shift_a", "cell_shift_b", "cell_shift_c"], - &[[0i32, 1, 0, 0, 0]], + [[0i32, 1, 0, 0, 0]], ); - let components = vec![Labels::new(["xyz"], &[[0i32], [1], [2]])]; - let properties = Labels::new(["distance"], &[[0i32]]); + let components = vec![Labels::new(["xyz"], [[0i32], [1], [2]])]; + let properties = Labels::new(["distance"], [[0i32]]); match dtype { "f32" => { @@ -459,9 +459,9 @@ use ndarray::{Array1, Array2}; } fn valid_custom_data(dtype: &str) -> TensorMap { - let keys = Labels::new(["key"], &[[0i32]]); - let samples = Labels::new(["sample"], &[[0i32]]); - let properties = Labels::new(["property"], &[[0i32]]); + let keys = Labels::new(["key"], [[0i32]]); + let samples = Labels::new(["sample"], [[0i32]]); + let properties = Labels::new(["property"], [[0i32]]); let block = match dtype { "f32" => { diff --git a/metatomic-core/tests/CMakeLists.txt b/metatomic-core/tests/CMakeLists.txt index 77e251ddc..29d432132 100644 --- a/metatomic-core/tests/CMakeLists.txt +++ b/metatomic-core/tests/CMakeLists.txt @@ -55,6 +55,16 @@ endif() enable_testing() add_subdirectory(test-plugins) +if (TARGET metatensor::shared) + get_target_property(METATENSOR_LOCATION metatensor::shared IMPORTED_LOCATION) + get_filename_component(METATENSOR_DIR ${METATENSOR_LOCATION} DIRECTORY) +elseif (TARGET metatensor) + get_target_property(METATENSOR_LOCATION metatensor LOCATION) + get_filename_component(METATENSOR_DIR ${METATENSOR_LOCATION} DIRECTORY) +else() + set(METATENSOR_DIR "") +endif() + file(GLOB ALL_TESTS *.cpp) foreach(_file_ ${ALL_TESTS}) get_filename_component(_name_ ${_file_} NAME_WE) @@ -71,7 +81,7 @@ foreach(_file_ ${ALL_TESTS}) NO_SYSTEM_FROM_IMPORTED ON ) - target_compile_definitions(${_name_} PRIVATE PLUGIN_DIR="${CMAKE_CURRENT_BINARY_DIR}/test-plugins") + target_compile_definitions(${_name_} PRIVATE PLUGIN_DIR="$") add_test( NAME ${_name_} @@ -79,11 +89,11 @@ foreach(_file_ ${ALL_TESTS}) ) if(WIN32) - # We need to set the path to allow access to metatomic.dll - # this does a similar job to the BUILD_RPATH above + # We need to set the path to allow access to metatomic.dll and + # metatensor.dll. This does a similar job to the BUILD_RPATH above. STRING(REPLACE ";" "\\;" PATH_STRING "$ENV{PATH}") set_tests_properties(${_name_} PROPERTIES - ENVIRONMENT "PATH=${PATH_STRING}\;$" + ENVIRONMENT "PATH=${PATH_STRING}\;$\;${METATENSOR_DIR}" ) endif() endforeach() diff --git a/metatomic-core/tests/utils/mod.rs b/metatomic-core/tests/utils/mod.rs index a04e8a194..3b23bca74 100644 --- a/metatomic-core/tests/utils/mod.rs +++ b/metatomic-core/tests/utils/mod.rs @@ -255,7 +255,7 @@ pub fn setup_torch_pip(python: &Path) -> PathBuf { /// Install metatensor in a Python virtualenv with pip, and return the /// CMAKE_PREFIX_PATH for the installed libmetatensor. pub fn setup_metatensor_pip(python: &Path) -> PathBuf { - pip_install(python, &["metatensor-core >=0.2.0,<0.3"], PipInstallOptions::default()); + pip_install(python, &["metatensor-core >=0.2.2,<0.3"], PipInstallOptions::default()); let mut cmd = Command::new(python); cmd.arg("-c"); @@ -275,7 +275,7 @@ pub fn setup_metatensor_pip(python: &Path) -> PathBuf { /// Install metatensor-torch in a Python virtualenv with pip, and return the /// CMAKE_PREFIX_PATH for the installed libmetatensor_torch. pub fn setup_metatensor_torch_pip(python: &Path) -> PathBuf { - pip_install(python, &["metatensor-torch >=0.9.0,<0.10"], PipInstallOptions::default()); + pip_install(python, &["metatensor-torch >=0.10.0,<0.11"], PipInstallOptions::default()); let mut cmd = Command::new(python); cmd.arg("-c"); diff --git a/python/metatomic_core/pyproject.toml b/python/metatomic_core/pyproject.toml index 9107f6805..62e501ba9 100644 --- a/python/metatomic_core/pyproject.toml +++ b/python/metatomic_core/pyproject.toml @@ -36,7 +36,7 @@ requires = [ "setuptools >=77", "packaging >=26", "cmake", - "metatensor-core >=0.2.0,<0.3", + "metatensor-core >=0.2.2,<0.3", ] build-backend = "setuptools.build_meta" diff --git a/python/metatomic_core/setup.py b/python/metatomic_core/setup.py index 35a2ef16c..905fb5c23 100644 --- a/python/metatomic_core/setup.py +++ b/python/metatomic_core/setup.py @@ -133,7 +133,7 @@ def create_version_number(version): authors = fd.read().splitlines() install_requires = [ - "metatensor-core >=0.2.0,<0.3", + "metatensor-core >=0.2.2,<0.3", ] setup( From 14280cbf55052d055b1457509a4d8c052f188dea Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Mon, 29 Jun 2026 15:51:06 +0200 Subject: [PATCH 26/43] Add a C API for systems --- metatomic-core/include/metatomic.h | 149 +++++++- metatomic-core/src/c_api/system.rs | 380 ++++++++++++++++++-- metatomic-core/src/system.rs | 3 + metatomic-core/tests/system.cpp | 558 +++++++++++++++++++++++++++++ 4 files changed, 1047 insertions(+), 43 deletions(-) create mode 100644 metatomic-core/tests/system.cpp diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index b1899af0d..397886926 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -110,7 +110,10 @@ typedef enum mta_status_t { } mta_status_t; /** - * TODO + * Kind of data always stored in a system. + * + * Other kinds of data can be stored with `mta_system_add_custom_data` and + * retrieved with `mta_system_get_custom_data`. */ typedef enum mta_system_data_kind { MTA_SYSTEM_DATA_TYPES = 0, @@ -120,7 +123,10 @@ typedef enum mta_system_data_kind { } mta_system_data_kind; /** - * TODO + * Opaque handle to an atomistic system. + * + * The system owns DLPack tensors for types, positions, cell, and PBC, as well + * as metatensor blocks for pair lists and tensor maps for custom data. */ typedef struct mta_system_t mta_system_t; @@ -404,7 +410,25 @@ enum mta_status_t mta_unit_conversion_factor(const char *from_unit, double *conversion); /** - * TODO + * Create a new system from raw DLPack tensors. + * + * This function **takes ownership** of `types`, `positions`, `cell`, and + * `pbc`. The caller must not use these tensors after calling this function. + * + * @param length_unit A null-terminated C string containing the length unit + * (e.g. "Angstrom", "nanometer"). Must not be null. + * @param types A DLPack managed tensor with shape `(n_atoms,)` and dtype + * `int32`. Ownership is transferred. + * @param positions A DLPack managed tensor with shape `(n_atoms, 3)` and + * dtype `float32` or `float64`. Ownership is transferred. + * @param cell A DLPack managed tensor with shape `(3, 3)` and the same dtype + * as `positions`. Ownership is transferred. + * @param pbc A DLPack managed tensor with shape `(3,)` and dtype `bool`. + * Ownership is transferred. + * @param system Output parameter, set to the newly created system handle. + * The caller takes ownership and must free it with `mta_system_free`. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_create(const char *length_unit, DLManagedTensorVersioned *types, @@ -414,64 +438,163 @@ enum mta_status_t mta_system_create(const char *length_unit, struct mta_system_t **system); /** - * TODO + * Free a system previously created by `mta_system_create`. + * + * If there are outstanding borrowed views (from `mta_system_get_data`), the + * system's data will remain alive until all views are released. + * + * @param system The system handle to free. Can be null, in which case this + * function is a no-op. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_free(struct mta_system_t *system); /** - * TODO + * Get the number of atoms in a system. + * + * @param system The system handle. Must not be null. + * @param size Output parameter, set to the number of atoms. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_size(const struct mta_system_t *system, uintptr_t *size); /** - * TODO + * Get a DLPack tensor from a system for the requested data. + * + * This function **returns a borrowed view** of the system's internal data. + * The returned `DLManagedTensorVersioned` has a custom deleter that decrements + * the system's reference count, keeping the system alive as long as the + * borrowed view exists. + * + * The caller is responsible for calling the deleter on the returned tensor + * when it is no longer needed. The tensor shares the data pointer with the + * system; do **not** modify it. + * + * @param system The system handle. Must not be null. + * @param request Which data to retrieve (types, positions, cell, or PBC). + * @param data Output parameter, set to a pointer to a newly allocated + * `DLManagedTensorVersioned` containing the requested data. The caller + * takes ownership and must call the deleter when done. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_get_data(const struct mta_system_t *system, enum mta_system_data_kind request, DLManagedTensorVersioned **data); /** - * TODO + * Get the length unit of a system. + * + * This function returns a new `mta_string_t` that the caller must free with + * `mta_string_free`. + * + * @param system The system handle. Must not be null. + * @param length_unit Output parameter, set to the length unit string. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_get_length_unit(const struct mta_system_t *system, mta_string_t *length_unit); /** - * TODO + * Add a pair list (neighbor list) to a system. + * + * This function **takes ownership** of `pairs`. The caller must not use the + * block after calling this function. + * + * @param system The system handle. Must not be null. + * @param options A JSON-serialized `PairListOptions` object. Must not be null. + * @param pairs A `mts_block_t` containing the pair data. Ownership is + * transferred. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_add_pairs(struct mta_system_t *system, const char *options, mts_block_t *pairs); /** - * TODO + * Get a pair list from a system. + * + * **Returns a borrowed view** of the pair list. The system must outlive the + * returned pointer. Do **not** free the returned block. + * + * @param system The system handle. Must not be null. + * @param options A JSON-serialized `PairListOptions` object identifying which + * pair list to retrieve. Must not be null. + * @param pairs Output parameter, set to a pointer to the pair list block, or + * NULL if no pair list matches the options. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_get_pairs(const struct mta_system_t *system, const char *options, const mts_block_t **pairs); /** - * TODO + * Get all pair list options known by a system. + * + * This function returns a new `mta_string_t` containing a JSON array of + * `PairListOptions` objects. The caller must free it with `mta_string_free`. + * + * @param system The system handle. Must not be null. + * @param pairs_options Output parameter, set to a JSON string containing an + * array of `PairListOptions` objects. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_known_pairs(const struct mta_system_t *system, mta_string_t *pairs_options); /** - * TODO + * Add custom data to a system. + * + * This function **takes ownership** of `data`. The caller must not use the + * tensor map after calling this function. + * + * @param system The system handle. Must not be null. + * @param name A null-terminated C string containing the name of the custom + * data. Must not be null. + * @param data A `mts_tensormap_t` containing the custom data. Ownership is + * transferred. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_add_custom_data(struct mta_system_t *system, const char *name, mts_tensormap_t *data); /** - * TODO + * Get custom data from a system by name. + * + * **Returns a borrowed view** of the custom data. The system must outlive the + * returned pointer. Do **not** free the returned tensor map. + * + * @param system The system handle. Must not be null. + * @param name A null-terminated C string containing the name of the custom + * data. Must not be null. + * @param data Output parameter, set to a pointer to the custom data tensor + * map, or an error if no data with the given name exists. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_get_custom_data(const struct mta_system_t *system, const char *name, const mts_tensormap_t **data); /** - * TODO + * Get all custom data names known by a system. + * + * **Returns a new** `mta_string_t` containing a JSON array of strings. The + * caller must free it with `mta_string_free`. + * + * @param system The system handle. Must not be null. + * @param names Output parameter, set to a JSON string containing an array of + * custom data names. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ enum mta_status_t mta_system_known_custom_data(const struct mta_system_t *system, mta_string_t *names); diff --git a/metatomic-core/src/c_api/system.rs b/metatomic-core/src/c_api/system.rs index 1697e5155..4f6ff1235 100644 --- a/metatomic-core/src/c_api/system.rs +++ b/metatomic-core/src/c_api/system.rs @@ -1,17 +1,40 @@ -use std::ffi::c_char; +use std::ffi::{c_char, CStr}; +use std::sync::Arc; use dlpk::sys::DLManagedTensorVersioned; +use dlpk::{DLPackTensor, DLPackVersion}; use metatensor::c_api::{mts_block_t, mts_tensormap_t}; +use metatensor::{TensorBlock, TensorMap}; -use crate::System; -use super::{mta_status_t, mta_string_t}; +use crate::{Error, PairListOptions, System}; +use super::{catch_unwind, mta_status_t, mta_string_t}; -/// TODO +/// Opaque handle to an atomistic system. +/// +/// The system owns DLPack tensors for types, positions, cell, and PBC, as well +/// as metatensor blocks for pair lists and tensor maps for custom data. #[allow(non_camel_case_types)] -pub struct mta_system_t(pub(crate) System); +pub struct mta_system_t(pub(crate) Arc); - -/// TODO +/// Create a new system from raw DLPack tensors. +/// +/// This function **takes ownership** of `types`, `positions`, `cell`, and +/// `pbc`. The caller must not use these tensors after calling this function. +/// +/// @param length_unit A null-terminated C string containing the length unit +/// (e.g. "Angstrom", "nanometer"). Must not be null. +/// @param types A DLPack managed tensor with shape `(n_atoms,)` and dtype +/// `int32`. Ownership is transferred. +/// @param positions A DLPack managed tensor with shape `(n_atoms, 3)` and +/// dtype `float32` or `float64`. Ownership is transferred. +/// @param cell A DLPack managed tensor with shape `(3, 3)` and the same dtype +/// as `positions`. Ownership is transferred. +/// @param pbc A DLPack managed tensor with shape `(3,)` and dtype `bool`. +/// Ownership is transferred. +/// @param system Output parameter, set to the newly created system handle. +/// The caller takes ownership and must free it with `mta_system_free`. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_create( length_unit: *const c_char, @@ -21,25 +44,73 @@ pub unsafe extern "C" fn mta_system_create( pbc: *mut DLManagedTensorVersioned, system: *mut *mut mta_system_t, ) -> mta_status_t { - todo!() + let unwind_wrapper = std::panic::AssertUnwindSafe(system); + catch_unwind(move || { + check_pointers_non_null!(length_unit, types, positions, cell, pbc, system); + + let length_unit = CStr::from_ptr(length_unit) + .to_str() + .map_err(|_| Error::InvalidParameter("length_unit is not valid UTF-8".into()))? + .to_string(); + + let types = DLPackTensor::from_ptr(types); + let positions = DLPackTensor::from_ptr(positions); + let cell = DLPackTensor::from_ptr(cell); + let pbc = DLPackTensor::from_ptr(pbc); + + let system_inner = System::new(length_unit, types, positions, cell, pbc)?; + + let _ = &unwind_wrapper; + *unwind_wrapper.0 = Box::into_raw(Box::new(mta_system_t(Arc::new(system_inner)))); + Ok(()) + }) } -/// TODO +/// Free a system previously created by `mta_system_create`. +/// +/// If there are outstanding borrowed views (from `mta_system_get_data`), the +/// system's data will remain alive until all views are released. +/// +/// @param system The system handle to free. Can be null, in which case this +/// function is a no-op. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_free(system: *mut mta_system_t) -> mta_status_t { - todo!() + catch_unwind(|| { + if system.is_null() { + return Ok(()); + } + + let _ = Box::from_raw(system); + Ok(()) + }) } -/// TODO +/// Get the number of atoms in a system. +/// +/// @param system The system handle. Must not be null. +/// @param size Output parameter, set to the number of atoms. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_size( system: *const mta_system_t, size: *mut usize, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(system, size); + + let system = &*system; + *size = system.0.size(); + Ok(()) + }) } -/// TODO +/// Kind of data always stored in a system. +/// +/// Other kinds of data can be stored with `mta_system_add_custom_data` and +/// retrieved with `mta_system_get_custom_data`. #[allow(non_camel_case_types)] #[repr(C)] #[non_exhaustive] @@ -50,82 +121,331 @@ pub enum mta_system_data_kind { MTA_SYSTEM_DATA_PBC = 3, } -/// TODO +/// Custom deleter for borrowed DLPack tensors returned by `mta_system_get_data`. +/// +/// Releases the `Arc` reference stored in `manager_ctx` and frees the +/// heap-allocated `DLManagedTensorVersioned`. +unsafe extern "C" fn borrowed_tensor_deleter( + tensor: *mut DLManagedTensorVersioned, +) { + let system: Arc = { + let ptr = (*tensor).manager_ctx as *const System; + Arc::from_raw(ptr) + }; + drop(system); + let _ = Box::from_raw(tensor); +} + +/// Get a DLPack tensor from a system for the requested data. +/// +/// This function **returns a borrowed view** of the system's internal data. +/// The returned `DLManagedTensorVersioned` has a custom deleter that decrements +/// the system's reference count, keeping the system alive as long as the +/// borrowed view exists. +/// +/// The caller is responsible for calling the deleter on the returned tensor +/// when it is no longer needed. The tensor shares the data pointer with the +/// system; do **not** modify it. +/// +/// @param system The system handle. Must not be null. +/// @param request Which data to retrieve (types, positions, cell, or PBC). +/// @param data Output parameter, set to a pointer to a newly allocated +/// `DLManagedTensorVersioned` containing the requested data. The caller +/// takes ownership and must call the deleter when done. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_get_data( system: *const mta_system_t, request: mta_system_data_kind, data: *mut *mut DLManagedTensorVersioned, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(system, data); + *data = std::ptr::null_mut(); + + let system = &*system; + let tensor_ref = match request { + mta_system_data_kind::MTA_SYSTEM_DATA_TYPES => system.0.types(), + mta_system_data_kind::MTA_SYSTEM_DATA_POSITIONS => system.0.positions(), + mta_system_data_kind::MTA_SYSTEM_DATA_CELL => system.0.cell(), + mta_system_data_kind::MTA_SYSTEM_DATA_PBC => system.0.pbc(), + }; + + // Clone the Arc to keep the system alive while the borrowed view exists + let system_arc = system.0.clone(); + let system_ptr = Arc::into_raw(system_arc); + + let packed = Box::new(DLManagedTensorVersioned { + version: DLPackVersion::current(), + manager_ctx: system_ptr as *mut std::ffi::c_void, + deleter: Some( + borrowed_tensor_deleter + as unsafe extern "C" fn(*mut DLManagedTensorVersioned), + ), + flags: dlpk::sys::DLPACK_FLAG_BITMASK_READ_ONLY, + dl_tensor: tensor_ref.raw.clone(), + }); + + *data = Box::into_raw(packed); + Ok(()) + }) } -/// TODO +/// Get the length unit of a system. +/// +/// This function returns a new `mta_string_t` that the caller must free with +/// `mta_string_free`. +/// +/// @param system The system handle. Must not be null. +/// @param length_unit Output parameter, set to the length unit string. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_get_length_unit( system: *const mta_system_t, length_unit: *mut mta_string_t, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(system, length_unit); + *length_unit = mta_string_t::null(); + + let system = &*system; + *length_unit = mta_string_t::new(system.0.length_unit()); + Ok(()) + }) } -/// TODO +/// Add a pair list (neighbor list) to a system. +/// +/// This function **takes ownership** of `pairs`. The caller must not use the +/// block after calling this function. +/// +/// @param system The system handle. Must not be null. +/// @param options A JSON-serialized `PairListOptions` object. Must not be null. +/// @param pairs A `mts_block_t` containing the pair data. Ownership is +/// transferred. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_add_pairs( system: *mut mta_system_t, options: *const c_char, pairs: *mut mts_block_t, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(system, options, pairs); + + let options_str = CStr::from_ptr(options) + .to_str() + .map_err(|_| Error::InvalidParameter("options is not valid UTF-8".into()))?; + + let options_json = json::parse(options_str) + .map_err(|e| Error::Serialization(format!("invalid JSON for PairListOptions: {e}")))?; + + let options = PairListOptions::try_from(&options_json)?; + + let pairs = TensorBlock::from_raw(pairs); + + let system = &mut *system; + let system = Arc::get_mut(&mut system.0).ok_or_else(|| { + Error::InvalidParameter( + "cannot modify system while there are outstanding borrowed views".into(), + ) + })?; + + system.add_pairs(options, pairs)?; + + Ok(()) + }) } -/// TODO +/// Get a pair list from a system. +/// +/// **Returns a borrowed view** of the pair list. The system must outlive the +/// returned pointer. Do **not** free the returned block. +/// +/// @param system The system handle. Must not be null. +/// @param options A JSON-serialized `PairListOptions` object identifying which +/// pair list to retrieve. Must not be null. +/// @param pairs Output parameter, set to a pointer to the pair list block, or +/// NULL if no pair list matches the options. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_get_pairs( system: *const mta_system_t, options: *const c_char, pairs: *mut *const mts_block_t, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(system, options, pairs); + *pairs = std::ptr::null(); + + let options_str = CStr::from_ptr(options) + .to_str() + .map_err(|_| Error::InvalidParameter("options is not valid UTF-8".into()))?; + + let options_json = json::parse(options_str) + .map_err(|e| Error::Serialization(format!("invalid JSON for PairListOptions: {e}")))?; + + let options = PairListOptions::try_from(&options_json)?; + + let system = &*system; + match system.0.get_pairs(&options) { + Some(block) => { + *pairs = block.as_ptr(); + } + None => { + return Err(Error::InvalidParameter( + "no pair list found for the given options".into(), + )); + } + } + + Ok(()) + }) } -/// TODO +/// Get all pair list options known by a system. +/// +/// This function returns a new `mta_string_t` containing a JSON array of +/// `PairListOptions` objects. The caller must free it with `mta_string_free`. +/// +/// @param system The system handle. Must not be null. +/// @param pairs_options Output parameter, set to a JSON string containing an +/// array of `PairListOptions` objects. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_known_pairs( system: *const mta_system_t, pairs_options: *mut mta_string_t, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(system, pairs_options); + *pairs_options = mta_string_t::null(); + + let system = &*system; + let known = system.0.known_pairs(); + let mut json_array = json::JsonValue::new_array(); + for options in known { + json_array.push(json::JsonValue::from(options.clone())).map_err(|_| { + Error::Internal("failed to build JSON array".into()) + })?; + } + + *pairs_options = mta_string_t::new(json::stringify(json_array)); + Ok(()) + }) } -/// TODO +/// Add custom data to a system. +/// +/// This function **takes ownership** of `data`. The caller must not use the +/// tensor map after calling this function. +/// +/// @param system The system handle. Must not be null. +/// @param name A null-terminated C string containing the name of the custom +/// data. Must not be null. +/// @param data A `mts_tensormap_t` containing the custom data. Ownership is +/// transferred. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_add_custom_data( system: *mut mta_system_t, name: *const c_char, data: *mut mts_tensormap_t, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(system, name, data); + + let name = CStr::from_ptr(name) + .to_str() + .map_err(|_| Error::InvalidParameter("name is not valid UTF-8".into()))? + .to_string(); + + let data = TensorMap::from_raw(data); + + let system = &mut *system; + let system = Arc::get_mut(&mut system.0).ok_or_else(|| { + Error::InvalidParameter( + "cannot modify system while there are outstanding borrowed views".into(), + ) + })?; + + system.add_custom_data(name, data, false)?; + + Ok(()) + }) } -/// TODO +/// Get custom data from a system by name. +/// +/// **Returns a borrowed view** of the custom data. The system must outlive the +/// returned pointer. Do **not** free the returned tensor map. +/// +/// @param system The system handle. Must not be null. +/// @param name A null-terminated C string containing the name of the custom +/// data. Must not be null. +/// @param data Output parameter, set to a pointer to the custom data tensor +/// map, or an error if no data with the given name exists. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_get_custom_data( system: *const mta_system_t, name: *const c_char, data: *mut *const mts_tensormap_t, ) -> mta_status_t { - todo!() + catch_unwind(|| { + check_pointers_non_null!(system, name, data); + *data = std::ptr::null_mut(); + + let name = CStr::from_ptr(name) + .to_str() + .map_err(|_| Error::InvalidParameter("name is not valid UTF-8".into()))?; + + let system = &*system; + let result = system.0.get_custom_data(name)?; + *data = result.as_ptr(); + + Ok(()) + }) } -/// TODO +/// Get all custom data names known by a system. +/// +/// **Returns a new** `mta_string_t` containing a JSON array of strings. The +/// caller must free it with `mta_string_free`. +/// +/// @param system The system handle. Must not be null. +/// @param names Output parameter, set to a JSON string containing an array of +/// custom data names. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[no_mangle] pub unsafe extern "C" fn mta_system_known_custom_data( system: *const mta_system_t, names: *mut mta_string_t, ) -> mta_status_t { - todo!() -} + catch_unwind(|| { + check_pointers_non_null!(system, names); + *names = mta_string_t::null(); + let system = &*system; + let known = system.0.known_custom_data(); + let mut json_array = json::JsonValue::new_array(); + for name in known { + json_array.push(name).map_err(|_| { + Error::Internal("failed to build JSON array".into()) + })?; + } + + *names = mta_string_t::new(json::stringify(json_array)); + Ok(()) + }) +} // TODO: mta_system_to(device, dtype) diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs index d014fcc80..578333d52 100644 --- a/metatomic-core/src/system.rs +++ b/metatomic-core/src/system.rs @@ -27,6 +27,9 @@ pub struct System { custom_data: HashMap, } +unsafe impl Send for System {} +unsafe impl Sync for System {} + impl System { /// Create a `System` from raw DLPack tensors pub fn new( diff --git a/metatomic-core/tests/system.cpp b/metatomic-core/tests/system.cpp new file mode 100644 index 000000000..fb67a2390 --- /dev/null +++ b/metatomic-core/tests/system.cpp @@ -0,0 +1,558 @@ +#include +#include + +#include + +#include +#include "metatomic.h" + + +template static DLManagedTensorVersioned* types_tensor(size_t n_atoms) { + std::vector type_data; + type_data.reserve(n_atoms); + for (size_t i = 0; i < n_atoms; i++) { + type_data.push_back(static_cast(i * 3 + 1)); + } + auto array = std::make_unique>( + std::vector{n_atoms}, + std::move(type_data) + ); + auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); + DLDevice cpu = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + return mts.as_dlpack(cpu, nullptr, version); +} + +template static DLManagedTensorVersioned* cell_tensor() { + auto array = std::make_unique>( + std::vector{3, 3}, + std::vector{ + T(10.0), T(0.0), T(0.0), + T(0.0), T(0.0), T(0.0), + T(0.0), T(0.0), T(10.0), + } + ); + auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); + DLDevice cpu = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + return mts.as_dlpack(cpu, nullptr, version); +} + +template static DLManagedTensorVersioned* positions_tensor(size_t n_atoms) { + std::vector position_data; + position_data.reserve(n_atoms * 3); + for (size_t i = 0; i < n_atoms; i++) { + position_data.push_back(static_cast(i * 3 + 1)); + position_data.push_back(static_cast(i * 3 + 2)); + position_data.push_back(static_cast(i * 3 + 3)); + } + auto array = std::make_unique>( + std::vector{n_atoms, 3}, + std::move(position_data) + ); + auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); + DLDevice cpu = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + return mts.as_dlpack(cpu, nullptr, version); +} + +template static DLManagedTensorVersioned* pbc_tensor() { + std::vector pbc_data = {1, 0, 1}; + auto array = std::make_unique>( + std::vector{3}, + std::move(pbc_data) + ); + auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); + DLDevice cpu = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + return mts.as_dlpack(cpu, nullptr, version); +} + + +/// SimpleDataArray doesn't compile (std::vector has no data() +/// method). We use SimpleDataArray and patch the dtype code +/// from kDLUInt to kDLBool. +template <> DLManagedTensorVersioned* pbc_tensor() { + std::vector pbc_data = {1, 0, 1}; + auto array = std::make_unique>( + std::vector{3}, + std::move(pbc_data) + ); + auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); + DLDevice cpu = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + auto* tensor = mts.as_dlpack(cpu, nullptr, version); + + tensor->dl_tensor.dtype.code = DLDataTypeCode::kDLBool; + + return tensor; +} + +static mts_block_t* pair_block() { + auto samples = metatensor::Labels( + {"first_atom", "second_atom", "cell_shift_a", "cell_shift_b", "cell_shift_c"}, + {{0, 1, 0, 0, 0}} + ); + + auto components = metatensor::Labels({"xyz"}, {{0}, {1}, {2}}); + std::vector components_list = { + components.as_mts_labels_t() + }; + + auto properties = metatensor::Labels({"distance"}, {{0}}); + + auto values = std::make_unique>( + std::vector{1, 3, 1}, + std::vector{1.5F, 2.5F, 3.5F} + ); + auto values_mts = metatensor::DataArrayBase::to_mts_array(std::move(values)); + + auto* block = mts_block( + std::move(values_mts).release(), + samples.as_mts_labels_t(), + components_list.data(), + components_list.size(), + properties.as_mts_labels_t() + ); + REQUIRE(block != nullptr); + return block; +} + +static mts_tensormap_t* custom_data() { + auto keys = metatensor::Labels({"key"}, {{0}}); + auto samples = metatensor::Labels({"sample"}, {{0}}); + auto properties = metatensor::Labels({"property"}, {{0}}); + + auto values = std::make_unique>( + std::vector{1, 1}, + std::vector{42.0F} + ); + auto values_mts = metatensor::DataArrayBase::to_mts_array(std::move(values)); + + auto* block = mts_block( + std::move(values_mts).release(), + samples.as_mts_labels_t(), + nullptr, + 0, + properties.as_mts_labels_t() + ); + REQUIRE(block != nullptr); + + std::vector blocks = {block}; + auto* tensormap = mts_tensormap( + keys.as_mts_labels_t(), + blocks.data(), + blocks.size() + ); + REQUIRE(tensormap != nullptr); + + return tensormap; +} + +TEST_CASE("system") { + SECTION("create and free") { + mta_system_t* system_f32 = nullptr; + auto status = mta_system_create( + "nm", + types_tensor(4), + positions_tensor(4), + cell_tensor(), + pbc_tensor(), + &system_f32 + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(system_f32 != nullptr); + + status = mta_system_free(system_f32); + CHECK(status == MTA_SUCCESS); + + mta_system_t* system_f64 = nullptr; + status = mta_system_create( + "nm", + types_tensor(4), + positions_tensor(4), + cell_tensor(), + pbc_tensor(), + &system_f64 + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(system_f64 != nullptr); + + status = mta_system_free(system_f64); + CHECK(status == MTA_SUCCESS); + + // free on null pointer is fine + status = mta_system_free(nullptr); + REQUIRE(status == MTA_SUCCESS); + } + + SECTION("errors") { + mta_system_t* system = nullptr; + + // wrong dtype for types (float instead of int32) + auto status = mta_system_create( + "Angstrom", + types_tensor(3), + positions_tensor(3), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status != MTA_SUCCESS); + CHECK(system == nullptr); + + const char* message = nullptr; + mta_last_error(&message, nullptr, nullptr); + CHECK(std::string(message) == "invalid parameter: `types` must be a tensor of 32-bit integers"); + + // wrong dtype for positions (int32 instead of float) + status = mta_system_create( + "Angstrom", + types_tensor(3), + positions_tensor(3), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status != MTA_SUCCESS); + CHECK(system == nullptr); + + mta_last_error(&message, nullptr, nullptr); + CHECK(std::string(message) == "invalid parameter: `positions` must be a tensor of 32 or 64-bit floating point data"); + + // wrong dtype for cell (int32 instead of float) + status = mta_system_create( + "Angstrom", + types_tensor(3), + positions_tensor(3), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status != MTA_SUCCESS); + CHECK(system == nullptr); + + mta_last_error(&message, nullptr, nullptr); + CHECK(std::string(message) == "invalid parameter: `cell` must have the same dtype as `positions`"); + + // wrong dtype for pbc (float instead of bool) + status = mta_system_create( + "Angstrom", + types_tensor(3), + positions_tensor(3), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status != MTA_SUCCESS); + CHECK(system == nullptr); + + mta_last_error(&message, nullptr, nullptr); + CHECK(std::string(message) == "invalid parameter: `pbc` must be a tensor of booleans"); + + // mismatched positions/type shapes + status = mta_system_create( + "Angstrom", + types_tensor(3), + positions_tensor(5), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status != MTA_SUCCESS); + CHECK(system == nullptr); + + mta_last_error(&message, nullptr, nullptr); + CHECK(std::string(message) == "invalid parameter: `positions` must be a (n_atoms x 3) tensor, got a tensor with shape [5, 3]"); + + + // wrong cell shape + auto* cell = cell_tensor(); + cell->dl_tensor.shape[0] = 9; + cell->dl_tensor.shape[1] = 1; + status = mta_system_create( + "Angstrom", + types_tensor(3), + positions_tensor(3), + cell, + pbc_tensor(), + &system + ); + CHECK(status != MTA_SUCCESS); + CHECK(system == nullptr); + + mta_last_error(&message, nullptr, nullptr); + CHECK(std::string(message) == "invalid parameter: `cell` must be a (3 x 3) tensor, got a tensor with shape [9, 1]"); + + + // wrong pbc shape + auto* pbc = pbc_tensor(); + pbc->dl_tensor.shape[0] = 2; + status = mta_system_create( + "Angstrom", + types_tensor(3), + positions_tensor(3), + cell_tensor(), + pbc, + &system + ); + CHECK(status != MTA_SUCCESS); + CHECK(system == nullptr); + + mta_last_error(&message, nullptr, nullptr); + CHECK(std::string(message) == "invalid parameter: `pbc` must contain 3 entries, got a tensor with shape [2]"); + } + + SECTION("size") { + mta_system_t* system = nullptr; + auto status = mta_system_create( + "nm", + types_tensor(4), + positions_tensor(4), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(system != nullptr); + + uintptr_t size = 0; + status = mta_system_size(system, &size); + CHECK(status == MTA_SUCCESS); + CHECK(size == 4); + + status = mta_system_free(system); + CHECK(status == MTA_SUCCESS); + } + + SECTION("length unit") { + mta_system_t* system = nullptr; + auto status = mta_system_create( + "nm", + types_tensor(4), + positions_tensor(4), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(system != nullptr); + + mta_string_t unit = nullptr; + status = mta_system_get_length_unit(system, &unit); + CHECK(status == MTA_SUCCESS); + CHECK(std::string(mta_string_view(unit)) == "nm"); + mta_string_free(unit); + + status = mta_system_free(system); + CHECK(status == MTA_SUCCESS); + } +} + +TEST_CASE("system data") { + mta_system_t* system = nullptr; + auto status = mta_system_create( + "nm", + types_tensor(4), + positions_tensor(4), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(system != nullptr); + + DLManagedTensorVersioned* data = nullptr; + + SECTION("types") { + status = mta_system_get_data( + system, MTA_SYSTEM_DATA_TYPES, &data + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(data != nullptr); + + CHECK(data->dl_tensor.ndim == 1); + CHECK(data->dl_tensor.shape[0] == 4); + CHECK(data->dl_tensor.dtype.code == kDLInt); + CHECK(data->dl_tensor.dtype.bits == 32); + + auto* types = reinterpret_cast(static_cast(data->dl_tensor.data) + data->dl_tensor.byte_offset); + CHECK(types[0] == 1); + CHECK(types[1] == 4); + CHECK(types[2] == 7); + CHECK(types[3] == 10); + } + + SECTION("positions") { + status = mta_system_get_data( + system, MTA_SYSTEM_DATA_POSITIONS, &data + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(data != nullptr); + + CHECK(data->dl_tensor.ndim == 2); + CHECK(data->dl_tensor.shape[0] == 4); + CHECK(data->dl_tensor.shape[1] == 3); + CHECK(data->dl_tensor.dtype.code == kDLFloat); + CHECK(data->dl_tensor.dtype.bits == 32); + + auto* positions = reinterpret_cast(static_cast(data->dl_tensor.data) + data->dl_tensor.byte_offset); + CHECK(positions[0] == 1.0F); + CHECK(positions[3] == 4.0F); + CHECK(positions[6] == 7.0F); + CHECK(positions[9] == 10.0F); + } + + SECTION("cell") { + status = mta_system_get_data( + system, MTA_SYSTEM_DATA_CELL, &data + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(data != nullptr); + + CHECK(data->dl_tensor.ndim == 2); + CHECK(data->dl_tensor.shape[0] == 3); + CHECK(data->dl_tensor.shape[1] == 3); + CHECK(data->dl_tensor.dtype.code == kDLFloat); + CHECK(data->dl_tensor.dtype.bits == 32); + + auto* cell = reinterpret_cast(static_cast(data->dl_tensor.data) + data->dl_tensor.byte_offset); + CHECK(cell[0] == 10.0F); + CHECK(cell[4] == 0.0F); + CHECK(cell[8] == 10.0F); + } + + SECTION("pbc") { + status = mta_system_get_data( + system, MTA_SYSTEM_DATA_PBC, &data + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(data != nullptr); + + CHECK(data->dl_tensor.ndim == 1); + CHECK(data->dl_tensor.shape[0] == 3); + CHECK(data->dl_tensor.dtype.code == kDLBool); + CHECK(data->dl_tensor.dtype.bits == 8); + + auto* pbc = reinterpret_cast(static_cast(data->dl_tensor.data) + data->dl_tensor.byte_offset); + CHECK(pbc[0] == true); + CHECK(pbc[1] == false); + CHECK(pbc[2] == true); + } + + data->deleter(data); + + status = mta_system_free(system); + CHECK(status == MTA_SUCCESS); +} + + +TEST_CASE("system pairs") { + mta_system_t* system = nullptr; + auto status = mta_system_create( + "nm", + types_tensor(4), + positions_tensor(4), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(system != nullptr); + + const auto* options_json = R"({ + "type": "metatomic_pair_options", + "cutoff": "0x00001000", + "full_list": true, + "strict": false, + "requestors": ["test"] + })"; + + auto* pairs = pair_block(); + status = mta_system_add_pairs(system, options_json, pairs); + CHECK(status == MTA_SUCCESS); + + const mts_block_t* recovered_pairs = nullptr; + status = mta_system_get_pairs(system, options_json, &recovered_pairs); + CHECK(status == MTA_SUCCESS); + // we get the same pointer back + CHECK(static_cast(recovered_pairs) == static_cast(pairs)); + + // Add a second block with different options + const auto* other_json = R"({ + "type": "metatomic_pair_options", + "cutoff": "0x00001000", + "full_list": true, + "strict": true, + "requestors": [] + })"; + + pairs = pair_block(); + status = mta_system_add_pairs(system, other_json, pairs); + CHECK(status == MTA_SUCCESS); + + // Check known pairs contains both + mta_string_t known = nullptr; + status = mta_system_known_pairs(system, &known); + CHECK(status == MTA_SUCCESS); + REQUIRE(known != nullptr); + + auto known_str = std::string(mta_string_view(known)); + mta_string_free(known); + + auto first = known_str.find("metatomic_pair_options"); + CHECK(first != std::string::npos); + known_str = known_str.substr(first + 1); + auto second = known_str.find("metatomic_pair_options"); + CHECK(second != std::string::npos); + + mta_system_free(system); +} + +TEST_CASE("system custom data") { + mta_system_t* system = nullptr; + auto status = mta_system_create( + "Angstrom", + types_tensor(4), + positions_tensor(4), + cell_tensor(), + pbc_tensor(), + &system + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(system != nullptr); + + auto* data = custom_data(); + status = mta_system_add_custom_data(system, "test::my_data", data); + CHECK(status == MTA_SUCCESS); + + const mts_tensormap_t* retrieved = nullptr; + status = mta_system_get_custom_data( + system, "test::my_data", &retrieved + ); + CHECK(status == MTA_SUCCESS); + CHECK(retrieved != nullptr); + CHECK(static_cast(retrieved) == static_cast(data)); + + retrieved = nullptr; + status = mta_system_get_custom_data( + system, "test::no_such_data", &retrieved + ); + CHECK(status != MTA_SUCCESS); + CHECK(retrieved == nullptr); + + data = custom_data(); + status = mta_system_add_custom_data(system, "test::other_data", data); + CHECK(status == MTA_SUCCESS); + + mta_string_t names = nullptr; + status = mta_system_known_custom_data(system, &names); + CHECK(status == MTA_SUCCESS); + CHECK(names != nullptr); + + auto names_str = std::string(mta_string_view(names)); + CHECK(names_str.find("test::my_data") != std::string::npos); + CHECK(names_str.find("test::other_data") != std::string::npos); + + mta_system_free(system); +} From ff38ac9b2f18c32d7ae67a922335969ebe195ce1 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Tue, 30 Jun 2026 10:43:46 +0200 Subject: [PATCH 27/43] Separate C++ and C API tests --- metatomic-core/tests/CMakeLists.txt | 30 ++++++---- metatomic-core/tests/cxx/misc.cpp | 64 +++++++++++++++++++++ metatomic-core/tests/cxx/plugins.cpp | 25 +++++++++ metatomic-core/tests/misc.cpp | 82 ++++++++++----------------- metatomic-core/tests/plugins.cpp | 84 ++++++++++------------------ 5 files changed, 168 insertions(+), 117 deletions(-) create mode 100644 metatomic-core/tests/cxx/misc.cpp create mode 100644 metatomic-core/tests/cxx/plugins.cpp diff --git a/metatomic-core/tests/CMakeLists.txt b/metatomic-core/tests/CMakeLists.txt index 29d432132..731904948 100644 --- a/metatomic-core/tests/CMakeLists.txt +++ b/metatomic-core/tests/CMakeLists.txt @@ -65,13 +65,11 @@ else() set(METATENSOR_DIR "") endif() -file(GLOB ALL_TESTS *.cpp) -foreach(_file_ ${ALL_TESTS}) - get_filename_component(_name_ ${_file_} NAME_WE) - add_executable(${_name_} ${_file_}) - target_link_libraries(${_name_} metatomic catch) +function(metatomic_add_test source target) + add_executable(${target} ${source}) + target_link_libraries(${target} metatomic catch) - set_target_properties(${_name_} PROPERTIES + set_target_properties(${target} PROPERTIES # Ensure that the binaries find the right shared library. # # Without this, when configuring with cmake before the library is built, @@ -81,19 +79,31 @@ foreach(_file_ ${ALL_TESTS}) NO_SYSTEM_FROM_IMPORTED ON ) - target_compile_definitions(${_name_} PRIVATE PLUGIN_DIR="$") + target_compile_definitions(${target} PRIVATE PLUGIN_DIR="$") add_test( - NAME ${_name_} - COMMAND ${TEST_COMMAND} $ + NAME ${target} + COMMAND ${TEST_COMMAND} $ ) if(WIN32) # We need to set the path to allow access to metatomic.dll and # metatensor.dll. This does a similar job to the BUILD_RPATH above. STRING(REPLACE ";" "\\;" PATH_STRING "$ENV{PATH}") - set_tests_properties(${_name_} PROPERTIES + set_tests_properties(${target} PROPERTIES ENVIRONMENT "PATH=${PATH_STRING}\;$\;${METATENSOR_DIR}" ) endif() +endfunction() + +file(GLOB ALL_TESTS *.cpp) +foreach(_file_ ${ALL_TESTS}) + get_filename_component(_name_ ${_file_} NAME_WE) + metatomic_add_test(${_file_} ${_name_}) +endforeach() + +file(GLOB ALL_CPP_TESTS cxx/*.cpp) +foreach(_file_ ${ALL_CPP_TESTS}) + get_filename_component(_name_ ${_file_} NAME_WE) + metatomic_add_test(${_file_} "cxx-${_name_}") endforeach() diff --git a/metatomic-core/tests/cxx/misc.cpp b/metatomic-core/tests/cxx/misc.cpp new file mode 100644 index 000000000..ed08aaf40 --- /dev/null +++ b/metatomic-core/tests/cxx/misc.cpp @@ -0,0 +1,64 @@ +#include + +#include "metatomic.hpp" + + +TEST_CASE("unit conversion factor") { + // same unit -> factor = 1.0 + auto factor = metatomic::unit_conversion_factor("m", "m"); + CHECK(factor == 1.0); + + // kJ/mol -> eV + factor = metatomic::unit_conversion_factor("kJ/mol", "eV"); + CHECK(factor == Approx(0.010364269656262174).epsilon(1e-15)); + + REQUIRE_THROWS_WITH( + metatomic::unit_conversion_factor("m", "kg"), + "invalid parameter: dimension mismatch in unit conversion: " + "'m' has dimension [L] but 'kg' has dimension [M]" + ); +} + + +TEST_CASE("metatdata formatting") { + std::string json =R"({ + "type": "metatomic_model_metadata", + "name": "name", + "description": "Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua. Ut enim ad minim veniam, quis nostrud exercitation.", + "authors": ["Short author", "Some extremely long author that will take more than one line in the printed output"], + "references": { + "architecture": ["ref-2", "ref-3"], + "model": ["a very long reference that will take more than one line in the printed output"], + "implementation": [] + }, + "extra": {} +})"; + + const auto* expected = R"(This is the name model +====================== + +Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor +incididunt ut labore et dolore magna aliqua. Ut enim ad minim veniam, quis +nostrud exercitation. + +Model authors +------------- + +- Short author +- Some extremely long author that will take more than one line in the printed + output + +Model references +---------------- + +Please cite the following references when using this model: +- about this specific model: + * a very long reference that will take more than one line in the printed + output +- about the architecture of this model: + * ref-2 + * ref-3 +)"; + + CHECK(metatomic::format_metadata(json) == expected); +} diff --git a/metatomic-core/tests/cxx/plugins.cpp b/metatomic-core/tests/cxx/plugins.cpp new file mode 100644 index 000000000..92be8995c --- /dev/null +++ b/metatomic-core/tests/cxx/plugins.cpp @@ -0,0 +1,25 @@ +#include + +#include "metatomic.hpp" + + +TEST_CASE("Load plugins") { + metatomic::load_plugin(PLUGIN_DIR "/test-c-plugin.so"); + + REQUIRE_THROWS_WITH( + metatomic::load_model("some_model", "{}", "test-c-plugin"), + "invalid parameter: failed to load model from 'some_model': plugin 'test-c-plugin' could not load the model" + ); + + REQUIRE_THROWS_WITH( + metatomic::load_model("some_model"), + "invalid parameter: failed to load model from 'some_model': tried the " + "following plugins, but none could load the model: test-c-plugin" + ); + + REQUIRE_THROWS_WITH( + metatomic::load_plugin(PLUGIN_DIR "/bad-abi-plugin.so"), + "invalid parameter: can not register plugin 'bad-abi-plugin': " + "plugin ABI version is 2, but metatomic expects 1" + ); +} diff --git a/metatomic-core/tests/misc.cpp b/metatomic-core/tests/misc.cpp index db9f21d40..8b3b7656f 100644 --- a/metatomic-core/tests/misc.cpp +++ b/metatomic-core/tests/misc.cpp @@ -3,7 +3,6 @@ #include #include "metatomic.h" -#include "metatomic.hpp" TEST_CASE("Version macros") { @@ -50,48 +49,29 @@ TEST_CASE("mta_string_t") { } TEST_CASE("unit conversion factor") { - SECTION("C API") { - double factor = 0.0; - - // same unit -> factor = 1.0 - auto status = mta_unit_conversion_factor("m", "m", &factor); - REQUIRE(status == MTA_SUCCESS); - CHECK(factor == 1.0); - - // kJ/mol -> eV - CHECK(mta_unit_conversion_factor("kJ/mol", "eV", &factor) == MTA_SUCCESS); - CHECK(factor == Approx(0.010364269656262174).epsilon(1e-15)); - - // dimension mismatch -> error - status = mta_unit_conversion_factor("m", "kg", &factor); - REQUIRE(status != MTA_SUCCESS); - - const char* error_msg = nullptr; - mta_last_error(&error_msg, nullptr, nullptr); - CHECK(std::string(error_msg) == - "invalid parameter: dimension mismatch in unit conversion: " - "'m' has dimension [L] but 'kg' has dimension [M]" - ); - } - - SECTION("C++ API") { - // same unit -> factor = 1.0 - auto factor = metatomic::unit_conversion_factor("m", "m"); - CHECK(factor == 1.0); - - // kJ/mol -> eV - factor = metatomic::unit_conversion_factor("kJ/mol", "eV"); - CHECK(factor == Approx(0.010364269656262174).epsilon(1e-15)); - - REQUIRE_THROWS_WITH( - metatomic::unit_conversion_factor("m", "kg"), - "invalid parameter: dimension mismatch in unit conversion: " - "'m' has dimension [L] but 'kg' has dimension [M]" - ); - } + double factor = 0.0; + + // same unit -> factor = 1.0 + auto status = mta_unit_conversion_factor("m", "m", &factor); + REQUIRE(status == MTA_SUCCESS); + CHECK(factor == 1.0); + + // kJ/mol -> eV + CHECK(mta_unit_conversion_factor("kJ/mol", "eV", &factor) == MTA_SUCCESS); + CHECK(factor == Approx(0.010364269656262174).epsilon(1e-15)); + + // dimension mismatch -> error + status = mta_unit_conversion_factor("m", "kg", &factor); + REQUIRE(status != MTA_SUCCESS); + + const char* error_msg = nullptr; + mta_last_error(&error_msg, nullptr, nullptr); + CHECK(std::string(error_msg) == + "invalid parameter: dimension mismatch in unit conversion: " + "'m' has dimension [L] but 'kg' has dimension [M]" + ); } - TEST_CASE("metatdata formatting") { std::string json =R"({ "type": "metatomic_model_metadata", @@ -105,6 +85,7 @@ TEST_CASE("metatdata formatting") { }, "extra": {} })"; + const auto* expected = R"(This is the name model ====================== @@ -131,17 +112,10 @@ Please cite the following references when using this model: * ref-3 )"; - SECTION("C API") { - auto* mta_string = mta_string_create(""); - REQUIRE(mta_string != nullptr); - auto status = mta_format_metadata(json.c_str(), &mta_string); - REQUIRE(status == MTA_SUCCESS); - CHECK(std::string(mta_string_view(mta_string)) == expected); - mta_string_free(mta_string); - } - - SECTION("C++ API") { - auto result = metatomic::format_metadata(json); - CHECK(result == expected); - } + auto* mta_string = mta_string_create(""); + REQUIRE(mta_string != nullptr); + auto status = mta_format_metadata(json.c_str(), &mta_string); + REQUIRE(status == MTA_SUCCESS); + CHECK(std::string(mta_string_view(mta_string)) == expected); + mta_string_free(mta_string); } diff --git a/metatomic-core/tests/plugins.cpp b/metatomic-core/tests/plugins.cpp index 2652d534c..70a2dd14f 100644 --- a/metatomic-core/tests/plugins.cpp +++ b/metatomic-core/tests/plugins.cpp @@ -1,71 +1,49 @@ #include #include "metatomic.h" -#include "metatomic.hpp" TEST_CASE("Load plugins") { - SECTION("C API") { - auto status = mta_load_plugin(PLUGIN_DIR "/test-c-plugin.so"); - CHECK(status == MTA_SUCCESS); + auto status = mta_load_plugin(PLUGIN_DIR "/test-c-plugin.so"); + CHECK(status == MTA_SUCCESS); - const char* error_message; - const char* error_origin; + const char* error_message; + const char* error_origin; - struct mta_model_t model; - status = mta_load_model("some_model", "{}", "test-c-plugin", &model); - CHECK(status == MTA_INVALID_PARAMETER_ERROR); + struct mta_model_t model; + status = mta_load_model("some_model", "{}", "test-c-plugin", &model); + CHECK(status == MTA_INVALID_PARAMETER_ERROR); - status = mta_last_error(&error_message, &error_origin, nullptr); - REQUIRE(status == MTA_SUCCESS); + status = mta_last_error(&error_message, &error_origin, nullptr); + REQUIRE(status == MTA_SUCCESS); - CHECK(std::string(error_origin) == "metatomic-core"); - CHECK(std::string(error_message) == ( - "invalid parameter: failed to load model from 'some_model': plugin 'test-c-plugin' could not load the model" - )); + CHECK(std::string(error_origin) == "metatomic-core"); + CHECK(std::string(error_message) == ( + "invalid parameter: failed to load model from 'some_model': plugin 'test-c-plugin' could not load the model" + )); - status = mta_load_model("some_model", "{}", nullptr, &model); - CHECK(status == MTA_INVALID_PARAMETER_ERROR); + status = mta_load_model("some_model", "{}", nullptr, &model); + CHECK(status == MTA_INVALID_PARAMETER_ERROR); - status = mta_last_error(&error_message, &error_origin, nullptr); - REQUIRE(status == MTA_SUCCESS); + status = mta_last_error(&error_message, &error_origin, nullptr); + REQUIRE(status == MTA_SUCCESS); - CHECK(std::string(error_origin) == "metatomic-core"); - CHECK(std::string(error_message) == ( - "invalid parameter: failed to load model from 'some_model': tried the " - "following plugins, but none could load the model: test-c-plugin" - )); + CHECK(std::string(error_origin) == "metatomic-core"); + CHECK(std::string(error_message) == ( + "invalid parameter: failed to load model from 'some_model': tried the " + "following plugins, but none could load the model: test-c-plugin" + )); - status = mta_load_plugin(PLUGIN_DIR "/bad-abi-plugin.so"); - CHECK(status == MTA_INVALID_PARAMETER_ERROR); + status = mta_load_plugin(PLUGIN_DIR "/bad-abi-plugin.so"); + CHECK(status == MTA_INVALID_PARAMETER_ERROR); - status = mta_last_error(&error_message, &error_origin, nullptr); - REQUIRE(status == MTA_SUCCESS); + status = mta_last_error(&error_message, &error_origin, nullptr); + REQUIRE(status == MTA_SUCCESS); - CHECK(std::string(error_origin) == "metatomic-core"); - CHECK(std::string(error_message) == ( - "invalid parameter: can not register plugin 'bad-abi-plugin': " - "plugin ABI version is 2, but metatomic expects 1" - )); - } - - SECTION("C++ API") { - REQUIRE_THROWS_WITH( - metatomic::load_model("some_model", "{}", "test-c-plugin"), - "invalid parameter: failed to load model from 'some_model': plugin 'test-c-plugin' could not load the model" - ); - - REQUIRE_THROWS_WITH( - metatomic::load_model("some_model"), - "invalid parameter: failed to load model from 'some_model': tried the " - "following plugins, but none could load the model: test-c-plugin" - ); - - REQUIRE_THROWS_WITH( - metatomic::load_plugin(PLUGIN_DIR "/bad-abi-plugin.so"), - "invalid parameter: can not register plugin 'bad-abi-plugin': " - "plugin ABI version is 2, but metatomic expects 1" - ); - } + CHECK(std::string(error_origin) == "metatomic-core"); + CHECK(std::string(error_message) == ( + "invalid parameter: can not register plugin 'bad-abi-plugin': " + "plugin ABI version is 2, but metatomic expects 1" + )); } From f3c161b8f06250366612efc11832626da21f9944 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Wed, 1 Jul 2026 14:53:21 +0200 Subject: [PATCH 28/43] Better fix to "no module named __pip__" in tests --- .github/workflows/rust-tests.yml | 5 +---- .github/workflows/torch-tests.yml | 9 +++------ metatomic-core/tests/utils/mod.rs | 16 +++------------- 3 files changed, 7 insertions(+), 23 deletions(-) diff --git a/.github/workflows/rust-tests.yml b/.github/workflows/rust-tests.yml index 9a29bf23b..ad17d814b 100644 --- a/.github/workflows/rust-tests.yml +++ b/.github/workflows/rust-tests.yml @@ -92,10 +92,7 @@ jobs: uses: actions/setup-python@v6 if: matrix.container == null with: - # Python 3.14.5 fails with "No module named pip.__main__; 'pip' is a - # package and cannot be directly executed" when using a venv, so we - # use 3.14.4 for now - python-version: "3.14.4" + python-version: "3.14" - name: Cache Rust dependencies uses: Leafwing-Studios/cargo-cache@v2.6.1 diff --git a/.github/workflows/torch-tests.yml b/.github/workflows/torch-tests.yml index af066c42b..93661740a 100644 --- a/.github/workflows/torch-tests.yml +++ b/.github/workflows/torch-tests.yml @@ -20,10 +20,7 @@ jobs: include: - os: ubuntu-24.04 torch-version: "2.13" - # Python 3.14.5 fails with "No module named pip.__main__; 'pip' is a - # package and cannot be directly executed" when using a venv, so we - # use 3.14.4 for now - python-version: "3.14.4" + python-version: "3.14" cargo-test-flags: --release do-valgrind: true @@ -36,12 +33,12 @@ jobs: - os: macos-15 torch-version: "2.13" - python-version: "3.14.4" + python-version: "3.14" cargo-test-flags: --release - os: windows-2022 torch-version: "2.13" - python-version: "3.14.4" + python-version: "3.14" cargo-test-flags: --release steps: - name: install dependencies in container diff --git a/metatomic-core/tests/utils/mod.rs b/metatomic-core/tests/utils/mod.rs index 3b23bca74..7a22f5e67 100644 --- a/metatomic-core/tests/utils/mod.rs +++ b/metatomic-core/tests/utils/mod.rs @@ -134,13 +134,13 @@ fn python_in_venv(venv_dir: &Path) -> PathBuf { python } -/// Create a fresh Python virtualenv using uv if available, else fallback to +/// Create a Python virtualenv using uv if available, else fallback to /// `python -m venv`, and return the path to the python executable in the venv pub fn create_python_venv(build_dir: PathBuf) -> PathBuf { if let Some(uv_bin) = find_uv() { let mut cmd = Command::new(&uv_bin); cmd.arg("venv"); - cmd.arg("--clear"); + cmd.arg("--allow-existing"); cmd.arg(&build_dir); run_command(cmd, "uv venv creation"); @@ -148,20 +148,10 @@ pub fn create_python_venv(build_dir: PathBuf) -> PathBuf { let mut cmd = Command::new(find_python()); cmd.arg("-m"); cmd.arg("venv"); + cmd.arg("--upgrade-deps"); cmd.arg(&build_dir); run_command(cmd, "python to create virtualenv with `venv`"); - - // update pip in case the system uses a very old one - let python = python_in_venv(&build_dir); - let mut cmd = Command::new(&python); - cmd.arg("-m"); - cmd.arg("pip"); - cmd.arg("install"); - cmd.arg("--upgrade"); - cmd.arg("pip"); - - run_command(cmd, "pip upgrade in virtualenv"); } python_in_venv(&build_dir) From e9126c21201c822cc2ac4cf96acec6493b5e6aa2 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Wed, 8 Jul 2026 11:09:17 +0200 Subject: [PATCH 29/43] Bump MSRV to 1.88 --- .github/workflows/rust-tests.yml | 25 +++- CONTRIBUTING.rst | 2 +- metatomic-core/CMakeLists.txt | 52 +------ metatomic-core/Cargo.toml | 17 +-- metatomic-core/cmake/detect_cargo.cmake | 181 ++++++++++++++++++++++++ metatomic-torch/Cargo.toml | 2 +- python/Cargo.toml | 2 +- 7 files changed, 207 insertions(+), 74 deletions(-) create mode 100644 metatomic-core/cmake/detect_cargo.cmake diff --git a/.github/workflows/rust-tests.yml b/.github/workflows/rust-tests.yml index ad17d814b..f9c2a1c4d 100644 --- a/.github/workflows/rust-tests.yml +++ b/.github/workflows/rust-tests.yml @@ -25,30 +25,33 @@ jobs: strategy: matrix: include: + # test our MSRV - os: ubuntu-24.04 - rust-version: stable + rust-version: 1.88 rust-target: x86_64-unknown-linux-gnu cxx: g++ cc: gcc + cargo: cargo cmake-generator: Unix Makefiles # check the build on a stock Ubuntu 22.04, which uses cmake 3.22, and - # with our minimal supported rust version + # using cargo/rustc from APT - os: ubuntu-24.04 - rust-version: 1.74 + rust-version: from APT container: ubuntu:22.04 rust-target: x86_64-unknown-linux-gnu extra-name: ", cmake 3.22" cxx: g++ cc: gcc + cargo: cargo-1.89 cmake-generator: Unix Makefiles - os: macos-15 rust-version: stable rust-target: aarch64-apple-darwin - extra-name: "" cxx: clang++ cc: clang + cargo: cargo cmake-generator: Unix Makefiles - os: windows-2022 @@ -57,6 +60,7 @@ jobs: extra-name: " / MSVC" cxx: cl.exe cc: cl.exe + cargo: cargo cmake-generator: Visual Studio 17 2022 - os: windows-2022 @@ -65,6 +69,7 @@ jobs: extra-name: " / MinGW" cxx: g++.exe cc: gcc.exe + cargo: cargo cmake-generator: MinGW Makefiles steps: - name: install dependencies in container @@ -72,7 +77,11 @@ jobs: run: | apt update apt install -y software-properties-common - apt install -y cmake make gcc g++ git curl python3-venv + apt install -y cmake make gcc g++ git curl python3-venv cargo-1.89 + + # for some reason, cargo-1.89 from APT tries to find `rustdoc` and + # not `rustdoc-1.89`, so we force it to use the correct one + echo "RUSTDOC=rustdoc-1.89" >> "$GITHUB_ENV" - uses: actions/checkout@v6 with: @@ -84,6 +93,7 @@ jobs: - name: setup rust uses: dtolnay/rust-toolchain@master + if: matrix.container == null with: toolchain: ${{ matrix.rust-version }} target: ${{ matrix.rust-target }} @@ -96,6 +106,7 @@ jobs: - name: Cache Rust dependencies uses: Leafwing-Studios/cargo-cache@v2.6.1 + if: matrix.container == null with: sweep-cache: true @@ -119,10 +130,10 @@ jobs: echo "CMAKE_CXX_COMPILER_LAUNCHER=sccache" >> $GITHUB_ENV - name: run tests - run: | - cargo test --package metatomic-core --target ${{ matrix.rust-target }} env: RUST_BACKTRACE: full + run: | + ${{ matrix.cargo }} test --package metatomic-core --target ${{ matrix.rust-target }} - name: check that the header was already up to date run: | diff --git a/CONTRIBUTING.rst b/CONTRIBUTING.rst index e13c2c6f1..f0f4dd5fd 100644 --- a/CONTRIBUTING.rst +++ b/CONTRIBUTING.rst @@ -19,7 +19,7 @@ on metatomic: - **the rust compiler**: you will need both ``rustc`` (the compiler) and ``cargo`` (associated build tool). You can install both using `rustup`_, or use a version provided by your operating system. We need at least Rust version - 1.74 to build metatomic. + 1.88 to build metatomic. - **Python**: you can install ``Python`` and ``pip`` on your operating system. We require a Python version of at least 3.9. - **tox**: a Python test runner, see https://tox.readthedocs.io/en/latest/. You diff --git a/metatomic-core/CMakeLists.txt b/metatomic-core/CMakeLists.txt index 6a2169173..73382e21b 100644 --- a/metatomic-core/CMakeLists.txt +++ b/metatomic-core/CMakeLists.txt @@ -117,46 +117,7 @@ else() endif() -find_program(CARGO_EXE "cargo" DOC "path to cargo (Rust build system)") -if (NOT CARGO_EXE) - message(FATAL_ERROR - "could not find cargo, please make sure the Rust compiler is installed \ - (see https://www.rust-lang.org/tools/install) or set CARGO_EXE" - ) -endif() - -execute_process( - COMMAND ${CARGO_EXE} "--version" "--verbose" - RESULT_VARIABLE CARGO_STATUS - OUTPUT_VARIABLE CARGO_VERSION_RAW -) - -if(CARGO_STATUS AND NOT CARGO_STATUS EQUAL 0) - message(FATAL_ERROR - "could not run cargo, please make sure the Rust compiler is installed \ - (see https://www.rust-lang.org/tools/install)" - ) -endif() - -set(REQUIRED_RUST_VERSION "1.74.0") -if (CARGO_VERSION_RAW MATCHES "cargo ([0-9]+\\.[0-9]+\\.[0-9]+).*") - set(CARGO_VERSION "${CMAKE_MATCH_1}") -else() - message(FATAL_ERROR "failed to determine cargo version, output was: ${CARGO_VERSION_RAW}") -endif() - -if (${CARGO_VERSION} VERSION_LESS ${REQUIRED_RUST_VERSION}) - message(FATAL_ERROR - "your Rust installation is too old (you have version ${CARGO_VERSION}), \ - at least ${REQUIRED_RUST_VERSION} is required" - ) -else() - if(NOT "${CACHED_LAST_CARGO_VERSION}" STREQUAL ${CARGO_VERSION}) - set(CACHED_LAST_CARGO_VERSION ${CARGO_VERSION} CACHE INTERNAL "Last version of cargo used in configuration") - message(STATUS "Using cargo version ${CARGO_VERSION} at ${CARGO_EXE}") - set(CARGO_VERSION_CHANGED TRUE) - endif() -endif() +include(cmake/detect_cargo.cmake) # ============================================================================ # # determine Cargo flags @@ -184,17 +145,6 @@ endif() set(CARGO_TARGET_DIR ${CMAKE_CURRENT_BINARY_DIR}/target) set(CARGO_BUILD_ARG "${CARGO_BUILD_ARG};--target-dir=${CARGO_TARGET_DIR}") -if (CARGO_VERSION_RAW MATCHES "host: ([a-zA-Z0-9_\\-]*)\n") - set(RUST_HOST_TARGET "${CMAKE_MATCH_1}") - if (RUST_HOST_TARGET MATCHES "([a-zA-Z0-9_]*)\\-") - set(RUST_HOST_ARCH "${CMAKE_MATCH_1}") - else() - message(FATAL_ERROR "failed to determine host CPU arch, target was: ${RUST_HOST_TARGET}") - endif() -else() - message(FATAL_ERROR "failed to determine host target, output was: ${CARGO_VERSION_RAW}") -endif() - if (WIN32) # on Windows, we need to use the same ABI in both CMake and cargo. If the # user did not explicitly request a target, we can try to set it ourself, diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index 4693e30dd..f4628ddef 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -3,7 +3,7 @@ name = "metatomic-core" version = "0.1.0" edition = "2021" publish = false -rust-version = "1.74" +rust-version = "1.88" exclude = [ "tests" ] @@ -14,26 +14,17 @@ name = "metatomic" bench = false [dependencies] -metatensor = { version = "0.4.1" } +metatensor = { version = "0.5.0" } once_cell = "1" -dlpk = { version = "0.3", features = ["ndarray"]} +dlpk = { version = "0.4", features = ["ndarray"]} json = "0.12" -libloading = "0.8" +libloading = "0.9" ndarray = "0.17" [build-dependencies] cbindgen = { version = "0.29", default-features = false } -# the last versions that supports Rust 1.74 -serde_spanned = "=1.0.1" -toml = "=0.9.6" -toml_datetime = "=0.7.1" -toml_parser = "=1.0.2" -toml_writer = "=1.0.2" -tempfile = "=3.24.0" -indexmap = "=2.11.4" - [dev-dependencies] lazy_static = "1" which = "8" diff --git a/metatomic-core/cmake/detect_cargo.cmake b/metatomic-core/cmake/detect_cargo.cmake new file mode 100644 index 000000000..0af2293e7 --- /dev/null +++ b/metatomic-core/cmake/detect_cargo.cmake @@ -0,0 +1,181 @@ +# This module finds a suitable cargo binary. It tries plain "cargo" first, then +# searches for versioned cargo binaries (e.g. cargo-1.82) commonly installed on +# Ubuntu. If a binary is found but too old, it continues searching for a newer +# one. +# +# Sets: +# CARGO_EXE - path to the chosen cargo binary +# CARGO_VERSION - parsed version string (e.g. 1.74.0) +# RUST_HOST_TARGET - host target triple (e.g. x86_64-unknown-linux-gnu) +# RUST_HOST_ARCH - host CPU architecture (e.g. x86_64) +# CACHED_LAST_CARGO_VERSION - cache variable for change detection +# CARGO_VERSION_CHANGED - true if the version differs from the last run + +set(REQUIRED_RUST_VERSION "1.88.0") + +# --------------------------------------------------------------------------- +# Helper: run cargo --version --verbose, extract version & host target +# --------------------------------------------------------------------------- +function(_try_cargo _exe _ok_var _version_var _host_target_var _host_arch_var) + execute_process( + COMMAND "${_exe}" "--version" "--verbose" + RESULT_VARIABLE _status + OUTPUT_VARIABLE _raw + ERROR_QUIET + ) + + if (NOT _status EQUAL 0) + set(${_ok_var} FALSE PARENT_SCOPE) + return() + endif() + + set(_ok TRUE) + set(_version "") + set(_host_target "") + + if (_raw MATCHES "cargo ([0-9]+\\.[0-9]+\\.[0-9]+)") + set(_version "${CMAKE_MATCH_1}") + else() + set(_ok FALSE) + endif() + + if (_raw MATCHES "host: ([a-zA-Z0-9_\\-]*)\n") + set(_host_target "${CMAKE_MATCH_1}") + else() + set(_ok FALSE) + endif() + + set(${_ok_var} ${_ok} PARENT_SCOPE) + set(${_version_var} "${_version}" PARENT_SCOPE) + set(${_host_target_var} "${_host_target}" PARENT_SCOPE) + + if (_host_target MATCHES "([a-zA-Z0-9_]*)\\-") + set(${_host_arch_var} "${CMAKE_MATCH_1}" PARENT_SCOPE) + else() + set(${_host_arch_var} "" PARENT_SCOPE) + endif() +endfunction() + +# --------------------------------------------------------------------------- +# Step 1: try plain "cargo" (or respect a pre-defined CARGO_EXE) +# --------------------------------------------------------------------------- +set(_cargo_found FALSE) +if (DEFINED CARGO_EXE AND NOT CARGO_EXE STREQUAL "CARGO_EXE-NOTFOUND") + _try_cargo("${CARGO_EXE}" _ok _ver _target _arch) + if (_ok AND ${_ver} VERSION_GREATER_EQUAL ${REQUIRED_RUST_VERSION}) + set(_cargo_found TRUE) + set(CARGO_VERSION "${_ver}") + set(RUST_HOST_TARGET "${_target}") + set(RUST_HOST_ARCH "${_arch}") + else() + # Cache is stale or binary changed; re-search below + message(STATUS "cargo at ${CARGO_EXE} is not usable, searching for alternatives...") + unset(CARGO_EXE) + unset(CARGO_EXE CACHE) + endif() +endif() + +if (NOT _cargo_found) + find_program(_cargo_vanilla "cargo") + if (_cargo_vanilla) + _try_cargo("${_cargo_vanilla}" _ok _ver _target _arch) + if (_ok AND ${_ver} VERSION_GREATER_EQUAL ${REQUIRED_RUST_VERSION}) + set(_cargo_found TRUE) + set(CARGO_EXE "${_cargo_vanilla}") + set(CARGO_VERSION "${_ver}") + set(RUST_HOST_TARGET "${_target}") + set(RUST_HOST_ARCH "${_arch}") + endif() + endif() +endif() + +# --------------------------------------------------------------------------- +# Step 2: search for versioned cargo-* binaries across PATH +# --------------------------------------------------------------------------- +if (NOT _cargo_found) + # Collect all directories to search + set(_search_dirs ${CMAKE_PROGRAM_PATH}) + + if (WIN32) + foreach(_dir IN LISTS $ENV{PATH}) + list(APPEND _search_dirs "${_dir}") + endforeach() + else() + string(REPLACE ":" ";" _sys_path "$ENV{PATH}") + list(APPEND _search_dirs ${_sys_path}) + endif() + + if (NOT "$ENV{HOME}" STREQUAL "") + list(APPEND _search_dirs "$ENV{HOME}/.cargo/bin") + endif() + + set(_cargo_candidates "") + foreach(_dir IN LISTS _search_dirs) + if (IS_DIRECTORY "${_dir}") + file(GLOB _bins "${_dir}/cargo-*") + list(APPEND _cargo_candidates ${_bins}) + endif() + endforeach() + + if (_cargo_candidates) + list(REMOVE_DUPLICATES _cargo_candidates) + endif() + + set(_best_exe "") + set(_best_version "0.0.0") + set(_best_target "") + set(_best_arch "") + + foreach(_bin IN LISTS _cargo_candidates) + _try_cargo("${_bin}" _ok _ver _target _arch) + if (_ok AND ${_ver} VERSION_GREATER_EQUAL ${REQUIRED_RUST_VERSION} + AND ${_ver} VERSION_GREATER ${_best_version}) + set(_best_exe "${_bin}") + set(_best_version "${_ver}") + set(_best_target "${_target}") + set(_best_arch "${_arch}") + endif() + endforeach() + + if (_best_exe) + set(_cargo_found TRUE) + set(CARGO_EXE "${_best_exe}") + set(CARGO_VERSION "${_best_version}") + set(RUST_HOST_TARGET "${_best_target}") + set(RUST_HOST_ARCH "${_best_arch}") + endif() +endif() + +# --------------------------------------------------------------------------- +# Final validation +# --------------------------------------------------------------------------- +if (NOT _cargo_found) + message(FATAL_ERROR + "could not find a suitable cargo binary (version >= ${REQUIRED_RUST_VERSION}).\n" + "Please install Rust from https://www.rust-lang.org/tools/install\n" + "or set CARGO_EXE to point to your cargo binary before calling CMake." + ) +endif() + +if (NOT RUST_HOST_TARGET) + message(FATAL_ERROR + "failed to determine host target from cargo --version --verbose" + ) +endif() + +if (NOT RUST_HOST_ARCH) + message(FATAL_ERROR + "failed to determine host CPU arch from target: ${RUST_HOST_TARGET}" + ) +endif() + +# --------------------------------------------------------------------------- +# Cache for change detection across CMake re-configures +# --------------------------------------------------------------------------- +if (NOT "${CACHED_LAST_CARGO_VERSION}" STREQUAL "${CARGO_VERSION}") + set(CACHED_LAST_CARGO_VERSION "${CARGO_VERSION}" + CACHE INTERNAL "Last version of cargo used in configuration" + ) + message(STATUS "Using cargo version ${CARGO_VERSION} at ${CARGO_EXE}") + set(CARGO_VERSION_CHANGED TRUE) +endif() diff --git a/metatomic-torch/Cargo.toml b/metatomic-torch/Cargo.toml index 3809a5a99..82672a105 100644 --- a/metatomic-torch/Cargo.toml +++ b/metatomic-torch/Cargo.toml @@ -3,7 +3,7 @@ name = "metatomic-torch" version = "0.0.0" edition = "2021" publish = false -rust-version = "1.74" +rust-version = "1.88" [lib] path = "lib.rs" diff --git a/python/Cargo.toml b/python/Cargo.toml index 2ca54178e..88607618f 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -3,7 +3,7 @@ name = "metatomic-python" version = "0.0.0" edition = "2021" publish = false -rust-version = "1.74" +rust-version = "1.88" [lib] path = "lib.rs" From 64e16dae90f3b644990994640ff6efa58c755c17 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Wed, 8 Jul 2026 11:22:15 +0200 Subject: [PATCH 30/43] Update to edition 2024 --- metatomic-core/Cargo.toml | 3 +- metatomic-core/src/c_api/model.rs | 10 ++- metatomic-core/src/c_api/plugin.rs | 30 ++++--- metatomic-core/src/c_api/status.rs | 28 +++--- metatomic-core/src/c_api/system.rs | 131 +++++++++++++++++------------ metatomic-core/src/c_api/utils.rs | 26 +++--- metatomic-core/src/model.rs | 107 +++++++++++++---------- metatomic-core/src/plugin.rs | 7 +- metatomic-core/src/quantities.rs | 16 ++-- metatomic-core/src/system.rs | 4 +- metatomic-core/src/units.rs | 4 +- metatomic-torch/Cargo.toml | 2 +- python/Cargo.toml | 2 +- 13 files changed, 209 insertions(+), 161 deletions(-) diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index f4628ddef..7cbc6d23f 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "metatomic-core" version = "0.1.0" -edition = "2021" +edition = "2024" publish = false rust-version = "1.88" exclude = [ @@ -15,7 +15,6 @@ bench = false [dependencies] metatensor = { version = "0.5.0" } -once_cell = "1" dlpk = { version = "0.4", features = ["ndarray"]} json = "0.12" libloading = "0.9" diff --git a/metatomic-core/src/c_api/model.rs b/metatomic-core/src/c_api/model.rs index 869db8ede..2f9aba2ab 100644 --- a/metatomic-core/src/c_api/model.rs +++ b/metatomic-core/src/c_api/model.rs @@ -200,7 +200,7 @@ impl mta_model_t { /// `requested_outputs_count` /// @return `MTA_SUCCESS` on success, another status code on error (the message /// is available through `mta_last_error`) -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_execute_model( model: mta_model_t, systems: *const *const mta_system_t, @@ -222,7 +222,7 @@ pub unsafe extern "C" fn mta_execute_model( /// metadata. The caller takes ownership and must free it with /// `mta_string_free`. /// @return `MTA_SUCCESS` on success, another status code on error -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_format_metadata( metadata: *const c_char, printed: *mut mta_string_t, @@ -230,7 +230,7 @@ pub unsafe extern "C" fn mta_format_metadata( catch_unwind(|| { check_pointers_non_null!(metadata, printed); - let metadata = std::ffi::CStr::from_ptr(metadata); + let metadata = unsafe { std::ffi::CStr::from_ptr(metadata) }; let metadata = metadata.to_str().map_err(|_| { Error::InvalidParameter("metadata is not valid UTF-8".into()) })?; @@ -241,7 +241,9 @@ pub unsafe extern "C" fn mta_format_metadata( let metadata = ModelMetadata::try_from(&metadata)?; - *printed = mta_string_t::new(metadata.print()); + unsafe { + *printed = mta_string_t::new(metadata.print()); + } Ok(()) }) } diff --git a/metatomic-core/src/c_api/plugin.rs b/metatomic-core/src/c_api/plugin.rs index 42ee45dee..ed056aa93 100644 --- a/metatomic-core/src/c_api/plugin.rs +++ b/metatomic-core/src/c_api/plugin.rs @@ -54,7 +54,7 @@ unsafe impl Send for mta_plugin_t {} /// @return `MTA_SUCCESS` if the plugin was registered successfully, or another /// status code if an error occurs. You can get more details about the error /// with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_register_plugin(plugin: mta_plugin_t) -> mta_status_t { catch_unwind(move || { let plugin = Plugin::new(plugin)?; @@ -73,14 +73,16 @@ pub unsafe extern "C" fn mta_register_plugin(plugin: mta_plugin_t) -> mta_status /// @return `MTA_SUCCESS` if the plugin was loaded successfully, or another /// status code if an error occurs. You can get more details about the /// error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_load_plugin(path: *const c_char) -> mta_status_t { catch_unwind(move || { check_pointers_non_null!(path); - let path = CStr::from_ptr(path).to_str().map_err(|_| { - Error::InvalidParameter("invalid UTF-8 in plugin path".into()) - })?; + let path = unsafe { CStr::from_ptr(path) } + .to_str() + .map_err(|_| { + Error::InvalidParameter("invalid UTF-8 in plugin path".into()) + })?; crate::plugin::load_plugin(path) }) @@ -110,7 +112,7 @@ pub unsafe extern "C" fn mta_load_plugin(path: *const c_char) -> mta_status_t { /// @return `MTA_SUCCESS` if the model was loaded successfully, or another /// status code if an error occurs. You can get more details about the /// error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_load_model( load_from: *const c_char, options_json: *const c_char, @@ -125,15 +127,16 @@ pub unsafe extern "C" fn mta_load_model( let plugin_name = if plugin_name.is_null() { None } else { - Some(CStr::from_ptr(plugin_name).to_str().map_err(|_| { + let cstr = unsafe { CStr::from_ptr(plugin_name) }; + Some(cstr.to_str().map_err(|_| { Error::InvalidParameter("invalid UTF-8 in plugin name".into()) })?) }; - let options_json = if options_json.is_null() { - CStr::from_bytes_with_nul(b"{}\0").expect("invalid CStr") + let options_json = if options_json.is_null() { + c"{}" } else { - CStr::from_ptr(options_json) + unsafe { CStr::from_ptr(options_json) } }; let options_str = options_json.to_str().map_err(|_| { @@ -157,10 +160,13 @@ pub unsafe extern "C" fn mta_load_model( } } - let loaded = crate::plugin::load_model(CStr::from_ptr(load_from), options_json, plugin_name)?; + let load_from = unsafe { CStr::from_ptr(load_from) }; + let loaded = crate::plugin::load_model(load_from, options_json, plugin_name)?; let _ = &unwind_wrapper; - *unwind_wrapper.0 = loaded.into_raw(); + unsafe { + *unwind_wrapper.0 = loaded.into_raw(); + } Ok(()) }) } diff --git a/metatomic-core/src/c_api/status.rs b/metatomic-core/src/c_api/status.rs index 48d658777..cefd7d347 100644 --- a/metatomic-core/src/c_api/status.rs +++ b/metatomic-core/src/c_api/status.rs @@ -121,7 +121,7 @@ impl From for mta_status_t { } /// Get last error message that was created on the current thread. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_last_error( message: *mut *const c_char, origin: *mut *const c_char, @@ -129,15 +129,17 @@ pub unsafe extern "C" fn mta_last_error( ) -> mta_status_t { let status = std::panic::catch_unwind(|| { LAST_ERROR.with(|last_error| { - let last_error = last_error.borrow(); - if !message.is_null() { - *message = last_error.message.as_ptr(); - } - if !origin.is_null() { - *origin = last_error.origin.as_ptr(); - } - if !data.is_null() { - *data = last_error.custom_data; + unsafe { + let last_error = last_error.borrow(); + if !message.is_null() { + *message = last_error.message.as_ptr(); + } + if !origin.is_null() { + *origin = last_error.origin.as_ptr(); + } + if !data.is_null() { + *data = last_error.custom_data; + } } }); }); @@ -171,7 +173,7 @@ pub unsafe extern "C" fn mta_last_error( } /// Set last error message for the current thread. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_set_last_error( message: *const c_char, origin: *const c_char, @@ -182,13 +184,13 @@ pub unsafe extern "C" fn mta_set_last_error( let message = if message.is_null() { CString::new("").expect("invalid C string") } else { - CString::from(CStr::from_ptr(message)) + unsafe { CString::from(CStr::from_ptr(message)) } }; let origin = if origin.is_null() { CString::new("").expect("invalid C string") } else { - CString::from(CStr::from_ptr(origin)) + unsafe { CString::from(CStr::from_ptr(origin)) } }; LAST_ERROR.with(|last_error| { diff --git a/metatomic-core/src/c_api/system.rs b/metatomic-core/src/c_api/system.rs index 4f6ff1235..b6f57b21a 100644 --- a/metatomic-core/src/c_api/system.rs +++ b/metatomic-core/src/c_api/system.rs @@ -35,7 +35,7 @@ pub struct mta_system_t(pub(crate) Arc); /// The caller takes ownership and must free it with `mta_system_free`. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_create( length_unit: *const c_char, types: *mut DLManagedTensorVersioned, @@ -48,20 +48,23 @@ pub unsafe extern "C" fn mta_system_create( catch_unwind(move || { check_pointers_non_null!(length_unit, types, positions, cell, pbc, system); - let length_unit = CStr::from_ptr(length_unit) - .to_str() - .map_err(|_| Error::InvalidParameter("length_unit is not valid UTF-8".into()))? - .to_string(); + unsafe { + let length_unit = CStr::from_ptr(length_unit) + .to_str() + .map_err(|_| Error::InvalidParameter("length_unit is not valid UTF-8".into()))? + .to_string(); - let types = DLPackTensor::from_ptr(types); - let positions = DLPackTensor::from_ptr(positions); - let cell = DLPackTensor::from_ptr(cell); - let pbc = DLPackTensor::from_ptr(pbc); - let system_inner = System::new(length_unit, types, positions, cell, pbc)?; + let types = DLPackTensor::from_ptr(types); + let positions = DLPackTensor::from_ptr(positions); + let cell = DLPackTensor::from_ptr(cell); + let pbc = DLPackTensor::from_ptr(pbc); - let _ = &unwind_wrapper; - *unwind_wrapper.0 = Box::into_raw(Box::new(mta_system_t(Arc::new(system_inner)))); + let system_inner = System::new(length_unit, types, positions, cell, pbc)?; + + let _ = &unwind_wrapper; + *unwind_wrapper.0 = Box::into_raw(Box::new(mta_system_t(Arc::new(system_inner)))); + } Ok(()) }) } @@ -75,14 +78,16 @@ pub unsafe extern "C" fn mta_system_create( /// function is a no-op. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_free(system: *mut mta_system_t) -> mta_status_t { catch_unwind(|| { if system.is_null() { return Ok(()); } - let _ = Box::from_raw(system); + unsafe { + std::mem::drop(Box::from_raw(system)); + } Ok(()) }) } @@ -93,7 +98,7 @@ pub unsafe extern "C" fn mta_system_free(system: *mut mta_system_t) -> mta_statu /// @param size Output parameter, set to the number of atoms. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_size( system: *const mta_system_t, size: *mut usize, @@ -101,8 +106,10 @@ pub unsafe extern "C" fn mta_system_size( catch_unwind(|| { check_pointers_non_null!(system, size); - let system = &*system; - *size = system.0.size(); + unsafe { + let system = &*system; + *size = system.0.size(); + } Ok(()) }) } @@ -129,11 +136,15 @@ unsafe extern "C" fn borrowed_tensor_deleter( tensor: *mut DLManagedTensorVersioned, ) { let system: Arc = { - let ptr = (*tensor).manager_ctx as *const System; - Arc::from_raw(ptr) + unsafe { + let ptr = (*tensor).manager_ctx as *const System; + Arc::from_raw(ptr) + } }; - drop(system); - let _ = Box::from_raw(tensor); + std::mem::drop(system); + unsafe { + std::mem::drop(Box::from_raw(tensor)); + } } /// Get a DLPack tensor from a system for the requested data. @@ -154,7 +165,7 @@ unsafe extern "C" fn borrowed_tensor_deleter( /// takes ownership and must call the deleter when done. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_get_data( system: *const mta_system_t, request: mta_system_data_kind, @@ -162,9 +173,11 @@ pub unsafe extern "C" fn mta_system_get_data( ) -> mta_status_t { catch_unwind(|| { check_pointers_non_null!(system, data); - *data = std::ptr::null_mut(); + unsafe { + *data = std::ptr::null_mut(); + } - let system = &*system; + let system = unsafe { &*system }; let tensor_ref = match request { mta_system_data_kind::MTA_SYSTEM_DATA_TYPES => system.0.types(), mta_system_data_kind::MTA_SYSTEM_DATA_POSITIONS => system.0.positions(), @@ -187,7 +200,9 @@ pub unsafe extern "C" fn mta_system_get_data( dl_tensor: tensor_ref.raw.clone(), }); - *data = Box::into_raw(packed); + unsafe { + *data = Box::into_raw(packed); + } Ok(()) }) } @@ -201,17 +216,18 @@ pub unsafe extern "C" fn mta_system_get_data( /// @param length_unit Output parameter, set to the length unit string. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_get_length_unit( system: *const mta_system_t, length_unit: *mut mta_string_t, ) -> mta_status_t { catch_unwind(|| { check_pointers_non_null!(system, length_unit); - *length_unit = mta_string_t::null(); - let system = &*system; - *length_unit = mta_string_t::new(system.0.length_unit()); + unsafe { + let system = &*system; + *length_unit = mta_string_t::new(system.0.length_unit()); + } Ok(()) }) } @@ -227,7 +243,7 @@ pub unsafe extern "C" fn mta_system_get_length_unit( /// transferred. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_add_pairs( system: *mut mta_system_t, options: *const c_char, @@ -236,7 +252,7 @@ pub unsafe extern "C" fn mta_system_add_pairs( catch_unwind(|| { check_pointers_non_null!(system, options, pairs); - let options_str = CStr::from_ptr(options) + let options_str = unsafe { CStr::from_ptr(options) } .to_str() .map_err(|_| Error::InvalidParameter("options is not valid UTF-8".into()))?; @@ -245,9 +261,9 @@ pub unsafe extern "C" fn mta_system_add_pairs( let options = PairListOptions::try_from(&options_json)?; - let pairs = TensorBlock::from_raw(pairs); + let pairs = unsafe { TensorBlock::from_raw(pairs) }; - let system = &mut *system; + let system = unsafe { &mut *system }; let system = Arc::get_mut(&mut system.0).ok_or_else(|| { Error::InvalidParameter( "cannot modify system while there are outstanding borrowed views".into(), @@ -272,7 +288,7 @@ pub unsafe extern "C" fn mta_system_add_pairs( /// NULL if no pair list matches the options. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_get_pairs( system: *const mta_system_t, options: *const c_char, @@ -280,9 +296,8 @@ pub unsafe extern "C" fn mta_system_get_pairs( ) -> mta_status_t { catch_unwind(|| { check_pointers_non_null!(system, options, pairs); - *pairs = std::ptr::null(); - let options_str = CStr::from_ptr(options) + let options_str = unsafe { CStr::from_ptr(options) } .to_str() .map_err(|_| Error::InvalidParameter("options is not valid UTF-8".into()))?; @@ -291,10 +306,12 @@ pub unsafe extern "C" fn mta_system_get_pairs( let options = PairListOptions::try_from(&options_json)?; - let system = &*system; + let system = unsafe { &*system }; match system.0.get_pairs(&options) { Some(block) => { - *pairs = block.as_ptr(); + unsafe { + *pairs = block.as_ptr(); + } } None => { return Err(Error::InvalidParameter( @@ -317,16 +334,15 @@ pub unsafe extern "C" fn mta_system_get_pairs( /// array of `PairListOptions` objects. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_known_pairs( system: *const mta_system_t, pairs_options: *mut mta_string_t, ) -> mta_status_t { catch_unwind(|| { check_pointers_non_null!(system, pairs_options); - *pairs_options = mta_string_t::null(); - let system = &*system; + let system = unsafe { &*system }; let known = system.0.known_pairs(); let mut json_array = json::JsonValue::new_array(); for options in known { @@ -335,7 +351,9 @@ pub unsafe extern "C" fn mta_system_known_pairs( })?; } - *pairs_options = mta_string_t::new(json::stringify(json_array)); + unsafe { + *pairs_options = mta_string_t::new(json::stringify(json_array)); + } Ok(()) }) } @@ -352,7 +370,7 @@ pub unsafe extern "C" fn mta_system_known_pairs( /// transferred. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_add_custom_data( system: *mut mta_system_t, name: *const c_char, @@ -361,14 +379,14 @@ pub unsafe extern "C" fn mta_system_add_custom_data( catch_unwind(|| { check_pointers_non_null!(system, name, data); - let name = CStr::from_ptr(name) + let name = unsafe { CStr::from_ptr(name) } .to_str() .map_err(|_| Error::InvalidParameter("name is not valid UTF-8".into()))? .to_string(); - let data = TensorMap::from_raw(data); + let data = unsafe { TensorMap::from_raw(data) }; - let system = &mut *system; + let system = unsafe { &mut *system }; let system = Arc::get_mut(&mut system.0).ok_or_else(|| { Error::InvalidParameter( "cannot modify system while there are outstanding borrowed views".into(), @@ -393,7 +411,7 @@ pub unsafe extern "C" fn mta_system_add_custom_data( /// map, or an error if no data with the given name exists. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_get_custom_data( system: *const mta_system_t, name: *const c_char, @@ -401,15 +419,17 @@ pub unsafe extern "C" fn mta_system_get_custom_data( ) -> mta_status_t { catch_unwind(|| { check_pointers_non_null!(system, name, data); - *data = std::ptr::null_mut(); - let name = CStr::from_ptr(name) + let name = unsafe { CStr::from_ptr(name) } .to_str() .map_err(|_| Error::InvalidParameter("name is not valid UTF-8".into()))?; - let system = &*system; + let system = unsafe { &*system }; let result = system.0.get_custom_data(name)?; - *data = result.as_ptr(); + + unsafe { + *data = result.as_ptr(); + } Ok(()) }) @@ -425,16 +445,15 @@ pub unsafe extern "C" fn mta_system_get_custom_data( /// custom data names. /// @return `MTA_SUCCESS` on success, or another status code if an error occurs. /// You can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_system_known_custom_data( system: *const mta_system_t, names: *mut mta_string_t, ) -> mta_status_t { catch_unwind(|| { check_pointers_non_null!(system, names); - *names = mta_string_t::null(); - let system = &*system; + let system = unsafe { &*system }; let known = system.0.known_custom_data(); let mut json_array = json::JsonValue::new_array(); for name in known { @@ -443,7 +462,9 @@ pub unsafe extern "C" fn mta_system_known_custom_data( })?; } - *names = mta_string_t::new(json::stringify(json_array)); + unsafe { + *names = mta_string_t::new(json::stringify(json_array)); + } Ok(()) }) } diff --git a/metatomic-core/src/c_api/utils.rs b/metatomic-core/src/c_api/utils.rs index 38e5222c4..22fabe38a 100644 --- a/metatomic-core/src/c_api/utils.rs +++ b/metatomic-core/src/c_api/utils.rs @@ -1,11 +1,11 @@ use std::ffi::{CString, c_char}; -use once_cell::sync::Lazy; +use std::sync::LazyLock; use super::{mta_status_t, catch_unwind}; use crate::Error; -static VERSION: Lazy = Lazy::new(|| { +static VERSION: LazyLock = LazyLock::new(|| { CString::new(env!("METATOMIC_FULL_VERSION")).expect("version contains NULL byte") }); @@ -13,7 +13,7 @@ static VERSION: Lazy = Lazy::new(|| { /// Get the runtime version of the metatomic library as a string. /// /// This version follows the `..[-]` format. -#[no_mangle] +#[unsafe(no_mangle)] pub extern "C" fn mta_version() -> *const c_char { return VERSION.as_ptr(); } @@ -80,7 +80,7 @@ impl mta_string_t { /// @param string A pointer to a null-terminated C string. Must not be null. /// @return A new `mta_string_t` containing a copy of `string`, or null if an /// error occurred. You can check the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_string_create( string: *const c_char, ) -> mta_string_t { @@ -90,7 +90,7 @@ pub unsafe extern "C" fn mta_string_create( catch_unwind(move || { check_pointers_non_null!(string); - let cstr = std::ffi::CStr::from_ptr(string); + let cstr = unsafe { std::ffi::CStr::from_ptr(string) }; let string = CString::from(cstr); let ptr = CString::into_raw(string); @@ -106,7 +106,7 @@ pub unsafe extern "C" fn mta_string_create( /// Free a `mta_string_t` previously created by `mta_string_create`. /// /// @param string A `mta_string_t` to free. Can be null, in which case this function is a no-op. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_string_free(string: mta_string_t) { catch_unwind(|| { if string.0.is_null() { @@ -114,7 +114,7 @@ pub unsafe extern "C" fn mta_string_free(string: mta_string_t) { } let ptr = string.0.cast::(); - let cstring = CString::from_raw(ptr); + let cstring = unsafe { CString::from_raw(ptr) }; std::mem::drop(cstring); Ok(()) @@ -127,7 +127,7 @@ pub unsafe extern "C" fn mta_string_free(string: mta_string_t) { /// /// @param string A `mta_string_t` containing the string to view. Must not be null. /// @return A pointer to the null-terminated C string inside `string` -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_string_view( string: mta_string_t, ) -> *const c_char { @@ -165,7 +165,7 @@ pub unsafe extern "C" fn mta_string_view( /// @param conversion A pointer to a `double` where the conversion factor will be stored. /// @return The status code of the operation. If this code is not `MTA_SUCCESS`, /// you can get more details about the error with `mta_last_error`. -#[no_mangle] +#[unsafe(no_mangle)] pub unsafe extern "C" fn mta_unit_conversion_factor( from_unit: *const c_char, to_unit: *const c_char, @@ -174,8 +174,8 @@ pub unsafe extern "C" fn mta_unit_conversion_factor( catch_unwind(|| { check_pointers_non_null!(from_unit, to_unit, conversion); - let from_cstr = std::ffi::CStr::from_ptr(from_unit); - let to_cstr = std::ffi::CStr::from_ptr(to_unit); + let from_cstr = unsafe { std::ffi::CStr::from_ptr(from_unit) }; + let to_cstr = unsafe { std::ffi::CStr::from_ptr(to_unit) }; let from_str = from_cstr.to_str().map_err(|_| { Error::InvalidParameter("from_unit is not valid UTF-8".into()) @@ -184,7 +184,9 @@ pub unsafe extern "C" fn mta_unit_conversion_factor( Error::InvalidParameter("to_unit is not valid UTF-8".into()) })?; - *conversion = crate::unit_conversion_factor(from_str, to_str)?; + unsafe { + *conversion = crate::unit_conversion_factor(from_str, to_str)?; + } Ok(()) }) diff --git a/metatomic-core/src/model.rs b/metatomic-core/src/model.rs index a9ab74b81..cfe7d09d5 100644 --- a/metatomic-core/src/model.rs +++ b/metatomic-core/src/model.rs @@ -175,14 +175,16 @@ mod tests { _data: *const c_void, out: *mut mta_string_t, ) -> mta_status_t { - *out = mta_string_t::new(r#"{ - "type": "metatomic_model_metadata", - "name": "test-model", - "authors": ["Alice"], - "description": "A test model", - "references": {"model": [], "architecture": [], "implementation": []}, - "extra": {} - }"#); + unsafe { + *out = mta_string_t::new(r#"{ + "type": "metatomic_model_metadata", + "name": "test-model", + "authors": ["Alice"], + "description": "A test model", + "references": {"model": [], "architecture": [], "implementation": []}, + "extra": {} + }"#); + } return mta_status_t::MTA_SUCCESS; } @@ -190,15 +192,23 @@ mod tests { _data: *const c_void, out: *mut mta_string_t, ) -> mta_status_t { - *out = mta_string_t::new(r#"{ - "type": "metatomic_model_capabilities", - "outputs": [{"type": "metatomic_quantity", "name": "energy", "unit": "eV", "gradients": [], "sample_kind": "system"}], - "atomic_types": [1, 6], - "interaction_range": 5.0, - "length_unit": "Angstrom", - "supported_devices": ["cpu"], - "dtype": "float32" - }"#); + unsafe { + *out = mta_string_t::new(r#"{ + "type": "metatomic_model_capabilities", + "outputs": [{ + "type": "metatomic_quantity", + "name": "energy", + "unit": "eV", + "gradients": [], + "sample_kind": "system" + }], + "atomic_types": [1, 6], + "interaction_range": 5.0, + "length_unit": "Angstrom", + "supported_devices": ["cpu"], + "dtype": "float32" + }"#); + } return mta_status_t::MTA_SUCCESS; } @@ -206,12 +216,14 @@ mod tests { _data: *const c_void, out: *mut mta_string_t, ) -> mta_status_t { - *out = mta_string_t::new(format!(r#"[{{ - "type": "metatomic_pair_options", - "cutoff": "{:#x}", - "full_list": true, - "strict": true - }}]"#, 3.5_f64.to_bits())); + unsafe { + *out = mta_string_t::new(format!(r#"[{{ + "type": "metatomic_pair_options", + "cutoff": "{:#x}", + "full_list": true, + "strict": true + }}]"#, 3.5_f64.to_bits())); + } return mta_status_t::MTA_SUCCESS; } @@ -219,13 +231,15 @@ mod tests { _data: *const c_void, out: *mut mta_string_t, ) -> mta_status_t { - *out = mta_string_t::new(r#"[{ - "type": "metatomic_quantity", - "name": "charge", - "unit": "e", - "gradients": [], - "sample_kind": "atom" - }]"#); + unsafe { + *out = mta_string_t::new(r#"[{ + "type": "metatomic_quantity", + "name": "charge", + "unit": "e", + "gradients": [], + "sample_kind": "atom" + }]"#); + } return mta_status_t::MTA_SUCCESS; } @@ -233,21 +247,24 @@ mod tests { _data: *const c_void, out: *mut mta_string_t, ) -> mta_status_t { - *out = mta_string_t::new(r#"[ - { - "type": "metatomic_quantity", - "name": "energy", - "unit": "eV", - "gradients": ["positions"], - "sample_kind": "system" - }, - { - "type": "metatomic_quantity", - "name": "custom::output", - "unit": "", - "gradients": [], - "sample_kind": "atom_pair" - }]"#); + unsafe { + *out = mta_string_t::new(r#"[ + { + "type": "metatomic_quantity", + "name": "energy", + "unit": "eV", + "gradients": ["positions"], + "sample_kind": "system" + }, + { + "type": "metatomic_quantity", + "name": "custom::output", + "unit": "", + "gradients": [], + "sample_kind": "atom_pair" + }]"# + ); + } return mta_status_t::MTA_SUCCESS; } diff --git a/metatomic-core/src/plugin.rs b/metatomic-core/src/plugin.rs index cd087965c..af8a9bbe8 100644 --- a/metatomic-core/src/plugin.rs +++ b/metatomic-core/src/plugin.rs @@ -1,8 +1,7 @@ use std::ffi::CStr; -use std::sync::Mutex; +use std::sync::{Mutex, LazyLock}; use libloading::Library; -use once_cell::sync::Lazy; use crate::c_api::{mta_model_t, mta_plugin_t, mta_register_plugin, mta_status_t}; use crate::{Error, Model}; @@ -15,10 +14,10 @@ use crate::{Error, Model}; pub const MTA_ABI_VERSION: i32 = 1; /// The list of registered plugins in the current process. -static PLUGINS: Lazy>> = Lazy::new(|| Mutex::new(Vec::new())); +static PLUGINS: LazyLock>> = LazyLock::new(|| Mutex::new(Vec::new())); /// Keep the loaded libraries alive for the entire process lifetime, to ensure /// that the plugin code is not unloaded while it's still in use. -static LIBRARIES: Lazy>> = Lazy::new(|| Mutex::new(Vec::new())); +static LIBRARIES: LazyLock>> = LazyLock::new(|| Mutex::new(Vec::new())); pub struct Plugin(mta_plugin_t); diff --git a/metatomic-core/src/quantities.rs b/metatomic-core/src/quantities.rs index cc313d555..85b726ccf 100644 --- a/metatomic-core/src/quantities.rs +++ b/metatomic-core/src/quantities.rs @@ -58,13 +58,12 @@ pub(crate) fn validate_quantity_name(name: &str) -> Result<(), Error> { ))); } - if let Some(variant) = variant { - if !is_valid_identifier(variant) { - return Err(Error::InvalidParameter(format!( - "invalid quantity variant '{}' in '{}': must be a valid identifier (alphanumeric or underscore, not starting with a digit)", - variant, name - ))); - } + if let Some(variant) = variant && !is_valid_identifier(variant) { + return Err(Error::InvalidParameter(format!( + "invalid quantity variant '{}' in '{}': must be a valid identifier \ + (alphanumeric or underscore, not starting with a digit)", + variant, name + ))); } if STANDARD_QUANTITIES.contains(&main_part) { @@ -75,7 +74,8 @@ pub(crate) fn validate_quantity_name(name: &str) -> Result<(), Error> { for component in &components { if !is_valid_identifier(component) { return Err(Error::InvalidParameter(format!( - "invalid quantity name component '{}' in '{}': must be a valid identifier (alphanumeric or underscore, not starting with a digit)", + "invalid quantity name component '{}' in '{}': must be a valid \ + identifier (alphanumeric or underscore, not starting with a digit)", component, name ))); } diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs index 578333d52..182c05d2e 100644 --- a/metatomic-core/src/system.rs +++ b/metatomic-core/src/system.rs @@ -1,5 +1,5 @@ use std::collections::{BTreeMap, HashMap, HashSet}; -use once_cell::sync::Lazy; +use std::sync::LazyLock; use dlpk::sys::{DLDataType, DLDevice, DLDeviceType}; use dlpk::{DLPackTensor, DLPackTensorRef}; @@ -8,7 +8,7 @@ use metatensor::{TensorBlock, TensorMap}; use crate::{Error, PairListOptions}; /// Names that can never be used as custom data in a system -static INVALID_DATA_NAMES: Lazy> = Lazy::new(|| { +static INVALID_DATA_NAMES: LazyLock> = LazyLock::new(|| { HashSet::from(["types", "type", "positions", "position", "cell", "neighbors", "neighbor", "pair", "pairs"]) }); diff --git a/metatomic-core/src/units.rs b/metatomic-core/src/units.rs index 4239b86be..4d5a908e8 100644 --- a/metatomic-core/src/units.rs +++ b/metatomic-core/src/units.rs @@ -1,6 +1,6 @@ use crate::Error; -use once_cell::sync::Lazy; +use std::sync::LazyLock; use std::collections::HashMap; use std::fmt; use std::ops::{Add, Sub}; @@ -139,7 +139,7 @@ struct UnitValue { /// All base units with SI factors and dimensions. /// Factors are expressed in SI base units (m, s, kg, C, K). /// Case-insensitive lookup: names are lowercased before searching. -static BASE_UNITS: Lazy> = Lazy::new(|| { +static BASE_UNITS: LazyLock> = LazyLock::new(|| { let mut map = HashMap::new(); // --- Temperature --- diff --git a/metatomic-torch/Cargo.toml b/metatomic-torch/Cargo.toml index 82672a105..6387cd4db 100644 --- a/metatomic-torch/Cargo.toml +++ b/metatomic-torch/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "metatomic-torch" version = "0.0.0" -edition = "2021" +edition = "2024" publish = false rust-version = "1.88" diff --git a/python/Cargo.toml b/python/Cargo.toml index 88607618f..3546f0179 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "metatomic-python" version = "0.0.0" -edition = "2021" +edition = "2024" publish = false rust-version = "1.88" From 2985872a3c8db0e3e92ad198e3ddaa03a25f4664 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Wed, 1 Jul 2026 14:51:20 +0200 Subject: [PATCH 31/43] Implement serialization for Systems Co-Authored-By: frostedoyster --- docs/src/core/reference/c/system.rst | 35 +- metatomic-core/Cargo.toml | 2 + metatomic-core/include/metatomic.h | 76 +++ metatomic-core/src/c_api/io.rs | 317 +++++++++++++ metatomic-core/src/c_api/mod.rs | 2 + metatomic-core/src/io/mod.rs | 34 ++ metatomic-core/src/io/npy_header.rs | 673 +++++++++++++++++++++++++++ metatomic-core/src/io/system.rs | 296 ++++++++++++ metatomic-core/src/io/tensor.rs | 333 +++++++++++++ metatomic-core/src/lib.rs | 19 +- metatomic-core/src/metadata.rs | 39 +- metatomic-core/src/system.rs | 59 ++- metatomic-core/tests/data/legacy.mta | Bin 0 -> 6060 bytes scripts/include/metatensor.h | 3 + 14 files changed, 1847 insertions(+), 41 deletions(-) create mode 100644 metatomic-core/src/c_api/io.rs create mode 100644 metatomic-core/src/io/mod.rs create mode 100644 metatomic-core/src/io/npy_header.rs create mode 100644 metatomic-core/src/io/system.rs create mode 100644 metatomic-core/src/io/tensor.rs create mode 100644 metatomic-core/tests/data/legacy.mta diff --git a/docs/src/core/reference/c/system.rst b/docs/src/core/reference/c/system.rst index 155245253..69895f256 100644 --- a/docs/src/core/reference/c/system.rst +++ b/docs/src/core/reference/c/system.rst @@ -5,17 +5,22 @@ System The following functions operate on :c:type:`mta_system_t`: -- :c:func:`mta_system_create`: TODO summary -- :c:func:`mta_system_free`: TODO summary -- :c:func:`mta_system_size`: TODO summary -- :c:func:`mta_system_get_data`: TODO summary -- :c:func:`mta_system_get_length_unit`: TODO summary -- :c:func:`mta_system_add_pairs`: TODO summary -- :c:func:`mta_system_get_pairs`: TODO summary -- :c:func:`mta_system_known_pairs`: TODO summary -- :c:func:`mta_system_add_custom_data`: TODO summary -- :c:func:`mta_system_get_custom_data`: TODO summary -- :c:func:`mta_system_known_custom_data`: TODO summary +- :c:func:`mta_system_create`: create a new system from types, positions, cell, and PBC data +- :c:func:`mta_system_free`: free a system handle +- :c:func:`mta_system_size`: get the number of atoms in a system +- :c:func:`mta_system_get_data`: get a borrowed DLPack tensor for some system data +- :c:func:`mta_system_get_length_unit`: get the length unit of a system +- :c:func:`mta_system_add_pairs`: add a pair list to a system +- :c:func:`mta_system_get_pairs`: get a borrowed view of a pair list from a system +- :c:func:`mta_system_known_pairs`: get all pair list options known by a system +- :c:func:`mta_system_add_custom_data`: add custom data to a system +- :c:func:`mta_system_get_custom_data`: get a borrowed view of custom data by name +- :c:func:`mta_system_known_custom_data`: get all custom data names known by a system + +- :c:func:`mta_save`: save a system to a file +- :c:func:`mta_save_buffer`: save a system to a buffer +- :c:func:`mta_load`: load a system from a file +- :c:func:`mta_load_buffer`: load a system from a buffer -------------------------------------------------------------------------------- @@ -40,3 +45,11 @@ The following functions operate on :c:type:`mta_system_t`: .. doxygenfunction:: mta_system_get_custom_data .. doxygenfunction:: mta_system_known_custom_data + +.. doxygenfunction:: mta_save + +.. doxygenfunction:: mta_save_buffer + +.. doxygenfunction:: mta_load + +.. doxygenfunction:: mta_load_buffer diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index 7cbc6d23f..934b6fe2b 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -19,6 +19,8 @@ dlpk = { version = "0.4", features = ["ndarray"]} json = "0.12" libloading = "0.9" ndarray = "0.17" +zip = { version = "8.6.0", default-features = false } +byteorder = {version = "1"} [build-dependencies] diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index 397886926..430d7b9f1 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -700,6 +700,82 @@ enum mta_status_t mta_load_model(const char *load_from, const char *plugin_name, struct mta_model_t *model); +/** + * Save a system to a file. + * + * The format consists of a zip archive containing NPY files for the system's + * data (types, positions, cell, pbc), a `info.json` file for metadata, and + * optional sub-directories for pair lists (`pairs//options.json` and + * `pairs//data.mts`) and custom data (`data/.mts`). + * + * @param path A null-terminated C string containing the file path. Must not be + * null. + * @param system The system to save. Must not be null. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. + */ +enum mta_status_t mta_save(const char *path, const struct mta_system_t *system); + +/** + * Save a system to an in-memory buffer. + * + * The buffer is grown as needed using the provided `realloc` callback. On + * success, `*buffer` points to the serialized data and `*buffer_count` + * contains the number of bytes written. + * + * @param buffer Pointer to the buffer pointer. On input, `*buffer` may be NULL + * (in which case `*buffer_count` must be 0). On output, `*buffer` is + * updated to point to the serialized data. + * @param buffer_count Pointer to the buffer size. On input, `*buffer_count` + * must contain the current allocation size. On output, it is set to the + * number of bytes written. + * @param realloc_user_data User data passed as the first argument to + * `realloc`. + * @param realloc Callback to grow the buffer. Must not be NULL. + * @param system The system to save. Must not be null. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. + */ +enum mta_status_t mta_save_buffer(uint8_t **buffer, + uintptr_t *buffer_count, + void *realloc_user_data, + mts_realloc_buffer_t realloc, + const struct mta_system_t *system); + +/** + * Load a system from a file. + * + * The file must have been written by `mta_save` and contain a valid metatomic + * system. + * + * @param path A null-terminated C string containing the file path. Must not be + * null. + * @param create_array Callback to allocate arrays for the system's data. Must + * not be NULL. + * @return A pointer to the newly allocated system. The caller takes ownership + * and must free it with `mta_system_free`. Returns NULL on error; use + * `mta_last_error` for details. + */ +struct mta_system_t *mta_load(const char *path, mts_create_array_callback_t create_array); + +/** + * Load a system from an in-memory buffer. + * + * The buffer must contain data serialized by `mta_save_buffer` (or the + * equivalent Rust function). + * + * @param buffer Pointer to the serialized data. Must not be NULL. + * @param buffer_size Number of bytes in `buffer`. + * @param create_array Callback to allocate arrays for the system's data. Must + * not be NULL. + * @return A pointer to the newly allocated system. The caller takes ownership + * and must free it with `mta_system_free`. Returns NULL on error; use + * `mta_last_error` for details. + */ +struct mta_system_t *mta_load_buffer(const uint8_t *buffer, + uintptr_t buffer_size, + mts_create_array_callback_t create_array); + #ifdef __cplusplus } // extern "C" #endif // __cplusplus diff --git a/metatomic-core/src/c_api/io.rs b/metatomic-core/src/c_api/io.rs new file mode 100644 index 000000000..ef47c115b --- /dev/null +++ b/metatomic-core/src/c_api/io.rs @@ -0,0 +1,317 @@ +use std::ffi::{c_char, c_void, CStr}; +use std::fs::File; +use std::io::{BufReader, Cursor}; +use std::sync::Arc; + +use metatensor::c_api::{mts_create_array_callback_t, mts_realloc_buffer_t}; + +use super::{catch_unwind, mta_status_t, mta_system_t}; +use crate::Error; + +/// Wrapper for an externally managed buffer, that can be grown to fit more data +struct ExternalBuffer { + data: *mut *mut u8, + writen: u64, + allocated: u64, + + realloc_user_data: *mut c_void, + realloc: unsafe extern "C" fn(*mut c_void, *mut u8, usize) -> *mut u8, + + current: u64, +} + +impl std::io::Write for ExternalBuffer { + #[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)] + fn write(&mut self, buf: &[u8]) -> std::io::Result { + let remaining_space = self.allocated.saturating_sub(self.current); + + if (remaining_space as usize) < buf.len() { + let required_size = self.current.saturating_add(buf.len() as u64); + let mut new_size = if self.allocated == 0 { 1024 } else { self.allocated }; + while new_size < required_size { + new_size = new_size.saturating_mul(2); + if new_size == 0 { + return Err(std::io::Error::new( + std::io::ErrorKind::OutOfMemory, + "requested allocation size overflow", + )); + } + } + + let new_ptr = unsafe { + (self.realloc)(self.realloc_user_data, *self.data, new_size as usize) + }; + + if new_ptr.is_null() { + return Err(std::io::Error::new( + std::io::ErrorKind::OutOfMemory, + "failed to allocate memory with the realloc callback" + )); + } + + unsafe { + *self.data = new_ptr; + } + + self.allocated = new_size; + } + + let mut output = unsafe { + let start = (*self.data).offset(self.current as isize); + // allocated >= current + buf.len() + std::slice::from_raw_parts_mut(start, buf.len()) + }; + + let count = output.write(buf).expect("failed to write to pre-allocated slice"); + assert_eq!(count, buf.len()); + self.current += count as u64; + + if self.current > self.writen { + self.writen = self.current; + } + return Ok(count); + } + + fn flush(&mut self) -> std::io::Result<()> { + return Ok(()); + } +} + + +#[allow(clippy::cast_sign_loss, clippy::cast_possible_wrap)] +impl std::io::Seek for ExternalBuffer { + fn seek(&mut self, pos: std::io::SeekFrom) -> std::io::Result { + match pos { + std::io::SeekFrom::Start(offset) => { + if offset > self.writen { + return Err(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, "tried to seek past the end of the buffer") + ); + } + + self.current = offset; + }, + + std::io::SeekFrom::End(offset) => { + if offset > 0 { + return Err(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, "tried to seek past the end of the buffer") + ); + } + + if -offset > self.writen as i64 { + return Err(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, "tried to seek past the beginning of the buffer") + ); + } + + self.current = (self.writen as i64 + offset) as u64; + }, + + std::io::SeekFrom::Current(offset) => { + let result = self.current as i64 + offset; + if result > self.writen as i64 { + return Err(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, "tried to seek past the end of the buffer") + ); + } + + if result < 0 { + return Err(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, "tried to seek past the beginning of the buffer") + ); + } + + self.current = result as u64; + }, + } + + return Ok(self.current); + } + + fn rewind(&mut self) -> std::io::Result<()> { + self.current = 0; + return Ok(()); + } + + fn stream_position(&mut self) -> std::io::Result { + return Ok(self.current); + } +} + + +/// Save a system to a file. +/// +/// The format consists of a zip archive containing NPY files for the system's +/// data (types, positions, cell, pbc), a `info.json` file for metadata, and +/// optional sub-directories for pair lists (`pairs//options.json` and +/// `pairs//data.mts`) and custom data (`data/.mts`). +/// +/// @param path A null-terminated C string containing the file path. Must not be +/// null. +/// @param system The system to save. Must not be null. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn mta_save( + path: *const c_char, + system: *const mta_system_t +) -> mta_status_t { + catch_unwind(|| { + check_pointers_non_null!(path, system); + + let path = unsafe { CStr::from_ptr(path) }.to_str() + .map_err(|_| Error::InvalidParameter("path is not valid UTF-8".into()))?; + + let file = File::create(path)?; + let system = unsafe { &*system }; + crate::io::save(file, &system.0)?; + + Ok(()) + }) +} + +/// Save a system to an in-memory buffer. +/// +/// The buffer is grown as needed using the provided `realloc` callback. On +/// success, `*buffer` points to the serialized data and `*buffer_count` +/// contains the number of bytes written. +/// +/// @param buffer Pointer to the buffer pointer. On input, `*buffer` may be NULL +/// (in which case `*buffer_count` must be 0). On output, `*buffer` is +/// updated to point to the serialized data. +/// @param buffer_count Pointer to the buffer size. On input, `*buffer_count` +/// must contain the current allocation size. On output, it is set to the +/// number of bytes written. +/// @param realloc_user_data User data passed as the first argument to +/// `realloc`. +/// @param realloc Callback to grow the buffer. Must not be NULL. +/// @param system The system to save. Must not be null. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. +#[unsafe(no_mangle)] +#[allow(clippy::cast_possible_truncation)] +pub unsafe extern "C" fn mta_save_buffer( + buffer: *mut *mut u8, + buffer_count: *mut usize, + realloc_user_data: *mut c_void, + realloc: mts_realloc_buffer_t, + system: *const mta_system_t, +) -> mta_status_t { + catch_unwind(|| { + check_pointers_non_null!(buffer, buffer_count, system); + + let realloc = if let Some(realloc) = realloc { + realloc + } else { + return Err(Error::InvalidParameter( + "realloc callback can not be NULL in mta_save_buffer".into() + )); + }; + + if unsafe { (*buffer).is_null() } { + // `ExternalBuffer.write` calls realloc with the current `*buffer` + // (which may be null) for the initial allocation. + unsafe { *buffer = std::ptr::null_mut(); } + } + + let system = unsafe { &*system }; + let mut external_buffer = ExternalBuffer { + data: buffer, + allocated: unsafe { *buffer_count } as u64, + writen: 0, + realloc_user_data, + realloc, + current: 0, + }; + + crate::io::save(&mut external_buffer, &system.0)?; + + unsafe { + *buffer_count = external_buffer.current as usize; + } + + Ok(()) + }) +} + +/// Load a system from a file. +/// +/// The file must have been written by `mta_save` and contain a valid metatomic +/// system. +/// +/// @param path A null-terminated C string containing the file path. Must not be +/// null. +/// @param create_array Callback to allocate arrays for the system's data. Must +/// not be NULL. +/// @return A pointer to the newly allocated system. The caller takes ownership +/// and must free it with `mta_system_free`. Returns NULL on error; use +/// `mta_last_error` for details. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn mta_load( + path: *const c_char, + create_array: mts_create_array_callback_t +) -> *mut mta_system_t { + let mut result = std::ptr::null_mut(); + let unwind_wrapper = std::panic::AssertUnwindSafe(&mut result); + let status = catch_unwind(move || { + check_pointers_non_null!(path); + + let path = unsafe { CStr::from_ptr(path) }.to_str() + .map_err(|_| Error::InvalidParameter("path is not valid UTF-8".into()))?; + + let file = BufReader::new(File::open(path)?); + let system = crate::io::load(file, create_array)?; + + let _ = &unwind_wrapper; + *unwind_wrapper.0 = Box::into_raw(Box::new(mta_system_t(Arc::new(system)))); + Ok(()) + }); + + if status != mta_status_t::MTA_SUCCESS { + return std::ptr::null_mut(); + } + + return result; +} + +/// Load a system from an in-memory buffer. +/// +/// The buffer must contain data serialized by `mta_save_buffer` (or the +/// equivalent Rust function). +/// +/// @param buffer Pointer to the serialized data. Must not be NULL. +/// @param buffer_size Number of bytes in `buffer`. +/// @param create_array Callback to allocate arrays for the system's data. Must +/// not be NULL. +/// @return A pointer to the newly allocated system. The caller takes ownership +/// and must free it with `mta_system_free`. Returns NULL on error; use +/// `mta_last_error` for details. +#[unsafe(no_mangle)] +pub unsafe extern "C" fn mta_load_buffer( + buffer: *const u8, + buffer_size: usize, + create_array: mts_create_array_callback_t +) -> *mut mta_system_t { + let mut result = std::ptr::null_mut(); + let unwind_wrapper = std::panic::AssertUnwindSafe(&mut result); + let status = catch_unwind(move || { + check_pointers_non_null!(buffer); + + let slice = unsafe { + std::slice::from_raw_parts(buffer, buffer_size) + }; + let cursor = Cursor::new(slice); + let system = crate::io::load(cursor, create_array)?; + + let _ = &unwind_wrapper; + *unwind_wrapper.0 = Box::into_raw(Box::new(mta_system_t(Arc::new(system)))); + Ok(()) + }); + + if status != mta_status_t::MTA_SUCCESS { + return std::ptr::null_mut(); + } + + return result; +} diff --git a/metatomic-core/src/c_api/mod.rs b/metatomic-core/src/c_api/mod.rs index 235b00296..dc705f81a 100644 --- a/metatomic-core/src/c_api/mod.rs +++ b/metatomic-core/src/c_api/mod.rs @@ -16,3 +16,5 @@ pub use self::model::mta_model_t; mod plugin; pub use self::plugin::{mta_plugin_t, mta_register_plugin, mta_load_plugin, mta_load_model}; + +mod io; diff --git a/metatomic-core/src/io/mod.rs b/metatomic-core/src/io/mod.rs new file mode 100644 index 000000000..b0cfa2530 --- /dev/null +++ b/metatomic-core/src/io/mod.rs @@ -0,0 +1,34 @@ +use crate::Error; + +mod npy_header; + +mod tensor; +mod system; + +pub use system::{load, save}; + +pub trait ReadAndSeek: std::io::Read + std::io::Seek {} +impl ReadAndSeek for T {} + +pub enum PathOrBuffer<'a> { + Path(&'a str), + Buffer(&'a mut dyn ReadAndSeek), +} + +/// Byte order for multi-byte values in NPY files. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum Endianness { + Little, + Big, + Native, +} + +// returns an error if the given reader contains any more data +fn check_for_extra_bytes(reader: &mut R) -> Result<(), Error> { + let extra = reader.read_to_end(&mut Vec::new())?; + if extra == 0 { + Ok(()) + } else { + Err(Error::Serialization(format!("found {} extra bytes after the expected end of data", extra))) + } +} diff --git a/metatomic-core/src/io/npy_header.rs b/metatomic-core/src/io/npy_header.rs new file mode 100644 index 000000000..3b789aafe --- /dev/null +++ b/metatomic-core/src/io/npy_header.rs @@ -0,0 +1,673 @@ +// This file was initially taken from https://github.com/jturner314/ndarray-npy, +// version 0.8.1. It is Copyright 2018–2021 Jim Turner and ndarray-npy +// developers, released under MIT and Apache Licenses. +use std::convert::TryFrom; +use std::sync::Arc; +use std::error::Error; +use std::fmt::Write as FmtWrite; +use std::io::Write as IoWrite; + +use byteorder::{ByteOrder, LittleEndian, ReadBytesExt}; + +/// Magic string to indicate npy format. +const MAGIC_STRING: &[u8] = b"\x93NUMPY"; + +/// The total header length (including magic string, version number, header +/// length value, array format description, padding, and final newline) must be +/// evenly divisible by this value. +// If this changes, update the docs of `ViewNpyExt` and `ViewMutNpyExt`. +const HEADER_DIVISOR: usize = 64; + +#[derive(Debug)] +pub enum ParseHeaderError { + MagicString, + Version { + major: u8, + minor: u8, + }, + /// Indicates that the `HEADER_LEN` doesn't fit in `usize`. + HeaderLengthOverflow(u32), + /// Indicates that the array format string contains non-ASCII characters. + /// This is an error for .npy format versions 1.0 and 2.0. + NonAscii, + /// Error parsing the array format string as UTF-8. This does not apply to + /// .npy format versions 1.0 and 2.0, which require the array format string + /// to be ASCII. + Utf8Parse(std::str::Utf8Error), + InvalidHeader(String), +} + +impl Error for ParseHeaderError { + fn source(&self) -> Option<&(dyn Error + 'static)> { + match self { + ParseHeaderError::Utf8Parse(err) => Some(err), + ParseHeaderError::MagicString | + ParseHeaderError::Version { .. } | + ParseHeaderError::HeaderLengthOverflow(_) | + ParseHeaderError::NonAscii | + ParseHeaderError::InvalidHeader(_) => None, + } + } +} + +impl std::fmt::Display for ParseHeaderError { + fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + match self { + ParseHeaderError::MagicString => write!(f, "start does not match magic string"), + ParseHeaderError::Version { major, minor } => write!(f, "unknown version number: {}.{}", major, minor), + ParseHeaderError::HeaderLengthOverflow(header_len) => write!(f, "HEADER_LEN {} does not fit in `usize`", header_len), + ParseHeaderError::NonAscii => write!(f, "non-ascii in array format string; this is not supported in .npy format versions 1.0 and 2.0"), + ParseHeaderError::Utf8Parse(err) => write!(f, "error parsing array format string as UTF-8: {}", err), + ParseHeaderError::InvalidHeader(value) => write!(f, "invalid header in file: {}", value), + } + } +} + +impl From for ParseHeaderError { + fn from(err: std::str::Utf8Error) -> ParseHeaderError { + ParseHeaderError::Utf8Parse(err) + } +} + +impl From for ParseHeaderError { + fn from(e: std::num::ParseIntError) -> Self { + ParseHeaderError::InvalidHeader(format!("failed to parse an integer: {}", e)) + } +} + +#[derive(Debug)] +pub enum ReadHeaderError { + Io(std::io::Error), + Parse(ParseHeaderError), +} + +impl Error for ReadHeaderError { + fn source(&self) -> Option<&(dyn Error + 'static)> { + match self { + ReadHeaderError::Io(err) => Some(err), + ReadHeaderError::Parse(err) => Some(err), + } + } +} + +impl std::fmt::Display for ReadHeaderError { + fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + match self { + ReadHeaderError::Io(err) => write!(f, "I/O error: {}", err), + ReadHeaderError::Parse(err) => write!(f, "error parsing header: {}", err), + } + } +} + +impl From for ReadHeaderError { + fn from(err: std::io::Error) -> ReadHeaderError { + ReadHeaderError::Io(err) + } +} + +impl From for ReadHeaderError { + fn from(err: ParseHeaderError) -> ReadHeaderError { + ReadHeaderError::Parse(err) + } +} + +#[derive(Clone, Copy)] +#[allow(non_camel_case_types)] +enum Version { + V1_0, + V2_0, + V3_0, +} + +impl Version { + /// Number of bytes taken up by version number (1 byte for major version, 1 + /// byte for minor version). + const VERSION_NUM_BYTES: usize = 2; + + fn from_bytes(bytes: &[u8]) -> Result { + debug_assert_eq!(bytes.len(), Self::VERSION_NUM_BYTES); + match (bytes[0], bytes[1]) { + (0x01, 0x00) => Ok(Version::V1_0), + (0x02, 0x00) => Ok(Version::V2_0), + (0x03, 0x00) => Ok(Version::V3_0), + (major, minor) => Err(ParseHeaderError::Version { major, minor }), + } + } + + /// Major version number. + fn major_version(self) -> u8 { + match self { + Version::V1_0 => 1, + Version::V2_0 => 2, + Version::V3_0 => 3, + } + } + + /// Major version number. + fn minor_version(self) -> u8 { + match self { + Version::V1_0 | Version::V2_0 | Version::V3_0 => 0, + } + } + + /// Number of bytes in representation of header length. + fn header_len_num_bytes(self) -> usize { + match self { + Version::V1_0 => 2, + Version::V2_0 | Version::V3_0 => 4, + } + } + + /// Read header length. + fn read_header_len(self, reader: &mut R) -> Result { + match self { + Version::V1_0 => Ok(usize::from(reader.read_u16::()?)), + Version::V2_0 | Version::V3_0 => { + let header_len: u32 = reader.read_u32::()?; + Ok(usize::try_from(header_len) + .map_err(|_| ParseHeaderError::HeaderLengthOverflow(header_len))?) + } + } + } + + /// Format header length as bytes for writing to file. + /// + /// Returns `None` if the value of `header_len` is too large for this .npy version. + fn format_header_len(self, header_len: usize) -> Option> { + match self { + Version::V1_0 => { + let header_len: u16 = u16::try_from(header_len).ok()?; + let mut out = vec![0; self.header_len_num_bytes()]; + LittleEndian::write_u16(&mut out, header_len); + Some(out) + } + Version::V2_0 | Version::V3_0 => { + let header_len: u32 = u32::try_from(header_len).ok()?; + let mut out = vec![0; self.header_len_num_bytes()]; + LittleEndian::write_u32(&mut out, header_len); + Some(out) + } + } + } + + /// Computes the total header length, formatted `HEADER_LEN` value, and + /// padding length for this .npy version. + /// + /// `unpadded_arr_format` is the Python literal describing the array + /// format, formatted as an ASCII string without any padding. + /// + /// Returns `None` if the total header length overflows `usize` or if the + /// value of `HEADER_LEN` is too large for this .npy version. + fn compute_lengths(self, unpadded_arr_format: &[u8]) -> Option { + /// Length of a '\n' char in bytes. + const NEWLINE_LEN: usize = 1; + + let prefix_len: usize = + MAGIC_STRING.len() + Version::VERSION_NUM_BYTES + self.header_len_num_bytes(); + let unpadded_total_len: usize = prefix_len + .checked_add(unpadded_arr_format.len())? + .checked_add(NEWLINE_LEN)?; + let padding_len: usize = HEADER_DIVISOR - unpadded_total_len % HEADER_DIVISOR; + let total_len: usize = unpadded_total_len.checked_add(padding_len)?; + let header_len: usize = total_len - prefix_len; + let formatted_header_len = self.format_header_len(header_len)?; + Some(HeaderLengthInfo { + total_len, + formatted_header_len, + }) + } +} + +struct HeaderLengthInfo { + /// Total header length (including magic string, version number, header + /// length value, array format description, padding, and final newline). + total_len: usize, + /// Formatted `HEADER_LEN` value. (This is the number of bytes in the array + /// format description, padding, and final newline.) + formatted_header_len: Vec, +} + +#[derive(Debug)] +pub enum WriteHeaderError { + Io(std::io::Error), + Format(String), +} + +impl Error for WriteHeaderError { + fn source(&self) -> Option<&(dyn Error + 'static)> { + match self { + WriteHeaderError::Io(err) => Some(err), + WriteHeaderError::Format(_) => None, + } + } +} + +impl std::fmt::Display for WriteHeaderError { + fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + match self { + WriteHeaderError::Io(err) => write!(f, "I/O error: {}", err), + WriteHeaderError::Format(err) => write!(f, "error formatting header: {}", err), + } + } +} + +impl From for WriteHeaderError { + fn from(err: std::io::Error) -> WriteHeaderError { + WriteHeaderError::Io(err) + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum DataType { + Scalar(String), + Compound(Vec<(String, String)>), +} + +impl std::fmt::Display for DataType { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + DataType::Scalar(v) => write!(f, "'{}'", v), + DataType::Compound(list) => { + write!(f, "[")?; + for (k, v) in list { + write!(f, "('{}', '{}'), ", k, v)?; + } + write!(f, "]") + } + } + } +} + +#[derive(Clone, Debug, PartialEq)] +pub struct Header { + pub type_descriptor: DataType, + pub fortran_order: bool, + pub shape: Vec, +} + +#[derive(Clone, Debug)] +struct HeaderParser { + data: Vec, + position: usize, +} + +impl HeaderParser { + fn done(&self) -> bool { + return self.position >= self.data.len(); + } + + fn current(&self) -> char { + return self.data[self.position]; + } + + fn advance(&mut self) -> char { + let value = self.current(); + self.position += 1; + return value; + } + + fn expects(&mut self, c: char) -> Result<(), ParseHeaderError> { + if self.current() == c { + self.advance(); + return Ok(()); + } else { + return Err(ParseHeaderError::InvalidHeader(format!( + "expected '{}', got '{}'", c, self.current() + ))); + } + } + + fn skip_whitespaces(&mut self) { + let mut c = self.current(); + while !self.done() && (c == ' ' || c == '\t' || c == '\x0C') { + self.advance(); + c = self.current(); + } + } + + fn parse_string(&mut self) -> Result { + let mut value = String::new(); + if self.current() == '\'' { + self.advance(); + while self.current() != '\'' { + value.push(self.advance()); + } + self.advance(); + + } else if self.current() == '"' { + self.advance(); + while self.current() != '"' { + value.push(self.advance()); + } + self.advance(); + } else { + return Err(ParseHeaderError::InvalidHeader(format!( + "expected a string, got '{}'", self.current() + ))); + } + + return Ok(value); + } + + fn parse_integer(&mut self) -> Result { + let mut value = String::new(); + loop { + if self.current().is_ascii_digit() { + value.push(self.advance()); + } else { + break; + } + } + + if value.is_empty() { + return Err(ParseHeaderError::InvalidHeader(format!( + "expected an integer, got '{}'", self.current() + ))); + } + + return Ok(value.parse()?); + } + + fn parse_data_type(&mut self) -> Result { + if self.current() == '\'' || self.current() == '"' { + let value = self.parse_string()?; + return Ok(DataType::Scalar(value)); + } else if self.current() == '[' { + self.advance(); + + let mut data_type = Vec::new(); + loop { + self.skip_whitespaces(); + self.expects('(')?; + self.skip_whitespaces(); + + let name = self.parse_string()?; + + self.skip_whitespaces(); + self.expects(',')?; + self.skip_whitespaces(); + + let value = self.parse_string()?; + + self.skip_whitespaces(); + self.expects(')')?; + self.skip_whitespaces(); + + data_type.push((name, value)); + + if self.current() == ',' { + self.advance(); + self.skip_whitespaces(); + } else { + self.expects(']')?; + break; + } + + if self.current() == ']' { + self.advance(); + break; + } + } + + return Ok(DataType::Compound(data_type)); + } else { + return Err(ParseHeaderError::InvalidHeader(format!( + "expected a string or a list, got '{}'", self.current() + ))); + } + } + + fn parse_bool(&mut self) -> Result { + if self.current() == 'T' { + self.advance(); + self.expects('r')?; + self.expects('u')?; + self.expects('e')?; + return Ok(true); + } else if self.current() == 'F' { + self.advance(); + self.expects('a')?; + self.expects('l')?; + self.expects('s')?; + self.expects('e')?; + return Ok(false); + } else { + return Err(ParseHeaderError::InvalidHeader(format!( + "expected a bool, got '{}'", self.current() + ))); + } + } + + fn parse_shape(&mut self) -> Result, ParseHeaderError> { + let mut shape = Vec::new(); + self.expects('(')?; + loop { + self.skip_whitespaces(); + shape.push(self.parse_integer()?); + self.skip_whitespaces(); + + + if self.current() == ',' { + self.advance(); + self.skip_whitespaces(); + } else { + self.expects(')')?; + break; + } + + if self.current() == ')' { + self.advance(); + break; + } + } + + return Ok(shape); + } + + fn parse(&mut self) -> Result { + let mut type_descriptor: Option = None; + let mut fortran_order: Option = None; + let mut shape: Option> = None; + + self.skip_whitespaces(); + self.expects('{')?; + self.skip_whitespaces(); + + loop { + let key = self.parse_string()?; + self.skip_whitespaces(); + self.expects(':')?; + self.skip_whitespaces(); + + if key == "descr" { + type_descriptor = Some(self.parse_data_type()?); + } else if key == "fortran_order" { + fortran_order = Some(self.parse_bool()?); + } else if key == "shape" { + shape = Some(self.parse_shape()?); + } else { + return Err(ParseHeaderError::InvalidHeader(format!( + "unknown key: '{}'", key + ))); + } + + self.skip_whitespaces(); + if self.current() == ',' { + self.advance(); + self.skip_whitespaces(); + } else { + self.expects('}')?; + break; + } + + if self.current() == '}' { + self.advance(); + break; + } + } + + match (type_descriptor, fortran_order, shape) { + (Some(type_descriptor), Some(fortran_order), Some(shape)) => Ok(Header { + type_descriptor, + fortran_order, + shape, + }), + (None, _, _) => Err(ParseHeaderError::InvalidHeader("missing 'descr' key".into())), + (_, None, _) => Err(ParseHeaderError::InvalidHeader("missing 'fortran_order' key".into())), + (_, _, None) => Err(ParseHeaderError::InvalidHeader("missing 'shape' key".into())), + } + } +} + +impl Header { + fn from_str(value: &str) -> Result { + let mut parser = HeaderParser { data: value.chars().collect(), position: 0 }; + return parser.parse(); + } + + pub fn from_reader(reader: &mut R) -> Result { + // Check for magic string. + let mut buf = vec![0; MAGIC_STRING.len()]; + reader.read_exact(&mut buf)?; + if buf != MAGIC_STRING { + return Err(ParseHeaderError::MagicString.into()); + } + + // Get version number. + let mut buf = [0; Version::VERSION_NUM_BYTES]; + reader.read_exact(&mut buf)?; + let version = Version::from_bytes(&buf)?; + + // Get `HEADER_LEN`. + let header_len = version.read_header_len(reader)?; + + // Parse the dictionary describing the array's format. + let mut buf = vec![0; header_len]; + reader.read_exact(&mut buf)?; + let without_newline = match buf.split_last() { + Some((&b'\n', rest)) => rest, + Some(_) | None => return Err(ParseHeaderError::InvalidHeader("missing new line".into()))?, + }; + let header_str = match version { + Version::V1_0 | Version::V2_0 => { + if without_newline.is_ascii() { + // ASCII strings are always valid UTF-8. + unsafe { std::str::from_utf8_unchecked(without_newline) } + } else { + return Err(ParseHeaderError::NonAscii.into()); + } + } + Version::V3_0 => { + std::str::from_utf8(without_newline).map_err(ParseHeaderError::from)? + } + }; + + Ok(Header::from_str(header_str)?) + } + + fn to_dict_literal(&self) -> String { + let mut result = String::new(); + write!(&mut result, "{{ 'descr': {}, ", self.type_descriptor).expect("failed to write"); + + let order = if self.fortran_order { + "True" + } else { + "False" + }; + write!(&mut result, "'fortran_order': {}, ", order).expect("failed to write"); + + write!(&mut result, "'shape': (").expect("failed to write"); + for s in &self.shape { + write!(&mut result, "{}, ", s).expect("failed to write"); + } + write!(&mut result, ") }}").expect("failed to write"); + return result; + } + + pub fn to_bytes(&self) -> Result, WriteHeaderError> { + // Metadata describing array's format as ASCII string. + let mut arr_format = Vec::new(); + + write!(&mut arr_format, "{}", self.to_dict_literal())?; + + // Determine appropriate version based on header length, and compute + // length information. + let (version, length_info) = [Version::V1_0, Version::V2_0] + .iter() + .find_map(|&version| Some((version, version.compute_lengths(&arr_format)?))) + .ok_or_else(|| WriteHeaderError::Format("header too long".into()))?; + + // Write the header. + let mut out = Vec::with_capacity(length_info.total_len); + out.extend_from_slice(MAGIC_STRING); + out.push(version.major_version()); + out.push(version.minor_version()); + out.extend_from_slice(&length_info.formatted_header_len); + out.extend_from_slice(&arr_format); + out.resize(length_info.total_len - 1, b' '); + out.push(b'\n'); + + // Verify the length of the header. + debug_assert_eq!(out.len(), length_info.total_len); + debug_assert_eq!(out.len() % HEADER_DIVISOR, 0); + + Ok(out) + } + + pub fn write(&self, mut writer: W) -> Result<(), WriteHeaderError> { + let bytes = self.to_bytes()?; + writer.write_all(&bytes)?; + Ok(()) + } +} + +/******************************************************************************/ + +impl From for crate::Error { + fn from(error: ReadHeaderError) -> Self { + match error { + ReadHeaderError::Io(e) => crate::Error::Io(Arc::new(e)), + ReadHeaderError::Parse(e) => crate::Error::Serialization(e.to_string()), + } + } +} + +impl From for crate::Error { + fn from(error: WriteHeaderError) -> Self { + match error { + WriteHeaderError::Io(e) => crate::Error::Io(Arc::new(e)), + WriteHeaderError::Format(e) => crate::Error::Serialization(e), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn npy_header_parsing() { + let header = " \t{'descr': [('a', '(reader: R, create_array: mts_create_array_callback_t) -> Result + where R: std::io::Read + std::io::Seek +{ + let mut archive = ZipArchive::new(reader).map_err(|e| ("", e))?; + + let mut length_unit = String::new(); + if let Some(index) = archive.index_for_name("info.json") { + let mut info_file = archive.by_index(index).map_err(|e| ("info.json", e))?; + let mut info_content = String::new(); + info_file.read_to_string(&mut info_content)?; + let info: JsonValue = json::parse(&info_content)?; + + + if info["format"].as_str() != Some("metatomic_system") { + return Err(Error::Serialization(format!( + "invalid format in info.json, expected 'metatomic_system', found {:?}", + info["format"] + ))); + } + + if info["version"].as_u8() != Some(1) { + return Err(Error::Serialization(format!( + "unsupported version in info.json, expected 1, found {:?}", + info["version"] + ))); + } + + if !info.has_key("length_unit") || !info["length_unit"].is_string() { + return Err(Error::Serialization( + "missing or invalid 'length_unit' field in info.json".into() + )); + } + length_unit = info["length_unit"].as_str().unwrap().to_string(); + } else { + // this is a legacy file from metatomic-torch + } + + let data_file = archive.by_name("types.npy").map_err(|e| ("types.npy", e))?; + let types = read_tensor(data_file, create_array)?; + + let data_file = archive.by_name("positions.npy").map_err(|e| ("positions.npy", e))?; + let position = read_tensor(data_file, create_array)?; + + let data_file = archive.by_name("cell.npy").map_err(|e| ("cell.npy", e))?; + let cell = read_tensor(data_file, create_array)?; + + let data_file = archive.by_name("pbc.npy").map_err(|e| ("pbc.npy", e))?; + let pbc = read_tensor(data_file, create_array)?; + + let mut system = System::new(length_unit, types, position, cell, pbc)?; + + let pairs_paths: Vec = archive.file_names() + .filter(|path| path.starts_with("pairs/") && path.ends_with("/options.json")) + .map(|path| path.to_string()) + .collect(); + + let mut buffer = Vec::new(); + for path in pairs_paths { + let options: PairListOptions = { + let mut options_file = archive.by_name(&path).map_err(|e| (&path, e))?; + let mut options_content = String::new(); + options_file.read_to_string(&mut options_content)?; + let options_json: &JsonValue = &json::parse(&options_content)?; + + options_json.try_into()? + }; + + let data_path = path.strip_suffix("/options.json").unwrap().to_string() + "/data.mts"; + let mut data_file = archive.by_name(&data_path).map_err(|e| (data_path, e))?; + + buffer.clear(); + data_file.read_to_end(&mut buffer)?; + + let pairs = metatensor::io::load_block_buffer_custom_array(&buffer, create_array)?; + + system.add_pairs(options, pairs)?; + } + + let data_paths: Vec = archive.file_names() + .filter(|path| path.starts_with("data/")) + .map(|path| path.to_string()) + .collect(); + + for path in data_paths { + let name = path.strip_prefix("data/").expect("data path should start with 'data/'") + .strip_suffix(".mts").expect("data path should end with '.mts'").to_string(); + + let mut data_file = archive.by_name(&path).map_err(|e| (&path, e))?; + + buffer.clear(); + data_file.read_to_end(&mut buffer)?; + + let data = metatensor::io::load_buffer_custom_array(&buffer, create_array)?; + + system.add_custom_data(name, data, /*override*/ true)?; + } + + return Ok(system); +} + +/// Save the given system to a file (or any other writer). +/// +/// The format consists of a zip archive containing NPY files for the system's +/// data (types, positions, cell, pbc), a `info.json` file for metadata, and +/// optional sub-directories for pair lists (`pairs//options.json` and +/// `pairs//data.mts`) and custom data (`data/.mts`). +/// +/// The recommended file extension is `.mta`. +pub fn save(writer: W, system: &System) -> Result<(), Error> { + let mut archive = ZipWriter::new(writer); + + let options = zip::write::FileOptions::<'_, ()>::default() + .with_alignment(16) + .compression_method(zip::CompressionMethod::Stored) + .large_file(true) + .last_modified_time(zip::DateTime::from_date_and_time(2000, 1, 1, 0, 0, 0).expect("invalid datetime")); + + archive.start_file("info.json", options).map_err(|e| ("info.json", e))?; + let info = json::object! { + "format": "metatomic_system", + "version": 1, + "length_unit": system.length_unit(), + }; + info.write(&mut archive)?; + + archive.start_file("types.npy", options).map_err(|e| ("types.npy", e))?; + write_tensor(&mut archive, system.types())?; + + archive.start_file("positions.npy", options).map_err(|e| ("positions.npy", e))?; + write_tensor(&mut archive, system.positions())?; + + archive.start_file("cell.npy", options).map_err(|e| ("cell.npy", e))?; + write_tensor(&mut archive, system.cell())?; + + archive.start_file("pbc.npy", options).map_err(|e| ("pbc.npy", e))?; + write_tensor(&mut archive, system.pbc())?; + + let mut buffer = Vec::new(); + for (i, &pairs_options) in system.known_pairs().iter().enumerate() { + let path = format!("pairs/{}/options.json", i); + archive.start_file(&path, options).map_err(|e| (path, e))?; + let json: JsonValue = pairs_options.clone().into(); + json.write(&mut archive)?; + + + let pairs_block = system.get_pairs(pairs_options).expect("pairs block should exist"); + buffer.clear(); + pairs_block.save_buffer(&mut buffer)?; + + let path = format!("pairs/{}/data.mts", i); + archive.start_file(&path, options).map_err(|e| (path, e))?; + archive.write_all(&buffer)?; + } + + for name in system.known_custom_data() { + let tensor = system.get_custom_data(name).expect("custom data should exist"); + buffer.clear(); + tensor.save_buffer(&mut buffer)?; + let path = format!("data/{}.mts", name); + archive.start_file(&path, options).map_err(|e| (path, e))?; + archive.write_all(&buffer)?; + } + + archive.finish().map_err(|e| ("", e))?; + + return Ok(()); +} + + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn load_legacy() { + let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("tests/data/legacy.mta"); + + let file = std::fs::File::open(&path).unwrap(); + let system = load(file, Some(metatensor::io::create_ndarray)).unwrap(); + + assert_eq!(system.length_unit(), ""); + + let types: ndarray::ArrayView1 = system.types().try_into().unwrap(); + let positions: ndarray::ArrayView2 = system.positions().try_into().unwrap(); + let cell: ndarray::ArrayView2 = system.cell().try_into().unwrap(); + let pbc: ndarray::ArrayView1 = system.pbc().try_into().unwrap(); + + assert_eq!(types, ndarray::arr1(&[1, 6, 7, 8])); + assert_eq!( + positions, + ndarray::arr2(&[[0.0, 0.0, 0.0], [1.0, 2.0, 3.0], [4.0, 5.0, 6.0], [7.0, 8.0, 9.0]]) + ); + assert_eq!( + cell, + ndarray::arr2(&[[6.0, 0.0, 0.0], [0.0, 4.3, 0.0], [0.0, 0.0, 0.0]]) + ); + assert_eq!(pbc, ndarray::arr1(&[true, true, false])); + + let options = PairListOptions { cutoff: 5.5, full_list: true, strict: true, requestors: vec![] }; + let pairs = system.get_pairs(&options).unwrap(); + assert_eq!(pairs.samples().names(), ["first_atom", "second_atom", "cell_shift_a", "cell_shift_b", "cell_shift_c"]); + assert_eq!(pairs.samples().count(), 28); + assert_eq!(pairs.values().shape().unwrap(), [28, 3, 1]); + + let options = system.known_pairs(); + assert_eq!(options.len(), 1); + // requestors are not used when looking up pairs, but are stored in the file + assert_eq!(options[0].requestors, ["some requestor", "another one with UTF8 Θµ"]); + + assert_eq!(system.known_custom_data(), vec!["custom::data"]); + let custom = system.get_custom_data("custom::data").unwrap(); + assert_eq!(custom.keys().count(), 2); + } + + #[test] + fn save_load_system() { + let system = crate::system::test_system(); + + let path = std::env::temp_dir().join(format!("system-{}.mta", std::process::id())); + { + let file = std::fs::File::create(&path).unwrap(); + save(file, &system).unwrap(); + } + + { + let file = std::fs::File::open(&path).unwrap(); + let mut archive = zip::ZipArchive::new(file).unwrap(); + assert!(archive.by_name("types.npy").is_ok()); + assert!(archive.by_name("positions.npy").is_ok()); + assert!(archive.by_name("cell.npy").is_ok()); + assert!(archive.by_name("pbc.npy").is_ok()); + assert!(archive.by_name("pairs/0/data.mts").is_ok()); + assert!(archive.by_name("data/custom::data/name.mts").is_ok()); + + let options_file = archive.by_name("pairs/0/options.json").unwrap(); + let options_json = std::io::read_to_string(options_file).unwrap(); + let options_json: JsonValue = json::parse(&options_json).unwrap(); + + assert_eq!(options_json["type"].as_str(), Some("metatomic_pair_options")); + assert_eq!(options_json["cutoff"].as_str(), Some(&*format!("0x{:x}", 3.5_f64.to_bits()))); + assert_eq!(options_json["full_list"].as_bool(), Some(true)); + assert_eq!(options_json["strict"].as_bool(), Some(false)); + } + + let loaded = { + let file = std::fs::File::open(&path).unwrap(); + let loaded = load(file, Some(metatensor::io::create_ndarray)).unwrap(); + std::fs::remove_file(&path).unwrap(); + loaded + }; + + assert_eq!(loaded.length_unit(), "Angstrom"); + + let types: ndarray::ArrayView1 = loaded.types().try_into().unwrap(); + let positions: ndarray::ArrayView2 = loaded.positions().try_into().unwrap(); + let cell: ndarray::ArrayView2 = loaded.cell().try_into().unwrap(); + let pbc: ndarray::ArrayView1 = loaded.pbc().try_into().unwrap(); + + assert_eq!(types, ndarray::arr1(&[1, 6, 8])); + assert_eq!( + positions, + ndarray::arr2(&[[0.0, 0.0, 0.0], [1.0, 0.0, 0.0], [2.0, 0.0, 0.0]]) + ); + assert_eq!( + cell, + ndarray::arr2(&[[10.0, 0.0, 0.0], [0.0, 10.0, 0.0], [0.0, 0.0, 10.0]]) + ); + assert_eq!(pbc, ndarray::arr1(&[true, true, true])); + + let options = PairListOptions { + cutoff: 3.5, + full_list: true, + strict: false, + requestors: vec![], + }; + assert!(loaded.get_pairs(&options).is_some()); + assert!(loaded.get_custom_data("custom::data/name").is_ok()); + } +} diff --git a/metatomic-core/src/io/tensor.rs b/metatomic-core/src/io/tensor.rs new file mode 100644 index 000000000..81c8987cd --- /dev/null +++ b/metatomic-core/src/io/tensor.rs @@ -0,0 +1,333 @@ +use byteorder::{BigEndian, LittleEndian, NativeEndian, ReadBytesExt, WriteBytesExt}; + +use dlpk::{DLDataType, DLDataTypeCode, DLDevice, DLPackTensor, DLPackTensorRef, DLPackVersion}; +use metatensor::MtsArray; +use metatensor::c_api::{MTS_SUCCESS, mts_array_t, mts_create_array_callback_t}; + +use crate::Error; + +use super::{Endianness, check_for_extra_bytes}; +use super::npy_header::{Header, DataType}; + +/// Parse an NPY type descriptor string (e.g. `" Result<(DLDataTypeCode, u8, Endianness), Error> { + if descr.len() < 3 { + return Err(Error::Serialization(format!("invalid type descriptor: {}", descr))); + } + + let endian = match &descr[0..1] { + "<" => Endianness::Little, + "=" | "|" => Endianness::Native, + ">" => Endianness::Big, + // not applicable for single-byte types + _ => return Err(Error::Serialization(format!("unknown endianness in type descriptor: {}", descr))), + }; + + let type_char = &descr[1..2]; + let size_str = &descr[2..]; + let size: u8 = size_str.parse().map_err(|_| { + Error::Serialization(format!("invalid size in type descriptor: {}", descr)) + })?; + + let (code, bits) = match (type_char, size) { + ("f", 4) => (DLDataTypeCode::kDLFloat, 32), + ("f", 8) => (DLDataTypeCode::kDLFloat, 64), + ("i", 1) => (DLDataTypeCode::kDLInt, 8), + ("i", 2) => (DLDataTypeCode::kDLInt, 16), + ("i", 4) => (DLDataTypeCode::kDLInt, 32), + ("i", 8) => (DLDataTypeCode::kDLInt, 64), + ("u", 1) => (DLDataTypeCode::kDLUInt, 8), + ("u", 2) => (DLDataTypeCode::kDLUInt, 16), + ("u", 4) => (DLDataTypeCode::kDLUInt, 32), + ("u", 8) => (DLDataTypeCode::kDLUInt, 64), + ("b", 1) => (DLDataTypeCode::kDLBool, 8), + ("c", 8) => (DLDataTypeCode::kDLComplex, 64), + ("c", 16) => (DLDataTypeCode::kDLComplex, 128), + ("f", 2) => (DLDataTypeCode::kDLFloat, 16), + _ => return Err(Error::Serialization(format!("unsupported type descriptor: {}", descr))), + }; + + Ok((code, bits, endian)) +} + + +fn read_as(reader: &mut R, tensor: dlpk::DLPackTensorRefMut<'_>, cb: impl Fn(&mut R, &mut T) -> Result<(), std::io::Error>) -> Result<(), Error> +where R: std::io::Read, + T: dlpk::DLPackPointerCast + 'static +{ + let mut view: ndarray::ArrayViewMutD = tensor.try_into() + .map_err(|e| Error::Serialization(format!("failed to convert DLPack to ndarray mutable view: {}", e)))?; + + for value in &mut view { + cb(reader, value)?; + } + + Ok(()) +} + +// Read a data array from the given reader, using numpy's NPY format +#[allow(clippy::too_many_lines)] +pub fn read_tensor(mut reader: R, create_array: mts_create_array_callback_t) -> Result + where R: std::io::Read +{ + let create_array = create_array.ok_or_else(|| Error::InvalidParameter("create_array callback is NULL".into()))?; + let header = super::npy_header::Header::from_reader(&mut reader)?; + if header.fortran_order { + return Err(Error::Serialization("data can not be loaded from fortran-order arrays".into())); + } + + let descr = if let super::npy_header::DataType::Scalar(s) = &header.type_descriptor { + s.as_str() + } else { + return Err(Error::Serialization("structured arrays are not supported".into())); + }; + + let (file_code, file_bits, endian) = npy_descr_to_dtype(descr)?; + + let dl_dtype = DLDataType { code: file_code, bits: file_bits, lanes: 1 }; + + let shape = header.shape; + let mut array = mts_array_t::null(); + let status = unsafe { + create_array(shape.as_ptr(), shape.len(), dl_dtype, &mut array) + }; + + let array = if status == MTS_SUCCESS { + MtsArray::from_raw(array) + } else { + // TODO: how can we propagate the error from the callback? + return Err(Error::Serialization("failed to create array".into())); + }; + + let device = DLDevice::cpu(); + let version = DLPackVersion::current(); + let mut dl_tensor = array.as_dlpack(device, None, version)?; + + let num_elements: usize = shape.iter().product(); + if num_elements == 0 { + check_for_extra_bytes(&mut reader)?; + return Ok(dl_tensor); + } + + let tensor = dl_tensor.as_mut(); + + // Endianness is handled inside each arm to avoid tripling the number of + // match arms (which inflates uncovered-line counts for big/native paths + // that are not exercised in tests on little-endian CI). + match (file_code, file_bits) { + // Standard Floats + (DLDataTypeCode::kDLFloat, 32) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => r.read_f32::()?, + Endianness::Big => r.read_f32::()?, + Endianness::Native => r.read_f32::()?, + }; + Ok(()) + }), + (DLDataTypeCode::kDLFloat, 64) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => r.read_f64::()?, + Endianness::Big => r.read_f64::()?, + Endianness::Native => r.read_f64::()?, + }; + Ok(()) + }), + + // Standard Ints + (DLDataTypeCode::kDLInt, 8) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = r.read_i8()?; + Ok(()) + }), + (DLDataTypeCode::kDLInt, 16) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => r.read_i16::()?, + Endianness::Big => r.read_i16::()?, + Endianness::Native => r.read_i16::()?, + }; + Ok(()) + }), + (DLDataTypeCode::kDLInt, 32) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => r.read_i32::()?, + Endianness::Big => r.read_i32::()?, + Endianness::Native => r.read_i32::()?, + }; + Ok(()) + }), + (DLDataTypeCode::kDLInt, 64) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => r.read_i64::()?, + Endianness::Big => r.read_i64::()?, + Endianness::Native => r.read_i64::()?, + }; + Ok(()) + }), + + // Unsigned Ints + (DLDataTypeCode::kDLUInt, 8) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = r.read_u8()?; + Ok(()) + }), + (DLDataTypeCode::kDLUInt, 16) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => r.read_u16::()?, + Endianness::Big => r.read_u16::()?, + Endianness::Native => r.read_u16::()?, + }; + Ok(()) + }), + (DLDataTypeCode::kDLUInt, 32) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => r.read_u32::()?, + Endianness::Big => r.read_u32::()?, + Endianness::Native => r.read_u32::()?, + }; + Ok(()) + }), + (DLDataTypeCode::kDLUInt, 64) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => r.read_u64::()?, + Endianness::Big => r.read_u64::()?, + Endianness::Native => r.read_u64::()?, + }; + Ok(()) + }), + + // Boolean (Read as u8) + (DLDataTypeCode::kDLBool, 8) => read_as::(&mut reader, tensor, |r: &mut R, v| { + *v = r.read_u8()? != 0; + Ok(()) + }), + + // Complex Numbers (Read as array of 2 floats) + (DLDataTypeCode::kDLComplex, 64) => read_as::<[f32; 2], _>(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => [r.read_f32::()?, r.read_f32::()?], + Endianness::Big => [r.read_f32::()?, r.read_f32::()?], + Endianness::Native => [r.read_f32::()?, r.read_f32::()?], + }; + Ok(()) + }), + (DLDataTypeCode::kDLComplex, 128) => read_as::<[f64; 2], _>(&mut reader, tensor, |r: &mut R, v| { + *v = match endian { + Endianness::Little => [r.read_f64::()?, r.read_f64::()?], + Endianness::Big => [r.read_f64::()?, r.read_f64::()?], + Endianness::Native => [r.read_f64::()?, r.read_f64::()?], + }; + Ok(()) + }), + + _ => Err(Error::Serialization(format!( + "unsupported dtype for reading: {:?} {} bits", file_code, file_bits + ))), + }?; + + check_for_extra_bytes(&mut reader)?; + Ok(dl_tensor) +} + +fn dlpack_to_npy_descr(code: DLDataTypeCode, bits: u8) -> Result { + let endian = if cfg!(target_endian = "little") { "<" } else { ">" }; + + let (type_char, type_size) = match (code, bits) { + (DLDataTypeCode::kDLInt, 8) => ("i", 1), + (DLDataTypeCode::kDLInt, 16) => ("i", 2), + (DLDataTypeCode::kDLInt, 32) => ("i", 4), + (DLDataTypeCode::kDLInt, 64) => ("i", 8), + (DLDataTypeCode::kDLUInt, 8) => ("u", 1), + (DLDataTypeCode::kDLUInt, 16) => ("u", 2), + (DLDataTypeCode::kDLUInt, 32) => ("u", 4), + (DLDataTypeCode::kDLUInt, 64) => ("u", 8), + (DLDataTypeCode::kDLFloat, 32) => ("f", 4), + (DLDataTypeCode::kDLFloat, 64) => ("f", 8), + (DLDataTypeCode::kDLBool, 8) => ("b", 1), + (DLDataTypeCode::kDLComplex, 64) => ("c", 8), + (DLDataTypeCode::kDLComplex, 128) => ("c", 16), + (DLDataTypeCode::kDLFloat, 16) => ("f", 2), + _ => return Err(Error::Serialization( + format!("unsupported DLPack dtype: code {:?}, bits {:?}", code, bits) + ) + ), + }; + + Ok(format!("{}{}{}", endian, type_char, type_size)) +} + + +fn write_as(writer: &mut W, tensor: dlpk::DLPackTensorRef<'_>, cb: impl Fn(&mut W, T) -> Result<(), std::io::Error>) -> Result<(), Error> +where W: std::io::Write, + T: Copy + dlpk::DLPackPointerCast + 'static +{ + let view: ndarray::ArrayViewD = tensor.try_into() + .map_err(|e| Error::Serialization(format!("failed to convert DLPack to ndarray view: {}", e)))?; + + for &value in &view { + cb(writer, value)?; + } + + Ok(()) +} + +// Write an array to the given writer, using numpy's NPY format +#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)] +pub fn write_tensor(writer: &mut W, tensor: DLPackTensorRef<'_>) -> Result<(), Error> { + let dtype = tensor.raw.dtype; + let (code, bits) = (dtype.code, dtype.bits); + + // Validate Lanes + if dtype.lanes != 1 { + return Err(Error::Serialization(format!( + "unsupported DLPack dtype: lanes != 1 ({})", dtype.lanes + ))); + } + + // Write Header + let tdesc = dlpack_to_npy_descr(code, bits)?; + let header = Header { + type_descriptor: DataType::Scalar(tdesc), + fortran_order: false, + shape: tensor.shape().iter().map(|&s| s as usize).collect(), + }; + + header.write(&mut *writer)?; + + // Get metadata for size and pointer for data + let num_elements: usize = header.shape.iter().product(); + if num_elements == 0 { + return Ok(()); + } + + match (code, bits) { + // Standard Floats + (DLDataTypeCode::kDLFloat, 32) => write_as::(writer, tensor, |w: &mut W, v| w.write_f32::(v)), + (DLDataTypeCode::kDLFloat, 64) => write_as::(writer, tensor, |w: &mut W, v| w.write_f64::(v)), + + // Standard Ints + (DLDataTypeCode::kDLInt, 8) => write_as::(writer, tensor, |w: &mut W, v| w.write_i8(v)), + (DLDataTypeCode::kDLInt, 16) => write_as::(writer, tensor, |w: &mut W, v| w.write_i16::(v)), + (DLDataTypeCode::kDLInt, 32) => write_as::(writer, tensor, |w: &mut W, v| w.write_i32::(v)), + (DLDataTypeCode::kDLInt, 64) => write_as::(writer, tensor, |w: &mut W, v| w.write_i64::(v)), + + // Unsigned Ints + (DLDataTypeCode::kDLUInt, 8) => write_as::(writer, tensor, |w: &mut W, v| w.write_u8(v)), + (DLDataTypeCode::kDLUInt, 16) => write_as::(writer, tensor, |w: &mut W, v| w.write_u16::(v)), + (DLDataTypeCode::kDLUInt, 32) => write_as::(writer, tensor, |w: &mut W, v| w.write_u32::(v)), + (DLDataTypeCode::kDLUInt, 64) => write_as::(writer, tensor, |w: &mut W, v| w.write_u64::(v)), + + // Boolean, stored as u8 + (DLDataTypeCode::kDLBool, 8) => write_as::(writer, tensor, |w: &mut W, v| w.write_u8(u8::from(v))), + + // Complex Numbers + (DLDataTypeCode::kDLComplex, 64) => write_as::<[f32; 2], _>(writer, tensor, |w: &mut W, v| { + w.write_f32::(v[0])?; + w.write_f32::(v[1]) + }), + (DLDataTypeCode::kDLComplex, 128) => write_as::<[f64; 2], _>(writer, tensor, |w: &mut W, v| { + w.write_f64::(v[0])?; + w.write_f64::(v[1]) + }), + + _ => Err(Error::Serialization(format!("unsupported dtype for writing: {:?} {} bits", code, bits))), + } +} diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index 9e1acb29b..18dd99cac 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -28,6 +28,8 @@ pub use self::quantities::{Quantity, SampleKind, Gradients}; mod system; pub use self::system::System; +mod io; + mod model; pub use self::model::Model; @@ -40,7 +42,7 @@ pub use self::units::unit_conversion_factor; /// The possible sources of error in metatomic #[derive(Debug, Clone)] pub enum Error { - /// Error while serializing data to or deserializing data from JSON + /// Error while serializing data to or deserializing data Serialization(String), /// Invalid parameters passed to a function InvalidParameter(String), @@ -122,3 +124,18 @@ impl From for Error { Error::Metatensor(error) } } + +impl From for Error { + fn from(error: json::Error) -> Self { + Error::Serialization(format!("json error: {}", error)) + } +} + +impl> From<(T, zip::result::ZipError)> for Error { + fn from((path, error): (T, zip::result::ZipError)) -> Self { + match error { + zip::result::ZipError::Io(e) => Error::Io(Arc::new(e)), + error => Error::Serialization(format!("{}: at '{}'", error, path.as_ref())), + } + } +} diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs index 3bf1ce00d..0ccdba8a2 100644 --- a/metatomic-core/src/metadata.rs +++ b/metatomic-core/src/metadata.rs @@ -69,20 +69,37 @@ impl<'a> TryFrom<&'a JsonValue> for PairListOptions { )); } - if value["type"].as_str() != Some("metatomic_pair_options") { - return Err(Error::Serialization( - "'type' in JSON for PairListOptions must be 'metatomic_pair_options'".into() - )); - } + let cutoff = if value.has_key("class") { + // this is the legacy format from metatomic-torch, which can be used + // to load serialized PairListOptions + if value["class"].as_str() != Some("NeighborListOptions") { + return Err(Error::Serialization( + "'class' in legacy JSON for PairListOptions must be 'NeighborListOptions'".into() + )); + } - let cutoff = value["cutoff"].as_str().ok_or_else(|| Error::Serialization( - "'cutoff' in JSON for PairListOptions must be a hex-encoded string".into() - ))?; - let bits = u64::from_str_radix(cutoff.strip_prefix("0x").unwrap_or(cutoff), 16) - .map_err(|_| Error::Serialization( + let cutoff_bits = value["cutoff"].as_u64().ok_or_else(|| Error::Serialization( + "'cutoff' in legacy JSON for PairListOptions must be an integer".into() + ))?; + + f64::from_bits(cutoff_bits) + } else { + if value["type"].as_str() != Some("metatomic_pair_options") { + return Err(Error::Serialization( + "'type' in JSON for PairListOptions must be 'metatomic_pair_options'".into() + )); + } + + let cutoff_str = value["cutoff"].as_str().ok_or_else(|| Error::Serialization( "'cutoff' in JSON for PairListOptions must be a hex-encoded string".into() ))?; - let cutoff = f64::from_bits(bits); + let cutoff_bits = u64::from_str_radix(cutoff_str.strip_prefix("0x").unwrap_or(cutoff_str), 16) + .map_err(|_| Error::Serialization( + "'cutoff' in JSON for PairListOptions must be a hex-encoded string".into() + ))?; + + f64::from_bits(cutoff_bits) + }; if !cutoff.is_finite() || cutoff <= 0.0 { return Err(Error::Serialization( diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs index 182c05d2e..b3aeb8e3e 100644 --- a/metatomic-core/src/system.rs +++ b/metatomic-core/src/system.rs @@ -155,21 +155,17 @@ impl System { )); } - // TODO: add TensorBlock::device/dtype and use them here - let values = pairs.values(); - let values_device = values.device()?; - if values_device != self.device() { + if pairs.device()? != self.device() { return Err(Error::InvalidParameter(format!( "`pairs` device ({}) does not match this system's device ({})", - values_device, self.device(), + pairs.device()?, self.device(), ))); } - let values_dtype = values.dtype()?; - if values_dtype != self.dtype() { + if pairs.dtype()? != self.dtype() { return Err(Error::InvalidParameter(format!( "`pairs` dtype ({}) does not match this system's dtype ({})", - values_dtype, self.dtype(), + pairs.dtype()?, self.dtype(), ))); } @@ -191,7 +187,8 @@ impl System { /// /// If `override_` is `true`, existing data with the same name will be /// replaced. - pub fn add_custom_data(&mut self, name: String, data: TensorMap, override_: bool) -> Result<(), Error> { + pub fn add_custom_data(&mut self, name: impl Into, data: TensorMap, override_: bool) -> Result<(), Error> { + let name = name.into(); if INVALID_DATA_NAMES.contains(name.to_lowercase().as_str()) { return Err(Error::InvalidParameter(format!( "custom data can not be named '{}'", name @@ -377,11 +374,14 @@ fn validate_cpu_system_data(system: &System) -> Result<(), Error> { return Ok(()); } +#[cfg(test)] +pub(crate) use tests::test_system; + #[cfg(test)] mod tests { use super::*; use metatensor::Labels; -use ndarray::{Array1, Array2}; + use ndarray::{Array1, Array2}; // ----------------------------------------------------------------------- // helpers to create DLPack tensors @@ -489,6 +489,29 @@ use ndarray::{Array1, Array2}; assert_eq!(error.to_string(), expected); } + pub(crate) fn test_system() -> System { + let mut system = System::new( + "Angstrom".into(), + tests::type_tensor(&[1, 6, 8]), + tests::positions_tensor(3, "f32"), + tests::cell_tensor(10.0, "f32"), + tests::pbc_tensor(&[true, true, true]), + ).unwrap(); + + system.add_custom_data("custom::data/name", valid_custom_data("f32"), true).unwrap(); + + let options = PairListOptions { + cutoff: 3.5, + full_list: true, + strict: false, + requestors: vec![], + }; + + system.add_pairs(options, valid_pair_block("f32")).unwrap(); + + return system; + } + #[test] fn system() { let system = System::new( @@ -679,17 +702,17 @@ use ndarray::{Array1, Array2}; ).unwrap(); let data = valid_custom_data("f32"); - system.add_custom_data("test::my_data".into(), data, false).unwrap(); + system.add_custom_data("test::my_data", data, false).unwrap(); assert_eq!(system.known_custom_data(), vec!["test::my_data"]); assert_eq!(system.get_custom_data("test::my_data").unwrap().keys().names(), ["key"]); assert_error( - system.add_custom_data("test::my_data".into(), valid_custom_data("f32"), false), + system.add_custom_data("test::my_data", valid_custom_data("f32"), false), "invalid parameter: custom data 'test::my_data' is already present in this system", ); let replacement = valid_custom_data("f32"); - system.add_custom_data("test::my_data".into(), replacement, true).unwrap(); + system.add_custom_data("test::my_data", replacement, true).unwrap(); assert_eq!(system.known_custom_data(), vec!["test::my_data"]); let mut system = System::new( @@ -699,8 +722,8 @@ use ndarray::{Array1, Array2}; cell_tensor(10.0, "f32"), pbc_tensor(&[true, true, true]), ).unwrap(); - system.add_custom_data("test::a".into(), valid_custom_data("f32"), false).unwrap(); - system.add_custom_data("test::b".into(), valid_custom_data("f32"), false).unwrap(); + system.add_custom_data("test::a", valid_custom_data("f32"), false).unwrap(); + system.add_custom_data("test::b", valid_custom_data("f32"), false).unwrap(); let mut names = system.known_custom_data(); names.sort_unstable(); assert_eq!(names, vec!["test::a", "test::b"]); @@ -733,20 +756,20 @@ use ndarray::{Array1, Array2}; } assert_error( - system.add_custom_data("my_data".into(), valid_custom_data("f32"), false), + system.add_custom_data("my_data", valid_custom_data("f32"), false), "invalid parameter: 'my_data' is not a standard quantity name; custom quantity names must use '::'", ); let keys = Labels::empty(vec!["key"]); let empty = TensorMap::new(keys, vec![]).unwrap(); assert_error( - system.add_custom_data("test::empty".into(), empty, false), + system.add_custom_data("test::empty", empty, false), "invalid parameter: custom data 'test::empty' has no blocks", ); let dtype_mismatch = valid_custom_data("f64"); assert_error( - system.add_custom_data("test::dtype".into(), dtype_mismatch, false), + system.add_custom_data("test::dtype", dtype_mismatch, false), "invalid parameter: dtype of custom data 'test::dtype' does not match this system dtype", ); } diff --git a/metatomic-core/tests/data/legacy.mta b/metatomic-core/tests/data/legacy.mta new file mode 100644 index 0000000000000000000000000000000000000000..1eee677e78f11523ef624cc9fc27094c6b58edba GIT binary patch literal 6060 zcmd5=-D@0G6u(JZv)yWmHc~$dPWGXjLP)Yn+H69l>q`-$HC2n&wua4SwpqK`S!ZUe zu`L+&p@I*76!k?2xS)SQ-n3OjEBYb`iZ2$_ClwZ{f)AqScjunjJ2x{Wn}~RpJ7?z3 zIcLuA-gD1AduOJ%Y!xDv5_IYRg{~ppLU(n?tN0bC<_*>AOK%)G_TbF%E^$_zv$FHH zS8}scR`y^ypB=QaWykg1Vr|xO=WX;KE>=C8`n`o>-KOV(@j+{B(AsBRur^5P(6iW^ z)*;nh2zW~IUd8(qzeDjoid%|j3NzC^se$hsC$0%0&}I^c2BTH7tJTeq3JuA>GAfLU z95ZN4(yaxfd(9)zKmYFDd#}8-z1w{H{_mGSik#U>(x0x8`^&fG+;hZ?2@ zFrGMi-#f3;Jz6phXw>IQ$#ZyqF1Icao0~uhqD%Vy&26jOUnH;4lKt+3-^GgS4UY^v z^3cpnr%H(@fWOPj~>ojCmSPUlX<$1%RXA*qQVN zO!Lu4@)v)l!9YtMu3l}p=8JxDXwml|jbQ}!59Ic?KKffxTJ(vt#VUb8Ty%+EVQmTI zm~So%j^^cp&uuspsw|kg@};EfmjmxO+oi6k=OkWG@0qh5l|5Ns2$eUKoj^;7)}gYb zZq?74 zH`UK}ljKADEU$R7yuxMppDZ7^#(so6X?(cfu)o&N=-RsZ_?r*&+S);D%PTI{tL??r z>JHps&WmeRFELIGWXqK9{Mn-KEFvQxq}^rTpf!;7?2=QPU)ztpdDdH~lu0L|puLz& zdQs9|oLWz1gEbxOVsbYEB!&!h$R7$O9wlMdVmZ|nk;8kFV3Yg!J@?_+?1O#oV?Fe- zFRY7a_=sl^Wbv%>p6o{-$h2oZvrODKSeCi|U1^*yJXS2j{mi+aIr{k?T)*dc>X4%z4E`igH>p&n`XJX&i5)j!DBEAOyN7noD1O{RN}hNoeY>Q@35e(q8>9 zL79ko5|()?LyEPM-6^0BcUg&LX1c3;XiKy>uxQ%Qq(I=ZXfETXT`HJ*NPNg@Xgib+ zSabO#S<=#=sfZxJ>tg znfTqPpJ@_%Qulv8*4jg5J{l%9{tx2qB#nw{xYVH1Z*giWzPu_-UJr9iisImmy%ZK# z7x6_R=1CDiI-O4bC?{oXtPW5KNZ z|1$aCVf-2(RkXG%q$sGNOK09=P<=WmEjEFQr-pefDoQxse=$MyhYur+gcC*c-P8}? zGMe2wnrO%a5N8|aNr38|0@~PgNkCCBpPawvNd}|ETh55L6$}N@io6OD3#5~H+sKk> z2c3$yR3UI@?YT=7;ZZby{`jH8Xtemt4X=$qUP#Pir;Qd= z$MLtZ=}W`yU_nK9KICoqh?WZGR9=IfkF&VB0U=i+0ix}SXnjmiF`3gefV{+?M6xop#B99oFZjA;aGaD^1nqBuq=0+bPbMVbZ>-|;wh zZ`BOk00&}S03l1bv=QNlh*Nhz{q8ZMu^j3r`nKHT{GP=XY(oMr!Ib&s5USiS&pP z;%%DBxT7@Ff?f~xw!cYtpV3UPg?9N99bA3jQ8TrvpGvfZ3aUMeD>RIMNWFgmJF%OF literal 0 HcmV?d00001 diff --git a/scripts/include/metatensor.h b/scripts/include/metatensor.h index fb8e88f0d..56c085abf 100644 --- a/scripts/include/metatensor.h +++ b/scripts/include/metatensor.h @@ -4,5 +4,8 @@ typedef struct mts_labels_t mts_labels_t; typedef struct mts_block_t mts_block_t; typedef struct mts_tensormap_t mts_tensormap_t; +typedef void (*mts_create_array_callback_t)(void*); +typedef void (*mts_realloc_buffer_t)(void*); + typedef struct DLManagedTensorVersioned DLManagedTensorVersioned; From b2ef098d033a0a466d06ddef12fd3efecb77cef1 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Wed, 8 Jul 2026 11:33:35 +0200 Subject: [PATCH 32/43] Remove double indirection (Box + Arc) in mta_system_t --- metatomic-core/src/c_api/io.rs | 7 ++- metatomic-core/src/c_api/system.rs | 74 ++++++++++++++++++------------ 2 files changed, 50 insertions(+), 31 deletions(-) diff --git a/metatomic-core/src/c_api/io.rs b/metatomic-core/src/c_api/io.rs index ef47c115b..b6170fb1a 100644 --- a/metatomic-core/src/c_api/io.rs +++ b/metatomic-core/src/c_api/io.rs @@ -263,8 +263,10 @@ pub unsafe extern "C" fn mta_load( let file = BufReader::new(File::open(path)?); let system = crate::io::load(file, create_array)?; + let system = Arc::new(mta_system_t(system)); + let _ = &unwind_wrapper; - *unwind_wrapper.0 = Box::into_raw(Box::new(mta_system_t(Arc::new(system)))); + *unwind_wrapper.0 = Arc::into_raw(system).cast_mut(); Ok(()) }); @@ -303,9 +305,10 @@ pub unsafe extern "C" fn mta_load_buffer( }; let cursor = Cursor::new(slice); let system = crate::io::load(cursor, create_array)?; + let system = Arc::new(mta_system_t(system)); let _ = &unwind_wrapper; - *unwind_wrapper.0 = Box::into_raw(Box::new(mta_system_t(Arc::new(system)))); + *unwind_wrapper.0 = Arc::into_raw(system).cast_mut(); Ok(()) }); diff --git a/metatomic-core/src/c_api/system.rs b/metatomic-core/src/c_api/system.rs index b6f57b21a..97d2b9f4e 100644 --- a/metatomic-core/src/c_api/system.rs +++ b/metatomic-core/src/c_api/system.rs @@ -14,7 +14,21 @@ use super::{catch_unwind, mta_status_t, mta_string_t}; /// The system owns DLPack tensors for types, positions, cell, and PBC, as well /// as metatensor blocks for pair lists and tensor maps for custom data. #[allow(non_camel_case_types)] -pub struct mta_system_t(pub(crate) Arc); +pub struct mta_system_t(pub(crate) System); + +impl mta_system_t { + /// Convert an mta_system_t into a pointer inside an Arc, to be + /// passed through the C API + fn into_raw(self) -> *mut mta_system_t { + Arc::into_raw(Arc::new(self)).cast_mut() + } + + /// Recover the Arc from a pointer created with + /// [`mta_system_t::into_raw`] + unsafe fn from_raw(ptr: *const mta_system_t) -> Arc { + unsafe { Arc::from_raw(ptr) } + } +} /// Create a new system from raw DLPack tensors. /// @@ -60,10 +74,10 @@ pub unsafe extern "C" fn mta_system_create( let cell = DLPackTensor::from_ptr(cell); let pbc = DLPackTensor::from_ptr(pbc); - let system_inner = System::new(length_unit, types, positions, cell, pbc)?; + let system = mta_system_t(System::new(length_unit, types, positions, cell, pbc)?); let _ = &unwind_wrapper; - *unwind_wrapper.0 = Box::into_raw(Box::new(mta_system_t(Arc::new(system_inner)))); + *unwind_wrapper.0 = system.into_raw(); } Ok(()) }) @@ -85,9 +99,8 @@ pub unsafe extern "C" fn mta_system_free(system: *mut mta_system_t) -> mta_statu return Ok(()); } - unsafe { - std::mem::drop(Box::from_raw(system)); - } + let system = unsafe { mta_system_t::from_raw(system.cast_const()) }; + std::mem::drop(system); Ok(()) }) } @@ -130,16 +143,13 @@ pub enum mta_system_data_kind { /// Custom deleter for borrowed DLPack tensors returned by `mta_system_get_data`. /// -/// Releases the `Arc` reference stored in `manager_ctx` and frees the -/// heap-allocated `DLManagedTensorVersioned`. +/// Releases the `Arc` reference stored in `manager_ctx` and +/// frees the heap-allocated `DLManagedTensorVersioned`. unsafe extern "C" fn borrowed_tensor_deleter( tensor: *mut DLManagedTensorVersioned, ) { - let system: Arc = { - unsafe { - let ptr = (*tensor).manager_ctx as *const System; - Arc::from_raw(ptr) - } + let system = unsafe { + mta_system_t::from_raw((*tensor).manager_ctx.cast()) }; std::mem::drop(system); unsafe { @@ -177,7 +187,13 @@ pub unsafe extern "C" fn mta_system_get_data( *data = std::ptr::null_mut(); } - let system = unsafe { &*system }; + // increase the reference count of the system so that it stays alive as + // long as the returned tensor is alive. We do this by creating a + // temporary Arc from the raw pointer, cloning it and storing the clone + // in the manager_ctx. + let system = unsafe { mta_system_t::from_raw(system) }; + let arc_clone = system.clone(); + let tensor_ref = match request { mta_system_data_kind::MTA_SYSTEM_DATA_TYPES => system.0.types(), mta_system_data_kind::MTA_SYSTEM_DATA_POSITIONS => system.0.positions(), @@ -185,21 +201,17 @@ pub unsafe extern "C" fn mta_system_get_data( mta_system_data_kind::MTA_SYSTEM_DATA_PBC => system.0.pbc(), }; - // Clone the Arc to keep the system alive while the borrowed view exists - let system_arc = system.0.clone(); - let system_ptr = Arc::into_raw(system_arc); - let packed = Box::new(DLManagedTensorVersioned { version: DLPackVersion::current(), - manager_ctx: system_ptr as *mut std::ffi::c_void, - deleter: Some( - borrowed_tensor_deleter - as unsafe extern "C" fn(*mut DLManagedTensorVersioned), - ), + manager_ctx: Arc::into_raw(arc_clone) as *mut std::ffi::c_void, + deleter: Some(borrowed_tensor_deleter), flags: dlpk::sys::DLPACK_FLAG_BITMASK_READ_ONLY, dl_tensor: tensor_ref.raw.clone(), }); + // do not drop the system, it is still owned by the caller. + std::mem::forget(system); + unsafe { *data = Box::into_raw(packed); } @@ -263,14 +275,16 @@ pub unsafe extern "C" fn mta_system_add_pairs( let pairs = unsafe { TensorBlock::from_raw(pairs) }; - let system = unsafe { &mut *system }; - let system = Arc::get_mut(&mut system.0).ok_or_else(|| { + let mut system = unsafe { mta_system_t::from_raw(system.cast_const()) }; + let system_mut = Arc::get_mut(&mut system).ok_or_else(|| { Error::InvalidParameter( "cannot modify system while there are outstanding borrowed views".into(), ) })?; + system_mut.0.add_pairs(options, pairs)?; - system.add_pairs(options, pairs)?; + // do not drop the system, it is still owned by the caller. + std::mem::forget(system); Ok(()) }) @@ -386,14 +400,16 @@ pub unsafe extern "C" fn mta_system_add_custom_data( let data = unsafe { TensorMap::from_raw(data) }; - let system = unsafe { &mut *system }; - let system = Arc::get_mut(&mut system.0).ok_or_else(|| { + let mut system = unsafe { mta_system_t::from_raw(system.cast_const()) }; + let system_mut = Arc::get_mut(&mut system).ok_or_else(|| { Error::InvalidParameter( "cannot modify system while there are outstanding borrowed views".into(), ) })?; + system_mut.0.add_custom_data(name, data, false)?; - system.add_custom_data(name, data, false)?; + // do not drop the system, it is still owned by the caller. + std::mem::forget(system); Ok(()) }) From 3b5e1a2e0e422b601d7cdba8631cc734444deabf Mon Sep 17 00:00:00 2001 From: Rocco Meli Date: Wed, 8 Jul 2026 19:32:20 +0200 Subject: [PATCH 33/43] C++ metadata classes with JSON serialization/deserialization (#277) --- metatomic-core/CMakeLists.txt | 3 + .../cmake/metatomic-config.in.cmake | 7 +- metatomic-core/cmake/nlohmann_json.cmake | 37 + metatomic-core/include/metatomic.hpp | 1 + metatomic-core/include/metatomic/metadata.hpp | 1227 +++++++++++++++++ metatomic-core/tests/CMakeLists.txt | 1 + metatomic-core/tests/cxx/metadata.cpp | 910 ++++++++++++ 7 files changed, 2184 insertions(+), 2 deletions(-) create mode 100644 metatomic-core/cmake/nlohmann_json.cmake create mode 100644 metatomic-core/include/metatomic/metadata.hpp create mode 100644 metatomic-core/tests/cxx/metadata.cpp diff --git a/metatomic-core/CMakeLists.txt b/metatomic-core/CMakeLists.txt index 73382e21b..128d52f91 100644 --- a/metatomic-core/CMakeLists.txt +++ b/metatomic-core/CMakeLists.txt @@ -116,6 +116,7 @@ else() find_package(metatensor ${REQUIRED_METATENSOR_VERSION} CONFIG REQUIRED) endif() +include(cmake/nlohmann_json.cmake) include(cmake/detect_cargo.cmake) @@ -396,6 +397,8 @@ else() target_link_libraries(metatomic::shared INTERFACE metatensor) endif() +target_link_libraries(metatomic::static INTERFACE nlohmann_json::nlohmann_json) +target_link_libraries(metatomic::shared INTERFACE nlohmann_json::nlohmann_json) if (BUILD_SHARED_LIBS) add_library(metatomic ALIAS metatomic::shared) diff --git a/metatomic-core/cmake/metatomic-config.in.cmake b/metatomic-core/cmake/metatomic-config.in.cmake index 90fca167a..15652cbac 100644 --- a/metatomic-core/cmake/metatomic-config.in.cmake +++ b/metatomic-core/cmake/metatomic-config.in.cmake @@ -15,6 +15,9 @@ enable_language(CXX) set(REQUIRED_METATENSOR_VERSION @REQUIRED_METATENSOR_VERSION@) find_package(metatensor ${REQUIRED_METATENSOR_VERSION} CONFIG REQUIRED) +# Find nlohmann_json dependency +find_dependency(nlohmann_json 3.11.0) + get_filename_component(METATOMIC_PREFIX_DIR "${CMAKE_CURRENT_LIST_DIR}/@PACKAGE_RELATIVE_PATH@" ABSOLUTE) if (WIN32) @@ -46,7 +49,7 @@ if (@METATOMIC_INSTALL_BOTH_STATIC_SHARED@ OR @BUILD_SHARED_LIBS@) ) target_compile_features(metatomic::shared INTERFACE cxx_std_17) - target_link_libraries(metatomic::shared INTERFACE metatensor) + target_link_libraries(metatomic::shared INTERFACE metatensor nlohmann_json::nlohmann_json) if (WIN32) if (NOT EXISTS ${METATOMIC_IMPLIB_LOCATION}) @@ -75,7 +78,7 @@ if (@METATOMIC_INSTALL_BOTH_STATIC_SHARED@ OR NOT @BUILD_SHARED_LIBS@) ) target_compile_features(metatomic::static INTERFACE cxx_std_17) - target_link_libraries(metatomic::static INTERFACE metatensor) + target_link_libraries(metatomic::static INTERFACE metatensor nlohmann_json::nlohmann_json) endif() # Export either the shared or static library as the metatomic target diff --git a/metatomic-core/cmake/nlohmann_json.cmake b/metatomic-core/cmake/nlohmann_json.cmake new file mode 100644 index 000000000..39fd371fc --- /dev/null +++ b/metatomic-core/cmake/nlohmann_json.cmake @@ -0,0 +1,37 @@ +# Find or fetch nlohmann JSON library +# +# This module first tries to find nlohmann_json via find_package. +# If that fails, it falls back to fetching it via FetchContent. +# +# After including this module, you can link against nlohmann_json::nlohmann_json + +# Guard against multiple inclusion +if(TARGET nlohmann_json::nlohmann_json) + return() +endif() + +if (POLICY CMP0135) + cmake_policy(SET CMP0135 NEW) # DOWNLOAD_EXTRACT_TIMESTAMP TRUE in FetchContent_Declare +endif() + +include(FetchContent) + +find_package(nlohmann_json 3.11.0 QUIET) + +if(nlohmann_json_FOUND) + message(STATUS "Found nlohmann_json via find_package: ${nlohmann_json_VERSION}") +else() + message(STATUS "nlohmann_json not found via find_package, fetching from GitHub") + + # Fetch the release tarball, which contains the CMake build files and headers + # but not the benchmark reports with very long filenames that break Windows. + FetchContent_Declare( + nlohmann_json + URL https://github.com/nlohmann/json/releases/download/v3.11.3/json.tar.xz + ) + + set(JSON_BuildTests OFF CACHE INTERNAL "") + set(JSON_Install ON CACHE INTERNAL "") + + FetchContent_MakeAvailable(nlohmann_json) +endif() diff --git a/metatomic-core/include/metatomic.hpp b/metatomic-core/include/metatomic.hpp index e41f09542..a714290e3 100644 --- a/metatomic-core/include/metatomic.hpp +++ b/metatomic-core/include/metatomic.hpp @@ -3,3 +3,4 @@ #include "metatomic/model.hpp" // IWYU pragma: export #include "metatomic/plugin.hpp" // IWYU pragma: export #include "metatomic/errors.hpp" // IWYU pragma: export +#include "metatomic/metadata.hpp" // IWYU pragma: export diff --git a/metatomic-core/include/metatomic/metadata.hpp b/metatomic-core/include/metatomic/metadata.hpp new file mode 100644 index 000000000..f539766cc --- /dev/null +++ b/metatomic-core/include/metatomic/metadata.hpp @@ -0,0 +1,1227 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include // std::move +#include // std::isfinite +#include // std::memcpy +#include // std::uint64_t, std::int64_t +#include // std::isxdigit + +#include +#include + +namespace metatomic{ + namespace detail { + + inline std::vector read_string_array( + const nlohmann::json& j, const std::string& key, const char* context + ) { + if (!j.contains(key) || !j[key].is_array()) { + throw metatomic::Error("'" + key + "' in " + context + " must be an array"); + } + + std::vector result; + for (const auto& item : j[key]) { + if (!item.is_string()) { + throw metatomic::Error("'" + key + "' in " + context + " must be an array of strings"); + } + result.push_back(item.get()); + } + return result; + } + + } // namespace detail + + /// Options for the calculation of a pair list (neighbor list) + struct PairListOptions{ + private: + /// Cutoff radius for this pair list in the length unit of the model + std::optional cutoff_; + /// Whether the list is a full list (contains both the pair `i -> j` and `j -> i`) + /// or a half list (contains only `i -> j`) + std::optional full_list_; + /// Whether the list guarantees that only atoms within the cutoff are + /// included (strict) or may also include pairs slightly beyond the cutoff + /// (non-strict) + bool strict_ = true; + /// List of strings describing who requested this pair list + std::vector requestors_; + + public: + /// Set the cutoff radius for this pair list. + /// + /// @throw metatomic::Error if the value is not a finite positive number. + void cutoff(double value) { + if (!std::isfinite(value) || value <= 0.0) { + throw metatomic::Error("cutoff must be a finite positive number"); + } + cutoff_ = value; + } + + /// Get the cutoff radius for this pair list. + /// + /// @throw metatomic::Error if the value has not been set. + double cutoff() const { + if (!cutoff_.has_value()) { + throw metatomic::Error("cutoff is not set in PairListOptions"); + } + return cutoff_.value(); + } + + /// Set whether this pair list is a full list. + /// + /// @throw metatomic::Error if the value has not been set. + void full_list(bool value) { + full_list_ = value; + } + + /// Get whether this pair list is a full list. + /// + /// @throw metatomic::Error if the value has not been set. + bool full_list() const { + if (!full_list_.has_value()) { + throw metatomic::Error("full_list is not set in PairListOptions"); + } + return full_list_.value(); + } + + /// Set whether this pair list is strict. + void strict(bool value) { + strict_ = value; + } + + /// Get whether this pair list is strict. + bool strict() const { + return strict_; + } + + /// Set the list of requestors for this pair list. + void requestors(std::vector value) { + requestors_ = std::move(value); + } + + /// Get the list of requestors for this pair list. + const std::vector& requestors() const { + return requestors_; + } + + /// Add a requestor to the list. + /// + /// Empty strings and duplicates are ignored, keeping first-seen order. + void add_requestor(const std::string& requestor) { + if (!requestor.empty() && std::find(requestors_.begin(), requestors_.end(), requestor) == requestors_.end()) { + requestors_.push_back(requestor); + } + } + + /// Clear the list of requestors. + void clear_requestors() { + requestors_.clear(); + } + + /// Check if two `PairListOptions` are equal. + /// + /// The list of requestors is ignored when checking for equality. + bool operator==(const PairListOptions& other) const { + return cutoff_ == other.cutoff_ && + full_list_ == other.full_list_ && + strict_ == other.strict_; + } + + /// Check if two `PairListOptions` are different. + /// + /// The list of requestors is ignored when checking for equality. + bool operator!=(const PairListOptions& other) const { + return !(*this == other); + } + + /// Create a default `PairListOptions`. The cutoff and full_list fields + /// must be set before the object can be used. + PairListOptions() = default; + + /// Create a `PairListOptions` with the given values. + /// + /// @param cutoff spherical cutoff radius for the pair list + /// @param full_list whether the list is a full list + /// @param strict whether the list is strict + /// @param requestors list of strings describing who requested this pair list + PairListOptions( + double cutoff, + bool full_list, + bool strict = true, + std::vector requestors = {} + ) { + this->cutoff(cutoff); + this->full_list(full_list); + this->strict(strict); + this->requestors(std::move(requestors)); + } + }; + + inline void to_json(nlohmann::json& j, const PairListOptions& p){ + // Store cutoff as hex-encoded bit pattern + // Floating-point round-trip conversions is exact + double cutoff = p.cutoff(); + uint64_t bits; + std::memcpy(&bits, &cutoff, sizeof(double)); + std::ostringstream oss; + oss << "0x" << std::hex << bits; + + j = nlohmann::json{ + {"type", "metatomic_pair_options"}, + {"cutoff", oss.str()}, + {"full_list", p.full_list()}, + {"strict", p.strict()}, + {"requestors", p.requestors()} + }; + } + + inline void from_json(const nlohmann::json& j, PairListOptions& p) { + if (!j.is_object()) { + throw metatomic::Error("invalid JSON data for PairListOptions, expected an object"); + } + + if (!j.contains("type") || !j["type"].is_string() || j["type"].get() != "metatomic_pair_options") { + throw metatomic::Error("'type' in JSON for PairListOptions must be 'metatomic_pair_options'"); + } + + // Cutoff is an hex-encoded string + if (!j.contains("cutoff") || !j["cutoff"].is_string()) { + throw metatomic::Error("'cutoff' in JSON for PairListOptions must be a hex-encoded string"); + } + std::string cutoff_str = j["cutoff"].get(); + + // Strip "0x" prefix if present + if (cutoff_str.size() >= 2 && cutoff_str[0] == '0' && cutoff_str[1] == 'x') { + cutoff_str = cutoff_str.substr(2); + } + + uint64_t bits; + try { + // std::isxdigit checks for hex digits + if (cutoff_str.empty() || !std::all_of(cutoff_str.begin(), cutoff_str.end(), [](unsigned char c) { return std::isxdigit(c); })) { + throw metatomic::Error("'cutoff' in JSON for PairListOptions must be a hex-encoded string"); + } + + std::size_t pos = 0; + bits = std::stoull(cutoff_str, &pos, 16); + if (pos != cutoff_str.size()) { + throw metatomic::Error("'cutoff' in JSON for PairListOptions must be a hex-encoded string"); + } + } catch (...) { + throw metatomic::Error("'cutoff' in JSON for PairListOptions must be a hex-encoded string"); + } + double cutoff; + std::memcpy(&cutoff, &bits, sizeof(double)); + + if (!std::isfinite(cutoff) || cutoff <= 0.0) { + throw metatomic::Error("'cutoff' in JSON for PairListOptions must be a finite positive number"); + } + + if (!j.contains("full_list") || !j["full_list"].is_boolean()) { + throw metatomic::Error("'full_list' in JSON for PairListOptions must be a boolean"); + } + bool full_list = j["full_list"].get(); + + if (!j.contains("strict") || !j["strict"].is_boolean()) { + throw metatomic::Error("'strict' in JSON for PairListOptions must be a boolean"); + } + bool strict = j["strict"].get(); + + p = PairListOptions(cutoff, full_list, strict, {}); + if (j.contains("requestors")) { + if (!j["requestors"].is_array()) { + throw metatomic::Error("'requestors' in JSON for PairListOptions must be an array"); + } + + for (const auto& requestor : j["requestors"]) { + if (!requestor.is_string()) { + throw metatomic::Error("'requestors' in JSON for PairListOptions must be an array of strings"); + } + p.add_requestor(requestor.get()); + } + } + } + + // Forward declarations + // The ModelMetadata::print function uses to_json + struct ModelMetadata; + void to_json(nlohmann::json&, const ModelMetadata&); + + struct ModelMetadata { + /// References for a model, divided into three categories: references about + /// the model as a whole, references about the architecture of the model, + /// and references about the implementation of the model. + struct References { + private: + /// The references about the model as a whole, e.g. a paper describing the + /// model or a website presenting it. + std::vector model_; + /// The references about the architecture of the model, e.g. papers + /// describing the mathematical form of the model. + std::vector architecture_; + /// The references about the implementation of the model, e.g. a link to + /// the source code repository or a paper describing the software. + std::vector implementation_; + + public: + /// Set the references about the model as a whole. + void model(std::vector value) { + model_ = std::move(value); + } + + /// Get the references about the model as a whole. + const std::vector& model() const { + return model_; + } + + /// Add a reference about the model as a whole. + void add_model(const std::string& reference) { + model_.push_back(reference); + } + + /// Clear the references about the model as a whole. + void clear_model() { + model_.clear(); + } + + /// Set the references about the architecture of the model. + void architecture(std::vector value) { + architecture_ = std::move(value); + } + + /// Get the references about the architecture of the model. + const std::vector& architecture() const { + return architecture_; + } + + /// Add a reference about the architecture of the model. + void add_architecture(const std::string& reference) { + architecture_.push_back(reference); + } + + /// Clear the references about the architecture of the model. + void clear_architecture() { + architecture_.clear(); + } + + /// Set the references about the implementation of the model. + void implementation(std::vector value) { + implementation_ = std::move(value); + } + + /// Get the references about the implementation of the model. + const std::vector& implementation() const { + return implementation_; + } + + /// Add a reference about the implementation of the model. + void add_implementation(const std::string& reference) { + implementation_.push_back(reference); + } + + /// Clear the references about the implementation of the model. + void clear_implementation() { + implementation_.clear(); + } + + /// Create a `References` with the given values. + /// + /// @param model references about the model as a whole + /// @param architecture references about the architecture of the model + /// @param implementation references about the implementation of the model + References( + std::vector model = {}, + std::vector architecture = {}, + std::vector implementation = {} + ) { + this->model(std::move(model)); + this->architecture(std::move(architecture)); + this->implementation(std::move(implementation)); + } + }; + + private: + std::string name_; + std::vector authors_; + std::string description_; + References references_; + // BTreeMap in Rust is an ordered map + std::map extra_; + + public: + /// Set the name of the model. + void name(std::string value) { + name_ = std::move(value); + } + + /// Get the name of the model. + const std::string& name() const { + return name_; + } + + /// Set the list of authors of the model. + void authors(std::vector value) { + authors_ = std::move(value); + } + + /// Get the list of authors of the model. + const std::vector& authors() const { + return authors_; + } + + /// Add an author to the list of authors. + void add_author(const std::string& author) { + authors_.push_back(author); + } + + /// Clear the list of authors. + void clear_authors() { + authors_.clear(); + } + + /// Set the description of the model. + void description(std::string value) { + description_ = std::move(value); + } + + /// Get the description of the model. + const std::string& description() const { + return description_; + } + + /// Set the references for the model. + void references(References value) { + references_ = std::move(value); + } + + /// Get the references for the model. + const References& references() const { + return references_; + } + + /// Add a reference to the given section. + /// + /// @param section reference section, one of "model", "architecture", or + /// "implementation" + /// @param reference the reference to add + /// @throw metatomic::Error if `section` is not one of the allowed values + void add_reference(const std::string& section, const std::string& reference) { + if (section == "model") { + references_.add_model(reference); + } else if (section == "architecture") { + references_.add_architecture(reference); + } else if (section == "implementation") { + references_.add_implementation(reference); + } else { + throw metatomic::Error( + "reference section must be 'model', 'architecture', or 'implementation', got '" + section + "'" + ); + } + } + + /// Clear a single reference section. + /// + /// @param section reference section, one of "model", "architecture", or + /// "implementation" + /// @throw metatomic::Error if `section` is not one of the allowed values + void clear_reference(const std::string& section) { + if (section == "model") { + references_.clear_model(); + } else if (section == "architecture") { + references_.clear_architecture(); + } else if (section == "implementation") { + references_.clear_implementation(); + } else { + throw metatomic::Error( + "reference section must be 'model', 'architecture', or 'implementation', got '" + section + "'" + ); + } + } + + /// Clear all references for the model. + void clear_references() { + references_.clear_model(); + references_.clear_architecture(); + references_.clear_implementation(); + } + + /// Set the extra metadata for the model. + void extra(std::map value) { + extra_ = std::move(value); + } + + /// Get the extra metadata for the model. + const std::map& extra() const { + return extra_; + } + + /// Add a key/value pair to the extra metadata. + /// + /// If the key already exists, its value is overwritten. + /// + /// @param key key for the extra metadata entry + /// @param value value for the extra metadata entry + void add_extra(const std::string& key, const std::string& value) { + extra_[key] = value; + } + + /// Clear the extra metadata. + void clear_extra() { + extra_.clear(); + } + + /// Create a `ModelMetadata` with the given values. + /// + /// @param name name of the model + /// @param authors list of authors of the model + /// @param description description of the model + /// @param references references for the model + /// @param extra extra metadata for the model + ModelMetadata( + std::string name = "", + std::vector authors = {}, + std::string description = "", + References references = {}, + std::map extra = {} + ) { + this->name(std::move(name)); + this->authors(std::move(authors)); + this->description(std::move(description)); + this->references(std::move(references)); + this->extra(std::move(extra)); + } + + /// Print the metadata as a human-readable string. + std::string print() const { + // Re-use C API to avoid re-implementing 'normalize_withespace' and 'wrap_80_chars' + mta_string_t mta_string; + nlohmann::json j; + + to_json(j, *this); + auto status = mta_format_metadata(j.dump().c_str(), &mta_string); + details::check_status(status); + + std::string output = mta_string_view(mta_string); + mta_string_free(mta_string); + + return output; + } + }; + + inline void to_json(nlohmann::json& j, const ModelMetadata::References& r) { + j = nlohmann::json{ + {"model", r.model()}, + {"architecture", r.architecture()}, + {"implementation", r.implementation()} + }; + } + + inline void from_json(const nlohmann::json& j, ModelMetadata::References& r) { + if (!j.is_object()) { + throw metatomic::Error("invalid JSON data for references in ModelMetadata, expected an object"); + } + + r = ModelMetadata::References( + detail::read_string_array(j, "model", "references of ModelMetadata"), + detail::read_string_array(j, "architecture", "references of ModelMetadata"), + detail::read_string_array(j, "implementation", "references of ModelMetadata") + ); + } + + inline void to_json(nlohmann::json& j, const ModelMetadata& m) { + j = nlohmann::json{ + {"type", "metatomic_model_metadata"}, + {"name", m.name()}, + {"authors", m.authors()}, + {"description", m.description()}, + {"references", m.references()}, + {"extra", m.extra()} + }; + } + + inline void from_json(const nlohmann::json& j, ModelMetadata& m) { + if (!j.is_object()) { + throw metatomic::Error("invalid JSON data for ModelMetadata, expected an object"); + } + + if (!j.contains("type") || !j["type"].is_string() || j["type"].get() != "metatomic_model_metadata") { + throw metatomic::Error("'type' in JSON for ModelMetadata must be 'metatomic_model_metadata'"); + } + + if (!j.contains("name") || !j["name"].is_string()) { + throw metatomic::Error("'name' in JSON for ModelMetadata must be a string"); + } + std::string name = j["name"].get(); + + auto authors = metatomic::detail::read_string_array(j, "authors", "JSON for ModelMetadata"); + + if (!j.contains("description") || !j["description"].is_string()) { + throw metatomic::Error("'description' in JSON for ModelMetadata must be a string"); + } + std::string description = j["description"].get(); + + if (!j.contains("references") || !j["references"].is_object()) { + throw metatomic::Error("invalid JSON data for references in ModelMetadata, expected an object"); + } + auto references = j["references"].get(); + + if (!j.contains("extra") || !j["extra"].is_object()) { + throw metatomic::Error("'extra' in JSON for ModelMetadata must be an object"); + } + std::map extra; + for (const auto& item : j["extra"].items()) { + if (!item.value().is_string()) { + throw metatomic::Error("'extra' in JSON for ModelMetadata must be an object with string values"); + } + extra[item.key()] = item.value().get(); + } + + // Validate authors content + for (const auto& author : authors) { + if (author.empty()) { + throw metatomic::Error("author can not be empty string in ModelMetadata"); + } + } + + // Validate references content + for (const auto& ref : references.model()) { + if (ref.empty()) { + throw metatomic::Error("reference can not be empty string (in 'model' section)"); + } + } + + for (const auto& ref : references.architecture()) { + if (ref.empty()) { + throw metatomic::Error("reference can not be empty string (in 'architecture' section)"); + } + } + + for (const auto& ref : references.implementation()) { + if (ref.empty()) { + throw metatomic::Error("reference can not be empty string (in 'implementation' section)"); + } + } + + m = ModelMetadata(name, authors, description, references, extra); + } + + /// The kind of samples a quantity can be associated with + enum class SampleKind { + /// The quantity is defined for each atom (e.g. atomic energy, charge, ...) + Atom, + /// The quantity is defined for the whole system (e.g. total energy, ...) + System, + /// The quantity is defined for each pair of atoms (e.g. hamiltonian elements, ...) + AtomPair, + }; + + /// The gradients a quantity can have + enum class Gradients { + /// Gradients with respect to atomic positions + Positions, + /// Gradients with respect to the strain (typically used for stress) + Strain, + }; + + /// A quantity that a model can use as input or output + struct Quantity { + private: + /// Name of the quantity, this can be a standard name from + /// https://docs.metatensor.org/metatomic/latest/quantities/index.html, or + /// a custom name of the form `::[/]` + std::optional name_; + /// Unit of the quantity + std::optional unit_; + /// Description of the quantity, used to provide more details about the + /// quantity, especially when a model defines multiple variants of the same + /// quantity. An empty string is treated as no description. + std::string description_; + /// List of explicit gradients for this quantity + std::vector gradients_; + /// The kind of samples this quantity is associated with + std::optional sample_kind_; + + public: + /// Set the name of this quantity. + void name(std::string value) { + name_ = std::move(value); + } + + /// Get the name of this quantity. + /// + /// @throw metatomic::Error if the value has not been set. + const std::string& name() const { + if (!name_.has_value()) { + throw metatomic::Error("name is not set in Quantity"); + } + return name_.value(); + } + + /// Set the unit of this quantity. + void unit(std::string value) { + unit_ = std::move(value); + } + + /// Get the unit of this quantity. + /// + /// @throw metatomic::Error if the value has not been set. + const std::string& unit() const { + if (!unit_.has_value()) { + throw metatomic::Error("unit is not set in Quantity"); + } + return unit_.value(); + } + + /// Set the description of this quantity. + void description(std::string value) { + description_ = std::move(value); + } + + /// Get the description of this quantity. + const std::string& description() const { + return description_; + } + + /// Set the list of explicit gradients for this quantity. + void gradients(std::vector value) { + gradients_ = std::move(value); + } + + /// Get the list of explicit gradients for this quantity. + const std::vector& gradients() const { + return gradients_; + } + + /// Add an explicit gradient to this quantity. + void add_gradient(Gradients gradient) { + gradients_.push_back(gradient); + } + + /// Clear the list of explicit gradients for this quantity. + void clear_gradients() { + gradients_.clear(); + } + + /// Set the kind of samples this quantity is associated with. + void sample_kind(const SampleKind& value) { + sample_kind_ = value; + } + + /// Get the kind of samples this quantity is associated with. + /// + /// @throw metatomic::Error if the value has not been set. + SampleKind sample_kind() const { + if (!sample_kind_.has_value()) { + throw metatomic::Error("sample_kind is not set in Quantity"); + } + return sample_kind_.value(); + } + + /// Create a default `Quantity`. The name, unit, and sample_kind fields + /// must be set before the object can be used. + Quantity() = default; + + /// Create a `Quantity` with the given values. + /// + /// @param name name of the quantity + /// @param unit unit of the quantity + /// @param sample_kind kind of samples this quantity is associated with + /// @param description description of the quantity + /// @param gradients list of explicit gradients for this quantity + Quantity( + std::string name, + std::string unit, + SampleKind sample_kind, + std::string description = "", + std::vector gradients = {} + ) { + this->name(std::move(name)); + this->unit(std::move(unit)); + this->sample_kind(sample_kind); + this->description(std::move(description)); + this->gradients(std::move(gradients)); + } + }; + + /// Capabilities of a model: which outputs it provides, which atoms it + /// supports, etc. + struct ModelCapabilities { + /// The data type of a model, used for all inputs and outputs. + enum class DType { + /// 32-bit floating point, following the IEEE 754 standard + Float32, + /// 64-bit floating point, following the IEEE 754 standard + Float64, + }; + + /// A device on which a model can run. + enum class Device { + CPU, + CUDA, + ROCM, + Metal, + }; + + using SampleKind = metatomic::SampleKind; ///< Alias for top-level `metatomic::SampleKind` + using Gradients = metatomic::Gradients; ///< Alias for top-level `metatomic::Gradients` + using Quantity = metatomic::Quantity; ///< Alias for top-level `metatomic::Quantity` + + private: + /// The outputs this model can provide + std::vector outputs_; + /// The atomic types this model supports. The meaning of the integers in + /// this list is up to the model, and is not required to be the atomic + /// numbers. + std::optional> atomic_types_; + /// The interaction range of the model (in the length unit of the model), + /// i.e. the maximum distance between two atoms for which the model's output + /// can depend on their relative position. + std::optional interaction_range_; + /// The length unit of the model, e.g. "angstrom" or "nanometer". This is + /// used to interpret the `interaction_range` and convert the inputs. + std::optional length_unit_; + /// The devices on which the model can run, e.g. `["cpu", "cuda"]`. + std::optional> supported_devices_; + /// The data type of the model, used for all inputs and outputs. + std::optional dtype_; + + public: + /// Set the list of outputs this model can provide. + void outputs(std::vector value) { + outputs_ = std::move(value); + } + + /// Get the list of outputs this model can provide. + const std::vector& outputs() const { + return outputs_; + } + + /// Add an output to the list of outputs this model can provide. + void add_output(const Quantity& output) { + outputs_.push_back(output); + } + + /// Clear the list of outputs this model can provide. + void clear_outputs() { + outputs_.clear(); + } + + /// Set the atomic types this model supports. + void atomic_types(std::vector value) { + atomic_types_ = std::move(value); + } + + /// Get the atomic types this model supports. + /// + /// @throw metatomic::Error if the value has not been set. + const std::vector& atomic_types() const { + if (!atomic_types_.has_value()) { + throw metatomic::Error("atomic_types is not set in ModelCapabilities"); + } + return atomic_types_.value(); + } + + /// Add an atomic type to the list of atomic types this model supports. + void add_atomic_type(int64_t atomic_type) { + if (!atomic_types_.has_value()) { + atomic_types_ = std::vector(); + } + atomic_types_->push_back(atomic_type); + } + + /// Clear the list of atomic types this model supports. + /// + /// If `atomic_types` has not been set, this function does nothing. + void clear_atomic_types() { + if (atomic_types_.has_value()) { + atomic_types_->clear(); + } + } + + /// Set the interaction range of the model. + /// + /// @throw metatomic::Error if the value is negative. + void interaction_range(double value) { + if (value < 0.0) { + throw metatomic::Error("interaction_range must be non-negative"); + } + interaction_range_ = value; + } + + /// Get the interaction range of the model. + /// + /// @throw metatomic::Error if the value has not been set. + double interaction_range() const { + if (!interaction_range_.has_value()) { + throw metatomic::Error("interaction_range is not set in ModelCapabilities"); + } + return interaction_range_.value(); + } + + /// Set the length unit of the model. + void length_unit(std::string value) { + length_unit_ = std::move(value); + } + + /// Get the length unit of the model. + /// + /// @throw metatomic::Error if the value has not been set. + const std::string& length_unit() const { + if (!length_unit_.has_value()) { + throw metatomic::Error("length_unit is not set in ModelCapabilities"); + } + return length_unit_.value(); + } + + /// Set the devices on which this model can run. + void supported_devices(std::vector value) { + supported_devices_ = std::move(value); + } + + /// Get the devices on which this model can run. + /// + /// @throw metatomic::Error if the value has not been set. + const std::vector& supported_devices() const { + if (!supported_devices_.has_value()) { + throw metatomic::Error("supported_devices is not set in ModelCapabilities"); + } + return supported_devices_.value(); + } + + /// Add a device to the list of devices on which this model can run. + void add_supported_device(Device device) { + if (!supported_devices_.has_value()) { + supported_devices_ = std::vector(); + } + supported_devices_->push_back(device); + } + + /// Clear the list of devices on which this model can run. + /// + /// If `supported_devices` has not been set, this function does nothing. + void clear_supported_devices() { + if (supported_devices_.has_value()) { + supported_devices_->clear(); + } + } + + /// Set the data type of the model. + void dtype(DType value) { + dtype_ = value; + } + + /// Get the data type of the model. + /// + /// @throw metatomic::Error if the value has not been set. + DType dtype() const { + if (!dtype_.has_value()) { + throw metatomic::Error("dtype is not set in ModelCapabilities"); + } + return dtype_.value(); + } + + /// Create a default `ModelCapabilities`. All fields must be set before + /// the object can be used. + ModelCapabilities() = default; + + /// Create a `ModelCapabilities` with the given values. + /// + /// @param atomic_types atomic types this model supports + /// @param interaction_range interaction range of the model + /// @param length_unit length unit of the model + /// @param supported_devices devices on which this model can run + /// @param dtype data type of the model + /// @param outputs outputs this model can provide + ModelCapabilities( + std::vector atomic_types, + double interaction_range, + std::string length_unit, + std::vector supported_devices, + DType dtype, + std::vector outputs = {} + ) { + this->atomic_types(std::move(atomic_types)); + this->interaction_range(interaction_range); + this->length_unit(std::move(length_unit)); + this->supported_devices(std::move(supported_devices)); + this->dtype(dtype); + this->outputs(std::move(outputs)); + } + }; + + inline void to_json(nlohmann::json& j, const ModelCapabilities::DType& dtype) { + switch (dtype) { + case ModelCapabilities::DType::Float32: + j = "float32"; + break; + case ModelCapabilities::DType::Float64: + j = "float64"; + break; + default: + throw metatomic::Error("invalid dtype in ModelCapabilities"); + } + } + + inline void from_json(const nlohmann::json& j, ModelCapabilities::DType& dtype) { + if (!j.is_string()) { + throw metatomic::Error("dtype in JSON for ModelCapabilities must be a string"); + } + + std::string s = j.get(); + if (s == "float32") { + dtype = ModelCapabilities::DType::Float32; + } else if (s == "float64") { + dtype = ModelCapabilities::DType::Float64; + } else { + throw metatomic::Error( + "invalid string for dtype in JSON for ModelCapabilities, expected 'float32' or 'float64'" + ); + } + } + + inline void to_json(nlohmann::json& j, const ModelCapabilities::Device& device) { + switch (device) { + case ModelCapabilities::Device::CPU: + j = "cpu"; + break; + case ModelCapabilities::Device::CUDA: + j = "cuda"; + break; + case ModelCapabilities::Device::ROCM: + j = "rocm"; + break; + case ModelCapabilities::Device::Metal: + j = "metal"; + break; + default: + throw metatomic::Error("invalid device in ModelCapabilities"); + } + } + + inline void from_json(const nlohmann::json& j, ModelCapabilities::Device& device) { + if (!j.is_string()) { + throw metatomic::Error("device in JSON for ModelCapabilities must be a string"); + } + + std::string s = j.get(); + if (s == "cpu") { + device = ModelCapabilities::Device::CPU; + } else if (s == "cuda") { + device = ModelCapabilities::Device::CUDA; + } else if (s == "rocm") { + device = ModelCapabilities::Device::ROCM; + } else if (s == "metal") { + device = ModelCapabilities::Device::Metal; + } else { + throw metatomic::Error( + "invalid string for device in JSON for ModelCapabilities, expected 'cpu', 'cuda', 'rocm', or 'metal'" + ); + } + } + + inline void to_json(nlohmann::json& j, const SampleKind& kind) { + switch (kind) { + case SampleKind::Atom: + j = "atom"; + break; + case SampleKind::System: + j = "system"; + break; + case SampleKind::AtomPair: + j = "atom_pair"; + break; + default: + throw metatomic::Error("invalid sample_kind in Quantity"); + } + } + + inline void from_json(const nlohmann::json& j, SampleKind& kind) { + if (!j.is_string()) { + throw metatomic::Error("'sample_kind' in JSON for Quantity must be a string"); + } + + std::string s = j.get(); + if (s == "atom") { + kind = SampleKind::Atom; + } else if (s == "system") { + kind = SampleKind::System; + } else if (s == "atom_pair") { + kind = SampleKind::AtomPair; + } else { + throw metatomic::Error( + "'sample_kind' in JSON for Quantity must be 'atom', 'system' or 'atom_pair', got '" + s + "'" + ); + } + } + + inline void to_json(nlohmann::json& j, const Gradients& gradients) { + switch (gradients) { + case Gradients::Positions: + j = "positions"; + break; + case Gradients::Strain: + j = "strain"; + break; + default: + throw metatomic::Error("invalid gradients in Quantity"); + } + } + + inline void from_json(const nlohmann::json& j, Gradients& gradients) { + if (!j.is_string()) { + throw metatomic::Error("'gradients' in JSON for Quantity must be a string"); + } + + std::string s = j.get(); + if (s == "positions") { + gradients = Gradients::Positions; + } else if (s == "strain") { + gradients = Gradients::Strain; + } else { + throw metatomic::Error( + "'gradients' in JSON for Quantity must be 'positions' or 'strain', got '" + s + "'" + ); + } + } + + inline void to_json(nlohmann::json& j, const Quantity& q) { + j = nlohmann::json{ + {"type", "metatomic_quantity"}, + {"name", q.name()}, + {"unit", q.unit()}, + {"gradients", q.gradients()}, + {"sample_kind", q.sample_kind()} + }; + + if (!q.description().empty()) { + j["description"] = q.description(); + } + } + + inline void from_json(const nlohmann::json& j, Quantity& q) { + if (!j.is_object()) { + throw metatomic::Error("invalid JSON data for Quantity, expected an object"); + } + + if (!j.contains("type") || !j["type"].is_string() || j["type"].get() != "metatomic_quantity") { + throw metatomic::Error("'type' in JSON for Quantity must be 'metatomic_quantity'"); + } + + if (!j.contains("name") || !j["name"].is_string()) { + throw metatomic::Error("'name' in JSON for Quantity must be a string"); + } + std::string name = j["name"].get(); + + if (!j.contains("unit") || !j["unit"].is_string()) { + throw metatomic::Error("'unit' in JSON for Quantity must be a string"); + } + std::string unit = j["unit"].get(); + + std::string description; + if (j.contains("description")) { + if (!j["description"].is_string()) { + throw metatomic::Error("'description' in JSON for Quantity must be a string"); + } + description = j["description"].get(); + } + + if (!j.contains("gradients") || !j["gradients"].is_array()) { + throw metatomic::Error("'gradients' in JSON for Quantity must be an array"); + } + std::vector gradients; + for (const auto& gradient : j["gradients"]) { + gradients.push_back(gradient.get()); + } + + if (!j.contains("sample_kind") || !j["sample_kind"].is_string()) { + throw metatomic::Error("'sample_kind' in JSON for Quantity must be a string"); + } + auto sample_kind = j["sample_kind"].get(); + + q = Quantity(name, unit, sample_kind, description, gradients); + } + + inline void to_json(nlohmann::json& j, const ModelCapabilities& c) { + j = nlohmann::json{ + {"type", "metatomic_model_capabilities"}, + {"outputs", c.outputs()}, + {"atomic_types", c.atomic_types()}, + {"interaction_range", c.interaction_range()}, + {"length_unit", c.length_unit()}, + {"supported_devices", c.supported_devices()}, + {"dtype", c.dtype()} + }; + } + + inline void from_json(const nlohmann::json& j, ModelCapabilities& c) { + if (!j.is_object()) { + throw metatomic::Error("invalid JSON data for ModelCapabilities, expected an object"); + } + + if (!j.contains("type") || !j["type"].is_string() || j["type"].get() != "metatomic_model_capabilities") { + throw metatomic::Error("'type' in JSON for ModelCapabilities must be 'metatomic_model_capabilities'"); + } + + if (!j.contains("outputs") || !j["outputs"].is_array()) { + throw metatomic::Error("'outputs' in JSON for ModelCapabilities must be an array"); + } + std::vector outputs; + for (const auto& output : j["outputs"]) { + outputs.push_back(output.get()); + } + + if (!j.contains("atomic_types") || !j["atomic_types"].is_array()) { + throw metatomic::Error("'atomic_types' in JSON for ModelCapabilities must be an array"); + } + std::vector atomic_types; + for (const auto& atomic_type : j["atomic_types"]) { + if (!atomic_type.is_number_integer()) { + throw metatomic::Error("'atomic_types' in JSON for ModelCapabilities must be an array of integers"); + } + atomic_types.push_back(atomic_type.get()); + } + + if (!j.contains("interaction_range") || !j["interaction_range"].is_number()) { + throw metatomic::Error("'interaction_range' in JSON for ModelCapabilities must be a number"); + } + double interaction_range = j["interaction_range"].get(); + if (interaction_range < 0.0) { + throw metatomic::Error("'interaction_range' in JSON for ModelCapabilities must be non-negative"); + } + + if (!j.contains("length_unit") || !j["length_unit"].is_string()) { + throw metatomic::Error("'length_unit' in JSON for ModelCapabilities must be a string"); + } + std::string length_unit = j["length_unit"].get(); + + // Validate that `length_unit` has the dimension of length by asking the + // C API for a conversion factor to meters. The call only succeeds when + // the dimensions match; otherwise `check_status` throws with the C API's + // dimension-mismatch message. + double conversion_factor = 0.0; + auto status = mta_unit_conversion_factor(length_unit.c_str(), "m", &conversion_factor); + metatomic::details::check_status(status); + + if (!j.contains("supported_devices") || !j["supported_devices"].is_array()) { + throw metatomic::Error("'supported_devices' in JSON for ModelCapabilities must be an array"); + } + std::vector supported_devices; + for (const auto& device : j["supported_devices"]) { + supported_devices.push_back(device.get()); + } + + if (!j.contains("dtype") || !j["dtype"].is_string()) { + throw metatomic::Error("dtype in JSON for ModelCapabilities must be a string"); + } + auto dtype = j["dtype"].get(); + + c = ModelCapabilities(atomic_types, interaction_range, length_unit, supported_devices, dtype, outputs); + } + +} // namespace metatomic diff --git a/metatomic-core/tests/CMakeLists.txt b/metatomic-core/tests/CMakeLists.txt index 731904948..efb9b1f7b 100644 --- a/metatomic-core/tests/CMakeLists.txt +++ b/metatomic-core/tests/CMakeLists.txt @@ -49,6 +49,7 @@ if (CMAKE_CXX_COMPILER_ID MATCHES "Clang") set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-unsafe-buffer-usage") set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-poison-system-directories") set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-allocator-wrappers") + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-covered-switch-default") endif() diff --git a/metatomic-core/tests/cxx/metadata.cpp b/metatomic-core/tests/cxx/metadata.cpp new file mode 100644 index 000000000..697f9814e --- /dev/null +++ b/metatomic-core/tests/cxx/metadata.cpp @@ -0,0 +1,910 @@ +#include + +#include "metatomic.hpp" + +TEST_CASE("JSON serialization C++ API") { + SECTION("PairListOptions"){ + double cutoff = 3.0; + std::string cutoff_hex = "0x4008000000000000"; + + SECTION("Constructor with default arguments") { + metatomic::PairListOptions p1(cutoff, true, false, {"model1", "model2"}); + + nlohmann::json j = p1; + + CHECK(j["cutoff"] == cutoff_hex); + CHECK(j["full_list"] == true); + CHECK(j["strict"] == false); + CHECK(j["requestors"].is_array()); + CHECK(j["requestors"].size() == 2); + CHECK(j["requestors"][0] == "model1"); + CHECK(j["requestors"][1] == "model2"); + + auto p2 = j.get(); + CHECK(p2.cutoff() == Approx(cutoff)); + CHECK(p2.full_list() == true); + CHECK(p2.strict() == false); + CHECK(p2.requestors().size() == 2); + CHECK(p2.requestors()[0] == "model1"); + CHECK(p2.requestors()[1] == "model2"); + } + + SECTION("Default constructor initialized with setters") { + metatomic::PairListOptions p1; + p1.cutoff(cutoff); + p1.full_list(true); + p1.strict(false); + p1.requestors({"model1", "model2"}); + + nlohmann::json j = p1; + + CHECK(j["cutoff"] == cutoff_hex); + CHECK(j["full_list"] == true); + CHECK(j["strict"] == false); + CHECK(j["requestors"].is_array()); + CHECK(j["requestors"].size() == 2); + CHECK(j["requestors"][0] == "model1"); + CHECK(j["requestors"][1] == "model2"); + + auto p2 = j.get(); + CHECK(p2.cutoff() == Approx(cutoff)); + CHECK(p2.full_list() == true); + CHECK(p2.strict() == false); + CHECK(p2.requestors().size() == 2); + CHECK(p2.requestors()[0] == "model1"); + CHECK(p2.requestors()[1] == "model2"); + } + + SECTION("add_requestor ignores empty strings and duplicates") { + metatomic::PairListOptions p1; + p1.cutoff(cutoff); + p1.full_list(true); + p1.add_requestor("model1"); + p1.add_requestor(""); + p1.add_requestor("model2"); + p1.add_requestor("model1"); + + auto requestors = p1.requestors(); + CHECK(requestors.size() == 2); + CHECK(requestors[0] == "model1"); + CHECK(requestors[1] == "model2"); + + nlohmann::json j = p1; + CHECK(j["requestors"].size() == 2); + CHECK(j["requestors"][0] == "model1"); + CHECK(j["requestors"][1] == "model2"); + } + + SECTION("clear_requestors empties the list") { + metatomic::PairListOptions p1; + p1.cutoff(cutoff); + p1.full_list(true); + p1.requestors({"model1", "model2"}); + p1.clear_requestors(); + + CHECK(p1.requestors().empty()); + + nlohmann::json j = p1; + CHECK(j["requestors"].is_array()); + CHECK(j["requestors"].size() == 0); + } + } + + SECTION("References") { + SECTION("Constructor") { + metatomic::ModelMetadata::References r1( + {"model ref 1", "model ref 2"}, + {"architecture ref 1"}, + {"implementation ref 1", "implementation ref 2"} + ); + + nlohmann::json j = r1; + + CHECK(j["model"].is_array()); + CHECK(j["model"].size() == 2); + CHECK(j["model"][0] == "model ref 1"); + CHECK(j["model"][1] == "model ref 2"); + + CHECK(j["architecture"].is_array()); + CHECK(j["architecture"].size() == 1); + CHECK(j["architecture"][0] == "architecture ref 1"); + + CHECK(j["implementation"].is_array()); + CHECK(j["implementation"].size() == 2); + CHECK(j["implementation"][0] == "implementation ref 1"); + CHECK(j["implementation"][1] == "implementation ref 2"); + + auto r2 = j.get(); + CHECK(r2.model()[0] == "model ref 1"); + CHECK(r2.model()[1] == "model ref 2"); + CHECK(r2.architecture().size() == 1); + CHECK(r2.architecture()[0] == "architecture ref 1"); + CHECK(r2.implementation().size() == 2); + CHECK(r2.implementation()[0] == "implementation ref 1"); + CHECK(r2.implementation()[1] == "implementation ref 2"); + } + + SECTION("Default constructor initialized with setters") { + metatomic::ModelMetadata::References r1; + r1.model({"model ref 1", "model ref 2"}); + r1.architecture({"architecture ref 1"}); + r1.implementation({"implementation ref 1", "implementation ref 2"}); + + nlohmann::json j = r1; + + CHECK(j["model"].size() == 2); + CHECK(j["model"][0] == "model ref 1"); + CHECK(j["architecture"].size() == 1); + CHECK(j["implementation"].size() == 2); + + auto r2 = j.get(); + CHECK(r2.model()[0] == "model ref 1"); + CHECK(r2.architecture().size() == 1); + CHECK(r2.implementation().size() == 2); + } + + SECTION("add and clear reference sections") { + metatomic::ModelMetadata::References r1; + r1.add_model("model ref 1"); + r1.add_model("model ref 2"); + r1.add_architecture("architecture ref 1"); + r1.add_implementation("implementation ref 1"); + r1.add_implementation("implementation ref 2"); + + CHECK(r1.model().size() == 2); + CHECK(r1.model()[0] == "model ref 1"); + CHECK(r1.model()[1] == "model ref 2"); + CHECK(r1.architecture().size() == 1); + CHECK(r1.architecture()[0] == "architecture ref 1"); + CHECK(r1.implementation().size() == 2); + CHECK(r1.implementation()[1] == "implementation ref 2"); + + r1.clear_model(); + CHECK(r1.model().empty()); + CHECK(r1.architecture().size() == 1); + + r1.clear_architecture(); + r1.clear_implementation(); + CHECK(r1.architecture().empty()); + CHECK(r1.implementation().empty()); + } + } + + SECTION("ModelMetadata") { + auto create_example = []() { + return metatomic::ModelMetadata( + "test-model", + {"Alice", "Bob"}, + "A test model", + metatomic::ModelMetadata::References( + {"doi:10.1234/test"}, + {"doi:10.1234/arch"}, + {"https://github.com/test"} + ), + std::map{ + {"key1", "value1"}, + {"key2", "value2"} + } + ); + }; + + auto create_example_with_setters = []() { + metatomic::ModelMetadata metadata; + metadata.name("test-model"); + metadata.authors({"Alice", "Bob"}); + metadata.description("A test model"); + metadata.references(metatomic::ModelMetadata::References( + {"doi:10.1234/test"}, + {"doi:10.1234/arch"}, + {"https://github.com/test"} + )); + metadata.extra(std::map{ + {"key1", "value1"}, + {"key2", "value2"} + }); + return metadata; + }; + + SECTION("JSON roundtrip conversion with constructor") { + auto m1 = create_example(); + nlohmann::json j = m1; + + CHECK(j["type"] == "metatomic_model_metadata"); + CHECK(j["name"] == "test-model"); + CHECK(j["authors"].is_array()); + CHECK(j["authors"].size() == 2); + CHECK(j["authors"][0] == "Alice"); + CHECK(j["authors"][1] == "Bob"); + CHECK(j["description"] == "A test model"); + CHECK(j["references"]["model"][0] == "doi:10.1234/test"); + CHECK(j["references"]["architecture"][0] == "doi:10.1234/arch"); + CHECK(j["references"]["implementation"][0] == "https://github.com/test"); + CHECK(j["extra"]["key1"] == "value1"); + CHECK(j["extra"]["key2"] == "value2"); + + auto m2 = j.get(); + CHECK(m2.name() == m1.name()); + CHECK(m2.authors() == m1.authors()); + CHECK(m2.description() == m1.description()); + CHECK(m2.references().model() == m1.references().model()); + CHECK(m2.references().architecture() == m1.references().architecture()); + CHECK(m2.references().implementation() == m1.references().implementation()); + CHECK(m2.extra() == m1.extra()); + } + + SECTION("JSON roundtrip conversion with default constructor and setters") { + auto m1 = create_example_with_setters(); + nlohmann::json j = m1; + + CHECK(j["type"] == "metatomic_model_metadata"); + CHECK(j["name"] == "test-model"); + CHECK(j["authors"].is_array()); + CHECK(j["authors"].size() == 2); + CHECK(j["authors"][0] == "Alice"); + CHECK(j["authors"][1] == "Bob"); + CHECK(j["description"] == "A test model"); + CHECK(j["references"]["model"][0] == "doi:10.1234/test"); + CHECK(j["references"]["architecture"][0] == "doi:10.1234/arch"); + CHECK(j["references"]["implementation"][0] == "https://github.com/test"); + CHECK(j["extra"]["key1"] == "value1"); + CHECK(j["extra"]["key2"] == "value2"); + + auto m2 = j.get(); + CHECK(m2.name() == m1.name()); + CHECK(m2.authors() == m1.authors()); + CHECK(m2.description() == m1.description()); + CHECK(m2.references().model() == m1.references().model()); + CHECK(m2.references().architecture() == m1.references().architecture()); + CHECK(m2.references().implementation() == m1.references().implementation()); + CHECK(m2.extra() == m1.extra()); + } + + SECTION("Invalid JSON data") { + auto m1 = create_example(); + nlohmann::json j = m1; + + CHECK_THROWS_WITH( + nlohmann::json("not an object").get(), + Catch::Matchers::StartsWith("invalid JSON data for ModelMetadata, expected an object") + ); + + { + auto j_copy = j; + j_copy["type"] = "something-else"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'type' in JSON for ModelMetadata must be 'metatomic_model_metadata'") + ); + } + + { + auto j_copy = j; + j_copy.erase("name"); + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'name' in JSON for ModelMetadata must be a string") + ); + } + + { + auto j_copy = j; + j_copy["name"] = 42; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'name' in JSON for ModelMetadata must be a string") + ); + } + + { + auto j_copy = j; + j_copy["authors"] = "Alice"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'authors' in JSON for ModelMetadata must be an array") + ); + } + + { + auto j_copy = j; + j_copy["authors"] = {"Alice", 42}; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'authors' in JSON for ModelMetadata must be an array of strings") + ); + } + + { + auto j_copy = j; + j_copy.erase("description"); + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'description' in JSON for ModelMetadata must be a string") + ); + } + + { + auto j_copy = j; + j_copy["extra"] = "not-an-object"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'extra' in JSON for ModelMetadata must be an object") + ); + } + + { + auto j_copy = j; + j_copy["extra"] = {{"key", 42}}; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'extra' in JSON for ModelMetadata must be an object with string values") + ); + } + + { + auto j_copy = j; + j_copy["references"] = "not-an-object"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("invalid JSON data for references in ModelMetadata, expected an object") + ); + } + } + + SECTION("Model metadata formatting") { + auto m1 = create_example(); + std::string output = m1.print(); + std::string expected = + "This is the test-model model\n" + "============================\n" + "\n" + "A test model\n" + "\n" + "Model authors\n" + "-------------\n" + "\n" + "- Alice\n" + "- Bob\n" + "\n" + "Model references\n" + "----------------\n" + "\n" + "Please cite the following references when using this model:\n" + "- about this specific model:\n" + " * doi:10.1234/test\n" + "- about the architecture of this model:\n" + " * doi:10.1234/arch\n" + "- about the implementation of this model:\n" + " * https://github.com/test\n"; + + CHECK(output == expected); + } + + SECTION("add and clear authors, references, and extra") { + metatomic::ModelMetadata m1; + m1.name("test-model"); + m1.add_author("Alice"); + m1.add_author("Bob"); + m1.add_reference("model", "doi:10.1234/test"); + m1.add_reference("architecture", "doi:10.1234/arch"); + m1.add_reference("implementation", "https://github.com/test"); + m1.add_extra("key1", "value1"); + m1.add_extra("key2", "value2"); + + CHECK(m1.authors().size() == 2); + CHECK(m1.authors()[0] == "Alice"); + CHECK(m1.authors()[1] == "Bob"); + CHECK(m1.references().model().size() == 1); + CHECK(m1.references().architecture().size() == 1); + CHECK(m1.references().implementation().size() == 1); + CHECK(m1.extra().size() == 2); + CHECK(m1.extra().at("key1") == "value1"); + CHECK(m1.extra().at("key2") == "value2"); + + nlohmann::json j = m1; + CHECK(j["authors"].size() == 2); + CHECK(j["references"]["model"].size() == 1); + CHECK(j["extra"].size() == 2); + + m1.clear_reference("model"); + CHECK(m1.references().model().empty()); + CHECK(m1.references().architecture().size() == 1); + CHECK(m1.references().implementation().size() == 1); + + m1.clear_authors(); + m1.clear_references(); + m1.clear_extra(); + CHECK(m1.authors().empty()); + CHECK(m1.references().model().empty()); + CHECK(m1.references().architecture().empty()); + CHECK(m1.references().implementation().empty()); + CHECK(m1.extra().empty()); + + CHECK_THROWS_WITH( + m1.add_reference("invalid", "ref"), + Catch::Matchers::StartsWith("reference section must be 'model', 'architecture', or 'implementation', got 'invalid'") + ); + + CHECK_THROWS_WITH( + m1.clear_reference("invalid"), + Catch::Matchers::StartsWith("reference section must be 'model', 'architecture', or 'implementation', got 'invalid'") + ); + } + } + + SECTION("DType") { + SECTION("JSON roundtrip conversion") { + auto dtype1 = metatomic::ModelCapabilities::DType::Float32; + nlohmann::json j = dtype1; + CHECK(j == "float32"); + auto dtype2 = j.get(); + CHECK(dtype2 == metatomic::ModelCapabilities::DType::Float32); + + auto dtype3 = metatomic::ModelCapabilities::DType::Float64; + nlohmann::json j2 = dtype3; + CHECK(j2 == "float64"); + auto dtype4 = j2.get(); + CHECK(dtype4 == metatomic::ModelCapabilities::DType::Float64); + } + + SECTION("Invalid JSON data") { + CHECK_THROWS_WITH( + nlohmann::json(42).get(), + Catch::Matchers::StartsWith("dtype in JSON for ModelCapabilities must be a string") + ); + + CHECK_THROWS_WITH( + nlohmann::json("float16").get(), + Catch::Matchers::StartsWith("invalid string for dtype in JSON for ModelCapabilities, expected 'float32' or 'float64'") + ); + } + } + + SECTION("Quantity") { + SECTION("JSON roundtrip conversion with description") { + metatomic::Quantity q1( + "energy", + "eV", + metatomic::SampleKind::System, + "total energy of the system", + {metatomic::Gradients::Positions} + ); + nlohmann::json j = q1; + + CHECK(j["type"] == "metatomic_quantity"); + CHECK(j["name"] == "energy"); + CHECK(j["unit"] == "eV"); + CHECK(j["description"] == "total energy of the system"); + CHECK(j["gradients"].is_array()); + CHECK(j["gradients"].size() == 1); + CHECK(j["gradients"][0] == "positions"); + CHECK(j["sample_kind"] == "system"); + + auto q2 = j.get(); + CHECK(q2.name() == q1.name()); + CHECK(q2.unit() == q1.unit()); + CHECK(q2.description() == q1.description()); + CHECK(q2.gradients() == q1.gradients()); + CHECK(q2.sample_kind() == q1.sample_kind()); + } + + SECTION("JSON roundtrip conversion without description") { + metatomic::Quantity q1( + "charge", + "e", + metatomic::SampleKind::Atom, + "", + {} + ); + nlohmann::json j = q1; + + CHECK(j["type"] == "metatomic_quantity"); + CHECK(j["name"] == "charge"); + CHECK(j["unit"] == "e"); + CHECK(!j.contains("description")); + CHECK(j["gradients"].is_array()); + CHECK(j["gradients"].size() == 0); + CHECK(j["sample_kind"] == "atom"); + + auto q2 = j.get(); + CHECK(q2.name() == q1.name()); + CHECK(q2.unit() == q1.unit()); + CHECK(q2.description().empty()); + CHECK(q2.gradients().empty()); + CHECK(q2.sample_kind() == q1.sample_kind()); + } + + SECTION("Default constructor initialized with setters") { + metatomic::Quantity q1; + q1.name("energy"); + q1.unit("eV"); + q1.sample_kind(metatomic::SampleKind::System); + q1.description("total energy of the system"); + q1.gradients({metatomic::Gradients::Positions}); + + nlohmann::json j = q1; + + CHECK(j["type"] == "metatomic_quantity"); + CHECK(j["name"] == "energy"); + CHECK(j["unit"] == "eV"); + CHECK(j["description"] == "total energy of the system"); + CHECK(j["gradients"].size() == 1); + CHECK(j["gradients"][0] == "positions"); + CHECK(j["sample_kind"] == "system"); + + auto q2 = j.get(); + CHECK(q2.name() == q1.name()); + CHECK(q2.unit() == q1.unit()); + CHECK(q2.description() == q1.description()); + CHECK(q2.gradients() == q1.gradients()); + CHECK(q2.sample_kind() == q1.sample_kind()); + } + + SECTION("add and clear gradients") { + metatomic::Quantity q1; + q1.name("energy"); + q1.unit("eV"); + q1.sample_kind(metatomic::SampleKind::System); + q1.add_gradient(metatomic::Gradients::Positions); + q1.add_gradient(metatomic::Gradients::Strain); + + CHECK(q1.gradients().size() == 2); + CHECK(q1.gradients()[0] == metatomic::Gradients::Positions); + CHECK(q1.gradients()[1] == metatomic::Gradients::Strain); + + nlohmann::json j = q1; + CHECK(j["gradients"].size() == 2); + CHECK(j["gradients"][0] == "positions"); + CHECK(j["gradients"][1] == "strain"); + + q1.clear_gradients(); + CHECK(q1.gradients().empty()); + } + + SECTION("Empty description is treated as no description") { + nlohmann::json j = { + {"type", "metatomic_quantity"}, + {"name", "charge"}, + {"unit", "e"}, + {"description", ""}, + {"gradients", nlohmann::json::array()}, + {"sample_kind", "atom"} + }; + + auto q = j.get(); + CHECK(q.name() == "charge"); + CHECK(q.unit() == "e"); + CHECK(q.description().empty()); + CHECK(q.gradients().empty()); + CHECK(q.sample_kind() == metatomic::SampleKind::Atom); + } + + SECTION("Invalid JSON data") { + CHECK_THROWS_WITH( + nlohmann::json("not an object").get(), + Catch::Matchers::StartsWith("invalid JSON data for Quantity, expected an object") + ); + + { + nlohmann::json j = {{"type", "wrong-type"}}; + CHECK_THROWS_WITH( + j.get(), + Catch::Matchers::StartsWith("'type' in JSON for Quantity must be 'metatomic_quantity'") + ); + } + + { + nlohmann::json j = { + {"type", "metatomic_quantity"}, + {"name", 42} + }; + CHECK_THROWS_WITH( + j.get(), + Catch::Matchers::StartsWith("'name' in JSON for Quantity must be a string") + ); + } + + { + nlohmann::json j = { + {"type", "metatomic_quantity"}, + {"name", "energy"}, + {"unit", "eV"}, + {"gradients", "positions"} + }; + CHECK_THROWS_WITH( + j.get(), + Catch::Matchers::StartsWith("'gradients' in JSON for Quantity must be an array") + ); + } + + { + nlohmann::json j = { + {"type", "metatomic_quantity"}, + {"name", "energy"}, + {"unit", "eV"}, + {"gradients", {"positions"}}, + {"sample_kind", "unknown"} + }; + CHECK_THROWS_WITH( + j.get(), + Catch::Matchers::StartsWith("'sample_kind' in JSON for Quantity must be 'atom', 'system' or 'atom_pair', got 'unknown'") + ); + } + } + } + + SECTION("ModelCapabilities") { + auto create_example = []() { + std::vector outputs = { + metatomic::Quantity( + "energy", + "eV", + metatomic::SampleKind::System, + "total energy", + {metatomic::Gradients::Positions} + ), + metatomic::Quantity( + "charge", + "e", + metatomic::SampleKind::Atom, + "", + {} + ) + }; + + return metatomic::ModelCapabilities( + {1, 6, 8}, + 5.0, + "Angstrom", + {metatomic::ModelCapabilities::Device::CPU, metatomic::ModelCapabilities::Device::CUDA}, + metatomic::ModelCapabilities::DType::Float32, + outputs + ); + }; + + auto create_example_with_setters = []() { + std::vector outputs = { + metatomic::Quantity( + "energy", + "eV", + metatomic::SampleKind::System, + "total energy", + {metatomic::Gradients::Positions} + ), + metatomic::Quantity( + "charge", + "e", + metatomic::SampleKind::Atom, + "", + {} + ) + }; + + metatomic::ModelCapabilities capabilities; + capabilities.atomic_types({1, 6, 8}); + capabilities.interaction_range(5.0); + capabilities.length_unit("Angstrom"); + capabilities.supported_devices({metatomic::ModelCapabilities::Device::CPU, metatomic::ModelCapabilities::Device::CUDA}); + capabilities.dtype(metatomic::ModelCapabilities::DType::Float32); + capabilities.outputs(outputs); + return capabilities; + }; + + SECTION("JSON roundtrip conversion with constructor") { + auto c1 = create_example(); + nlohmann::json j = c1; + + CHECK(j["type"] == "metatomic_model_capabilities"); + CHECK(j["outputs"].is_array()); + CHECK(j["outputs"].size() == 2); + CHECK(j["outputs"][0]["name"] == "energy"); + CHECK(j["outputs"][1]["name"] == "charge"); + CHECK(j["atomic_types"].is_array()); + CHECK(j["atomic_types"].size() == 3); + CHECK(j["atomic_types"][0] == 1); + CHECK(j["atomic_types"][1] == 6); + CHECK(j["atomic_types"][2] == 8); + CHECK(j["interaction_range"] == Approx(5.0)); + CHECK(j["length_unit"] == "Angstrom"); + CHECK(j["supported_devices"].is_array()); + CHECK(j["supported_devices"].size() == 2); + CHECK(j["supported_devices"][0] == "cpu"); + CHECK(j["supported_devices"][1] == "cuda"); + CHECK(j["dtype"] == "float32"); + + auto c2 = j.get(); + CHECK(c2.outputs().size() == c1.outputs().size()); + CHECK(c2.outputs()[0].name() == c1.outputs()[0].name()); + CHECK(c2.outputs()[1].name() == c1.outputs()[1].name()); + CHECK(c2.atomic_types() == c1.atomic_types()); + CHECK(c2.interaction_range() == Approx(c1.interaction_range())); + CHECK(c2.length_unit() == c1.length_unit()); + CHECK(c2.supported_devices() == c1.supported_devices()); + CHECK(c2.dtype() == c1.dtype()); + } + + SECTION("JSON roundtrip conversion with default constructor and setters") { + auto c1 = create_example_with_setters(); + nlohmann::json j = c1; + + CHECK(j["type"] == "metatomic_model_capabilities"); + CHECK(j["outputs"].is_array()); + CHECK(j["outputs"].size() == 2); + CHECK(j["outputs"][0]["name"] == "energy"); + CHECK(j["outputs"][1]["name"] == "charge"); + CHECK(j["atomic_types"].is_array()); + CHECK(j["atomic_types"].size() == 3); + CHECK(j["atomic_types"][0] == 1); + CHECK(j["atomic_types"][1] == 6); + CHECK(j["atomic_types"][2] == 8); + CHECK(j["interaction_range"] == Approx(5.0)); + CHECK(j["length_unit"] == "Angstrom"); + CHECK(j["supported_devices"].is_array()); + CHECK(j["supported_devices"].size() == 2); + CHECK(j["supported_devices"][0] == "cpu"); + CHECK(j["supported_devices"][1] == "cuda"); + CHECK(j["dtype"] == "float32"); + + auto c2 = j.get(); + CHECK(c2.outputs().size() == c1.outputs().size()); + CHECK(c2.outputs()[0].name() == c1.outputs()[0].name()); + CHECK(c2.outputs()[1].name() == c1.outputs()[1].name()); + CHECK(c2.atomic_types() == c1.atomic_types()); + CHECK(c2.interaction_range() == Approx(c1.interaction_range())); + CHECK(c2.length_unit() == c1.length_unit()); + CHECK(c2.supported_devices() == c1.supported_devices()); + CHECK(c2.dtype() == c1.dtype()); + } + + SECTION("add and clear outputs, atomic types, and supported devices") { + metatomic::ModelCapabilities c1; + c1.interaction_range(5.0); + c1.length_unit("Angstrom"); + c1.dtype(metatomic::ModelCapabilities::DType::Float32); + + c1.add_output(metatomic::Quantity( + "energy", + "eV", + metatomic::SampleKind::System, + "total energy", + {metatomic::Gradients::Positions} + )); + c1.add_output(metatomic::Quantity( + "charge", + "e", + metatomic::SampleKind::Atom, + "", + {} + )); + + c1.add_atomic_type(1); + c1.add_atomic_type(6); + c1.add_atomic_type(8); + + c1.add_supported_device(metatomic::ModelCapabilities::Device::CPU); + c1.add_supported_device(metatomic::ModelCapabilities::Device::CUDA); + + CHECK(c1.outputs().size() == 2); + CHECK(c1.outputs()[0].name() == "energy"); + CHECK(c1.outputs()[1].name() == "charge"); + CHECK(c1.atomic_types().size() == 3); + CHECK(c1.atomic_types()[0] == 1); + CHECK(c1.atomic_types()[1] == 6); + CHECK(c1.atomic_types()[2] == 8); + CHECK(c1.supported_devices().size() == 2); + CHECK(c1.supported_devices()[0] == metatomic::ModelCapabilities::Device::CPU); + CHECK(c1.supported_devices()[1] == metatomic::ModelCapabilities::Device::CUDA); + + nlohmann::json j = c1; + CHECK(j["outputs"].size() == 2); + CHECK(j["atomic_types"].size() == 3); + CHECK(j["supported_devices"].size() == 2); + + c1.clear_outputs(); + c1.clear_atomic_types(); + c1.clear_supported_devices(); + CHECK(c1.outputs().empty()); + CHECK(c1.atomic_types().empty()); + CHECK(c1.supported_devices().empty()); + } + + SECTION("Invalid JSON data") { + auto c1 = create_example(); + nlohmann::json j = c1; + + CHECK_THROWS_WITH( + nlohmann::json("not an object").get(), + Catch::Matchers::StartsWith("invalid JSON data for ModelCapabilities, expected an object") + ); + + { + auto j_copy = j; + j_copy["type"] = "something-else"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'type' in JSON for ModelCapabilities must be 'metatomic_model_capabilities'") + ); + } + + { + auto j_copy = j; + j_copy["outputs"] = "energy"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'outputs' in JSON for ModelCapabilities must be an array") + ); + } + + { + auto j_copy = j; + j_copy["atomic_types"] = "1"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'atomic_types' in JSON for ModelCapabilities must be an array") + ); + } + + { + auto j_copy = j; + j_copy["atomic_types"] = {1, "x"}; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'atomic_types' in JSON for ModelCapabilities must be an array of integers") + ); + } + + { + auto j_copy = j; + j_copy.erase("interaction_range"); + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'interaction_range' in JSON for ModelCapabilities must be a number") + ); + } + + { + auto j_copy = j; + j_copy["interaction_range"] = -1.0; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'interaction_range' in JSON for ModelCapabilities must be non-negative") + ); + } + + { + auto j_copy = j; + j_copy["length_unit"] = "eV"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("invalid parameter: dimension mismatch") + ); + } + + { + auto j_copy = j; + j_copy["supported_devices"] = "cpu"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("'supported_devices' in JSON for ModelCapabilities must be an array") + ); + } + + { + auto j_copy = j; + j_copy["supported_devices"] = {"cpu", "wat"}; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("invalid string for device in JSON for ModelCapabilities, expected 'cpu', 'cuda', 'rocm', or 'metal'") + ); + } + + { + auto j_copy = j; + j_copy["dtype"] = "float16"; + CHECK_THROWS_WITH( + j_copy.get(), + Catch::Matchers::StartsWith("invalid string for dtype in JSON for ModelCapabilities, expected 'float32' or 'float64'") + ); + } + } + } +} From 981789c27f3644c5c615a189ed392734852a1c17 Mon Sep 17 00:00:00 2001 From: frostedoyster Date: Wed, 8 Jul 2026 19:27:29 +0200 Subject: [PATCH 34/43] Implement System in the C++ API --- docs/src/core/reference/cxx/model.rst | 11 + docs/src/core/reference/cxx/system.rst | 6 + metatomic-core/include/metatomic/metadata.hpp | 18 +- metatomic-core/include/metatomic/model.hpp | 6 +- metatomic-core/include/metatomic/system.hpp | 313 ++++++++++++++++++ metatomic-core/include/metatomic/utils.hpp | 86 +++++ metatomic-core/tests/cxx/system.cpp | 299 +++++++++++++++++ 7 files changed, 727 insertions(+), 12 deletions(-) create mode 100644 metatomic-core/tests/cxx/system.cpp diff --git a/docs/src/core/reference/cxx/model.rst b/docs/src/core/reference/cxx/model.rst index 75338c89d..5411064fa 100644 --- a/docs/src/core/reference/cxx/model.rst +++ b/docs/src/core/reference/cxx/model.rst @@ -1,2 +1,13 @@ Model ===== + +.. TODO: Model classes + +.. doxygenclass:: metatomic::Quantity + :members: + +.. doxygenclass:: metatomic::ModelMetadata + :members: + +.. doxygenclass:: metatomic::ModelCapabilities + :members: diff --git a/docs/src/core/reference/cxx/system.rst b/docs/src/core/reference/cxx/system.rst index 3dcbaeea1..c6471bfce 100644 --- a/docs/src/core/reference/cxx/system.rst +++ b/docs/src/core/reference/cxx/system.rst @@ -1,2 +1,8 @@ System ====== + +.. doxygenclass:: metatomic::System + :members: + +.. doxygenclass:: metatomic::PairListOptions + :members: diff --git a/metatomic-core/include/metatomic/metadata.hpp b/metatomic-core/include/metatomic/metadata.hpp index f539766cc..617312d60 100644 --- a/metatomic-core/include/metatomic/metadata.hpp +++ b/metatomic-core/include/metatomic/metadata.hpp @@ -15,7 +15,7 @@ #include #include -namespace metatomic{ +namespace metatomic { namespace detail { inline std::vector read_string_array( @@ -38,7 +38,7 @@ namespace metatomic{ } // namespace detail /// Options for the calculation of a pair list (neighbor list) - struct PairListOptions{ + class PairListOptions final { private: /// Cutoff radius for this pair list in the length unit of the model std::optional cutoff_; @@ -250,14 +250,15 @@ namespace metatomic{ // Forward declarations // The ModelMetadata::print function uses to_json - struct ModelMetadata; + class ModelMetadata; void to_json(nlohmann::json&, const ModelMetadata&); - struct ModelMetadata { + class ModelMetadata final { + public: /// References for a model, divided into three categories: references about /// the model as a whole, references about the architecture of the model, /// and references about the implementation of the model. - struct References { + class References final { private: /// The references about the model as a whole, e.g. a paper describing the /// model or a website presenting it. @@ -630,7 +631,7 @@ namespace metatomic{ }; /// A quantity that a model can use as input or output - struct Quantity { + class Quantity final { private: /// Name of the quantity, this can be a standard name from /// https://docs.metatensor.org/metatomic/latest/quantities/index.html, or @@ -751,7 +752,8 @@ namespace metatomic{ /// Capabilities of a model: which outputs it provides, which atoms it /// supports, etc. - struct ModelCapabilities { + class ModelCapabilities final { + public: /// The data type of a model, used for all inputs and outputs. enum class DType { /// 32-bit floating point, following the IEEE 754 standard @@ -768,7 +770,7 @@ namespace metatomic{ Metal, }; - using SampleKind = metatomic::SampleKind; ///< Alias for top-level `metatomic::SampleKind` + using SampleKind = metatomic::SampleKind; ///< Alias for top-level `metatomic::SampleKind` using Gradients = metatomic::Gradients; ///< Alias for top-level `metatomic::Gradients` using Quantity = metatomic::Quantity; ///< Alias for top-level `metatomic::Quantity` diff --git a/metatomic-core/include/metatomic/model.hpp b/metatomic-core/include/metatomic/model.hpp index fc4df3911..640bedb5d 100644 --- a/metatomic-core/include/metatomic/model.hpp +++ b/metatomic-core/include/metatomic/model.hpp @@ -4,6 +4,7 @@ #include #include +#include namespace metatomic { /// Render model metadata as a human-readable string. @@ -16,9 +17,6 @@ namespace metatomic { auto status = mta_format_metadata(metadata.c_str(), &printed); details::check_status(status); - auto result = std::string(mta_string_view(printed)); - mta_string_free(printed); - - return result; + return details::string_from_mta(printed); } } // namespace metatomic diff --git a/metatomic-core/include/metatomic/system.hpp b/metatomic-core/include/metatomic/system.hpp index 1cae91bdf..e6cc37d81 100644 --- a/metatomic-core/include/metatomic/system.hpp +++ b/metatomic-core/include/metatomic/system.hpp @@ -1,7 +1,320 @@ #pragma once +#include +#include +#include + #include +#include + +#include +#include +#include namespace metatomic { + /// A `System` contains all the information about an atomistic system, and is + /// used as the input of atomistic models. + /// + /// This is a RAII wrapper around the `mta_system_t` type from the C API. It + /// can either own the underlying system (in which case it is freed with the + /// `System`), or be a non-owning view into a system owned elsewhere (for + /// example a system passed to a model by the runtime). + class System final { + public: + /// Create a new `System` from DLPack tensors. + /// + /// Ownership of all four tensors is transferred to the new `System`. + /// + /// @param length_unit unit of length used by `positions` and `cell` + /// @param types tensor with shape `(n_atoms,)` of atomic types + /// @param positions tensor with shape `(n_atoms, 3)` of atomic positions + /// @param cell tensor with shape `(3, 3)` of the unit cell vectors + /// @param pbc tensor with shape `(3,)` of periodic boundary conditions + /// + /// The dtype and layout required for each tensor are validated by + /// `mta_system_create`; see the C API documentation for details. + System( + const std::string& length_unit, + DLPackTensor types, + DLPackTensor positions, + DLPackTensor cell, + DLPackTensor pbc + ) { + auto status = mta_system_create( + length_unit.c_str(), + types.release(), + positions.release(), + cell.release(), + pbc.release(), + &system_ + ); + details::check_status(status); + details::check_pointer(system_); + } + + ~System() { + if (!is_view_) { + // `mta_system_free` is a no-op on a null pointer + mta_system_free(system_); + } + } + + /// `System` is not copy-constructible + System(const System&) = delete; + /// `System` is not copy-assignable + System& operator=(const System&) = delete; + + /// `System` is move-constructible + System(System&& other) noexcept { + *this = std::move(other); + } + + /// `System` is move-assignable + System& operator=(System&& other) noexcept { + if (!is_view_) { + mta_system_free(system_); + } + + system_ = other.system_; + is_view_ = other.is_view_; + + other.system_ = nullptr; + other.is_view_ = true; + + return *this; + } + + /// Get the number of atoms in this system. + size_t size() const { + uintptr_t size = 0; + auto status = mta_system_size(system_, &size); + details::check_status(status); + return static_cast(size); + } + + /// Get the unit of length used by the positions and cell of this system. + std::string length_unit() const { + mta_string_t length_unit = nullptr; + auto status = mta_system_get_length_unit(system_, &length_unit); + details::check_status(status); + return details::string_from_mta(length_unit); + } + + /// Get the atomic types of all atoms in this system, as a tensor with + /// shape `(n_atoms,)`. + /// + /// @see `data` for the meaning of the returned tensor. + DLPackTensor types() const { + return this->data(MTA_SYSTEM_DATA_TYPES); + } + + /// Get the positions of all atoms in this system, as a tensor with shape + /// `(n_atoms, 3)`. + /// + /// @see `data` for the meaning of the returned tensor. + DLPackTensor positions() const { + return this->data(MTA_SYSTEM_DATA_POSITIONS); + } + + /// Get the unit cell of this system, as a tensor with shape `(3, 3)`. + /// + /// @see `data` for the meaning of the returned tensor. + DLPackTensor cell() const { + return this->data(MTA_SYSTEM_DATA_CELL); + } + + /// Get the periodic boundary conditions of this system, as a tensor with + /// shape `(3,)`. + /// + /// @see `data` for the meaning of the returned tensor. + DLPackTensor pbc() const { + return this->data(MTA_SYSTEM_DATA_PBC); + } + + /// Add a pair list (i.e. neighbor list) to this system. + /// + /// Ownership of `pairs` is transferred to this `System`. + /// + /// @param options options describing the pair list + /// @param pairs pairs data, stored as a metatensor block + void add_pairs(const PairListOptions& options, metatensor::TensorBlock pairs) { + nlohmann::json j = options; + this->add_pairs(j.dump(), std::move(pairs)); + } + + /// Add a pair list (i.e. neighbor list) to this system. + /// + /// Ownership of `pairs` is transferred to this `System`. + /// + /// @param options_json JSON-serialized `PairListOptions` describing the + /// pair list + /// @param pairs pairs data, stored as a metatensor block + void add_pairs(const std::string& options_json, metatensor::TensorBlock pairs) { + auto status = mta_system_add_pairs(system_, options_json.c_str(), pairs.release()); + details::check_status(status); + } + + /// Get a previously stored pair list matching the given `options_json`. + /// + /// The returned block is a non-owning view into data owned by this + /// `System`, and is only valid for as long as this `System` is alive. + /// + /// @param options options identifying the pair list to retrieve + metatensor::TensorBlock pairs(const PairListOptions& options) const { + nlohmann::json j = options; + return this->pairs(j.dump()); + } + + /// Get a previously stored pair list matching the given `options_json`. + /// + /// The returned block is a non-owning view into data owned by this + /// `System`, and is only valid for as long as this `System` is alive. + /// + /// @param options_json JSON-serialized `PairListOptions` identifying + /// the pair list to retrieve + metatensor::TensorBlock pairs(const std::string& options_json) const { + const mts_block_t* pairs = nullptr; + auto status = mta_system_get_pairs(system_, options_json.c_str(), &pairs); + details::check_status(status); + details::check_pointer(pairs); + return metatensor::TensorBlock::unsafe_view_from_ptr(const_cast(pairs)); + } + + /// Get the options of all pair lists registered with this `System` + std::vector known_pairs() const { + mta_string_t options = nullptr; + auto status = mta_system_known_pairs(system_, &options); + details::check_status(status); + nlohmann::json j = nlohmann::json::parse(mta_string_view(options)); + mta_string_free(options); + return j.get>(); + } + + /// Get the options of all pair lists registered with this `System`, as + /// a JSON-serialized array of `PairListOptions`. + std::vector known_pairs_json() const { + mta_string_t options = nullptr; + auto status = mta_system_known_pairs(system_, &options); + details::check_status(status); + nlohmann::json j = nlohmann::json::parse(mta_string_view(options)); + mta_string_free(options); + return j.get>(); + } + + /// Add custom data to this system, stored under the given `name`. + /// + /// Ownership of `data` is transferred to this `System`. + /// + /// @param name name used to identify the custom data + /// @param data custom data, stored as a metatensor tensor map + void add_custom_data(const std::string& name, metatensor::TensorMap data) { + auto status = mta_system_add_custom_data(system_, name.c_str(), data.release()); + details::check_status(status); + } + + /// Get the custom data previously stored under the given `name`. + /// + /// The returned tensor map is a non-owning view into data owned by this + /// `System`, and is only valid for as long as this `System` is alive. + /// + /// @param name name of the custom data to retrieve + metatensor::TensorMap custom_data(const std::string& name) const { + const mts_tensormap_t* data = nullptr; + auto status = mta_system_get_custom_data(system_, name.c_str(), &data); + details::check_status(status); + details::check_pointer(data); + return metatensor::TensorMap::unsafe_view_from_ptr(const_cast(data)); + } + + /// Get the names of all custom data registered with this `System` + std::vector known_custom_data() const { + mta_string_t names = nullptr; + auto status = mta_system_known_custom_data(system_, &names); + details::check_status(status); + nlohmann::json j = nlohmann::json::parse(mta_string_view(names)); + mta_string_free(names); + return j.get>(); + } + + /// Get the raw `mta_system_t` pointer backing this `System`. + /// + /// The `System` keeps ownership of the pointer, which is only valid for + /// as long as this `System` is alive. + mta_system_t* as_mta_system_t() & { + return system_; + } + + /// Get the raw `mta_system_t` pointer backing this `System`. + /// + /// The `System` keeps ownership of the pointer, which is only valid for + /// as long as this `System` is alive. + const mta_system_t* as_mta_system_t() const & { + return system_; + } + + /// Getting the raw pointer from a temporary `System` is forbidden, as it + /// would immediately dangle. + mta_system_t* as_mta_system_t() && = delete; + + /// Create an owning `System` from a raw `mta_system_t` pointer, taking + /// ownership of it. The system will be freed when the `System` is + /// destroyed. + /// + /// This is an advanced function, and the caller is responsible for + /// ensuring that `system` was allocated by the C API and is not used + /// anywhere else. + static System unsafe_from_ptr(mta_system_t* system) { + return System(system, /*is_view*/ false); + } + + /// Create a non-owning `System` view from a raw `mta_system_t` pointer. + /// The system will *not* be freed when the `System` is destroyed, and + /// must outlive it. + /// + /// This is an advanced function, mainly useful to wrap the systems given + /// to a model by the runtime. + static System unsafe_view_from_ptr(const mta_system_t* system) { + return System(const_cast(system), /*is_view*/ true); + } + + /// Release the raw `mta_system_t` pointer from this `System` without + /// freeing it, transferring ownership back to the caller. + mta_system_t* release() { + this->check_not_view("release"); + auto* system = system_; + system_ = nullptr; + is_view_ = true; + return system; + } + + private: + /// Wrap an existing `mta_system_t` pointer, see `unsafe_from_ptr` and + /// `unsafe_view_from_ptr`. + explicit System(mta_system_t* system, bool is_view): + system_(system), is_view_(is_view) {} + + void check_not_view(const char* method_name) const { + if (is_view_) { + throw Error( + "can not call System::" + std::string(method_name) + + " on this system since it is a view of a system owned elsewhere." + ); + } + } + + /// Get one of the always-present data tensors of this system. + /// + /// The returned `DLPackTensor` is a view sharing its data with the + /// system, which is kept alive for as long as the view exists. + DLPackTensor data(mta_system_data_kind request) const { + DLManagedTensorVersioned* data = nullptr; + auto status = mta_system_get_data(system_, request, &data); + details::check_status(status); + details::check_pointer(data); + return DLPackTensor(data); + } + mta_system_t* system_ = nullptr; + bool is_view_ = false; + }; } // namespace metatomic diff --git a/metatomic-core/include/metatomic/utils.hpp b/metatomic-core/include/metatomic/utils.hpp index 38f6aaf79..87cbd30de 100644 --- a/metatomic-core/include/metatomic/utils.hpp +++ b/metatomic-core/include/metatomic/utils.hpp @@ -1,11 +1,97 @@ #pragma once +#include #include #include #include namespace metatomic { + /// RAII wrapper around a DLPack `DLManagedTensorVersioned*`. + /// + /// This owns the managed tensor and calls its deleter when the wrapper is + /// destroyed. It can be used to move ownership of DLPack tensors across the + /// metatomic C++ API. + class DLPackTensor final { + public: + /// Create an empty wrapper, not owning any tensor. + DLPackTensor() = default; + + /// Take ownership of an existing DLPack managed tensor. + explicit DLPackTensor(DLManagedTensorVersioned* tensor): tensor_(tensor) {} + + /// The managed tensor is freed through its own deleter on destruction. + ~DLPackTensor() = default; + + /// `DLPackTensor` is not copy-constructible + DLPackTensor(const DLPackTensor&) = delete; + /// `DLPackTensor` is not copy-assignable + DLPackTensor& operator=(const DLPackTensor&) = delete; + + /// `DLPackTensor` is move-constructible + DLPackTensor(DLPackTensor&&) noexcept = default; + /// `DLPackTensor` is move-assignable + DLPackTensor& operator=(DLPackTensor&&) noexcept = default; + + /// Check whether this wrapper currently owns a tensor. + explicit operator bool() const { + return static_cast(tensor_); + } + + /// Access the underlying `DLManagedTensorVersioned` without transferring + /// ownership. The pointer stays owned by this `DLPackTensor`. + DLManagedTensorVersioned* operator->() const { + return tensor_.get(); + } + + /// Get the underlying `DLManagedTensorVersioned` pointer. It stays owned + /// by this `DLPackTensor`, and is only valid for as long as it is alive. + DLManagedTensorVersioned* as_dlpack() const { + return tensor_.get(); + } + + /// Release the underlying `DLManagedTensorVersioned` without calling its + /// deleter, transferring ownership back to the caller. + DLManagedTensorVersioned* release() { + return tensor_.release(); + } + + private: + /// Deleter implementing the DLPack ownership protocol: invoke the managed + /// tensor's own `deleter` callback if it has one. + struct Deleter { + void operator()(DLManagedTensorVersioned* tensor) const noexcept { + if (tensor->deleter != nullptr) { + tensor->deleter(tensor); + } + } + }; + + std::unique_ptr tensor_; + }; + + namespace details { + /// Take ownership of an `mta_string_t` returned by the C API, copy its + /// contents into an owned `std::string`, and free the C string. + /// + /// The `unique_ptr` guard frees the C string on return, including if the + /// copy into the `std::string` throws. A null `mta_string_t` (as produced + /// by an empty output) yields an empty string. + inline std::string string_from_mta(mta_string_t string) { + struct Deleter { + void operator()(mta_string_t ptr) const noexcept { + mta_string_free(ptr); + } + }; + std::unique_ptr, Deleter> owned(string); + + if (string == nullptr) { + return std::string(); + } + return std::string(mta_string_view(string)); + } + } // namespace details + /// Get the multiplicative conversion factor to use to convert from /// `from_unit` to `to_unit`. Both units are parsed as expressions /// (e.g. `kJ / mol / A^2`, `(eV * u)^(1/2)`) and their dimensions must diff --git a/metatomic-core/tests/cxx/system.cpp b/metatomic-core/tests/cxx/system.cpp new file mode 100644 index 000000000..d21715ae2 --- /dev/null +++ b/metatomic-core/tests/cxx/system.cpp @@ -0,0 +1,299 @@ +#include +#include +#include +#include +#include + +#include + +#include +#include "metatomic.hpp" + + +// Helpers building the DLPack tensors used to create a `System`, wrapped in the +// RAII `metatomic::DLPackTensor`. These mirror the ones used in the C API tests. + +template +static metatomic::DLPackTensor types_tensor(size_t n_atoms) { + auto type_data = std::vector(); + type_data.reserve(n_atoms); + for (size_t i=0; i(i * 3 + 1)); + } + + auto array = std::make_unique>( + std::vector{n_atoms}, std::move(type_data) + ); + auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); + + DLDevice cpu = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + return metatomic::DLPackTensor(mts.as_dlpack(cpu, nullptr, version)); +} + +template +static metatomic::DLPackTensor positions_tensor(size_t n_atoms) { + auto position_data = std::vector(); + position_data.reserve(n_atoms * 3); + for (size_t i=0; i(i * 3 + 1)); + position_data.push_back(static_cast(i * 3 + 2)); + position_data.push_back(static_cast(i * 3 + 3)); + } + + auto array = std::make_unique>( + std::vector{n_atoms, 3}, std::move(position_data) + ); + auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); + + DLDevice cpu = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + return metatomic::DLPackTensor(mts.as_dlpack(cpu, nullptr, version)); +} + +template +static metatomic::DLPackTensor cell_tensor() { + // the `y` row is zero to match the non-periodic `y` direction in `pbc` + // (metatomic requires the cell vector of non-periodic directions to be zero) + auto array = std::make_unique>( + std::vector{3, 3}, + std::vector{ + T(10.0), T(0.0), T(0.0), + T(0.0), T(0.0), T(0.0), + T(0.0), T(0.0), T(10.0), + } + ); + auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); + + DLDevice cpu = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + return metatomic::DLPackTensor(mts.as_dlpack(cpu, nullptr, version)); +} + +static metatomic::DLPackTensor pbc_tensor() { + // `SimpleDataArray` does not compile (`std::vector` has no + // `data()` method), so we use `uint8_t` and patch the dtype code to + // `kDLBool`. + auto array = std::make_unique>( + std::vector{3}, std::vector{1, 0, 1} + ); + auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); + + DLDevice cpu = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + auto* tensor = mts.as_dlpack(cpu, nullptr, version); + tensor->dl_tensor.dtype.code = DLDataTypeCode::kDLBool; + + return metatomic::DLPackTensor(tensor); +} + +static metatomic::System test_system(size_t n_atoms = 4) { + return metatomic::System( + "nm", + types_tensor(n_atoms), + positions_tensor(n_atoms), + cell_tensor(), + pbc_tensor() + ); +} + +static metatensor::TensorBlock pair_block() { + auto samples = metatensor::Labels( + {"first_atom", "second_atom", "cell_shift_a", "cell_shift_b", "cell_shift_c"}, + {{0, 1, 0, 0, 0}} + ); + auto components = std::vector{ + metatensor::Labels({"xyz"}, {{0}, {1}, {2}}) + }; + auto properties = metatensor::Labels({"distance"}, {{0}}); + + auto values = std::make_unique>( + std::vector{1, 3, 1}, std::vector{1.5F, 2.5F, 3.5F} + ); + + return metatensor::TensorBlock(std::move(values), samples, components, properties); +} + +static metatensor::TensorMap custom_data() { + auto keys = metatensor::Labels({"key"}, {{0}}); + + auto samples = metatensor::Labels({"sample"}, {{0}}); + auto properties = metatensor::Labels({"property"}, {{0}}); + auto values = std::make_unique>( + std::vector{1, 1}, std::vector{42.0F} + ); + auto block = metatensor::TensorBlock(std::move(values), samples, {}, properties); + + auto blocks = std::vector(); + blocks.push_back(std::move(block)); + return metatensor::TensorMap(keys, std::move(blocks)); +} + + +TEST_CASE("System basics") { + auto system = test_system(4); + + CHECK(system.size() == 4); + CHECK(system.length_unit() == "nm"); +} + +TEST_CASE("System construction errors") { + // wrong dtype for `types` (float instead of int32) + REQUIRE_THROWS_WITH( + metatomic::System( + "Angstrom", + types_tensor(3), + positions_tensor(3), + cell_tensor(), + pbc_tensor() + ), + "invalid parameter: `types` must be a tensor of 32-bit integers" + ); +} + +TEST_CASE("System data") { + auto system = test_system(4); + + SECTION("types") { + auto types = system.types(); + REQUIRE(static_cast(types)); + CHECK(types->dl_tensor.ndim == 1); + CHECK(types->dl_tensor.shape[0] == 4); + CHECK(types->dl_tensor.dtype.code == kDLInt); + CHECK(types->dl_tensor.dtype.bits == 32); + + auto* data = reinterpret_cast( + static_cast(types->dl_tensor.data) + types->dl_tensor.byte_offset + ); + CHECK(data[0] == 1); + CHECK(data[3] == 10); + } + + SECTION("positions") { + auto positions = system.positions(); + REQUIRE(static_cast(positions)); + CHECK(positions->dl_tensor.ndim == 2); + CHECK(positions->dl_tensor.shape[0] == 4); + CHECK(positions->dl_tensor.shape[1] == 3); + CHECK(positions->dl_tensor.dtype.code == kDLFloat); + + auto* data = reinterpret_cast( + static_cast(positions->dl_tensor.data) + positions->dl_tensor.byte_offset + ); + CHECK(data[0] == 1.0F); + CHECK(data[9] == 10.0F); + } + + SECTION("cell") { + auto cell = system.cell(); + REQUIRE(static_cast(cell)); + CHECK(cell->dl_tensor.ndim == 2); + CHECK(cell->dl_tensor.shape[0] == 3); + CHECK(cell->dl_tensor.shape[1] == 3); + } + + SECTION("pbc") { + auto pbc = system.pbc(); + REQUIRE(static_cast(pbc)); + CHECK(pbc->dl_tensor.ndim == 1); + CHECK(pbc->dl_tensor.shape[0] == 3); + CHECK(pbc->dl_tensor.dtype.code == kDLBool); + + auto* data = reinterpret_cast( + static_cast(pbc->dl_tensor.data) + pbc->dl_tensor.byte_offset + ); + CHECK(data[0] == true); + CHECK(data[1] == false); + CHECK(data[2] == true); + } +} + +TEST_CASE("System pairs") { + auto system = test_system(4); + + auto options = metatomic::PairListOptions(); + options.cutoff(1.0); + options.full_list(true); + options.strict(false); + options.add_requestor("test"); + + system.add_pairs(options, pair_block()); + + const auto* options_json = R"({ + "type": "metatomic_pair_options", + "cutoff": "0x40364ccccccccccd", + "full_list": false, + "strict": true, + "requestors": [""] + })"; + + system.add_pairs(options_json, pair_block()); + + auto pairs = system.pairs(options); + CHECK(pairs.samples().count() == 1); + CHECK(pairs.properties().size() == 1); + + auto known = system.known_pairs(); + CHECK(known.size() == 2); + CHECK(known[0].cutoff() == 1.0); + CHECK(known[0].full_list() == true); + CHECK(known[0].strict() == false); + CHECK(known[0].requestors().size() == 1); + CHECK(known[0].requestors()[0] == "test"); + + CHECK(known[1].cutoff() == 22.3); + CHECK(known[1].full_list() == false); + CHECK(known[1].strict() == true); + CHECK(known[1].requestors().size() == 0); +} + +TEST_CASE("System custom data") { + auto system = test_system(4); + + system.add_custom_data("test::my_data", custom_data()); + + auto data = system.custom_data("test::my_data"); + CHECK(data.keys().count() == 1); + + // retrieving unknown data throws + REQUIRE_THROWS(system.custom_data("test::no_such_data")); + + system.add_custom_data("test::other_data", custom_data()); + auto names = system.known_custom_data(); + std::sort(names.begin(), names.end()); + CHECK(names.size() == 2); + CHECK(names[0] == "test::my_data"); + CHECK(names[1] == "test::other_data"); +} + +TEST_CASE("System ownership") { + SECTION("move") { + auto system = test_system(4); + auto* ptr = system.as_mta_system_t(); + + auto moved = std::move(system); + CHECK(moved.as_mta_system_t() == ptr); + CHECK(moved.size() == 4); + } + + SECTION("release / unsafe_from_ptr round-trip") { + auto system = test_system(4); + auto* raw = system.release(); + REQUIRE(raw != nullptr); + + auto owned = metatomic::System::unsafe_from_ptr(raw); + CHECK(owned.size() == 4); + } + + SECTION("unsafe_view_from_ptr does not free") { + auto system = test_system(4); + + { + auto view = metatomic::System::unsafe_view_from_ptr(system.as_mta_system_t()); + CHECK(view.size() == 4); + } + + // the original system is still usable after the view is destroyed + CHECK(system.size() == 4); + } +} From d37d963b12d819dd8fbd401f5269c909ffbd45f7 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Mon, 29 Jun 2026 11:17:42 +0200 Subject: [PATCH 35/43] Group all code accessing values in a new `kernels` module --- metatomic-core/src/kernels/cpu.rs | 60 ++++++++++++++++++++++++++ metatomic-core/src/kernels/mod.rs | 60 ++++++++++++++++++++++++++ metatomic-core/src/lib.rs | 2 + metatomic-core/src/system.rs | 72 +++++++++++++------------------ 4 files changed, 152 insertions(+), 42 deletions(-) create mode 100644 metatomic-core/src/kernels/cpu.rs create mode 100644 metatomic-core/src/kernels/mod.rs diff --git a/metatomic-core/src/kernels/cpu.rs b/metatomic-core/src/kernels/cpu.rs new file mode 100644 index 000000000..c054518b0 --- /dev/null +++ b/metatomic-core/src/kernels/cpu.rs @@ -0,0 +1,60 @@ +use dlpk::DLPackTensorRef; +use ndarray::{ArrayView1, ArrayView2, ArrayViewD}; + +use crate::Error; + +/// Check that the values of an i32 DLPack tensor match the expected reference. +/// +/// The tensor is converted to an ndarray view and compared element-wise and +/// shape-wise against `reference`. The `description` is used verbatim in the +/// error message on mismatch. +/// +/// # Parameters +/// - `tensor`: DLPack tensor with i32 data type +/// - `reference`: expected values with the same shape as the tensor +pub(crate) fn is_equal_i32( + tensor: DLPackTensorRef<'_>, + reference: ArrayViewD<'_, i32>, +) -> Result { + let values: ArrayViewD = tensor.try_into()?; + return Ok(values == reference); +} + +macro_rules! validate_cell { + ($T: ty, $pbc: expr, $cell: expr) => { + let pbc_array: ArrayView1 = $pbc.try_into()?; + let cell_array: ArrayView2<$T> = $cell.try_into()?; + for i in 0..3 { + if !pbc_array[i] && !cell_array.row(i).iter().all(|&x| x == 0.0) { + return Err(Error::InvalidParameter(format!( + "invalid cell: for non-periodic dimensions, the corresponding \ + cell vector must be zero, but cell[{}] contains non-zero values", + i + ))); + } + } + }; +} + +/// Validate that cell vectors are zero for non-periodic dimensions on CPU. +/// +/// Converts the DLPack tensors to ndarray views and checks that for every +/// dimension where `pbc` is false, the corresponding row of `cell` contains +/// only zeros. +/// +/// # Parameters +/// - `pbc`: 1D boolean tensor of length 3 (periodic boundary condition flags) +/// - `cell`: 3x3 tensor (unit cell vectors as rows) +pub(crate) fn validate_cell_pbc( + pbc: DLPackTensorRef<'_>, + cell: DLPackTensorRef<'_>, +) -> Result<(), Error> { + let dtype = cell.dtype(); + if dtype.bits == 32 { + validate_cell!(f32, pbc, cell); + } else { + assert_eq!(dtype.bits, 64); + validate_cell!(f64, pbc, cell); + } + return Ok(()); +} diff --git a/metatomic-core/src/kernels/mod.rs b/metatomic-core/src/kernels/mod.rs new file mode 100644 index 000000000..704ad07b3 --- /dev/null +++ b/metatomic-core/src/kernels/mod.rs @@ -0,0 +1,60 @@ +use dlpk::sys::DLDeviceType; +use dlpk::DLPackTensorRef; +use ndarray::ArrayViewD; + +use crate::Error; + +mod cpu; + +/// Check that the values of an i32 DLPack tensor match the expected reference. +/// +/// This dispatches to the appropriate backend based on the device of `tensor`. +/// +/// # Parameters +/// - `tensor`: DLPack tensor with i32 data type +/// - `reference`: expected values with the same shape as the tensor +pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_, i32>) -> Result { + match tensor.device().device_type { + DLDeviceType::kDLCPU + | DLDeviceType::kDLCUDAHost + | DLDeviceType::kDLROCMHost => { + cpu::is_equal_i32(tensor, reference) + } + _ => { + eprintln!( + "is_equal_i32 for non-CPU devices is not implemented, \ + got data on device: {:?}", tensor.device() + ); + Ok(true) + } + } +} + +/// Validate that cell vectors are zero for non-periodic dimensions. +/// +/// This dispatches to the appropriate backend based on the device of `pbc`. +/// +/// # Parameters +/// - `pbc`: 1D boolean tensor of length 3 (periodic boundary condition flags) +/// - `cell`: 3x3 tensor (unit cell vectors as rows) +pub(crate) fn validate_cell_pbc(pbc: DLPackTensorRef<'_>, cell: DLPackTensorRef<'_>) -> Result<(), Error> { + debug_assert!( + pbc.device() == cell.device(), + "pbc and cell must be on the same device" + ); + + match pbc.device().device_type { + DLDeviceType::kDLCPU + | DLDeviceType::kDLCUDAHost + | DLDeviceType::kDLROCMHost => { + cpu::validate_cell_pbc(pbc, cell) + } + _ => { + eprintln!( + "Cell/PBC validation for non-CPU devices is not implemented, \ + got data on device: {:?}", pbc.device() + ); + Ok(()) + } + } +} diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index 18dd99cac..954693f54 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -25,6 +25,8 @@ pub use self::metadata::{Device, DType, ModelCapabilities, ModelMetadata, PairLi mod quantities; pub use self::quantities::{Quantity, SampleKind, Gradients}; +mod kernels; + mod system; pub use self::system::System; diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs index b3aeb8e3e..26b11d844 100644 --- a/metatomic-core/src/system.rs +++ b/metatomic-core/src/system.rs @@ -1,7 +1,7 @@ use std::collections::{BTreeMap, HashMap, HashSet}; use std::sync::LazyLock; -use dlpk::sys::{DLDataType, DLDevice, DLDeviceType}; +use dlpk::sys::{DLDataType, DLDevice}; use dlpk::{DLPackTensor, DLPackTensorRef}; use metatensor::{TensorBlock, TensorMap}; @@ -51,9 +51,7 @@ impl System { custom_data: HashMap::new(), }; - if system.device().device_type == DLDeviceType::kDLCPU { - validate_cpu_system_data(&system)?; - } + crate::kernels::validate_cell_pbc(system.pbc(), system.cell())?; return Ok(system); } @@ -121,12 +119,22 @@ impl System { )); } - #[allow(clippy::collapsible_if)] - if components[0].device().device_type == DLDeviceType::kDLCPU { - if components[0][0] != [0] || components[0][1] != [1] || components[0][2] != [2] { + { + let mts_array = components[0].values(); + let dl_tensor = mts_array.as_dlpack( + components[0].device(), + None, + dlpk::sys::DLPackVersion::current(), + )?; + let reference = ndarray::ArrayViewD::::from_shape( + ndarray::IxDyn(&[3usize, 1]), + &[0i32, 1, 2], + ).unwrap(); + + if !crate::kernels::is_equal_i32(dl_tensor.as_ref(), reference)? { return Err(Error::InvalidParameter( - "invalid components for `pairs`: the 'xyz' \ - component should contain [0, 1, 2]".into() + "invalid components for `pairs`: the 'xyz' component should \ + contain [[0], [1], [2]]".into() )); } } @@ -139,9 +147,19 @@ impl System { )); } - #[allow(clippy::collapsible_if)] - if properties.device().device_type == DLDeviceType::kDLCPU { - if properties[0] != [0] { + { + let mts_array = properties.values(); + let dl_tensor = mts_array.as_dlpack( + properties.device(), + None, + dlpk::sys::DLPackVersion::current(), + )?; + let reference = ndarray::ArrayViewD::::from_shape( + ndarray::IxDyn(&[1usize, 1]), + &[0i32], + ).unwrap(); + + if !crate::kernels::is_equal_i32(dl_tensor.as_ref(), reference)? { return Err(Error::InvalidParameter( "invalid properties for `pairs`: the 'distance' property \ should contain [0]".into() @@ -344,36 +362,6 @@ fn validate_system_tensors( return Ok(()); } -fn validate_cpu_system_data(system: &System) -> Result<(), Error> { - let pbc_array: ndarray::ArrayView1 = system.pbc().try_into()?; - - if system.dtype().bits == 32 { - let cell_array: ndarray::ArrayView2 = system.cell().try_into()?; - for i in 0..3 { - if !pbc_array[i] && !cell_array.row(i).iter().all(|&x| x == 0.0) { - return Err(Error::InvalidParameter(format!( - "invalid cell: for non-periodic dimensions, the corresponding \ - cell vector must be zero, but cell[{}] contains non-zero values", - i - ))); - } - } - } else { - let cell_array: ndarray::ArrayView2 = system.cell().try_into()?; - for i in 0..3 { - if !pbc_array[i] && !cell_array.row(i).iter().all(|&x| x == 0.0) { - return Err(Error::InvalidParameter(format!( - "invalid cell: for non-periodic dimensions, the corresponding \ - cell vector must be zero, but cell[{}] contains non-zero values", - i - ))); - } - } - } - - return Ok(()); -} - #[cfg(test)] pub(crate) use tests::test_system; From 27ec8452c49dd879747d78aa7fc2325283ec4448 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Mon, 29 Jun 2026 17:14:21 +0200 Subject: [PATCH 36/43] Add cuda kernels for system validation --- metatomic-core/Cargo.toml | 5 + metatomic-core/src/kernels/cuda.rs | 285 +++++++++++++++++++++ metatomic-core/src/kernels/cuda_kernels.cu | 95 +++++++ metatomic-core/src/kernels/mod.rs | 23 +- 4 files changed, 398 insertions(+), 10 deletions(-) create mode 100644 metatomic-core/src/kernels/cuda.rs create mode 100644 metatomic-core/src/kernels/cuda_kernels.cu diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index 934b6fe2b..090b5c0ae 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -19,9 +19,14 @@ dlpk = { version = "0.4", features = ["ndarray"]} json = "0.12" libloading = "0.9" ndarray = "0.17" + +# For serialization of the systems zip = { version = "8.6.0", default-features = false } byteorder = {version = "1"} +# For custom kernels +cudarc = {version = "0.19", default-features = false, features=["std", "cuda-13030", "driver", "nvrtc", "dynamic-loading"]} + [build-dependencies] cbindgen = { version = "0.29", default-features = false } diff --git a/metatomic-core/src/kernels/cuda.rs b/metatomic-core/src/kernels/cuda.rs new file mode 100644 index 000000000..ccca5ca6c --- /dev/null +++ b/metatomic-core/src/kernels/cuda.rs @@ -0,0 +1,285 @@ +use std::collections::hash_map::Entry; +use std::collections::HashMap; +use std::sync::{Arc, Mutex, LazyLock}; + +use cudarc::driver::safe::DeviceRepr; +use cudarc::driver::safe::{ + CudaContext, CudaFunction, CudaModule, CudaStream, LaunchConfig, PushKernelArg, +}; +use cudarc::nvrtc::compile_ptx; +use dlpk::DLPackTensorRef; +use ndarray::ArrayViewD; + + +use crate::Error; + +// CUDA kernel source compiled at runtime via NVRTC for the exact GPU +const KERNEL_SRC: &str = include_str!("cuda_kernels.cu"); + +const MAX_NDIM: usize = 7; + +/// Multi-dimensional strided index matching the CUDA `StridedNDIndex` struct. +/// Supports up to 7 dimensions. +/// +/// WARNING: any change here needs to be reflected in the CUDA source +#[repr(C)] +struct StridedNDIndex { + ndim: i64, + shape: [i64; MAX_NDIM], + strides: [i64; MAX_NDIM], +} + +unsafe impl DeviceRepr for StridedNDIndex {} + +impl StridedNDIndex { + fn from_dlpack(tensor: &DLPackTensorRef<'_>) -> Self { + return StridedNDIndex::from_shape_strides(tensor.shape(), tensor.strides()) + } + + #[allow(clippy::cast_possible_wrap)] + fn from_ndarray(array: &ArrayViewD<'_, T>) -> Self { + let shape = array.shape().iter().map(|&s| s as i64).collect::>(); + let strides = array.strides().iter().map(|&s| s as i64).collect::>(); + return Self::from_shape_strides(&shape, Some(&strides)); + } + + #[allow(clippy::cast_possible_wrap)] + fn from_shape_strides(shape: &[i64], strides: Option<&[i64]>) -> Self { + let ndim = shape.len(); + assert!(ndim <= MAX_NDIM, "StridedNDIndex only supports up to {MAX_NDIM} dimensions, got {ndim}"); + let mut shape_arr = [0i64; MAX_NDIM]; + let mut strides_arr = [0i64; MAX_NDIM]; + + // Contiguous fallback strides (row-major / C-contiguous) + let mut acc: i64 = 1; + for i in (0..ndim).rev() { + shape_arr[i] = shape[i]; + strides_arr[i] = acc; + acc *= shape[i]; + } + + if let Some(strides) = strides { + strides_arr[..ndim].copy_from_slice(&strides[..ndim]); + } + StridedNDIndex { ndim: ndim as i64, shape: shape_arr, strides: strides_arr } + } +} + +/// Zero-cost wrapper to pass an existing device pointer as a CUDA kernel +/// argument. +/// +/// Does NOT own the memory — the caller (DLPack tensor) is responsible for +/// lifetime and must ensure the pointer remains valid for the duration of the +/// kernel launch. +/// +/// The `#[repr(transparent)]` wrapper over `cudarc::driver::sys::CUdeviceptr` +/// is passed to `PushKernelArg::arg()` which pushes the address of this struct +/// on the host stack. CUDA reads 8 bytes from that address as the kernel +/// parameter value, giving the kernel the correct device pointer. +#[repr(transparent)] +struct DevicePtrArg { + ptr: cudarc::driver::sys::CUdeviceptr, +} + +unsafe impl DeviceRepr for DevicePtrArg {} + +/// Per-device cached resources: context, module, and kernel function handles. +struct CudaKernelCache { + ctx: Arc, + module: Arc, + is_equal_i32: CudaFunction, + validate_cell_pbc_f32: CudaFunction, + validate_cell_pbc_f64: CudaFunction, +} + +impl CudaKernelCache { + fn new(device_id: usize) -> Result { + let ctx = CudaContext::new(device_id) + .map_err(|e| Error::Internal(format!("CudaContext::new({device_id}): {e}")))?; + let ptx = compile_ptx(KERNEL_SRC) + .map_err(|e| Error::Internal(format!("NVRTC compile failed: {e}")))?; + let module = ctx + .load_module(ptx) + .map_err(|e| Error::Internal(format!("PTX load failed: {e}")))?; + let is_equal_i32 = module + .load_function("is_equal_i32") + .map_err(|e| Error::Internal(format!("load_function(is_equal_i32): {e}")))?; + let validate_cell_pbc_f32 = module + .load_function("validate_cell_pbc_f32") + .map_err(|e| Error::Internal(format!("load_function(validate_cell_pbc_f32): {e}")))?; + let validate_cell_pbc_f64 = module + .load_function("validate_cell_pbc_f64") + .map_err(|e| Error::Internal(format!("load_function(validate_cell_pbc_f64): {e}")))?; + Ok(Self { + ctx, + module, + is_equal_i32, + validate_cell_pbc_f32, + validate_cell_pbc_f64, + }) + } +} + +static CUDA_CACHE: LazyLock>> = LazyLock::new(|| Mutex::new(HashMap::new())); + +fn get_or_init(device_id: usize) -> Result, Error> { + let mut cache = CUDA_CACHE.lock().expect("failed to lock CUDA_CACHE"); + let entry = match cache.entry(device_id) { + Entry::Occupied(entry) => entry.into_mut(), + Entry::Vacant(entry) => entry.insert(CudaKernelCache::new(device_id)?), + }; + Ok(entry.ctx.default_stream()) +} + +/// Extract a `CUdeviceptr` from a DLPack tensor's raw `data` + `byte_offset`. +/// +/// # Safety +/// +/// The returned `CUdeviceptr` is only valid as long as the DLPack tensor's +/// backing memory is alive. The caller must ensure the tensor is not dropped +/// before the kernel finishes execution. +unsafe fn dlpack_to_device_ptr(tensor: &DLPackTensorRef<'_>) -> cudarc::driver::sys::CUdeviceptr { + debug_assert!( + tensor.device().device_type == dlpk::sys::DLDeviceType::kDLCUDA, + "dlpack_to_device_ptr called on non-CUDA tensor" + ); + let raw_ptr = tensor.raw.data as u64; + (raw_ptr + tensor.raw.byte_offset) as cudarc::driver::sys::CUdeviceptr +} + +/// Check that the values of a CUDA-resident i32 DLPack tensor match an expected +/// reference array. +/// +/// The comparison is performed entirely on-device: the existing GPU pointer from +/// `tensor` is wrapped as a `DevicePtrArg`, the reference is uploaded to the GPU, +/// and a single-element result flag (`0` = ok, `1` = mismatch) is read back. +#[allow(clippy::cast_sign_loss, clippy::cast_possible_truncation)] +pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_, i32>) -> Result { + debug_assert!( + tensor.device().device_type == dlpk::sys::DLDeviceType::kDLCUDA, + "is_equal_i32 called on non-CUDA tensor" + ); + debug_assert!(tensor.device().device_id >= 0, "is_equal_i32 called on invalid device_id"); + + let device_id = tensor.device().device_id as usize; + let stream = get_or_init(device_id)?; + let cache = CUDA_CACHE.lock().expect("failed to lock CUDA_CACHE"); + let entry = &cache[&device_id]; + + let n_elements: i64 = tensor.shape().iter().product(); + + // Build strided index from the DLPack tensor (preserves actual strides) + let values_idx = StridedNDIndex::from_dlpack(&tensor); + + // Build strided index from the ndarray view (preserves actual strides) + let reference_idx = StridedNDIndex::from_ndarray(&reference); + + // Wrap the existing GPU-allocated tensor pointer + let tensor_ptr = unsafe { DevicePtrArg { ptr: dlpack_to_device_ptr(&tensor) } }; + + // Upload reference values to GPU + let ref_dev = stream.clone_htod(reference.as_slice().expect("reference should be contiguous")) + .map_err(|e| Error::Internal(format!("clone_htod reference: {e}")))?; + + // Allocate result flag (initialized to 0 = no mismatch) + let mut result = stream.alloc_zeros::(1) + .map_err(|e| Error::Internal(format!("alloc_zeros: {e}")))?; + + unsafe { + stream.launch_builder(&entry.is_equal_i32) + .arg(&tensor_ptr) + .arg(&values_idx) + .arg(&ref_dev) + .arg(&reference_idx) + .arg(&n_elements) + .arg(&mut result) + .launch(LaunchConfig::for_num_elems(u32::try_from(n_elements).expect("tensor too large for CUDA kernel"))) + .map_err(|e| Error::Internal(format!("kernel launch (is_equal_i32): {e}")))?; + } + + stream.synchronize() + .map_err(|e| Error::Internal(format!("device sync: {e}")))?; + + let host = stream.clone_dtoh(&result) + .map_err(|e| Error::Internal(format!("clone_dtoh result: {e}")))?; + + return Ok(host[0] == 0); +} + +/// Validate that cell vectors are zero for non-periodic dimensions, on CUDA device. +#[allow(clippy::cast_sign_loss)] +pub(crate) fn validate_cell_pbc( + pbc: DLPackTensorRef<'_>, + cell: DLPackTensorRef<'_>, +) -> Result<(), Error> { + debug_assert!( + pbc.device().device_type == dlpk::sys::DLDeviceType::kDLCUDA, + "validate_cell_pbc called on non-CUDA tensor" + ); + debug_assert!(pbc.device().device_id >= 0, "validate_cell_pbc called on invalid device_id"); + debug_assert!(cell.device() == pbc.device(), "pbc and cell must be on the same device"); + + + let device_id = pbc.device().device_id as usize; + let stream = get_or_init(device_id)?; + let cache = CUDA_CACHE.lock().expect("failed to lock CUDA_CACHE"); + let entry = &cache[&device_id]; + + let pbc_ptr = unsafe { DevicePtrArg { ptr: dlpack_to_device_ptr(&pbc) } }; + let cell_ptr = unsafe { DevicePtrArg { ptr: dlpack_to_device_ptr(&cell) } }; + + let pbc_idx = StridedNDIndex::from_dlpack(&pbc); + let cell_idx = StridedNDIndex::from_dlpack(&cell); + + let mut result = stream.alloc_zeros::(1) + .map_err(|e| Error::Internal(format!("alloc_zeros: {e}")))?; + + if cell.dtype().bits == 32 { + unsafe { + stream.launch_builder(&entry.validate_cell_pbc_f32) + .arg(&pbc_ptr) + .arg(&pbc_idx) + .arg(&cell_ptr) + .arg(&cell_idx) + .arg(&mut result) + .launch(LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (3, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| Error::Internal(format!("kernel launch (f32): {e}")))?; + } + } else { + assert_eq!(cell.dtype().bits, 64, "validate_cell_pbc: unsupported cell dtype"); + unsafe { + stream.launch_builder(&entry.validate_cell_pbc_f64) + .arg(&pbc_ptr) + .arg(&pbc_idx) + .arg(&cell_ptr) + .arg(&cell_idx) + .arg(&mut result) + .launch(LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (3, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| Error::Internal(format!("kernel launch (f64): {e}")))?; + } + } + + stream.synchronize() + .map_err(|e| Error::Internal(format!("device sync: {e}")))?; + + let host = stream.clone_dtoh(&result) + .map_err(|e| Error::Internal(format!("clone_dtoh result: {e}")))?; + + if host[0] != 0 { + let dim = host[0] - 1; + return Err(Error::InvalidParameter(format!( + "invalid cell: for non-periodic dimensions, the corresponding \ + cell vector must be zero, but cell[{}] contains non-zero values", + dim + ))); + } + Ok(()) +} diff --git a/metatomic-core/src/kernels/cuda_kernels.cu b/metatomic-core/src/kernels/cuda_kernels.cu new file mode 100644 index 000000000..641a55767 --- /dev/null +++ b/metatomic-core/src/kernels/cuda_kernels.cu @@ -0,0 +1,95 @@ +//////////////////////////////////////////////////////////////////////////////// +//////////////////////////////////////////////////////////////////////////////// + +#define MAX_NDIM 7 + +/// Multi-dimensional strided index (up to MAX_NDIM dimensions). +/// Decomposes a flat linear index into multi-dimensional coordinates from the +/// shape, then computes the strided memory offset using the stride array. +/// +/// WARNING: any change here needs to be reflected in the Rust source +struct StridedNDIndex { + int64_t ndim; + int64_t shape[MAX_NDIM]; + int64_t strides[MAX_NDIM]; + + /// Get the offset from the start of the array for a given flat index + __device__ int64_t offset(int64_t flat_idx) const { + int64_t off = 0; + for (int d = this->ndim - 1; d >= 0; d--) { + int64_t coord = flat_idx % this->shape[d]; + flat_idx /= this->shape[d]; + off += coord * this->strides[d]; + } + return off; + } +}; + +//////////////////////////////////////////////////////////////////////////////// +//////////////////////////////////////////////////////////////////////////////// + +extern "C" __global__ void is_equal_i32( + const int* values, + StridedNDIndex values_idx, + const int* reference, + StridedNDIndex reference_idx, + int64_t n, + int* mismatch +) { + int64_t i = blockIdx.x * blockDim.x + threadIdx.x; + if (i < n) { + int64_t value_offset = values_idx.offset(i); + int64_t reference_offset = reference_idx.offset(i); + if (values[value_offset] != reference[reference_offset]) { + atomicMax(mismatch, 1); + } + } +} + +//////////////////////////////////////////////////////////////////////////////// +//////////////////////////////////////////////////////////////////////////////// + +template +__device__ void validate_cell_pbc_impl( + const bool* pbc, + StridedNDIndex pbc_idx, + const T* cell, + StridedNDIndex cell_idx, + int* mismatch_idx +) { + int i = threadIdx.x; + if (i < 3) { + if (!pbc[pbc_idx.offset(i)]) { + if ( + cell[cell_idx.offset(i * 3 + 0)] != T(0) || + cell[cell_idx.offset(i * 3 + 1)] != T(0) || + cell[cell_idx.offset(i * 3 + 2)] != T(0) + ) { + atomicMax(mismatch_idx, i + 1); + } + } + } +} + +extern "C" __global__ void validate_cell_pbc_f32( + const bool* pbc, + StridedNDIndex pbc_idx, + const float* cell, + StridedNDIndex cell_idx, + int* mismatch_idx +) { + validate_cell_pbc_impl(pbc, pbc_idx, cell, cell_idx, mismatch_idx); +} + +extern "C" __global__ void validate_cell_pbc_f64( + const bool* pbc, + StridedNDIndex pbc_idx, + const double* cell, + StridedNDIndex cell_idx, + int* mismatch_idx +) { + validate_cell_pbc_impl(pbc, pbc_idx, cell, cell_idx, mismatch_idx); +} + +//////////////////////////////////////////////////////////////////////////////// +//////////////////////////////////////////////////////////////////////////////// diff --git a/metatomic-core/src/kernels/mod.rs b/metatomic-core/src/kernels/mod.rs index 704ad07b3..ba0c39991 100644 --- a/metatomic-core/src/kernels/mod.rs +++ b/metatomic-core/src/kernels/mod.rs @@ -5,6 +5,7 @@ use ndarray::ArrayViewD; use crate::Error; mod cpu; +mod cuda; /// Check that the values of an i32 DLPack tensor match the expected reference. /// @@ -15,15 +16,16 @@ mod cpu; /// - `reference`: expected values with the same shape as the tensor pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_, i32>) -> Result { match tensor.device().device_type { - DLDeviceType::kDLCPU - | DLDeviceType::kDLCUDAHost - | DLDeviceType::kDLROCMHost => { + DLDeviceType::kDLCPU | DLDeviceType::kDLCUDAHost | DLDeviceType::kDLROCMHost => { cpu::is_equal_i32(tensor, reference) } + DLDeviceType::kDLCUDA | DLDeviceType::kDLCUDAManaged => { + cuda::is_equal_i32(tensor, reference) + } _ => { eprintln!( - "is_equal_i32 for non-CPU devices is not implemented, \ - got data on device: {:?}", tensor.device() + "is_equal_i32 for device {:?} is not implemented", + tensor.device() ); Ok(true) } @@ -44,15 +46,16 @@ pub(crate) fn validate_cell_pbc(pbc: DLPackTensorRef<'_>, cell: DLPackTensorRef< ); match pbc.device().device_type { - DLDeviceType::kDLCPU - | DLDeviceType::kDLCUDAHost - | DLDeviceType::kDLROCMHost => { + DLDeviceType::kDLCPU | DLDeviceType::kDLCUDAHost | DLDeviceType::kDLROCMHost => { cpu::validate_cell_pbc(pbc, cell) } + DLDeviceType::kDLCUDA | DLDeviceType::kDLCUDAManaged => { + cuda::validate_cell_pbc(pbc, cell) + } _ => { eprintln!( - "Cell/PBC validation for non-CPU devices is not implemented, \ - got data on device: {:?}", pbc.device() + "Cell/PBC validation for device {:?} is not implemented", + pbc.device() ); Ok(()) } From c7ae5b287b3f4d21172035797c0cfd47c1956041 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Tue, 30 Jun 2026 11:15:51 +0200 Subject: [PATCH 37/43] Add custom metal kernels for system validation --- metatomic-core/CMakeLists.txt | 7 + metatomic-core/Cargo.toml | 4 + .../cmake/metatomic-config.in.cmake | 10 +- metatomic-core/src/kernels/cuda.rs | 48 +-- metatomic-core/src/kernels/cuda_kernels.cu | 2 +- metatomic-core/src/kernels/metal.rs | 290 ++++++++++++++++++ .../src/kernels/metal_kernels.metal | 76 +++++ metatomic-core/src/kernels/mod.rs | 80 +++++ 8 files changed, 468 insertions(+), 49 deletions(-) create mode 100644 metatomic-core/src/kernels/metal.rs create mode 100644 metatomic-core/src/kernels/metal_kernels.metal diff --git a/metatomic-core/CMakeLists.txt b/metatomic-core/CMakeLists.txt index 128d52f91..e7dd396a7 100644 --- a/metatomic-core/CMakeLists.txt +++ b/metatomic-core/CMakeLists.txt @@ -400,6 +400,13 @@ endif() target_link_libraries(metatomic::static INTERFACE nlohmann_json::nlohmann_json) target_link_libraries(metatomic::shared INTERFACE nlohmann_json::nlohmann_json) +if(APPLE) + target_link_libraries(metatomic::static INTERFACE + "-framework Metal" "-framework CoreGraphics" "-framework CoreFoundation" "-framework Foundation" objc + ) +endif() + + if (BUILD_SHARED_LIBS) add_library(metatomic ALIAS metatomic::shared) else() diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index 090b5c0ae..db3482f68 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -27,6 +27,10 @@ byteorder = {version = "1"} # For custom kernels cudarc = {version = "0.19", default-features = false, features=["std", "cuda-13030", "driver", "nvrtc", "dynamic-loading"]} +[target.'cfg(target_os = "macos")'.dependencies] +objc2-metal = "0.3" +objc2 = "0.6" +objc2-foundation = "0.3" [build-dependencies] cbindgen = { version = "0.29", default-features = false } diff --git a/metatomic-core/cmake/metatomic-config.in.cmake b/metatomic-core/cmake/metatomic-config.in.cmake index 15652cbac..4bccc34ba 100644 --- a/metatomic-core/cmake/metatomic-config.in.cmake +++ b/metatomic-core/cmake/metatomic-config.in.cmake @@ -78,7 +78,15 @@ if (@METATOMIC_INSTALL_BOTH_STATIC_SHARED@ OR NOT @BUILD_SHARED_LIBS@) ) target_compile_features(metatomic::static INTERFACE cxx_std_17) - target_link_libraries(metatomic::static INTERFACE metatensor nlohmann_json::nlohmann_json) + + target_link_libraries(metatomic::static INTERFACE metatensor) + target_link_libraries(metatomic::static INTERFACE nlohmann_json::nlohmann_json) + + if(APPLE) + target_link_libraries(metatomic::static INTERFACE + "-framework Metal" "-framework CoreGraphics" "-framework CoreFoundation" "-framework Foundation" objc + ) + endif() endif() # Export either the shared or static library as the metatomic target diff --git a/metatomic-core/src/kernels/cuda.rs b/metatomic-core/src/kernels/cuda.rs index ccca5ca6c..291d55dc1 100644 --- a/metatomic-core/src/kernels/cuda.rs +++ b/metatomic-core/src/kernels/cuda.rs @@ -12,59 +12,13 @@ use ndarray::ArrayViewD; use crate::Error; +use super::StridedNDIndex; // CUDA kernel source compiled at runtime via NVRTC for the exact GPU const KERNEL_SRC: &str = include_str!("cuda_kernels.cu"); -const MAX_NDIM: usize = 7; - -/// Multi-dimensional strided index matching the CUDA `StridedNDIndex` struct. -/// Supports up to 7 dimensions. -/// -/// WARNING: any change here needs to be reflected in the CUDA source -#[repr(C)] -struct StridedNDIndex { - ndim: i64, - shape: [i64; MAX_NDIM], - strides: [i64; MAX_NDIM], -} - unsafe impl DeviceRepr for StridedNDIndex {} -impl StridedNDIndex { - fn from_dlpack(tensor: &DLPackTensorRef<'_>) -> Self { - return StridedNDIndex::from_shape_strides(tensor.shape(), tensor.strides()) - } - - #[allow(clippy::cast_possible_wrap)] - fn from_ndarray(array: &ArrayViewD<'_, T>) -> Self { - let shape = array.shape().iter().map(|&s| s as i64).collect::>(); - let strides = array.strides().iter().map(|&s| s as i64).collect::>(); - return Self::from_shape_strides(&shape, Some(&strides)); - } - - #[allow(clippy::cast_possible_wrap)] - fn from_shape_strides(shape: &[i64], strides: Option<&[i64]>) -> Self { - let ndim = shape.len(); - assert!(ndim <= MAX_NDIM, "StridedNDIndex only supports up to {MAX_NDIM} dimensions, got {ndim}"); - let mut shape_arr = [0i64; MAX_NDIM]; - let mut strides_arr = [0i64; MAX_NDIM]; - - // Contiguous fallback strides (row-major / C-contiguous) - let mut acc: i64 = 1; - for i in (0..ndim).rev() { - shape_arr[i] = shape[i]; - strides_arr[i] = acc; - acc *= shape[i]; - } - - if let Some(strides) = strides { - strides_arr[..ndim].copy_from_slice(&strides[..ndim]); - } - StridedNDIndex { ndim: ndim as i64, shape: shape_arr, strides: strides_arr } - } -} - /// Zero-cost wrapper to pass an existing device pointer as a CUDA kernel /// argument. /// diff --git a/metatomic-core/src/kernels/cuda_kernels.cu b/metatomic-core/src/kernels/cuda_kernels.cu index 641a55767..3fc47ad1b 100644 --- a/metatomic-core/src/kernels/cuda_kernels.cu +++ b/metatomic-core/src/kernels/cuda_kernels.cu @@ -7,7 +7,7 @@ /// Decomposes a flat linear index into multi-dimensional coordinates from the /// shape, then computes the strided memory offset using the stride array. /// -/// WARNING: any change here needs to be reflected in the Rust source +/// WARNING: any change here needs to be reflected in the Rust and Metal sources. struct StridedNDIndex { int64_t ndim; int64_t shape[MAX_NDIM]; diff --git a/metatomic-core/src/kernels/metal.rs b/metatomic-core/src/kernels/metal.rs new file mode 100644 index 000000000..3ce2d0a06 --- /dev/null +++ b/metatomic-core/src/kernels/metal.rs @@ -0,0 +1,290 @@ +use std::collections::{HashMap, hash_map::Entry}; +use std::ptr::NonNull; +use std::sync::Mutex; +use std::sync::LazyLock; + +use objc2::rc::Retained; +use objc2::runtime::ProtocolObject; +use objc2_foundation::ns_string; + +use objc2_metal::{ + MTLBuffer, MTLCommandBuffer, MTLCommandEncoder, MTLCommandQueue, + MTLComputeCommandEncoder, MTLComputePipelineState, + MTLCreateSystemDefaultDevice, MTLCompileOptions, + MTLDevice, MTLLibrary, MTLResourceOptions, MTLSize, +}; + +use dlpk::DLPackTensorRef; +use ndarray::ArrayViewD; + +use crate::Error; +use super::StridedNDIndex; + +const KERNEL_SRC: &str = include_str!("metal_kernels.metal"); + +/// Cached metal ressources: device, command queue, and pipeline states for kernels. +struct MetalKernelCache { + device: Retained>, + queue: Retained>, + is_equal_i32: Retained>, + validate_cell_pbc_f32: Retained>, +} + +impl MetalKernelCache { + fn new(device_id: usize) -> Result { + let device = MTLCreateSystemDefaultDevice() + .ok_or_else(|| Error::Internal(format!("no Metal device found for id {device_id}")))?; + + let library = device + .newLibraryWithSource_options_error( + ns_string!(KERNEL_SRC), + Some(&MTLCompileOptions::new()), + ) + .map_err(|e| Error::Internal(format!("MSL compile failed: {e}")))?; + + let is_equal_i32 = make_pipeline(&device, &library, "is_equal_i32")?; + let validate_cell_pbc_f32 = make_pipeline(&device, &library, "validate_cell_pbc_f32")?; + + let queue = device + .newCommandQueue() + .ok_or_else(|| Error::Internal("failed to create command queue".into()))?; + + Ok(Self { + device, + queue, + is_equal_i32, + validate_cell_pbc_f32, + }) + } +} + +fn make_pipeline( + device: &ProtocolObject, + library: &ProtocolObject, + name: &str, +) -> Result>, Error> { + use objc2_foundation::NSString; + + let ns_name = NSString::from_str(name); + let function = library + .newFunctionWithName(&ns_name) + .ok_or_else(|| Error::Internal(format!("get_function({name}): not found")))?; + + device + .newComputePipelineStateWithFunction_error(&function) + .map_err(|e| Error::Internal(format!("pipeline state ({name}): {e}"))) +} + +static METAL_CACHE: LazyLock>> = LazyLock::new(|| Mutex::new(HashMap::new())); + +fn get_or_init(cache: &mut HashMap, device_id: usize) -> Result<&MetalKernelCache, Error> { + let entry = match cache.entry(device_id) { + Entry::Occupied(entry) => entry.into_mut(), + Entry::Vacant(entry) => entry.insert(MetalKernelCache::new(device_id)?), + }; + Ok(entry) +} + +/// Compute the byte span of a DLPack tensor's data (including gaps from +/// strides). +#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)] +fn tensor_num_bytes(tensor: &DLPackTensorRef<'_>) -> usize { + let elem_size = tensor.dtype().bits as usize / 8; + let shape = tensor.shape(); + match tensor.strides() { + None => shape.iter().map(|&s| s as usize).product::() * elem_size, + Some(strides) => { + let max_idx: i64 = shape.iter() + .zip(strides.iter()) + .map(|(&s, &st)| (s - 1) * st) + .sum(); + (max_idx as usize + 1) * elem_size + } + } +} + +/// Extract a raw pointer to the tensor's data, accounting for byte_offset. +/// +/// # Safety +/// +/// The returned pointer is only valid as long as the DLPack tensor's backing +/// memory is alive. +#[allow(clippy::cast_possible_truncation)] +fn dlpack_data_ptr(tensor: &DLPackTensorRef<'_>) -> *const std::ffi::c_void { + unsafe { + tensor.raw.data.cast::().add(tensor.raw.byte_offset as usize).cast() + } +} + +/// Check that the values of a Metal-resident i32 DLPack tensor match an expected +/// reference array. +#[allow(clippy::cast_sign_loss, clippy::cast_possible_truncation)] +pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_, i32>) -> Result { + let device_id = tensor.device().device_id as usize; + let mut lock = METAL_CACHE.lock().expect("failed to lock METAL_CACHE"); + let cache = get_or_init(&mut lock, device_id)?; + + let n_elements: usize = tensor.shape().iter().map(|&s| s as usize).product(); + let ref_bytes = n_elements * std::mem::size_of::(); + + let ref_ptr: *const std::ffi::c_void = reference.as_slice().expect("reference should be contiguous").as_ptr().cast(); + + // Build strided indices for values and reference + let values_idx = StridedNDIndex::from_dlpack(&tensor); + let reference_idx = StridedNDIndex::from_ndarray(&reference); + + let values_buf = unsafe { + cache.device.newBufferWithBytes_length_options( + NonNull::new(dlpack_data_ptr(&tensor).cast_mut()).expect("values pointer must not be null"), + tensor_num_bytes(&tensor), + MTLResourceOptions::empty(), + ).expect("failed to create values buffer") + }; + let ref_buf = unsafe { + cache.device.newBufferWithBytes_length_options( + NonNull::new(ref_ptr.cast_mut()).expect("reference pointer must not be null"), + ref_bytes, + MTLResourceOptions::empty(), + ).expect("failed to create reference buffer") + }; + let result_buf = unsafe { + cache.device.newBufferWithBytes_length_options( + NonNull::from(&0i32).cast(), + std::mem::size_of::(), + MTLResourceOptions::empty(), + ).expect("failed to create result buffer") + }; + + objc2::rc::autoreleasepool(|_| { + let cmd_buf = cache.queue.commandBuffer().expect("failed to create command buffer"); + let encoder = cmd_buf.computeCommandEncoder().expect("failed to create compute encoder"); + + encoder.setComputePipelineState(&cache.is_equal_i32); + unsafe { + encoder.setBuffer_offset_atIndex(Some(&*values_buf), 0, 0); + + encoder.setBytes_length_atIndex( + NonNull::from(&values_idx).cast(), + std::mem::size_of::(), + 1, + ); + + encoder.setBuffer_offset_atIndex(Some(&*ref_buf), 0, 2); + + encoder.setBytes_length_atIndex( + NonNull::from(&reference_idx).cast(), + std::mem::size_of::(), + 3, + ); + + encoder.setBytes_length_atIndex( + NonNull::from(&(n_elements as u32)).cast(), + std::mem::size_of::(), + 4, + ); + + encoder.setBuffer_offset_atIndex(Some(&*result_buf), 0, 5); + } + + let tg_size = 32; + let tg_count = n_elements.div_ceil(tg_size); + encoder.dispatchThreadgroups_threadsPerThreadgroup( + MTLSize { width: tg_count, height: 1, depth: 1 }, + MTLSize { width: tg_size, height: 1, depth: 1 }, + ); + encoder.endEncoding(); + cmd_buf.commit(); + cmd_buf.waitUntilCompleted(); + }); + + let result = unsafe { + *result_buf.contents().as_ptr().cast::() + }; + return Ok(result == 0); +} + +/// Validate that cell vectors are zero for non-periodic dimensions on Metal. +#[allow(clippy::cast_sign_loss)] +pub(crate) fn validate_cell_pbc( + pbc: DLPackTensorRef<'_>, + cell: DLPackTensorRef<'_>, +) -> Result<(), Error> { + let device_id = pbc.device().device_id as usize; + let mut lock = METAL_CACHE.lock().expect("failed to lock METAL_CACHE"); + let cache = get_or_init(&mut lock, device_id)?; + + let pbc_idx = StridedNDIndex::from_dlpack(&pbc); + let cell_idx = StridedNDIndex::from_dlpack(&cell); + + let pbc_buf = unsafe { + cache.device.newBufferWithBytes_length_options( + NonNull::new(dlpack_data_ptr(&pbc).cast_mut()).expect("pbc pointer must not be null"), + tensor_num_bytes(&pbc), + MTLResourceOptions::empty(), + ).expect("failed to create pbc buffer") + }; + let cell_buf = unsafe { + cache.device.newBufferWithBytes_length_options( + NonNull::new(dlpack_data_ptr(&cell).cast_mut()).expect("cell pointer must not be null"), + tensor_num_bytes(&cell), + MTLResourceOptions::empty(), + ).expect("failed to create cell buffer") + }; + let result_buf = unsafe { + cache.device.newBufferWithBytes_length_options( + NonNull::from(&0i32).cast(), + std::mem::size_of::(), + MTLResourceOptions::empty(), + ).expect("failed to create result buffer") + }; + + objc2::rc::autoreleasepool(|_| { + let cmd_buf = cache.queue.commandBuffer().expect("failed to create command buffer"); + let encoder = cmd_buf.computeCommandEncoder().expect("failed to create compute encoder"); + + assert!(cell.dtype().bits == 32, "only float32 is supported on Metal"); + + encoder.setComputePipelineState(&cache.validate_cell_pbc_f32); + unsafe { + encoder.setBuffer_offset_atIndex(Some(&*pbc_buf), 0, 0); + + encoder.setBytes_length_atIndex( + NonNull::from(&pbc_idx).cast(), + std::mem::size_of::(), + 1, + ); + + encoder.setBuffer_offset_atIndex(Some(&*cell_buf), 0, 2); + + encoder.setBytes_length_atIndex( + NonNull::from(&cell_idx).cast(), + std::mem::size_of::(), + 3, + ); + + encoder.setBuffer_offset_atIndex(Some(&*result_buf), 0, 4); + } + + encoder.dispatchThreadgroups_threadsPerThreadgroup( + MTLSize { width: 1, height: 1, depth: 1 }, + MTLSize { width: 3, height: 1, depth: 1 }, + ); + encoder.endEncoding(); + cmd_buf.commit(); + cmd_buf.waitUntilCompleted(); + }); + + let result = unsafe { + *result_buf.contents().as_ptr().cast::() + }; + + if result != 0 { + let dim = result - 1; + return Err(Error::InvalidParameter(format!( + "invalid cell: for non-periodic dimensions, the corresponding \ + cell vector must be zero, but cell[{}] contains non-zero values", + dim + ))); + } + Ok(()) +} diff --git a/metatomic-core/src/kernels/metal_kernels.metal b/metatomic-core/src/kernels/metal_kernels.metal new file mode 100644 index 000000000..aff14a43e --- /dev/null +++ b/metatomic-core/src/kernels/metal_kernels.metal @@ -0,0 +1,76 @@ +#include +using namespace metal; + +// --------------------------------------------------------------------------- +// Multi-dimensional strided index helper (up to MAX_NDIM dimensions). +// +// Decomposes a flat linear index into multi-dimensional coordinates based on +// the shape and then computes the strided memory offset using the stride +// array. +// +// WARNING: the layout of this struct must match both the CUDA +// (cuda_kernels.cu) and Rust (kernels/mod.rs) definitions. +// --------------------------------------------------------------------------- +constant long MAX_NDIM [[maybe_unused]] = 7; + +struct StridedNDIndex { + long ndim; + long shape[MAX_NDIM]; + long strides[MAX_NDIM]; + + long offset(long flat_idx) const { + long off = 0; + for (int d = ndim - 1; d >= 0; d--) { + long coord = flat_idx % shape[d]; + flat_idx /= shape[d]; + off += coord * strides[d]; + } + return off; + } +}; + +//////////////////////////////////////////////////////////////////////////////// +//////////////////////////////////////////////////////////////////////////////// + +kernel void is_equal_i32( + [[buffer(0)]] device const int* values, + [[buffer(1)]] constant StridedNDIndex& values_idx, + [[buffer(2)]] device const int* reference, + [[buffer(3)]] constant StridedNDIndex& reference_idx, + [[buffer(4)]] constant uint& n, + [[buffer(5)]] device atomic_int* mismatch, + [[thread_position_in_grid]] uint gid +) { + if (gid < n) { + long v_off = values_idx.offset(gid); + long r_off = reference_idx.offset(gid); + if (values[v_off] != reference[r_off]) { + atomic_fetch_max_explicit(mismatch, 1, memory_order_relaxed); + } + } +} + +//////////////////////////////////////////////////////////////////////////////// +//////////////////////////////////////////////////////////////////////////////// + +/// Validate cell vectors against PBC flags (f32 only on Metal). +kernel void validate_cell_pbc_f32( + [[buffer(0)]] device const bool* pbc, + [[buffer(1)]] constant StridedNDIndex& pbc_idx, + [[buffer(2)]] device const float* cell, + [[buffer(3)]] constant StridedNDIndex& cell_idx, + [[buffer(4)]] device atomic_int* mismatch_idx, + [[thread_position_in_threadgroup]] uint tid +) { + if (tid < 3) { + if (!pbc[pbc_idx.offset(tid)]) { + if ( + cell[cell_idx.offset(tid * 3 + 0)] != 0.0f || + cell[cell_idx.offset(tid * 3 + 1)] != 0.0f || + cell[cell_idx.offset(tid * 3 + 2)] != 0.0f + ) { + atomic_fetch_max_explicit(mismatch_idx, int(tid + 1), memory_order_relaxed); + } + } + } +} diff --git a/metatomic-core/src/kernels/mod.rs b/metatomic-core/src/kernels/mod.rs index ba0c39991..f31869166 100644 --- a/metatomic-core/src/kernels/mod.rs +++ b/metatomic-core/src/kernels/mod.rs @@ -7,6 +7,66 @@ use crate::Error; mod cpu; mod cuda; +#[cfg(target_os = "macos")] +mod metal; + +const MAX_NDIM: usize = 7; + +/// Multi-dimensional strided index (up to MAX_NDIM dimensions). +/// +/// Decomposes a flat linear index into multi-dimensional coordinates from the +/// shape, then computes the strided memory offset using the stride array. +/// +/// WARNING: any change here needs to be reflected in the CUDA and Metal sources. +#[repr(C)] +pub(crate) struct StridedNDIndex { + pub(crate) ndim: i64, + pub(crate) shape: [i64; MAX_NDIM], + pub(crate) strides: [i64; MAX_NDIM], +} + +#[allow(clippy::cast_possible_wrap)] +impl StridedNDIndex { + /// Create a `StridedNDIndex` from a DLPack tensor's shape and strides. + pub(crate) fn from_dlpack(tensor: &DLPackTensorRef<'_>) -> Self { + Self::from_shape_strides(tensor.shape(), tensor.strides()) + } + + /// Create a `StridedNDIndex` from an ndarray view's shape and strides. + pub(crate) fn from_ndarray(array: &ArrayViewD<'_, T>) -> Self { + let shape: Vec = array.shape().iter().map(|&s| s as i64).collect(); + let strides: Vec = array.strides().iter().map(|&s| s as i64).collect(); + Self::from_shape_strides(&shape, Some(&strides)) + } + + /// Create a `StridedNDIndex` from shape and optional strides. + /// + /// If strides is `None`, the strides are computed as if the array were + /// contiguous (row-major / C-contiguous). + pub(crate) fn from_shape_strides(shape: &[i64], strides: Option<&[i64]>) -> Self { + let ndim = shape.len(); + assert!( + ndim <= MAX_NDIM, + "StridedNDIndex only supports up to {MAX_NDIM} dimensions, got {ndim}" + ); + let mut shape_arr = [0i64; MAX_NDIM]; + let mut strides_arr = [0i64; MAX_NDIM]; + + // Contiguous fallback strides (row-major / C-contiguous) + let mut acc: i64 = 1; + for i in (0..ndim).rev() { + shape_arr[i] = shape[i]; + strides_arr[i] = acc; + acc *= shape[i]; + } + + if let Some(strides) = strides { + strides_arr[..ndim].copy_from_slice(&strides[..ndim]); + } + StridedNDIndex { ndim: ndim as i64, shape: shape_arr, strides: strides_arr } + } +} + /// Check that the values of an i32 DLPack tensor match the expected reference. /// /// This dispatches to the appropriate backend based on the device of `tensor`. @@ -22,6 +82,16 @@ pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_ DLDeviceType::kDLCUDA | DLDeviceType::kDLCUDAManaged => { cuda::is_equal_i32(tensor, reference) } + DLDeviceType::kDLMetal => { + #[cfg(target_os = "macos")] { + metal::is_equal_i32(tensor, reference) + } + #[cfg(not(target_os = "macos"))] { + Err(Error::Internal( + "Metal backend is only available on macOS".into(), + )) + } + } _ => { eprintln!( "is_equal_i32 for device {:?} is not implemented", @@ -52,6 +122,16 @@ pub(crate) fn validate_cell_pbc(pbc: DLPackTensorRef<'_>, cell: DLPackTensorRef< DLDeviceType::kDLCUDA | DLDeviceType::kDLCUDAManaged => { cuda::validate_cell_pbc(pbc, cell) } + DLDeviceType::kDLMetal => { + #[cfg(target_os = "macos")] { + metal::validate_cell_pbc(pbc, cell) + } + #[cfg(not(target_os = "macos"))] { + Err(Error::Internal( + "Metal backend is only available on macOS".into(), + )) + } + } _ => { eprintln!( "Cell/PBC validation for device {:?} is not implemented", From b57a44af40af795d5b35b2e2fa434a952c9c982b Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Tue, 14 Jul 2026 11:51:45 +0200 Subject: [PATCH 38/43] Cache reference values on device This way we don't have to keep uploading small values to the device over and over again --- metatomic-core/src/kernels/cpu.rs | 5 +-- metatomic-core/src/kernels/cuda.rs | 31 +++++++++--------- metatomic-core/src/kernels/metal.rs | 49 ++++++++++++++++++++--------- metatomic-core/src/kernels/mod.rs | 29 +++++++++++++++-- metatomic-core/src/system.rs | 31 ++++++++++++------ 5 files changed, 102 insertions(+), 43 deletions(-) diff --git a/metatomic-core/src/kernels/cpu.rs b/metatomic-core/src/kernels/cpu.rs index c054518b0..1225a9042 100644 --- a/metatomic-core/src/kernels/cpu.rs +++ b/metatomic-core/src/kernels/cpu.rs @@ -2,6 +2,7 @@ use dlpk::DLPackTensorRef; use ndarray::{ArrayView1, ArrayView2, ArrayViewD}; use crate::Error; +use super::ReferenceValue; /// Check that the values of an i32 DLPack tensor match the expected reference. /// @@ -14,10 +15,10 @@ use crate::Error; /// - `reference`: expected values with the same shape as the tensor pub(crate) fn is_equal_i32( tensor: DLPackTensorRef<'_>, - reference: ArrayViewD<'_, i32>, + reference: &ReferenceValue, ) -> Result { let values: ArrayViewD = tensor.try_into()?; - return Ok(values == reference); + return Ok(values == reference.cpu.view()); } macro_rules! validate_cell { diff --git a/metatomic-core/src/kernels/cuda.rs b/metatomic-core/src/kernels/cuda.rs index 291d55dc1..b1fcb8334 100644 --- a/metatomic-core/src/kernels/cuda.rs +++ b/metatomic-core/src/kernels/cuda.rs @@ -8,11 +8,9 @@ use cudarc::driver::safe::{ }; use cudarc::nvrtc::compile_ptx; use dlpk::DLPackTensorRef; -use ndarray::ArrayViewD; - use crate::Error; -use super::StridedNDIndex; +use super::{ReferenceValue, StridedNDIndex}; // CUDA kernel source compiled at runtime via NVRTC for the exact GPU const KERNEL_SRC: &str = include_str!("cuda_kernels.cu"); @@ -104,11 +102,12 @@ unsafe fn dlpack_to_device_ptr(tensor: &DLPackTensorRef<'_>) -> cudarc::driver:: /// Check that the values of a CUDA-resident i32 DLPack tensor match an expected /// reference array. /// -/// The comparison is performed entirely on-device: the existing GPU pointer from -/// `tensor` is wrapped as a `DevicePtrArg`, the reference is uploaded to the GPU, -/// and a single-element result flag (`0` = ok, `1` = mismatch) is read back. +/// The comparison is performed entirely on-device: the existing GPU pointer +/// from `tensor` is wrapped as a `DevicePtrArg`, the reference is uploaded to +/// the GPU (and cached for subsequent calls), and a single-element result flag +/// (`0` = ok, `1` = mismatch) is read back. #[allow(clippy::cast_sign_loss, clippy::cast_possible_truncation)] -pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_, i32>) -> Result { +pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: &ReferenceValue) -> Result { debug_assert!( tensor.device().device_type == dlpk::sys::DLDeviceType::kDLCUDA, "is_equal_i32 called on non-CUDA tensor" @@ -125,15 +124,17 @@ pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_ // Build strided index from the DLPack tensor (preserves actual strides) let values_idx = StridedNDIndex::from_dlpack(&tensor); - // Build strided index from the ndarray view (preserves actual strides) - let reference_idx = StridedNDIndex::from_ndarray(&reference); - // Wrap the existing GPU-allocated tensor pointer let tensor_ptr = unsafe { DevicePtrArg { ptr: dlpack_to_device_ptr(&tensor) } }; - // Upload reference values to GPU - let ref_dev = stream.clone_htod(reference.as_slice().expect("reference should be contiguous")) - .map_err(|e| Error::Internal(format!("clone_htod reference: {e}")))?; + // Upload reference values to GPU (cached after first call) + let (ref_dev, reference_idx) = reference.cuda.get_or_init(|| { + let slice = stream + .clone_htod(reference.cpu.as_slice().expect("reference should be contiguous")) + .expect("clone_htod reference failed"); + let idx = StridedNDIndex::from_ndarray(&reference.cpu.view()); + (slice, idx) + }); // Allocate result flag (initialized to 0 = no mismatch) let mut result = stream.alloc_zeros::(1) @@ -143,8 +144,8 @@ pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_ stream.launch_builder(&entry.is_equal_i32) .arg(&tensor_ptr) .arg(&values_idx) - .arg(&ref_dev) - .arg(&reference_idx) + .arg(ref_dev) + .arg(reference_idx) .arg(&n_elements) .arg(&mut result) .launch(LaunchConfig::for_num_elems(u32::try_from(n_elements).expect("tensor too large for CUDA kernel"))) diff --git a/metatomic-core/src/kernels/metal.rs b/metatomic-core/src/kernels/metal.rs index 3ce2d0a06..67339d855 100644 --- a/metatomic-core/src/kernels/metal.rs +++ b/metatomic-core/src/kernels/metal.rs @@ -15,10 +15,23 @@ use objc2_metal::{ }; use dlpk::DLPackTensorRef; -use ndarray::ArrayViewD; use crate::Error; -use super::StridedNDIndex; +use super::{ReferenceValue, StridedNDIndex}; + +// Small wrapper around MTLBuffer to implement Send and Sync, since the data is +// read-only after initialization. +pub(crate) struct MetalBuffer(Retained>); + +unsafe impl Send for MetalBuffer {} +unsafe impl Sync for MetalBuffer {} + +impl std::ops::Deref for MetalBuffer { + type Target = ProtocolObject; + fn deref(&self) -> &Self::Target { + &self.0 + } +} const KERNEL_SRC: &str = include_str!("metal_kernels.metal"); @@ -119,7 +132,7 @@ fn dlpack_data_ptr(tensor: &DLPackTensorRef<'_>) -> *const std::ffi::c_void { /// Check that the values of a Metal-resident i32 DLPack tensor match an expected /// reference array. #[allow(clippy::cast_sign_loss, clippy::cast_possible_truncation)] -pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_, i32>) -> Result { +pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: &ReferenceValue) -> Result { let device_id = tensor.device().device_id as usize; let mut lock = METAL_CACHE.lock().expect("failed to lock METAL_CACHE"); let cache = get_or_init(&mut lock, device_id)?; @@ -127,11 +140,26 @@ pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_ let n_elements: usize = tensor.shape().iter().map(|&s| s as usize).product(); let ref_bytes = n_elements * std::mem::size_of::(); - let ref_ptr: *const std::ffi::c_void = reference.as_slice().expect("reference should be contiguous").as_ptr().cast(); - - // Build strided indices for values and reference + // Build strided index for the values let values_idx = StridedNDIndex::from_dlpack(&tensor); - let reference_idx = StridedNDIndex::from_ndarray(&reference); + + // Upload reference values to Metal (cached after first call) + let (ref_buf, reference_idx) = reference.metal.get_or_init(|| { + let ref_bytes = reference.cpu.len() * std::mem::size_of::(); + let ref_ptr: *const std::ffi::c_void = reference.cpu.as_slice() + .expect("reference should be contiguous") + .as_ptr() + .cast(); + let buf = unsafe { + cache.device.newBufferWithBytes_length_options( + NonNull::new(ref_ptr.cast_mut()).expect("reference pointer must not be null"), + ref_bytes, + MTLResourceOptions::empty(), + ).expect("failed to create reference buffer") + }; + let idx = StridedNDIndex::from_ndarray(&reference.cpu.view()); + (MetalBuffer(buf), idx) + }); let values_buf = unsafe { cache.device.newBufferWithBytes_length_options( @@ -140,13 +168,6 @@ pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_ MTLResourceOptions::empty(), ).expect("failed to create values buffer") }; - let ref_buf = unsafe { - cache.device.newBufferWithBytes_length_options( - NonNull::new(ref_ptr.cast_mut()).expect("reference pointer must not be null"), - ref_bytes, - MTLResourceOptions::empty(), - ).expect("failed to create reference buffer") - }; let result_buf = unsafe { cache.device.newBufferWithBytes_length_options( NonNull::from(&0i32).cast(), diff --git a/metatomic-core/src/kernels/mod.rs b/metatomic-core/src/kernels/mod.rs index f31869166..6c586d8c0 100644 --- a/metatomic-core/src/kernels/mod.rs +++ b/metatomic-core/src/kernels/mod.rs @@ -1,6 +1,9 @@ +use std::sync::OnceLock; + +use cudarc::driver::CudaSlice; use dlpk::sys::DLDeviceType; use dlpk::DLPackTensorRef; -use ndarray::ArrayViewD; +use ndarray::{ArrayD, ArrayViewD}; use crate::Error; @@ -67,6 +70,28 @@ impl StridedNDIndex { } } +/// Store and cache reference values for different backends (CPU, CUDA, Metal). +pub struct ReferenceValue { + /// The reference values stored on the CPU, always there + pub(crate) cpu: ArrayD, + /// Reference values stored on CUDA, intialized on first use from the CPU values + pub(crate) cuda: OnceLock<(CudaSlice, StridedNDIndex)>, + #[cfg(target_os = "macos")] + /// Reference values stored on Metal, intialized on first use from the CPU values + pub(crate) metal: OnceLock<(metal::MetalBuffer, StridedNDIndex)>, +} + +impl ReferenceValue { + pub(crate) fn new(cpu: ArrayD) -> Self { + Self { + cpu, + cuda: OnceLock::new(), + #[cfg(target_os = "macos")] + metal: OnceLock::new(), + } + } +} + /// Check that the values of an i32 DLPack tensor match the expected reference. /// /// This dispatches to the appropriate backend based on the device of `tensor`. @@ -74,7 +99,7 @@ impl StridedNDIndex { /// # Parameters /// - `tensor`: DLPack tensor with i32 data type /// - `reference`: expected values with the same shape as the tensor -pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: ArrayViewD<'_, i32>) -> Result { +pub(crate) fn is_equal_i32(tensor: DLPackTensorRef<'_>, reference: &ReferenceValue) -> Result { match tensor.device().device_type { DLDeviceType::kDLCPU | DLDeviceType::kDLCUDAHost | DLDeviceType::kDLROCMHost => { cpu::is_equal_i32(tensor, reference) diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs index 26b11d844..740cece00 100644 --- a/metatomic-core/src/system.rs +++ b/metatomic-core/src/system.rs @@ -6,12 +6,31 @@ use dlpk::{DLPackTensor, DLPackTensorRef}; use metatensor::{TensorBlock, TensorMap}; use crate::{Error, PairListOptions}; +use crate::kernels::ReferenceValue; /// Names that can never be used as custom data in a system static INVALID_DATA_NAMES: LazyLock> = LazyLock::new(|| { HashSet::from(["types", "type", "positions", "position", "cell", "neighbors", "neighbor", "pair", "pairs"]) }); +static XYZ_REFERENCE: LazyLock> = LazyLock::new(|| { + ReferenceValue::new( + ndarray::ArrayD::from_shape_vec( + ndarray::IxDyn(&[3usize, 1]), + vec![0i32, 1, 2], + ).unwrap() + ) +}); + +static DISTANCE_REFERENCE: LazyLock> = LazyLock::new(|| { + ReferenceValue::new( + ndarray::ArrayD::from_shape_vec( + ndarray::IxDyn(&[1usize, 1]), + vec![0i32], + ).unwrap() + ) +}); + /// Storage for an atomistic system. /// /// This owns the raw DLPack tensors and metatensor objects used at FFI @@ -126,12 +145,8 @@ impl System { None, dlpk::sys::DLPackVersion::current(), )?; - let reference = ndarray::ArrayViewD::::from_shape( - ndarray::IxDyn(&[3usize, 1]), - &[0i32, 1, 2], - ).unwrap(); - if !crate::kernels::is_equal_i32(dl_tensor.as_ref(), reference)? { + if !crate::kernels::is_equal_i32(dl_tensor.as_ref(), &XYZ_REFERENCE)? { return Err(Error::InvalidParameter( "invalid components for `pairs`: the 'xyz' component should \ contain [[0], [1], [2]]".into() @@ -154,12 +169,8 @@ impl System { None, dlpk::sys::DLPackVersion::current(), )?; - let reference = ndarray::ArrayViewD::::from_shape( - ndarray::IxDyn(&[1usize, 1]), - &[0i32], - ).unwrap(); - if !crate::kernels::is_equal_i32(dl_tensor.as_ref(), reference)? { + if !crate::kernels::is_equal_i32(dl_tensor.as_ref(), &DISTANCE_REFERENCE)? { return Err(Error::InvalidParameter( "invalid properties for `pairs`: the 'distance' property \ should contain [0]".into() From 8ff51e071041d5c6f4865683f7654df6dd080e8f Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Wed, 15 Jul 2026 10:53:55 +0200 Subject: [PATCH 39/43] Remove cargo caching in CI, it interacts badly with our setup We create virtualenvs inside target/tmp, and these get corrupted by the caching setup --- .github/workflows/python-tests.yml | 5 ----- .github/workflows/rust-tests.yml | 6 ------ .github/workflows/torch-tests.yml | 5 ----- 3 files changed, 16 deletions(-) diff --git a/.github/workflows/python-tests.yml b/.github/workflows/python-tests.yml index 282f915d4..584fd6693 100644 --- a/.github/workflows/python-tests.yml +++ b/.github/workflows/python-tests.yml @@ -62,11 +62,6 @@ jobs: with: toolchain: stable - - name: Cache Rust dependencies - uses: Leafwing-Studios/cargo-cache@v2.6.1 - with: - sweep-cache: true - - name: Setup sccache if: ${{ !env.ACT }} uses: mozilla-actions/sccache-action@v0.0.10 diff --git a/.github/workflows/rust-tests.yml b/.github/workflows/rust-tests.yml index f9c2a1c4d..36eb59bff 100644 --- a/.github/workflows/rust-tests.yml +++ b/.github/workflows/rust-tests.yml @@ -104,12 +104,6 @@ jobs: with: python-version: "3.14" - - name: Cache Rust dependencies - uses: Leafwing-Studios/cargo-cache@v2.6.1 - if: matrix.container == null - with: - sweep-cache: true - - name: install valgrind if: matrix.do-valgrind run: | diff --git a/.github/workflows/torch-tests.yml b/.github/workflows/torch-tests.yml index 93661740a..d9dd09dc3 100644 --- a/.github/workflows/torch-tests.yml +++ b/.github/workflows/torch-tests.yml @@ -70,11 +70,6 @@ jobs: with: toolchain: stable - - name: Cache Rust dependencies - uses: Leafwing-Studios/cargo-cache@v2.6.1 - with: - sweep-cache: true - - name: install valgrind if: matrix.do-valgrind run: | From be1fbdd5e73d2626c551d64fa7d6721cb526d4d3 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Thu, 16 Jul 2026 15:31:41 +0200 Subject: [PATCH 40/43] Use an output parameter in mta_load{_buffer} This is consistent with mta_system_create --- metatomic-core/include/metatomic.h | 25 ++- metatomic-core/src/c_api/io.rs | 66 +++---- metatomic-core/src/c_api/system.rs | 4 +- metatomic-core/tests/system.cpp | 293 +++++++++++++++++++++++++++++ 4 files changed, 337 insertions(+), 51 deletions(-) diff --git a/metatomic-core/include/metatomic.h b/metatomic-core/include/metatomic.h index 430d7b9f1..9dbe221d8 100644 --- a/metatomic-core/include/metatomic.h +++ b/metatomic-core/include/metatomic.h @@ -752,11 +752,14 @@ enum mta_status_t mta_save_buffer(uint8_t **buffer, * null. * @param create_array Callback to allocate arrays for the system's data. Must * not be NULL. - * @return A pointer to the newly allocated system. The caller takes ownership - * and must free it with `mta_system_free`. Returns NULL on error; use - * `mta_last_error` for details. + * @param system Output parameter, set to the newly created system handle. + * The caller takes ownership and must free it with `mta_system_free`. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ -struct mta_system_t *mta_load(const char *path, mts_create_array_callback_t create_array); +enum mta_status_t mta_load(const char *path, + mts_create_array_callback_t create_array, + struct mta_system_t **system); /** * Load a system from an in-memory buffer. @@ -768,13 +771,15 @@ struct mta_system_t *mta_load(const char *path, mts_create_array_callback_t crea * @param buffer_size Number of bytes in `buffer`. * @param create_array Callback to allocate arrays for the system's data. Must * not be NULL. - * @return A pointer to the newly allocated system. The caller takes ownership - * and must free it with `mta_system_free`. Returns NULL on error; use - * `mta_last_error` for details. + * @param system Output parameter, set to the newly created system handle. + * The caller takes ownership and must free it with `mta_system_free`. + * @return `MTA_SUCCESS` on success, or another status code if an error occurs. + * You can get more details about the error with `mta_last_error`. */ -struct mta_system_t *mta_load_buffer(const uint8_t *buffer, - uintptr_t buffer_size, - mts_create_array_callback_t create_array); +enum mta_status_t mta_load_buffer(const uint8_t *buffer, + uintptr_t buffer_size, + mts_create_array_callback_t create_array, + struct mta_system_t **system); #ifdef __cplusplus } // extern "C" diff --git a/metatomic-core/src/c_api/io.rs b/metatomic-core/src/c_api/io.rs index b6170fb1a..ba803f7f7 100644 --- a/metatomic-core/src/c_api/io.rs +++ b/metatomic-core/src/c_api/io.rs @@ -1,7 +1,6 @@ use std::ffi::{c_char, c_void, CStr}; use std::fs::File; use std::io::{BufReader, Cursor}; -use std::sync::Arc; use metatensor::c_api::{mts_create_array_callback_t, mts_realloc_buffer_t}; @@ -244,37 +243,31 @@ pub unsafe extern "C" fn mta_save_buffer( /// null. /// @param create_array Callback to allocate arrays for the system's data. Must /// not be NULL. -/// @return A pointer to the newly allocated system. The caller takes ownership -/// and must free it with `mta_system_free`. Returns NULL on error; use -/// `mta_last_error` for details. +/// @param system Output parameter, set to the newly created system handle. +/// The caller takes ownership and must free it with `mta_system_free`. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[unsafe(no_mangle)] pub unsafe extern "C" fn mta_load( path: *const c_char, - create_array: mts_create_array_callback_t -) -> *mut mta_system_t { - let mut result = std::ptr::null_mut(); - let unwind_wrapper = std::panic::AssertUnwindSafe(&mut result); - let status = catch_unwind(move || { + create_array: mts_create_array_callback_t, + system: *mut *mut mta_system_t, +) -> mta_status_t { + catch_unwind(move || { check_pointers_non_null!(path); let path = unsafe { CStr::from_ptr(path) }.to_str() .map_err(|_| Error::InvalidParameter("path is not valid UTF-8".into()))?; let file = BufReader::new(File::open(path)?); - let system = crate::io::load(file, create_array)?; + let new_system = mta_system_t(crate::io::load(file, create_array)?); - let system = Arc::new(mta_system_t(system)); + unsafe { + *system = mta_system_t::into_raw(new_system); + } - let _ = &unwind_wrapper; - *unwind_wrapper.0 = Arc::into_raw(system).cast_mut(); Ok(()) - }); - - if status != mta_status_t::MTA_SUCCESS { - return std::ptr::null_mut(); - } - - return result; + }) } /// Load a system from an in-memory buffer. @@ -286,35 +279,30 @@ pub unsafe extern "C" fn mta_load( /// @param buffer_size Number of bytes in `buffer`. /// @param create_array Callback to allocate arrays for the system's data. Must /// not be NULL. -/// @return A pointer to the newly allocated system. The caller takes ownership -/// and must free it with `mta_system_free`. Returns NULL on error; use -/// `mta_last_error` for details. +/// @param system Output parameter, set to the newly created system handle. +/// The caller takes ownership and must free it with `mta_system_free`. +/// @return `MTA_SUCCESS` on success, or another status code if an error occurs. +/// You can get more details about the error with `mta_last_error`. #[unsafe(no_mangle)] pub unsafe extern "C" fn mta_load_buffer( buffer: *const u8, buffer_size: usize, - create_array: mts_create_array_callback_t -) -> *mut mta_system_t { - let mut result = std::ptr::null_mut(); - let unwind_wrapper = std::panic::AssertUnwindSafe(&mut result); - let status = catch_unwind(move || { + create_array: mts_create_array_callback_t, + system: *mut *mut mta_system_t, +) -> mta_status_t { + catch_unwind(move || { check_pointers_non_null!(buffer); let slice = unsafe { std::slice::from_raw_parts(buffer, buffer_size) }; let cursor = Cursor::new(slice); - let system = crate::io::load(cursor, create_array)?; - let system = Arc::new(mta_system_t(system)); + let new_system = mta_system_t(crate::io::load(cursor, create_array)?); - let _ = &unwind_wrapper; - *unwind_wrapper.0 = Arc::into_raw(system).cast_mut(); - Ok(()) - }); - - if status != mta_status_t::MTA_SUCCESS { - return std::ptr::null_mut(); - } + unsafe { + *system = mta_system_t::into_raw(new_system); + } - return result; + Ok(()) + }) } diff --git a/metatomic-core/src/c_api/system.rs b/metatomic-core/src/c_api/system.rs index 97d2b9f4e..ec78bd5fe 100644 --- a/metatomic-core/src/c_api/system.rs +++ b/metatomic-core/src/c_api/system.rs @@ -19,13 +19,13 @@ pub struct mta_system_t(pub(crate) System); impl mta_system_t { /// Convert an mta_system_t into a pointer inside an Arc, to be /// passed through the C API - fn into_raw(self) -> *mut mta_system_t { + pub(crate) fn into_raw(self) -> *mut mta_system_t { Arc::into_raw(Arc::new(self)).cast_mut() } /// Recover the Arc from a pointer created with /// [`mta_system_t::into_raw`] - unsafe fn from_raw(ptr: *const mta_system_t) -> Arc { + pub(crate) unsafe fn from_raw(ptr: *const mta_system_t) -> Arc { unsafe { Arc::from_raw(ptr) } } } diff --git a/metatomic-core/tests/system.cpp b/metatomic-core/tests/system.cpp index fb67a2390..638b31154 100644 --- a/metatomic-core/tests/system.cpp +++ b/metatomic-core/tests/system.cpp @@ -1,5 +1,9 @@ +#include +#include #include +#include #include +#include #include @@ -556,3 +560,292 @@ TEST_CASE("system custom data") { mta_system_free(system); } + +/// Build a system containing all kinds of data (basic data, pairs, and custom +/// data) for use in serialization round-trip tests. +static mta_system_t* full_test_system() { + mta_system_t* system = nullptr; + auto status = mta_system_create( + "Angstrom", + types_tensor(4), + positions_tensor(4), + cell_tensor(), + pbc_tensor(), + &system + ); + REQUIRE(status == MTA_SUCCESS); + REQUIRE(system != nullptr); + + const auto* pairs_options_json = R"({ + "type": "metatomic_pair_options", + "cutoff": "0x00001000", + "full_list": true, + "strict": false, + "requestors": ["test"] + })"; + status = mta_system_add_pairs(system, pairs_options_json, pair_block()); + CHECK(status == MTA_SUCCESS); + + status = mta_system_add_custom_data(system, "test::my_data", custom_data()); + CHECK(status == MTA_SUCCESS); + + return system; +} + +/// Check that the given system contains the data expected from +/// `full_test_system`, independently of how it was loaded back. +static void check_full_system_data(const mta_system_t* system) { + uintptr_t size = 0; + CHECK(mta_system_size(system, &size) == MTA_SUCCESS); + CHECK(size == 4); + + mta_string_t unit = nullptr; + CHECK(mta_system_get_length_unit(system, &unit) == MTA_SUCCESS); + CHECK(std::string(mta_string_view(unit)) == "Angstrom"); + mta_string_free(unit); + + DLManagedTensorVersioned* data = nullptr; + + // types + CHECK(mta_system_get_data(system, MTA_SYSTEM_DATA_TYPES, &data) == MTA_SUCCESS); + REQUIRE(data != nullptr); + CHECK(data->dl_tensor.ndim == 1); + CHECK(data->dl_tensor.shape[0] == 4); + CHECK(data->dl_tensor.dtype.code == kDLInt); + CHECK(data->dl_tensor.dtype.bits == 32); + { + auto* types = reinterpret_cast( + static_cast(data->dl_tensor.data) + data->dl_tensor.byte_offset + ); + CHECK(types[0] == 1); + CHECK(types[1] == 4); + CHECK(types[2] == 7); + CHECK(types[3] == 10); + } + data->deleter(data); + + // positions + CHECK(mta_system_get_data(system, MTA_SYSTEM_DATA_POSITIONS, &data) == MTA_SUCCESS); + REQUIRE(data != nullptr); + CHECK(data->dl_tensor.ndim == 2); + CHECK(data->dl_tensor.shape[0] == 4); + CHECK(data->dl_tensor.shape[1] == 3); + CHECK(data->dl_tensor.dtype.code == kDLFloat); + CHECK(data->dl_tensor.dtype.bits == 32); + { + auto* positions = reinterpret_cast( + static_cast(data->dl_tensor.data) + data->dl_tensor.byte_offset + ); + CHECK(positions[0] == 1.0F); + CHECK(positions[3] == 4.0F); + CHECK(positions[6] == 7.0F); + CHECK(positions[9] == 10.0F); + } + data->deleter(data); + + // cell + CHECK(mta_system_get_data(system, MTA_SYSTEM_DATA_CELL, &data) == MTA_SUCCESS); + REQUIRE(data != nullptr); + CHECK(data->dl_tensor.ndim == 2); + CHECK(data->dl_tensor.shape[0] == 3); + CHECK(data->dl_tensor.shape[1] == 3); + CHECK(data->dl_tensor.dtype.code == kDLFloat); + CHECK(data->dl_tensor.dtype.bits == 32); + { + auto* cell = reinterpret_cast( + static_cast(data->dl_tensor.data) + data->dl_tensor.byte_offset + ); + CHECK(cell[0] == 10.0F); + CHECK(cell[4] == 0.0F); + CHECK(cell[8] == 10.0F); + } + data->deleter(data); + + // pbc + CHECK(mta_system_get_data(system, MTA_SYSTEM_DATA_PBC, &data) == MTA_SUCCESS); + REQUIRE(data != nullptr); + CHECK(data->dl_tensor.ndim == 1); + CHECK(data->dl_tensor.shape[0] == 3); + CHECK(data->dl_tensor.dtype.code == kDLBool); + CHECK(data->dl_tensor.dtype.bits == 8); + { + auto* pbc = reinterpret_cast( + static_cast(data->dl_tensor.data) + data->dl_tensor.byte_offset + ); + CHECK(pbc[0] == 1); + CHECK(pbc[1] == 0); + CHECK(pbc[2] == 1); + } + data->deleter(data); + + // known pairs survive the round-trip + mta_string_t known = nullptr; + CHECK(mta_system_known_pairs(system, &known) == MTA_SUCCESS); + REQUIRE(known != nullptr); + { + auto known_str = std::string(mta_string_view(known)); + CHECK(known_str.find("metatomic_pair_options") != std::string::npos); + CHECK(known_str.find("\"full_list\":true") != std::string::npos); + } + mta_string_free(known); + + // the pairs block can be retrieved + const auto* pairs_options_json = R"({ + "type": "metatomic_pair_options", + "cutoff": "0x00001000", + "full_list": true, + "strict": false, + "requestors": [] + })"; + const mts_block_t* pairs = nullptr; + CHECK(mta_system_get_pairs(system, pairs_options_json, &pairs) == MTA_SUCCESS); + CHECK(pairs != nullptr); + + // custom data survives the round-trip + mta_string_t names = nullptr; + CHECK(mta_system_known_custom_data(system, &names) == MTA_SUCCESS); + REQUIRE(names != nullptr); + { + auto names_str = std::string(mta_string_view(names)); + CHECK(names_str.find("test::my_data") != std::string::npos); + } + mta_string_free(names); + + const mts_tensormap_t* retrieved = nullptr; + CHECK(mta_system_get_custom_data(system, "test::my_data", &retrieved) == MTA_SUCCESS); + CHECK(retrieved != nullptr); +} + +/// `DataArrayBase` storing boolean data as `uint8_t` (since +/// `std::vector` has no `data()` method, `SimpleDataArray` can not +/// be used). This class reports its dtype as `kDLBool` so that the metatensor +/// serialization code correctly handles it. +class BoolDataArray: public metatensor::SimpleDataArray { +public: + using SimpleDataArray::SimpleDataArray; + + DLDataType dtype() const override { + DLDataType dtype; + dtype.code = DLDataTypeCode::kDLBool; + dtype.bits = 8; + dtype.lanes = 1; + return dtype; + } + + DLManagedTensorVersioned* as_dlpack( + DLDevice device, + const int64_t* stream, + DLPackVersion max_version + ) override { + auto* managed = SimpleDataArray::as_dlpack(device, stream, max_version); + managed->dl_tensor.dtype.code = DLDataTypeCode::kDLBool; + return managed; + } + + std::unique_ptr copy(DLDevice device) const override { + if (device.device_type != kDLCPU) { + throw metatensor::Error("BoolDataArray only supports copying to CPU"); + } + return std::unique_ptr(new BoolDataArray(*this)); + } + + std::unique_ptr create( + std::vector shape, + metatensor::MtsArray fill_value + ) const override { + DLDevice cpu_device = {kDLCPU, 0}; + DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; + auto fill_dlpack = fill_value.as_dlpack_array(cpu_device, nullptr, version); + + if (!fill_dlpack.shape().empty()) { + throw metatensor::Error("`fill_value` must be a single scalar"); + } + + auto scalar = fill_dlpack.data()[0]; + return std::unique_ptr(new BoolDataArray(std::move(shape), scalar)); + } +}; + +/// `mts_realloc_buffer_t` callback backed by a `std::vector`. +static uint8_t* vector_realloc(void* user_data, uint8_t* /*ptr*/, uintptr_t new_size) { + auto* buffer = static_cast*>(user_data); + buffer->resize(new_size, 0); + return buffer->data(); +} + +/// `mts_create_array_callback_t` that delegates to +/// `metatensor::details::default_create_array`, but handles `kDLBool` by +/// creating a `BoolDataArray` (`SimpleDataArray` does not compile since +/// `std::vector` has no `data()` method). Can be removed once +/// https://github.com/metatensor/metatensor/pull/1164 is released. +static mts_status_t create_array_with_bool( + const uintptr_t* shape_ptr, + uintptr_t shape_count, + DLDataType dtype, + mts_array_t* array +) { + if (dtype.code == kDLBool && dtype.bits == 8 && dtype.lanes == 1) { + auto shape = std::vector(); + for (uintptr_t i = 0; i < shape_count; i++) { + shape.push_back(shape_ptr[i]); + } + auto cxx_array = std::make_unique(shape); + *array = metatensor::DataArrayBase::to_mts_array(std::move(cxx_array)).release(); + return MTS_SUCCESS; + } + + return metatensor::details::default_create_array(shape_ptr, shape_count, dtype, array); +} + +TEST_CASE("system serialization") { + SECTION("save and load to a file") { + auto* system = full_test_system(); + + auto path = (std::filesystem::temp_directory_path() / "metatomic-test-system.mta").string(); + + CHECK(mta_save(path.c_str(), system) == MTA_SUCCESS); + + mta_system_t* loaded = nullptr; + auto status = mta_load( + path.c_str(), + create_array_with_bool, + &loaded + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(loaded != nullptr); + + check_full_system_data(loaded); + + CHECK(mta_system_free(loaded) == MTA_SUCCESS); + CHECK(mta_system_free(system) == MTA_SUCCESS); + std::remove(path.c_str()); + } + + SECTION("save and load to an in-memory buffer") { + auto* system = full_test_system(); + + std::vector buffer; + uint8_t* ptr = buffer.data(); + uintptr_t size = buffer.size(); + + auto status = mta_save_buffer( + &ptr, &size, &buffer, vector_realloc, system + ); + CHECK(status == MTA_SUCCESS); + buffer.resize(size); + + mta_system_t* loaded = nullptr; + status = mta_load_buffer( + buffer.data(), buffer.size(), + create_array_with_bool, + &loaded + ); + CHECK(status == MTA_SUCCESS); + REQUIRE(loaded != nullptr); + + check_full_system_data(loaded); + + CHECK(mta_system_free(loaded) == MTA_SUCCESS); + CHECK(mta_system_free(system) == MTA_SUCCESS); + } +} From 75ef3ca62bdecf88ac2d3a272f16e14d052d0110 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Fri, 10 Jul 2026 19:45:53 +0200 Subject: [PATCH 41/43] Refactory quantity name parsing, keep the separate components accessible --- metatomic-core/src/lib.rs | 4 +- metatomic-core/src/metadata.rs | 17 +- metatomic-core/src/model.rs | 8 +- metatomic-core/src/quantity/mod.rs | 2 + .../src/{ => quantity}/quantities.rs | 230 ++++++++++++------ metatomic-core/src/system.rs | 9 +- 6 files changed, 176 insertions(+), 94 deletions(-) create mode 100644 metatomic-core/src/quantity/mod.rs rename metatomic-core/src/{ => quantity}/quantities.rs (64%) diff --git a/metatomic-core/src/lib.rs b/metatomic-core/src/lib.rs index 954693f54..03da987ff 100644 --- a/metatomic-core/src/lib.rs +++ b/metatomic-core/src/lib.rs @@ -22,8 +22,8 @@ use crate::c_api::mta_status_t; pub use self::metadata::{Device, DType, ModelCapabilities, ModelMetadata, PairListOptions}; -mod quantities; -pub use self::quantities::{Quantity, SampleKind, Gradients}; +mod quantity; +pub use self::quantity::{QuantityName, Quantity, SampleKind, Gradients}; mod kernels; diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs index 0ccdba8a2..b173b8b8b 100644 --- a/metatomic-core/src/metadata.rs +++ b/metatomic-core/src/metadata.rs @@ -909,20 +909,21 @@ Please cite the following references when using this model: } mod model_capabilities { + use crate::QuantityName; use super::super::*; fn example() -> ModelCapabilities { ModelCapabilities { outputs: vec![ Quantity { - name: "energy".into(), + name: QuantityName::new("energy".into()).unwrap(), unit: "eV".into(), description: Some("total energy".into()), gradients: vec![crate::Gradients::Positions], sample_kind: crate::SampleKind::System, }, Quantity { - name: "charge".into(), + name: QuantityName::new("custom::charge/with_variant".into()).unwrap(), unit: "e".into(), description: None, gradients: vec![], @@ -944,7 +945,7 @@ Please cite the following references when using this model: assert_eq!(json["type"].as_str(), Some("metatomic_model_capabilities")); assert_eq!(json["outputs"][0]["name"].as_str(), Some("energy")); - assert_eq!(json["outputs"][1]["name"].as_str(), Some("charge")); + assert_eq!(json["outputs"][1]["name"].as_str(), Some("custom::charge/with_variant")); assert_eq!(json["atomic_types"][0].as_i64(), Some(1)); assert_eq!(json["atomic_types"][1].as_i64(), Some(6)); assert_eq!(json["atomic_types"][2].as_i64(), Some(8)); @@ -956,8 +957,14 @@ Please cite the following references when using this model: let parsed = ModelCapabilities::try_from(&json).unwrap(); assert_eq!(parsed.outputs.len(), 2); - assert_eq!(parsed.outputs[0].name, "energy"); - assert_eq!(parsed.outputs[1].name, "charge"); + assert_eq!(parsed.outputs[0].name.namespace(), None); + assert_eq!(parsed.outputs[0].name.base(), "energy"); + assert_eq!(parsed.outputs[0].name.variant(), None); + + assert_eq!(parsed.outputs[1].name.namespace(), Some("custom")); + assert_eq!(parsed.outputs[1].name.base(), "charge"); + assert_eq!(parsed.outputs[1].name.variant(), Some("with_variant")); + assert_eq!(parsed.atomic_types, vec![1, 6, 8]); assert_eq!(parsed.interaction_range.to_bits(), 5.0_f64.to_bits()); assert_eq!(parsed.length_unit, "Angstrom"); diff --git a/metatomic-core/src/model.rs b/metatomic-core/src/model.rs index cfe7d09d5..b413f3430 100644 --- a/metatomic-core/src/model.rs +++ b/metatomic-core/src/model.rs @@ -293,7 +293,7 @@ mod tests { fn capabilities() { let capabilities = test_model().capabilities().unwrap(); assert_eq!(capabilities.outputs.len(), 1); - assert_eq!(capabilities.outputs[0].name, "energy"); + assert_eq!(capabilities.outputs[0].name.full(), "energy"); assert_eq!(capabilities.atomic_types, vec![1, 6]); assert_eq!(capabilities.interaction_range.to_bits(), 5.0_f64.to_bits()); assert_eq!(capabilities.length_unit, "Angstrom"); @@ -312,7 +312,7 @@ mod tests { fn requested_inputs() { let inputs = test_model().requested_inputs().unwrap(); assert_eq!(inputs.len(), 1); - assert_eq!(inputs[0].name, "charge"); + assert_eq!(inputs[0].name.full(), "charge"); assert_eq!(inputs[0].unit, "e"); } @@ -320,10 +320,10 @@ mod tests { fn supported_outputs() { let outputs = test_model().supported_outputs().unwrap(); assert_eq!(outputs.len(), 2); - assert_eq!(outputs[0].name, "energy"); + assert_eq!(outputs[0].name.full(), "energy"); assert_eq!(outputs[0].unit, "eV"); - assert_eq!(outputs[1].name, "custom::output"); + assert_eq!(outputs[1].name.full(), "custom::output"); assert_eq!(outputs[1].unit, ""); } } diff --git a/metatomic-core/src/quantity/mod.rs b/metatomic-core/src/quantity/mod.rs new file mode 100644 index 000000000..ef6e319f3 --- /dev/null +++ b/metatomic-core/src/quantity/mod.rs @@ -0,0 +1,2 @@ +mod quantities; +pub use quantities::{QuantityName, Quantity, SampleKind, Gradients}; diff --git a/metatomic-core/src/quantities.rs b/metatomic-core/src/quantity/quantities.rs similarity index 64% rename from metatomic-core/src/quantities.rs rename to metatomic-core/src/quantity/quantities.rs index 85b726ccf..cc4ca70fb 100644 --- a/metatomic-core/src/quantities.rs +++ b/metatomic-core/src/quantity/quantities.rs @@ -29,68 +29,130 @@ fn is_valid_identifier(s: &str) -> bool { s.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') } -/// Validate a quantity name. +/// The name of a quantity, which can be either a standard name or a custom name +/// with an optional variant. /// -/// The name can be either a standard name or a custom name with the form -/// `::`, where the namespace can itself contain `::` to define -/// sub-namespaces. -/// -/// Both standard and custom names can also define a variant with the form -/// `/` or `::/`. -/// -/// All components (namespace, name, variant) must be non-empty if they are -/// present, and must be valid identifiers (alphanumeric + underscore, not -/// starting with a digit). -pub(crate) fn validate_quantity_name(name: &str) -> Result<(), Error> { - if STANDARD_QUANTITIES.contains(&name) { - return Ok(()); - } +/// This struct enforces that the name is either a known standard name or a +/// custom name. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct QuantityName { + /// The full name of the quantity, including namespace and variant if present + full: String, + /// Optional namespace for custom quantity names. Standard quantity names do + /// not have a namespace. + namespace: Option, + /// The base name of the quantity + base: String, + /// Optional variant of the quantity, i.e. `pbe0` in `energy/pbe0` + variant: Option, +} - let (main_part, variant) = if let Some(pos) = name.find('/') { - (&name[..pos], Some(&name[pos + 1..])) - } else { - (name, None) - }; +impl QuantityName { + /// Parse and validate a quantity name. + /// + /// The name can be either a standard name or a custom name with the form + /// `::`, where the namespace can itself contain `::` to + /// define sub-namespaces. + /// + /// Both standard and custom names can also define a variant with the form + /// `/` or `::/`. + /// + /// All components (namespace, name, variant) must be non-empty if they are + /// present, and must be valid identifiers (alphanumeric + underscore, not + /// starting with a digit). + pub fn new(name: String) -> Result { + let (main_part, variant) = if let Some(pos) = name.find('/') { + (&name[..pos], Some(name[pos + 1..].to_string())) + } else { + (&*name, None) + }; - if main_part.is_empty() { - return Err(Error::InvalidParameter(format!( - "quantity name cannot be empty in '{}'", name - ))); - } + let (namespace, base) = match main_part.rsplit_once("::") { + Some((ns, base)) => (Some(ns.to_string()), base.to_string()), + None => (None, main_part.to_string()), + }; - if let Some(variant) = variant && !is_valid_identifier(variant) { - return Err(Error::InvalidParameter(format!( - "invalid quantity variant '{}' in '{}': must be a valid identifier \ - (alphanumeric or underscore, not starting with a digit)", - variant, name - ))); - } + if let Some(ref ns) = namespace { + for component in ns.split("::") { + if !is_valid_identifier(component) { + return Err(Error::InvalidParameter(format!( + "invalid namespace '{}' in '{}': must be a valid \ + identifier (alphanumeric or underscore, not starting with a digit)", + ns, name + ))); + } + } + } - if STANDARD_QUANTITIES.contains(&main_part) { - return Ok(()); - } + if base.is_empty() { + return Err(Error::InvalidParameter(format!( + "quantity name cannot be empty in '{}'", name + ))); + } - let components: Vec<_> = main_part.split("::").collect(); - for component in &components { - if !is_valid_identifier(component) { + if !is_valid_identifier(&base) { return Err(Error::InvalidParameter(format!( - "invalid quantity name component '{}' in '{}': must be a valid \ - identifier (alphanumeric or underscore, not starting with a digit)", - component, name + "invalid quantity name '{}' in '{}': \ + must be a valid identifier (alphanumeric or underscore, not starting with a digit)", + base, name ))); } + + if let Some(ref variant) = variant && !is_valid_identifier(variant) { + return Err(Error::InvalidParameter(format!( + "invalid quantity variant '{}' in '{}': \ + must be a valid identifier (alphanumeric or underscore, not starting with a digit)", + variant, name + ))); + } + + if namespace.is_none() && !STANDARD_QUANTITIES.contains(&&*base) { + return Err(Error::InvalidParameter(format!( + "'{}' is not a standard quantity name; custom quantity names must use '::'", + name + ))); + } + + return Ok(QuantityName { + full: name, + namespace, + base, + variant, + }) + } + + /// Is this a custom quantity name? + pub fn is_custom(&self) -> bool { + self.namespace.is_some() + } + + /// Get the base name of this quantity + pub fn base(&self) -> &str { + &self.base } - if components.len() == 1 { - return Err(Error::InvalidParameter(format!( - "'{}' is not a standard quantity name; custom quantity names must use '::'", - name - ))); + /// Get the namespace of this quantity, if any + pub fn namespace(&self) -> Option<&str> { + self.namespace.as_deref() } - Ok(()) + /// Get the variant of this quantity, if any + pub fn variant(&self) -> Option<&str> { + self.variant.as_deref() + } + + /// Get the full name of this quantity, including namespace and variant if + /// present + pub fn full(&self) -> &str { + &self.full + } } +impl std::fmt::Display for QuantityName { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.full()) + } +} /// Different kind of samples a quantity can be associated with #[derive(Debug, Clone, PartialEq)] @@ -132,6 +194,16 @@ impl<'a> TryFrom<&'a JsonValue> for SampleKind { } } +impl std::fmt::Display for SampleKind { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + SampleKind::Atom => write!(f, "atom"), + SampleKind::AtomPair => write!(f, "atompair"), + SampleKind::System => write!(f, "system"), + } + } +} + /// Different gradients that a quantity can have #[derive(Debug, Clone, PartialEq)] pub enum Gradients { @@ -174,7 +246,7 @@ pub struct Quantity { /// Name of the quantity, this can be a standard name from /// , or /// a custom name of the form `::[/]` - pub name: String, + pub name: QuantityName, /// Unit of the quantity pub unit: String, /// Description of the quantity, used to provide more details about the @@ -193,7 +265,7 @@ impl From for JsonValue { fn from(value: Quantity) -> Self { let mut result = JsonValue::new_object(); result["type"] = "metatomic_quantity".into(); - result["name"] = value.name.into(); + result["name"] = value.name.full().into(); result["unit"] = value.unit.into(); if let Some(description) = value.description { result["description"] = description.into(); @@ -224,7 +296,7 @@ impl<'a> TryFrom<&'a JsonValue> for Quantity { let name = value["name"].as_str().ok_or_else(|| Error::Serialization( "'name' in JSON for Quantity must be a string".into() ))?; - validate_quantity_name(name)?; + let name = QuantityName::new(name.to_string())?; let unit = value["unit"].as_str().ok_or_else(|| Error::Serialization( "'unit' in JSON for Quantity must be a string".into() @@ -249,7 +321,7 @@ impl<'a> TryFrom<&'a JsonValue> for Quantity { let sample_kind = SampleKind::try_from(&value["sample_kind"])?; Ok(Quantity { - name: name.to_string(), + name: name, unit: unit.to_string(), description, gradients, @@ -265,7 +337,7 @@ mod tests { fn example() -> Quantity { Quantity { - name: "energy".into(), + name: QuantityName::new("energy".into()).unwrap(), unit: "eV".into(), description: Some("total energy of the system".into()), gradients: vec![Gradients::Positions], @@ -285,7 +357,7 @@ mod tests { assert_eq!(json["sample_kind"].as_str(), Some("atom")); let parsed = Quantity::try_from(&json).unwrap(); - assert_eq!(parsed.name, "energy"); + assert_eq!(parsed.name.base, "energy"); assert_eq!(parsed.unit, "eV"); assert_eq!(parsed.gradients, vec![Gradients::Positions]); assert!(matches!(parsed.sample_kind, SampleKind::Atom)); @@ -301,7 +373,7 @@ mod tests { vec![Gradients::Positions, Gradients::Strain], ] { let quantity = Quantity { - name: "test_ns::test".into(), + name: QuantityName::new("test_ns::test".into()).unwrap(), unit: "unit".into(), description: Some("Hello".to_string()), gradients: grads.clone(), @@ -372,7 +444,7 @@ mod tests { #[test] fn validate_names() { for name in STANDARD_QUANTITIES { - assert!(validate_quantity_name(name).is_ok(), "expected '{}' to be valid", name); + QuantityName::new(name.to_string()).unwrap(); } let custom = [ @@ -383,7 +455,7 @@ mod tests { "_ns::_name", ]; for name in custom { - assert!(validate_quantity_name(name).is_ok(), "expected '{}' to be valid", name); + QuantityName::new(name.to_string()).unwrap(); } let variants = [ @@ -392,49 +464,49 @@ mod tests { "ns1::ns2::energy/some_variant", ]; for name in variants { - assert!(validate_quantity_name(name).is_ok(), "expected '{}' to be valid", name); + QuantityName::new(name.to_string()).unwrap(); } - let error = validate_quantity_name("").expect_err("expected an error"); + let error = QuantityName::new(String::new()).expect_err("expected an error"); assert_eq!(error.to_string(), "invalid parameter: quantity name cannot be empty in ''"); - let error = validate_quantity_name("not_a_standard_name").expect_err("expected an error"); + let error = QuantityName::new("not_a_standard_name".into()).expect_err("expected an error"); assert_eq!(error.to_string(), "invalid parameter: 'not_a_standard_name' is not a standard quantity name; custom quantity names must use '::'"); - let error = validate_quantity_name("/variant").expect_err("expected an error"); + let error = QuantityName::new("/variant".into()).expect_err("expected an error"); assert_eq!(error.to_string(), "invalid parameter: quantity name cannot be empty in '/variant'"); - let error = validate_quantity_name("name/").expect_err("expected an error"); + let error = QuantityName::new("name/".into()).expect_err("expected an error"); assert_eq!(error.to_string(), "invalid parameter: invalid quantity variant '' in 'name/': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); - let error = validate_quantity_name("::energy").expect_err("expected an error"); - assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '' in '::energy': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + let error = QuantityName::new("::energy".into()).expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid namespace '' in '::energy': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); - let error = validate_quantity_name("ns::").expect_err("expected an error"); - assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '' in 'ns::': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + let error = QuantityName::new("ns::".into()).expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: quantity name cannot be empty in 'ns::'"); - let error = validate_quantity_name("ns::/variant").expect_err("expected an error"); - assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '' in 'ns::/variant': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + let error = QuantityName::new("ns::/variant".into()).expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: quantity name cannot be empty in 'ns::/variant'"); - let error = validate_quantity_name("::").expect_err("expected an error"); - assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '' in '::': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + let error = QuantityName::new("::".into()).expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid namespace '' in '::': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); - let error = validate_quantity_name("123name").expect_err("expected an error"); - assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '123name' in '123name': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + let error = QuantityName::new("123name".into()).expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name '123name' in '123name': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); - let error = validate_quantity_name("my_ns::123name").expect_err("expected an error"); - assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component '123name' in 'my_ns::123name': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + let error = QuantityName::new("my_ns::123name".into()).expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name '123name' in 'my_ns::123name': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); - let error = validate_quantity_name("my_ns::name/123variant").expect_err("expected an error"); + let error = QuantityName::new("my_ns::name/123variant".into()).expect_err("expected an error"); assert_eq!(error.to_string(), "invalid parameter: invalid quantity variant '123variant' in 'my_ns::name/123variant': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); - let error = validate_quantity_name("has spaces").expect_err("expected an error"); - assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component 'has spaces' in 'has spaces': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + let error = QuantityName::new("has spaces".into()).expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name 'has spaces' in 'has spaces': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); - let error = validate_quantity_name("my_ns::name/has spaces").expect_err("expected an error"); + let error = QuantityName::new("my_ns::name/has spaces".into()).expect_err("expected an error"); assert_eq!(error.to_string(), "invalid parameter: invalid quantity variant 'has spaces' in 'my_ns::name/has spaces': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); - let error = validate_quantity_name("has-dash").expect_err("expected an error"); - assert_eq!(error.to_string(), "invalid parameter: invalid quantity name component 'has-dash' in 'has-dash': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); + let error = QuantityName::new("has-dash".into()).expect_err("expected an error"); + assert_eq!(error.to_string(), "invalid parameter: invalid quantity name 'has-dash' in 'has-dash': must be a valid identifier (alphanumeric or underscore, not starting with a digit)"); } } diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs index 740cece00..a189e449d 100644 --- a/metatomic-core/src/system.rs +++ b/metatomic-core/src/system.rs @@ -5,8 +5,8 @@ use dlpk::sys::{DLDataType, DLDevice}; use dlpk::{DLPackTensor, DLPackTensorRef}; use metatensor::{TensorBlock, TensorMap}; -use crate::{Error, PairListOptions}; use crate::kernels::ReferenceValue; +use crate::{Error, PairListOptions, QuantityName}; /// Names that can never be used as custom data in a system static INVALID_DATA_NAMES: LazyLock> = LazyLock::new(|| { @@ -224,9 +224,10 @@ impl System { ))); } - crate::quantities::validate_quantity_name(&name)?; + // validate the quantity name + let name = QuantityName::new(name)?; - if !override_ && self.custom_data.contains_key(&name) { + if !override_ && self.custom_data.contains_key(name.full()) { return Err(Error::InvalidParameter(format!( "custom data '{}' is already present in this system", name @@ -259,7 +260,7 @@ impl System { ))); } - self.custom_data.insert(name, data); + self.custom_data.insert(name.full().to_string(), data); return Ok(()); } From f810c3dfe53f5c0ab669067874db60654f86b027 Mon Sep 17 00:00:00 2001 From: GardevoirX Date: Thu, 4 Jun 2026 14:26:04 +0200 Subject: [PATCH 42/43] Add function to check that Quantity match the expected layout Co-Authored-By: Guillaume Fraux --- docs/src/quantities/mass.rst | 2 +- docs/src/quantities/non_conservative.rst | 2 +- docs/src/quantities/velocity.rst | 2 +- metatomic-core/src/kernels/mod.rs | 8 + metatomic-core/src/metadata.rs | 11 + metatomic-core/src/quantity/charge.rs | 390 ++++++++++++ metatomic-core/src/quantity/checks.rs | 400 +++++++++++++ metatomic-core/src/quantity/energy.rs | 557 ++++++++++++++++++ metatomic-core/src/quantity/feature.rs | 305 ++++++++++ metatomic-core/src/quantity/heat_flux.rs | 361 ++++++++++++ metatomic-core/src/quantity/mass.rs | 316 ++++++++++ metatomic-core/src/quantity/mod.rs | 85 +++ metatomic-core/src/quantity/momentum.rs | 375 ++++++++++++ .../src/quantity/non_conservative_force.rs | 376 ++++++++++++ .../src/quantity/non_conservative_stress.rs | 369 ++++++++++++ metatomic-core/src/quantity/position.rs | 377 ++++++++++++ metatomic-core/src/quantity/quantities.rs | 18 +- .../src/quantity/spin_multiplicity.rs | 302 ++++++++++ metatomic-core/src/quantity/velocity.rs | 376 ++++++++++++ metatomic-core/src/system.rs | 257 ++++---- metatomic-core/tests/system.cpp | 2 +- 21 files changed, 4768 insertions(+), 123 deletions(-) create mode 100644 metatomic-core/src/quantity/charge.rs create mode 100644 metatomic-core/src/quantity/checks.rs create mode 100644 metatomic-core/src/quantity/energy.rs create mode 100644 metatomic-core/src/quantity/feature.rs create mode 100644 metatomic-core/src/quantity/heat_flux.rs create mode 100644 metatomic-core/src/quantity/mass.rs create mode 100644 metatomic-core/src/quantity/momentum.rs create mode 100644 metatomic-core/src/quantity/non_conservative_force.rs create mode 100644 metatomic-core/src/quantity/non_conservative_stress.rs create mode 100644 metatomic-core/src/quantity/position.rs create mode 100644 metatomic-core/src/quantity/spin_multiplicity.rs create mode 100644 metatomic-core/src/quantity/velocity.rs diff --git a/docs/src/quantities/mass.rst b/docs/src/quantities/mass.rst index b2a46b7ab..ed5ab88d8 100644 --- a/docs/src/quantities/mass.rst +++ b/docs/src/quantities/mass.rst @@ -37,7 +37,7 @@ following metadata: - the ``"mass"`` quantity must not have any components * - properties - - ``"mass`` + - ``"mass"`` - The ``"mass"`` quantity must have a single property dimension named ``"mass"``, with a single entry set to ``0``. diff --git a/docs/src/quantities/non_conservative.rst b/docs/src/quantities/non_conservative.rst index 29e5ee24c..9085553d2 100644 --- a/docs/src/quantities/non_conservative.rst +++ b/docs/src/quantities/non_conservative.rst @@ -133,7 +133,7 @@ and must have the following metadata: * - keys - ``"_"`` - the keys must have a single dimension named ``"_"``, with a single entry - set to ``0``. The ``"non_conservative_force"`` quantity is always + set to ``0``. The ``"non_conservative_stress"`` quantity is always represented as a :py:class:`metatensor.torch.TensorMap` with a single block. diff --git a/docs/src/quantities/velocity.rst b/docs/src/quantities/velocity.rst index 4ad546662..f868ce0dc 100644 --- a/docs/src/quantities/velocity.rst +++ b/docs/src/quantities/velocity.rst @@ -36,7 +36,7 @@ following metadata: * - components - ``"xyz"`` - The ``"velocity"`` quantity must have a single component dimension named - ``"xyz"``, with three entries set to ``0``, ``1``, and ``2``. The position + ``"xyz"``, with three entries set to ``0``, ``1``, and ``2``. The velocity is always a 3D vector, and the order of the components is ``x, y, z``. * - properties diff --git a/metatomic-core/src/kernels/mod.rs b/metatomic-core/src/kernels/mod.rs index 6c586d8c0..7fe7788d5 100644 --- a/metatomic-core/src/kernels/mod.rs +++ b/metatomic-core/src/kernels/mod.rs @@ -81,6 +81,14 @@ pub struct ReferenceValue { pub(crate) metal: OnceLock<(metal::MetalBuffer, StridedNDIndex)>, } +impl std::fmt::Debug for ReferenceValue { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ReferenceValue") + .field("cpu", &self.cpu) + .finish() + } +} + impl ReferenceValue { pub(crate) fn new(cpu: ArrayD) -> Self { Self { diff --git a/metatomic-core/src/metadata.rs b/metatomic-core/src/metadata.rs index b173b8b8b..34f9df3eb 100644 --- a/metatomic-core/src/metadata.rs +++ b/metatomic-core/src/metadata.rs @@ -3,6 +3,7 @@ use std::fmt::Write; use json::JsonValue; +use crate::metadata::DType::Float32; use crate::{Error, Quantity}; use crate::units::validate_unit; @@ -459,6 +460,16 @@ pub enum DType { Float64, } +impl std::fmt::Display for DType { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + if *self == Float32 { + write!(f, "float32") + } else { + write!(f, "float64") + } + } +} + impl From for JsonValue { fn from(value: DType) -> Self { match value { diff --git a/metatomic-core/src/quantity/charge.rs b/metatomic-core/src/quantity/charge.rs new file mode 100644 index 000000000..636b423dd --- /dev/null +++ b/metatomic-core/src/quantity/charge.rs @@ -0,0 +1,390 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, SINGLE_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; + + +/// Check the layout of the "charge" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "charge"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::System, SampleKind::Atom])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[])?; + + let expected_properties = ExpectedLabels { + names: &["charge"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]" + }; + checks::it_should_have_expected_labels(&context, "properties", &block.properties(), expected_properties)?; + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("charge".into()).unwrap(), + unit: "e".into(), + description: None, + gradients: vec![], + sample_kind: SampleKind::Atom, + } + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2]], + ); + let properties = Labels::new(["charge"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(); + TensorBlock::new(values, &samples, &[], &properties).unwrap() + } + + fn valid_charge() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_charge(), &[system(3)], None).unwrap(); + + // Also check SampleKind::System + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.5]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["charge"], [[0]]) + ).unwrap(); + + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &charge, &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + // Empty systems slice, per-atom output + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &charge, &[], None).unwrap(); + + // System with 0 atoms, per-atom output + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &charge, &[system(0)], None).unwrap(); + + // Empty systems slice, per-system output + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 1], vec![]).unwrap(), + &Labels::new( + ["system"], + Array2::::from_shape_vec((0, 1), vec![]).unwrap(), + ), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &charge, &[], None).unwrap(); + } + + #[test] + fn selected_atoms() { + // Per-atom output with selected_atoms across multiple systems + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &charge, &systems, Some(&selected_atoms)).unwrap(); + + // Per-system values with selected_atoms across multiple systems + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, 1], vec![5.0, 6.0]).unwrap(), + &Labels::new(["system"], [[0], [1]]), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &charge, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![4, 1], vec![1.0; 4]).unwrap(), + &Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2], [1, 0]], + ), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &charge, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::AtomPair; + let err = check(&request, &valid_charge(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'charge': expected one of [system, atom], got 'atom_pair'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let charge = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'charge': expected a single block, but found 0 blocks" + ); + + let charge = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'charge': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let charge = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'charge': expected a single block with key '_', but found key names [foo]" + ); + + let charge = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'charge': expected a single block with key value 0" + ); + } + + #[test] + fn wrong_property() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["wrong"], [[0]]), + ).unwrap(); + + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties for 'charge': expected names [charge], got [wrong]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["charge"], [[1]]), + ).unwrap(); + + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties values for 'charge': expected [[0]]" + ); + } + + #[test] + fn has_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: components for 'charge' should be empty" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["charge"], [[0]]) + ).unwrap(); + + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'charge': expected [system, atom], got [system]" + ); + + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0]]), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&request, &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'charge': expected [system], got [system, atom]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![0.1, 0.2, 0.3]).unwrap(), + &Labels::new(["sample", "system", "atom"], [[0, 0, 0]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let charge = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &charge, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'charge': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + // samples that are not in the selected_atoms + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &charge, &[system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'charge', they do not match the `systems` and `selected_atoms`" + ); + + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, 1], vec![3.0, 4.0]).unwrap(), + // systems that are not in the selected_atoms + &Labels::new(["system"], [[0], [1]]), + &[], + &Labels::new(["charge"], [[0]]), + ).unwrap(); + let charge = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&request, &charge, &[system(3), system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'charge', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/quantity/checks.rs b/metatomic-core/src/quantity/checks.rs new file mode 100644 index 000000000..8dbdc3b21 --- /dev/null +++ b/metatomic-core/src/quantity/checks.rs @@ -0,0 +1,400 @@ +use std::collections::BTreeSet; +use std::sync::LazyLock; + +use metatensor::{Labels, TensorBlockRef, TensorMap}; + +use super::Quantity; + +use crate::{Error, SampleKind, System}; +use crate::kernels::{is_equal_i32, ReferenceValue}; + + +pub(super) static XYZ_LABELS_REFERENCE: LazyLock> = LazyLock::new(|| { + ReferenceValue::new( + ndarray::ArrayD::from_shape_vec( + ndarray::IxDyn(&[3usize, 1]), + vec![0i32, 1, 2], + ).unwrap() + ) +}); + +pub(super) static SINGLE_LABELS_REFERENCE: LazyLock> = LazyLock::new(|| { + ReferenceValue::new( + ndarray::ArrayD::from_shape_vec( + ndarray::IxDyn(&[1usize, 1]), + vec![0i32], + ).unwrap() + ) +}); + +/// Check that the `sample_kind` is one of the valid kinds for the given quantity. +pub(super) fn it_should_have_valid_sample_kind( + context: &str, + sample_kind: SampleKind, + valid_kinds: &[SampleKind] +) -> Result<(), Error> { + if !valid_kinds.contains(&sample_kind) { + return Err(Error::InvalidParameter(format!( + "invalid sample_kind for {}: expected one of [{}], got '{}'", + context, + valid_kinds.iter().map(|k| k.to_string()).collect::>().join(", "), + sample_kind + ))); + } + + return Ok(()); +} + +/// Ensure the TensorMap has a single block with the expected key +pub(super) fn it_should_have_a_single_block(context: &str, value: &TensorMap) -> Result<(), Error> { + let keys = value.keys(); + if keys.count() != 1 { + return Err(Error::InvalidParameter(format!( + "invalid {}: expected a single block, but found {} blocks", + context, + keys.count() + ))); + } + + if keys.names() != ["_"] { + return Err(Error::InvalidParameter(format!( + "invalid {}: expected a single block with key '_', but found key names [{}]", + context, + keys.names().join(", ") + ))); + } + + let values = keys.values().as_dlpack(dlpk::DLDevice::cpu(), None, dlpk::DLPackVersion::current())?; + if !is_equal_i32(values.as_ref(), &SINGLE_LABELS_REFERENCE)? { + return Err(Error::InvalidParameter(format!( + "invalid {}: expected a single block with key value 0", + context, + ))); + } + + Ok(()) +} + +/// Validate the values for "system" samples +#[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)] +fn validate_system_samples( + context: &str, + samples: &Labels, + systems: &[System], + selected_atoms: Option<&Labels>, +) -> Result<(), Error> { + let values = if let Some(selected) = selected_atoms { + // only include the systems that are present in the selected_atoms + let mut values = BTreeSet::new(); + for [system_i, _] in selected.iter_fixed_size::<2>() { + values.insert(system_i.i32()); + } + ndarray::Array2::from_shape_vec( + (values.len(), 1), + values.into_iter().collect() + ).expect("created invalid array for system samples") + } else { + ndarray::Array2::from_shape_vec( + (systems.len(), 1), + (0..systems.len()).map(|s| s as i32).collect() + ).expect("created invalid array for system samples") + }; + + let expected = Labels::new_assume_unique(["system"], values); + + if expected.union(samples, None, None)?.count() != expected.count() { + return Err(Error::InvalidParameter(format!( + "invalid samples for {}, they do not match the \ + `systems` and `selected_atoms`", + context, + // TODO: add Labels::print to metatensor and use it here + ))); + } + + return Ok(()); +} + +/// Validate the values for "atom" samples +#[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)] +fn validate_atom_samples( + context: &str, + samples: &Labels, + systems: &[System], + selected_atoms: Option<&Labels>, +) -> Result<(), Error> { + let total_atoms: usize = systems.iter().map(|s| s.size()).sum(); + let mut values = ndarray::Array2::from_elem((total_atoms, 2), 0); + + let mut index = 0; + for (system_i, system) in systems.iter().enumerate() { + for atom_i in 0..system.size() { + values[[index, 0]] = system_i as i32; + values[[index, 1]] = atom_i as i32; + index += 1; + } + } + let mut expected = Labels::new_assume_unique(["system", "atom"], values); + if let Some(selected) = selected_atoms { + expected = expected.intersection(selected, None, None)?; + } + + if expected.union(samples, None, None)?.count() != expected.count() { + return Err(Error::InvalidParameter(format!( + "invalid samples for {}, they do not match the \ + `systems` and `selected_atoms`", + context, + // TODO: add Labels::print to metatensor and use it here + ))); + } + + return Ok(()); +} + +/// Validate the values for "atom_pair" samples +#[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap, clippy::cast_sign_loss)] +fn validate_atom_pair_samples( + context: &str, + samples: &Labels, + systems: &[System], + selected_atoms: Option<&Labels>, +) -> Result<(), Error> { + for [system, first_atom, second_atom, _, _, _] in samples.iter_fixed_size::<6>() { + let system = system.i32(); + let first_atom = first_atom.i32(); + let second_atom = second_atom.i32(); + + if system < 0 || system >= systems.len() as i32 { + return Err(Error::InvalidParameter(format!( + "invalid system index in samples for {}: {} is out of bounds", + context, + system + ))); + } + + let n_atoms = systems[system as usize].size() as i32; + if first_atom < 0 || first_atom >= n_atoms { + return Err(Error::InvalidParameter(format!( + "invalid first_atom index in samples for {}: {} is out of bounds for system {}", + context, + first_atom, + system + ))); + } + if second_atom < 0 || second_atom >= n_atoms { + return Err(Error::InvalidParameter(format!( + "invalid second_atom index in samples for {}: {} is out of bounds for system {}", + context, + second_atom, + system + ))); + } + } + + return Ok(()); +} + +/// Validates that the sample labels match the expected structure based on the +/// sample_kind and the systems/selected_atoms provided. +pub(super) fn it_should_have_valid_samples( + context: &str, + sample_kind: SampleKind, + block: TensorBlockRef<'_>, + systems: &[System], + selected_atoms: Option<&Labels>, +) -> Result<(), Error> { + let expected_samples_names: &[&str] = match sample_kind { + SampleKind::System => &["system"], + SampleKind::Atom => &["system", "atom"], + SampleKind::AtomPair => &[ + "system", + "first_atom", + "second_atom", + "cell_shift_a", + "cell_shift_b", + "cell_shift_c", + ], + }; + + let samples = block.samples(); + if samples.names() != expected_samples_names { + return Err(Error::InvalidParameter(format!( + "invalid sample names for {}: expected [{}], got [{}]", + context, + expected_samples_names.join(", "), + samples.names().join(", ") + ))); + } + + // Check if the samples entries match the systems and selected_atoms + match sample_kind { + SampleKind::System => validate_system_samples(context, &samples, systems, selected_atoms), + SampleKind::Atom => validate_atom_samples(context, &samples, systems, selected_atoms), + SampleKind::AtomPair => validate_atom_pair_samples(context, &samples, systems, selected_atoms), + } +} + +#[derive(Debug, Clone, Copy)] +pub(super) struct ExpectedLabels<'a> { + /// Expected names of the labels + pub names: &'a [&'a str], + /// Expected values of the labels + pub values: &'a ReferenceValue, + /// Message to display if the values do not match, showing the expected values + pub values_message: &'a str, +} + +pub(super) fn it_should_have_expected_labels( + context: &str, + labels_kind: &str, + labels: &Labels, + expected: ExpectedLabels<'_>, +) -> Result<(), Error> { + + if labels.names() != expected.names { + return Err(Error::InvalidParameter(format!( + "invalid {} for {}: expected names [{}], got [{}]", + labels_kind, + context, + expected.names.join(", "), + labels.names().join(", ") + ))); + } + + let values = labels.values().as_dlpack(dlpk::DLDevice::cpu(), None, dlpk::DLPackVersion::current())?; + if !is_equal_i32(values.as_ref(), expected.values)? { + return Err(Error::InvalidParameter(format!( + "invalid {} values for {}: expected {}", + labels_kind, + context, + expected.values_message + ))); + } + Ok(()) +} + +pub(super) fn it_should_have_expected_components( + context: &str, + block: TensorBlockRef<'_>, + expected: &[ExpectedLabels<'_>], +) -> Result<(), Error> { + let components = block.components(); + if components.len() != expected.len() { + if expected.is_empty() { + return Err(Error::InvalidParameter(format!( + "components for {} should be empty", + context + ))); + } else { + return Err(Error::InvalidParameter(format!( + "invalid components for {}: expected {} component(s), got {}", + context, + expected.len(), + components.len() + ))); + } + } + + for (component, &expected) in components.iter().zip(expected) { + it_should_have_expected_labels(context, "components", component, expected)?; + } + + return Ok(()); +} + +pub(super) fn it_should_have_expected_gradients( + context: &str, + request: &Quantity, + block: TensorBlockRef<'_>, + potential_gradients: &[&str], +) -> Result<(), Error> { + if potential_gradients.is_empty() && block.gradients().len() > 0 { + return Err(Error::InvalidParameter(format!( + "invalid gradients for {}: expected no gradients, but found \ + gradients with respect to [{}]", + context, + block.gradient_list().join(", ") + ))); + } + + for (parameter, gradient) in block.gradients() { + if !potential_gradients.contains(¶meter) { + return Err(Error::InvalidParameter(format!( + "invalid gradient '{}' for {}: expected one of [{}]", + parameter, + context, + potential_gradients.join(", ") + ))); + } + + match parameter { + "strain" => { + if !request.gradients.contains(&super::Gradients::Strain) { + return Err(Error::InvalidParameter(format!( + "invalid gradient 'strain' for {}: these gradients were not requested", + context + ))); + } + + let context = format!("strain gradient of {}", context); + if gradient.samples().names() != ["sample"] { + return Err(Error::InvalidParameter(format!( + "invalid samples for {}: expected samples names ['sample'], got [{}]", + context, + gradient.samples().names().join(", ") + ))); + } + + it_should_have_expected_components( + &context, + gradient, + &[ + ExpectedLabels { + names: &["xyz_1"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]", + }, + ExpectedLabels { + names: &["xyz_2"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]", + }, + ] + )?; + }, + "positions" => { + if !request.gradients.contains(&super::Gradients::Positions) { + return Err(Error::InvalidParameter(format!( + "invalid gradient 'positions' for {}: these gradients were not requested", + context + ))); + } + + let context = format!("positions gradient of {}", context); + if gradient.samples().names() != ["sample", "system", "atom"] { + return Err(Error::InvalidParameter(format!( + "invalid samples for {}: expected samples names ['sample', 'system', 'atom'], got [{}]", + context, + gradient.samples().names().join(", ") + ))); + } + + it_should_have_expected_components( + &context, + gradient, + &[ + ExpectedLabels { + names: &["xyz"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]", + }, + ] + )?; + }, + _ => { + unreachable!("got unknown gradient parameter for {}: {}", context, parameter); + } + } + } + + return Ok(()); +} diff --git a/metatomic-core/src/quantity/energy.rs b/metatomic-core/src/quantity/energy.rs new file mode 100644 index 000000000..9a4e2920b --- /dev/null +++ b/metatomic-core/src/quantity/energy.rs @@ -0,0 +1,557 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, SINGLE_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; +use crate::kernels::ReferenceValue; + + +/// Check the layout of one of the energy-related quantities ("energy", +/// "energy_ensemble", "energy_uncertainty"). +#[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)] +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + let name = &request.name; + assert!(!name.is_custom()); + assert!(name.base() == "energy" || name.base() == "energy_ensemble" || name.base() == "energy_uncertainty"); + + let context = format!("'{}'", name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::System, SampleKind::Atom])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[])?; + + if name.base() == "energy" || name.base() == "energy_uncertainty" { + checks::it_should_have_expected_labels( + &context, + "properties", + &block.properties(), + ExpectedLabels { + names: &["energy"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]", + } + )?; + } else { + let n_ensemble_members = *block.values().shape()?.last().expect("energy block has an empty shape"); + let reference = ReferenceValue::new(ndarray::ArrayD::from_shape_vec( + vec![n_ensemble_members, 1], (0..n_ensemble_members as i32).collect() + ).expect("created invalid array for energy_ensemble properties")); + checks::it_should_have_expected_labels( + &context, + "properties", + &block.properties(), + ExpectedLabels { + names: &["energy"], + values: &reference, + values_message: "[[0, ..., n]]", + } + )?; + } + + checks::it_should_have_expected_gradients(&context, request, block, &["strain", "positions"])?; + return Ok(()); +} + +#[cfg(test)] +mod tests { + // use a macro to generate the test code for all three energy-related quantities + macro_rules! energy_tests { + ($base_name: ident) => { + mod $base_name { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Gradients, Quantity, QuantityName, SampleKind, System}; + + use super::super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new(String::from(stringify!($base_name))).unwrap(), + unit: "eV".into(), + description: None, + gradients: vec![Gradients::Positions, Gradients::Strain], + sample_kind: SampleKind::Atom, + } + } + + fn n_properties() -> usize { + if stringify!($base_name) == "energy_ensemble" { 2 } else { 1 } + } + + fn property_labels() -> Labels { + if stringify!($base_name) == "energy_ensemble" { + Labels::new(["energy"], [[0], [1]]) + } else { + Labels::new(["energy"], [[0]]) + } + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2]], + ); + let n_props = n_properties(); + let values = ArrayD::::from_shape_vec( + vec![3, n_props], + vec![1.0; 3 * n_props], + ).unwrap(); + TensorBlock::new(values, &samples, &[], &property_labels()).unwrap() + } + + fn with_gradients(block: &mut TensorBlock) { + let n_props = n_properties(); + let props = property_labels(); + + let pos_gradient = TensorBlock::new( + ArrayD::::from_shape_vec( + vec![1, 3, n_props], + vec![0.1; 3 * n_props], + ).unwrap(), + &Labels::new(["sample", "system", "atom"], [[0, 0, 0]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &props, + ).unwrap(); + block.add_gradient("positions", pos_gradient).unwrap(); + + let strain_gradient = TensorBlock::new( + ArrayD::::from_shape_vec( + vec![1, 3, 3, n_props], + vec![0.1; 9 * n_props], + ).unwrap(), + &Labels::new(["sample"], [[0]]), + &[ + Labels::new(["xyz_1"], [[0], [1], [2]]), + Labels::new(["xyz_2"], [[0], [1], [2]]), + ], + &props, + ).unwrap(); + block.add_gradient("strain", strain_gradient).unwrap(); + } + + fn valid_energy() -> TensorMap { + let mut block = valid_block(); + with_gradients(&mut block); + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![block]).unwrap() + } + + fn system_energy() -> TensorMap { + let samples = Labels::new(["system"], [[0]]); + let n_props = n_properties(); + let values = ArrayD::::from_shape_vec( + vec![1, n_props], + vec![2.0; n_props], + ).unwrap(); + let mut block = TensorBlock::new(values, &samples, &[], &property_labels()).unwrap(); + with_gradients(&mut block); + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![block]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_energy(), &[system(3)], None).unwrap(); + + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + check(&request, &system_energy(), &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + // Empty systems slice, per-atom output + let n_props = n_properties(); + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, n_props], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[], + &property_labels(), + ).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &energy, &[], None).unwrap(); + + // System with 0 atoms, per-atom output + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, n_props], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[], + &property_labels(), + ).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &energy, &[system(0)], None).unwrap(); + + // Empty systems slice, per-system output + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, n_props], vec![]).unwrap(), + &Labels::new( + ["system"], + Array2::::from_shape_vec((0, 1), vec![]).unwrap(), + ), + &[], + &property_labels(), + ).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &energy, &[], None).unwrap(); + } + + #[test] + fn selected_atoms() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let n_props = n_properties(); + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, n_props], vec![1.0; 3 * n_props]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]), + &[], + &property_labels(), + ).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &energy, &systems, Some(&selected_atoms)).unwrap(); + + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, n_props], vec![2.0; 2 * n_props]).unwrap(), + &Labels::new(["system"], [[0], [1]]), + &[], + &property_labels(), + ).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &energy, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let samples = Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2], [1, 0]], + ); + let n_props = n_properties(); + let values = ArrayD::::from_shape_vec( + vec![4, n_props], + vec![1.0; 4 * n_props], + ).unwrap(); + let mut block = TensorBlock::new(values, &samples, &[], &property_labels()).unwrap(); + with_gradients(&mut block); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &energy, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::AtomPair; + let err = check(&request, &valid_energy(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid sample_kind for '{}': expected one of [system, atom], got 'atom_pair'", + stringify!($base_name) + ) + ); + } + + #[test] + fn wrong_number_of_blocks() { + let energy = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid '{}': expected a single block, but found 0 blocks", + stringify!($base_name) + ) + ); + + let energy = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid '{}': expected a single block, but found 2 blocks", + stringify!($base_name) + ) + ); + } + + #[test] + fn wrong_key() { + let energy = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid '{}': expected a single block with key '_', but found key names [foo]", + stringify!($base_name) + ) + ); + + let energy = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid '{}': expected a single block with key value 0", + stringify!($base_name) + ) + ); + } + + #[test] + fn wrong_property() { + let samples = Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]); + let n_props = n_properties(); + let values = ArrayD::::from_shape_vec( + vec![3, n_props], + vec![1.0; 3 * n_props], + ).unwrap(); + + let props_wrong = if stringify!($base_name) == "energy_ensemble" { + Labels::new(["wrong"], [[0], [1]]) + } else { + Labels::new(["wrong"], [[0]]) + }; + + let block = TensorBlock::new(values, &samples, &[], &props_wrong).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid properties for '{}': expected names [energy], got [wrong]", + stringify!($base_name) + ) + ); + + let props_wrong = if stringify!($base_name) == "energy_ensemble" { + Labels::new(["energy"], [[1], [0]]) + } else { + Labels::new(["energy"], [[1]]) + }; + let values = ArrayD::::from_shape_vec( + vec![3, n_props], + vec![1.0; 3 * n_props], + ).unwrap(); + + let block = TensorBlock::new(values, &samples, &[], &props_wrong).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], None).unwrap_err(); + + let expected_msg = if stringify!($base_name) == "energy_ensemble" { + format!("invalid parameter: invalid properties values for '{}': expected [[0, ..., n]]", stringify!($base_name)) + } else { + format!("invalid parameter: invalid properties values for '{}': expected [[0]]", stringify!($base_name)) + }; + assert_eq!(err.to_string(), expected_msg); + } + + #[test] + fn has_components() { + let n_props = n_properties(); + let props = property_labels(); + let values = ArrayD::::from_shape_vec( + vec![3, 3, n_props], + vec![1.0; 9 * n_props], + ).unwrap(); + let block = TensorBlock::new( + values, + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &props, + ).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: components for '{}' should be empty", + stringify!($base_name) + ) + ); + } + + #[test] + fn wrong_sample_names() { + let n_props = n_properties(); + let props = property_labels(); + let values = ArrayD::::from_shape_vec( + vec![1, n_props], + vec![1.0; n_props], + ).unwrap(); + let block = TensorBlock::new( + values, + &Labels::new(["system"], [[0]]), + &[], + &props, + ).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid sample names for '{}': expected [system, atom], got [system]", + stringify!($base_name) + ) + ); + } + + #[test] + fn gradients_dummy() { + let n_props = n_properties(); + let props = property_labels(); + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec( + vec![3, n_props], + vec![1.0; 3 * n_props], + ).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &props, + ).unwrap(); + + let dummy_gradient = TensorBlock::new( + ArrayD::::from_shape_vec( + vec![1, n_props], + vec![0.1; n_props], + ).unwrap(), + &Labels::new(["sample"], [[0]]), + &[], + &props, + ).unwrap(); + block.add_gradient("dummy", dummy_gradient).unwrap(); + + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid gradient 'dummy' for '{}': expected one of [strain, positions]", + stringify!($base_name) + ) + ); + } + + #[test] + fn gradients_position_not_requested() { + let mut request = valid_request(); + request.gradients = vec![]; + + let n_props = n_properties(); + let props = property_labels(); + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec( + vec![3, n_props], + vec![1.0; 3 * n_props], + ).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &props, + ).unwrap(); + + let pos_gradient = TensorBlock::new( + ArrayD::::from_shape_vec( + vec![1, 3, n_props], + vec![0.1; 3 * n_props], + ).unwrap(), + &Labels::new(["sample", "system", "atom"], [[0, 0, 0]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &props, + ).unwrap(); + block.add_gradient("positions", pos_gradient).unwrap(); + + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&request, &energy, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid gradient 'positions' for '{}': these gradients were not requested", + stringify!($base_name) + ) + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let n_props = n_properties(); + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, n_props], vec![1.0; 3 * n_props]).unwrap(), + // samples that are not in the selected_atoms + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &property_labels(), + ).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &energy, &[system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid samples for '{}', they do not match the `systems` and `selected_atoms`", + stringify!($base_name) + ) + ); + + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, n_props], vec![2.0; 2 * n_props]).unwrap(), + // systems that are not in the selected_atoms + &Labels::new(["system"], [[0], [1]]), + &[], + &property_labels(), + ).unwrap(); + let energy = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&request, &energy, &[system(3), system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + format!( + "invalid parameter: invalid samples for '{}', they do not match the `systems` and `selected_atoms`", + stringify!($base_name) + ) + ); + } + } + }; + } + + energy_tests!(energy); + energy_tests!(energy_ensemble); + energy_tests!(energy_uncertainty); +} diff --git a/metatomic-core/src/quantity/feature.rs b/metatomic-core/src/quantity/feature.rs new file mode 100644 index 000000000..c43f6051e --- /dev/null +++ b/metatomic-core/src/quantity/feature.rs @@ -0,0 +1,305 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks; + +use crate::{Error, SampleKind, System}; + + +/// Check the layout of the "feature" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "feature"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::System, SampleKind::Atom])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[])?; + // no check on properties, as they can be anything for "feature" + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + Ok(()) +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("feature".into()).unwrap(), + unit: String::new(), + description: None, + gradients: vec![], + sample_kind: SampleKind::Atom, + } + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2]], + ); + let properties = Labels::new(["anything_goes_here"], [[-42], [5]]); + let values = ArrayD::::from_shape_vec(vec![3, 2], vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap(); + TensorBlock::new(values, &samples, &[], &properties).unwrap() + } + + fn valid_feature() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_feature(), &[system(3)], None).unwrap(); + + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.5]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["something_else"], [[0]]) + ).unwrap(); + let feature = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &feature, &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + // Empty systems slice, per-atom output + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[], + &Labels::new(["feature"], [[0]]), + ).unwrap(); + let feature = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &feature, &[], None).unwrap(); + + // System with 0 atoms, per-atom output + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[], + &Labels::new(["feature"], [[0]]), + ).unwrap(); + let feature = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &feature, &[system(0)], None).unwrap(); + + // Empty systems slice, per-system output + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 1], vec![]).unwrap(), + &Labels::new( + ["system"], + Array2::::from_shape_vec((0, 1), vec![]).unwrap(), + ), + &[], + &Labels::new(["feature"], [[0]]), + ).unwrap(); + let feature = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &feature, &[], None).unwrap(); + } + + #[test] + fn selected_atoms() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]), + &[], + &Labels::new(["feature"], [[0]]), + ).unwrap(); + let feature = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &feature, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![4, 1], vec![1.0; 4]).unwrap(), + &Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2], [1, 0]], + ), + &[], + &Labels::new(["feature"], [[0]]), + ).unwrap(); + + let feature = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &feature, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::AtomPair; + let err = check(&request, &valid_feature(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'feature': expected one of [system, atom], got 'atom_pair'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let feature = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &feature, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'feature': expected a single block, but found 0 blocks" + ); + + let feature = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &feature, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'feature': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let feature = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &feature, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'feature': expected a single block with key '_', but found key names [foo]" + ); + + let feature = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &feature, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'feature': expected a single block with key value 0" + ); + } + + #[test] + fn has_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &Labels::new(["feature"], [[0]]), + ).unwrap(); + + let feature = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &feature, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: components for 'feature' should be empty" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["feature"], [[0]]) + ).unwrap(); + + let feature = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &feature, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'feature': expected [system, atom], got [system]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["feature"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![0.1, 0.2, 0.3]).unwrap(), + &Labels::new(["sample", "system", "atom"], [[0, 0, 0]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &Labels::new(["feature"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let feature = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &feature, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'feature': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + // samples that are not in the selected_atoms + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["feature"], [[0]]), + ).unwrap(); + let feature = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &feature, &[system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'feature', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/quantity/heat_flux.rs b/metatomic-core/src/quantity/heat_flux.rs new file mode 100644 index 000000000..a6d609195 --- /dev/null +++ b/metatomic-core/src/quantity/heat_flux.rs @@ -0,0 +1,361 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, SINGLE_LABELS_REFERENCE, XYZ_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; + + +/// Check the layout of the "heat_flux" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "heat_flux"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::System])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[ + ExpectedLabels { + names: &["xyz"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]" + } + ])?; + + let expected_properties = ExpectedLabels { + names: &["heat_flux"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]" + }; + checks::it_should_have_expected_labels(&context, "properties", &block.properties(), expected_properties)?; + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("heat_flux".into()).unwrap(), + unit: "eV/ps".into(), + description: None, + gradients: vec![], + sample_kind: SampleKind::System, + } + } + + fn valid_xyz_component() -> Labels { + Labels::new(["xyz"], [[0], [1], [2]]) + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new(["system"], [[0]]); + let properties = Labels::new(["heat_flux"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(); + TensorBlock::new(values, &samples, &[valid_xyz_component()], &properties).unwrap() + } + + fn valid_heat_flux() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_heat_flux(), &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + // Empty systems slice, per-system output + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system"], + Array2::::from_shape_vec((0, 1), vec![]).unwrap(), + ), + &[valid_xyz_component()], + &Labels::new(["heat_flux"], [[0]]), + ).unwrap(); + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &heat_flux, &[], None).unwrap(); + } + + #[test] + fn selected_atoms() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, 3, 1], vec![1.0; 6]).unwrap(), + &Labels::new(["system"], [[0], [1]]), + &[valid_xyz_component()], + &Labels::new(["heat_flux"], [[0]]), + ).unwrap(); + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &heat_flux, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let samples = Labels::new( + ["system"], + [[0], [1]], + ); + let properties = Labels::new(["heat_flux"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![2, 3, 1], vec![1.0; 6]).unwrap(); + let block = TensorBlock::new(values, &samples, &[valid_xyz_component()], &properties).unwrap(); + + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &heat_flux, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::Atom; + let err = check(&request, &valid_heat_flux(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'heat_flux': expected one of [system], got 'atom'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let heat_flux = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'heat_flux': expected a single block, but found 0 blocks" + ); + + let heat_flux = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'heat_flux': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let heat_flux = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'heat_flux': expected a single block with key '_', but found key names [foo]" + ); + + let heat_flux = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'heat_flux': expected a single block with key value 0" + ); + } + + #[test] + fn wrong_property() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[valid_xyz_component()], + &Labels::new(["wrong"], [[0]]), + ).unwrap(); + + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties for 'heat_flux': expected names [heat_flux], got [wrong]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[valid_xyz_component()], + &Labels::new(["heat_flux"], [[1]]), + ).unwrap(); + + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties values for 'heat_flux': expected [[0]]" + ); + } + + #[test] + fn missing_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["heat_flux"], [[0]]), + ).unwrap(); + + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'heat_flux': expected 1 component(s), got 0" + ); + } + + #[test] + fn wrong_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[Labels::new(["abc"], [[0], [1], [2]])], + &Labels::new(["heat_flux"], [[0]]), + ).unwrap(); + + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'heat_flux': expected names [xyz], got [abc]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[Labels::new(["xyz"], [[1], [2], [3]])], + &Labels::new(["heat_flux"], [[0]]), + ).unwrap(); + + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components values for 'heat_flux': expected [[0], [1], [2]]" + ); + } + + #[test] + fn extra_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system"], [[0]]), + &[ + valid_xyz_component(), + Labels::new(["abc"], [[0], [1], [2]]), + ], + &Labels::new(["heat_flux"], [[0]]), + ).unwrap(); + + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'heat_flux': expected 1 component(s), got 2" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0]]), + &[valid_xyz_component()], + &Labels::new(["heat_flux"], [[0]]) + ).unwrap(); + + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'heat_flux': expected [system], got [system, atom]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[valid_xyz_component()], + &Labels::new(["heat_flux"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![0.1, 0.2, 0.3]).unwrap(), + &Labels::new(["sample"], [[0]]), + &[valid_xyz_component()], + &Labels::new(["heat_flux"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let heat_flux = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &heat_flux, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'heat_flux': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, 3, 1], vec![1.0; 6]).unwrap(), + // systems that are not in the selected_atoms + &Labels::new(["system"], [[0], [1]]), + &[valid_xyz_component()], + &Labels::new(["heat_flux"], [[0]]), + ).unwrap(); + let heat_flux = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &heat_flux, &[system(3), system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'heat_flux', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/quantity/mass.rs b/metatomic-core/src/quantity/mass.rs new file mode 100644 index 000000000..63a0a34bc --- /dev/null +++ b/metatomic-core/src/quantity/mass.rs @@ -0,0 +1,316 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, SINGLE_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; + + +/// Check the layout of the "mass" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "mass"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::Atom])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[])?; + + let expected_properties = ExpectedLabels { + names: &["mass"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]" + }; + checks::it_should_have_expected_labels(&context, "properties", &block.properties(), expected_properties)?; + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("mass".into()).unwrap(), + unit: "dalton".into(), + description: None, + gradients: vec![], + sample_kind: SampleKind::Atom, + } + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2]], + ); + let properties = Labels::new(["mass"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(); + TensorBlock::new(values, &samples, &[], &properties).unwrap() + } + + fn valid_mass() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_mass(), &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + // Empty systems slice, per-atom output + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[], + &Labels::new(["mass"], [[0]]), + ).unwrap(); + let mass = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &mass, &[], None).unwrap(); + + // System with 0 atoms, per-atom output + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[], + &Labels::new(["mass"], [[0]]), + ).unwrap(); + let mass = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &mass, &[system(0)], None).unwrap(); + } + + #[test] + fn selected_atoms() { + // Per-atom output with selected_atoms across multiple systems + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]), + &[], + &Labels::new(["mass"], [[0]]), + ).unwrap(); + let mass = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &mass, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![4, 1], vec![1.0; 4]).unwrap(), + &Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2], [1, 0]], + ), + &[], + &Labels::new(["mass"], [[0]]), + ).unwrap(); + + let mass = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &mass, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let err = check(&request, &valid_mass(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'mass': expected one of [atom], got 'system'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let mass = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &mass, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'mass': expected a single block, but found 0 blocks" + ); + + let mass = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &mass, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'mass': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let mass = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &mass, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'mass': expected a single block with key '_', but found key names [foo]" + ); + + let mass = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &mass, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'mass': expected a single block with key value 0" + ); + } + + #[test] + fn wrong_property() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["wrong"], [[0]]), + ).unwrap(); + + let mass = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &mass, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties for 'mass': expected names [mass], got [wrong]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["mass"], [[1]]), + ).unwrap(); + + let mass = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &mass, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties values for 'mass': expected [[0]]" + ); + } + + #[test] + fn has_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &Labels::new(["mass"], [[0]]), + ).unwrap(); + + let mass = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &mass, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: components for 'mass' should be empty" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["mass"], [[0]]) + ).unwrap(); + + let mass = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &mass, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'mass': expected [system, atom], got [system]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["mass"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![0.1, 0.2, 0.3]).unwrap(), + &Labels::new(["sample", "system", "atom"], [[0, 0, 0]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &Labels::new(["mass"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let mass = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &mass, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'mass': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + // samples that are not in the selected_atoms + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["mass"], [[0]]), + ).unwrap(); + let mass = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &mass, &[system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'mass', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/quantity/mod.rs b/metatomic-core/src/quantity/mod.rs index ef6e319f3..6f686255f 100644 --- a/metatomic-core/src/quantity/mod.rs +++ b/metatomic-core/src/quantity/mod.rs @@ -1,2 +1,87 @@ +use metatensor::{Labels, TensorMap}; + +use crate::{Error, System}; + + mod quantities; pub use quantities::{QuantityName, Quantity, SampleKind, Gradients}; + +mod checks; + +mod energy; +mod feature; +mod non_conservative_force; +mod non_conservative_stress; +mod position; +mod momentum; +mod velocity; +mod mass; +mod charge; +mod heat_flux; +mod spin_multiplicity; + + +/// Check that the provided `TensorMap` matches the expected layout for the +/// given `Quantity`. +/// +/// Only standard quantities are checked, custom quantities are only validated +/// for device/dtype compatibility. +/// +/// `selected_atoms` can change the expected samples, and should be provided if +/// the `TensorMap` was computed for a subset of atoms. +pub fn check_quantity( + quantity: &Quantity, + values: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels>, +) -> Result<(), Error> { + assert!(!systems.is_empty(), "systems must contain at least one system"); + debug_assert!(systems.iter().all(|s| s.dtype() == systems[0].dtype()), "all systems must have the same dtype"); + debug_assert!(systems.iter().all(|s| s.device() == systems[0].device()), "all systems must have the same device"); + + if !values.keys().is_empty() { + if values.device()? != systems[0].device() { + return Err(Error::InvalidParameter(format!( + "invalid device for quantity '{}': expected {}, got {}", + quantity.name, + systems[0].device(), + values.device()? + ))); + } + if values.dtype()? != systems[0].dtype() { + return Err(Error::InvalidParameter(format!( + "invalid dtype for quantity '{}': expected {}, got {}", + quantity.name, + systems[0].dtype(), + values.dtype()? + ))); + } + } + + if quantity.name.is_custom() { + // nothing to check + return Ok(()); + } + + match quantity.name.base() { + "energy" | "energy_ensemble" | "energy_uncertainty" => energy::check(quantity, values, systems, selected_atoms)?, + "feature" => feature::check(quantity, values, systems, selected_atoms)?, + "non_conservative_force" => non_conservative_force::check(quantity, values, systems, selected_atoms)?, + "non_conservative_stress" => non_conservative_stress::check(quantity, values, systems, selected_atoms)?, + "position" => position::check(quantity, values, systems, selected_atoms)?, + "momentum" => momentum::check(quantity, values, systems, selected_atoms)?, + "mass" => mass::check(quantity, values, systems, selected_atoms)?, + "velocity" => velocity::check(quantity, values, systems, selected_atoms)?, + "charge" => charge::check(quantity, values, systems, selected_atoms)?, + "heat_flux" => heat_flux::check(quantity, values, systems, selected_atoms)?, + "spin_multiplicity" => spin_multiplicity::check(quantity, values, systems, selected_atoms)?, + _ => { + return Err(Error::Internal(format!( + "invalid quantity name '{}': unknown standard quantity", + quantity.name + ))); + } + } + + Ok(()) +} diff --git a/metatomic-core/src/quantity/momentum.rs b/metatomic-core/src/quantity/momentum.rs new file mode 100644 index 000000000..0904f7ff9 --- /dev/null +++ b/metatomic-core/src/quantity/momentum.rs @@ -0,0 +1,375 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, SINGLE_LABELS_REFERENCE, XYZ_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; + + +/// Check the layout of the "momentum" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "momentum"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::Atom])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[ + ExpectedLabels { + names: &["xyz"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]", + }, + ])?; + + let expected_properties = ExpectedLabels { + names: &["momentum"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]" + }; + checks::it_should_have_expected_labels(&context, "properties", &block.properties(), expected_properties)?; + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("momentum".into()).unwrap(), + unit: "Angstrom*amu/ps".into(), + description: None, + gradients: vec![], + sample_kind: SampleKind::Atom, + } + } + + fn valid_xyz_component() -> Labels { + Labels::new(["xyz"], [[0], [1], [2]]) + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2]], + ); + let properties = Labels::new(["momentum"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(); + TensorBlock::new(values, &samples, &[valid_xyz_component()], &properties).unwrap() + } + + fn valid_momentum() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_momentum(), &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[valid_xyz_component()], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &momentum, &[], None).unwrap(); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[valid_xyz_component()], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &momentum, &[system(0)], None).unwrap(); + } + + #[test] + fn selected_atoms() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]), + &[valid_xyz_component()], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &momentum, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![4, 3, 1], vec![1.0; 12]).unwrap(), + &Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2], [1, 0]], + ), + &[valid_xyz_component()], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &momentum, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let err = check(&request, &valid_momentum(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'momentum': expected one of [atom], got 'system'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let momentum = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'momentum': expected a single block, but found 0 blocks" + ); + + let momentum = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'momentum': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let momentum = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'momentum': expected a single block with key '_', but found key names [foo]" + ); + + let momentum = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'momentum': expected a single block with key value 0" + ); + } + + #[test] + fn wrong_property() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["wrong"], [[0]]), + ).unwrap(); + + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties for 'momentum': expected names [momentum], got [wrong]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["momentum"], [[1]]), + ).unwrap(); + + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties values for 'momentum': expected [[0]]" + ); + } + + #[test] + fn missing_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0; 3]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'momentum': expected 1 component(s), got 0" + ); + } + + #[test] + fn wrong_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["abc"], [[0], [1], [2]])], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'momentum': expected names [xyz], got [abc]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["xyz"], [[1], [2], [3]])], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components values for 'momentum': expected [[0], [1], [2]]" + ); + } + + #[test] + fn extra_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 3, 1], vec![1.0; 27]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[ + valid_xyz_component(), + Labels::new(["abc"], [[0], [1], [2]]), + ], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'momentum': expected 1 component(s), got 2" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[valid_xyz_component()], + &Labels::new(["momentum"], [[0]]) + ).unwrap(); + + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'momentum': expected [system, atom], got [system]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![0.1, 0.2, 0.3]).unwrap(), + &Labels::new(["sample", "system", "atom"], [[0, 0, 0]]), + &[valid_xyz_component()], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let momentum = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &momentum, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'momentum': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + // samples that are not in the selected_atoms + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["momentum"], [[0]]), + ).unwrap(); + let momentum = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &momentum, &[system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'momentum', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/quantity/non_conservative_force.rs b/metatomic-core/src/quantity/non_conservative_force.rs new file mode 100644 index 000000000..dd99a8106 --- /dev/null +++ b/metatomic-core/src/quantity/non_conservative_force.rs @@ -0,0 +1,376 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, XYZ_LABELS_REFERENCE, SINGLE_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; + + + +/// Check the layout of the "non_conservative_force" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "non_conservative_force"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::Atom])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[ + ExpectedLabels { + names: &["xyz"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]", + }, + ])?; + + let expected_properties = ExpectedLabels { + names: &["non_conservative_force"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]" + }; + checks::it_should_have_expected_labels(&context, "properties", &block.properties(), expected_properties)?; + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("non_conservative_force".into()).unwrap(), + unit: "eV/Angstrom".into(), + description: None, + gradients: vec![], + sample_kind: SampleKind::Atom, + } + } + + fn valid_xyz_component() -> Labels { + Labels::new(["xyz"], [[0], [1], [2]]) + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2]], + ); + let properties = Labels::new(["non_conservative_force"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(); + TensorBlock::new(values, &samples, &[valid_xyz_component()], &properties).unwrap() + } + + fn valid_non_conservative_force() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_non_conservative_force(), &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[valid_xyz_component()], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &non_conservative_force, &[], None).unwrap(); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[valid_xyz_component()], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap(); + } + + #[test] + fn selected_atoms() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]), + &[valid_xyz_component()], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &non_conservative_force, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![4, 3, 1], vec![1.0; 12]).unwrap(), + &Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2], [1, 0]], + ), + &[valid_xyz_component()], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &non_conservative_force, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let err = check(&request, &valid_non_conservative_force(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'non_conservative_force': expected one of [atom], got 'system'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let non_conservative_force = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'non_conservative_force': expected a single block, but found 0 blocks" + ); + + let non_conservative_force = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'non_conservative_force': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let non_conservative_force = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'non_conservative_force': expected a single block with key '_', but found key names [foo]" + ); + + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'non_conservative_force': expected a single block with key value 0" + ); + } + + #[test] + fn wrong_property() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["wrong"], [[0]]), + ).unwrap(); + + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties for 'non_conservative_force': expected names [non_conservative_force], got [wrong]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["non_conservative_force"], [[1]]), + ).unwrap(); + + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties values for 'non_conservative_force': expected [[0]]" + ); + } + + #[test] + fn missing_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0; 3]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'non_conservative_force': expected 1 component(s), got 0" + ); + } + + #[test] + fn wrong_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["abc"], [[0], [1], [2]])], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'non_conservative_force': expected names [xyz], got [abc]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["xyz"], [[1], [2], [3]])], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components values for 'non_conservative_force': expected [[0], [1], [2]]" + ); + } + + #[test] + fn extra_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 3, 1], vec![1.0; 27]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[ + valid_xyz_component(), + Labels::new(["abc"], [[0], [1], [2]]), + ], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'non_conservative_force': expected 1 component(s), got 2" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[valid_xyz_component()], + &Labels::new(["non_conservative_force"], [[0]]) + ).unwrap(); + + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'non_conservative_force': expected [system, atom], got [system]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![0.1, 0.2, 0.3]).unwrap(), + &Labels::new(["sample", "system", "atom"], [[0, 0, 0]]), + &[valid_xyz_component()], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let non_conservative_force = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &non_conservative_force, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'non_conservative_force': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + // samples that are not in the selected_atoms + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["non_conservative_force"], [[0]]), + ).unwrap(); + let non_conservative_force = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_force, &[system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'non_conservative_force', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/quantity/non_conservative_stress.rs b/metatomic-core/src/quantity/non_conservative_stress.rs new file mode 100644 index 000000000..4386ff700 --- /dev/null +++ b/metatomic-core/src/quantity/non_conservative_stress.rs @@ -0,0 +1,369 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, XYZ_LABELS_REFERENCE, SINGLE_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; + +/// Check the layout of the "non_conservative_stress" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "non_conservative_stress"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::System])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[ + ExpectedLabels { + names: &["xyz_1"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]", + }, + ExpectedLabels { + names: &["xyz_2"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]", + }, + ])?; + + let expected_properties = ExpectedLabels { + names: &["non_conservative_stress"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]" + }; + checks::it_should_have_expected_labels(&context, "properties", &block.properties(), expected_properties)?; + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("non_conservative_stress".into()).unwrap(), + unit: "eV/Angstrom^3".into(), + description: None, + gradients: vec![], + sample_kind: SampleKind::System, + } + } + + fn valid_xyz_components() -> Vec { + vec![ + Labels::new(["xyz_1"], [[0], [1], [2]]), + Labels::new(["xyz_2"], [[0], [1], [2]]), + ] + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new(["system"], [[0]]); + let properties = Labels::new(["non_conservative_stress"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![1, 3, 3, 1], vec![1.0; 9]).unwrap(); + TensorBlock::new(values, &samples, &valid_xyz_components(), &properties).unwrap() + } + + fn valid_non_conservative_stress() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_non_conservative_stress(), &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + // Empty systems slice, per-system output + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system"], + Array2::::from_shape_vec((0, 1), vec![]).unwrap(), + ), + &valid_xyz_components(), + &Labels::new(["non_conservative_stress"], [[0]]), + ).unwrap(); + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &non_conservative_stress, &[], None).unwrap(); + } + + #[test] + fn selected_atoms() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, 3, 3, 1], vec![1.0; 18]).unwrap(), + &Labels::new(["system"], [[0], [1]]), + &valid_xyz_components(), + &Labels::new(["non_conservative_stress"], [[0]]), + ).unwrap(); + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &non_conservative_stress, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let samples = Labels::new( + ["system"], + [[0], [1]], + ); + let properties = Labels::new(["non_conservative_stress"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![2, 3, 3, 1], vec![1.0; 18]).unwrap(); + let block = TensorBlock::new(values, &samples, &valid_xyz_components(), &properties).unwrap(); + + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &non_conservative_stress, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::Atom; + let err = check(&request, &valid_non_conservative_stress(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'non_conservative_stress': expected one of [system], got 'atom'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let non_conservative_stress = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'non_conservative_stress': expected a single block, but found 0 blocks" + ); + + let non_conservative_stress = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'non_conservative_stress': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let non_conservative_stress = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'non_conservative_stress': expected a single block with key '_', but found key names [foo]" + ); + + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'non_conservative_stress': expected a single block with key value 0" + ); + } + + #[test] + fn wrong_property() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system"], [[0]]), + &valid_xyz_components(), + &Labels::new(["wrong"], [[0]]), + ).unwrap(); + + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties for 'non_conservative_stress': expected names [non_conservative_stress], got [wrong]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system"], [[0]]), + &valid_xyz_components(), + &Labels::new(["non_conservative_stress"], [[1]]), + ).unwrap(); + + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties values for 'non_conservative_stress': expected [[0]]" + ); + } + + #[test] + fn missing_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["non_conservative_stress"], [[0]]), + ).unwrap(); + + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'non_conservative_stress': expected 2 component(s), got 0" + ); + } + + #[test] + fn wrong_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system"], [[0]]), + &[Labels::new(["abc"], [[0], [1], [2]]), Labels::new(["xyz_2"], [[0], [1], [2]])], + &Labels::new(["non_conservative_stress"], [[0]]), + ).unwrap(); + + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'non_conservative_stress': expected names [xyz_1], got [abc]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system"], [[0]]), + &[Labels::new(["xyz_1"], [[1], [2], [3]]), Labels::new(["xyz_2"], [[0], [1], [2]])], + &Labels::new(["non_conservative_stress"], [[0]]), + ).unwrap(); + + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components values for 'non_conservative_stress': expected [[0], [1], [2]]" + ); + } + + #[test] + fn extra_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 3, 3, 1], vec![1.0; 27]).unwrap(), + &Labels::new(["system"], [[0]]), + &[ + Labels::new(["xyz_1"], [[0], [1], [2]]), + Labels::new(["xyz_2"], [[0], [1], [2]]), + Labels::new(["abc"], [[0], [1], [2]]), + ], + &Labels::new(["non_conservative_stress"], [[0]]), + ).unwrap(); + + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'non_conservative_stress': expected 2 component(s), got 3" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0]]), + &valid_xyz_components(), + &Labels::new(["non_conservative_stress"], [[0]]) + ).unwrap(); + + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'non_conservative_stress': expected [system], got [system, atom]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system"], [[0]]), + &valid_xyz_components(), + &Labels::new(["non_conservative_stress"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["sample"], [[0]]), + &valid_xyz_components(), + &Labels::new(["non_conservative_stress"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let non_conservative_stress = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &non_conservative_stress, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'non_conservative_stress': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, 3, 3, 1], vec![1.0; 18]).unwrap(), + // systems that are not in the selected_atoms + &Labels::new(["system"], [[0], [1]]), + &valid_xyz_components(), + &Labels::new(["non_conservative_stress"], [[0]]), + ).unwrap(); + let non_conservative_stress = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &non_conservative_stress, &[system(3), system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'non_conservative_stress', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/quantity/position.rs b/metatomic-core/src/quantity/position.rs new file mode 100644 index 000000000..8d691326f --- /dev/null +++ b/metatomic-core/src/quantity/position.rs @@ -0,0 +1,377 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, SINGLE_LABELS_REFERENCE, XYZ_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; + +/// Check the layout of the "position" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "position"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::Atom])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[ + ExpectedLabels { + names: &["xyz"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]", + }, + ])?; + + let expected_properties = ExpectedLabels { + names: &["position"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]" + }; + checks::it_should_have_expected_labels(&context, "properties", &block.properties(), expected_properties)?; + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("position".into()).unwrap(), + unit: "Angstrom".into(), + description: None, + gradients: vec![], + sample_kind: SampleKind::Atom, + } + } + + fn valid_xyz_component() -> Labels { + Labels::new(["xyz"], [[0], [1], [2]]) + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2]], + ); + let properties = Labels::new(["position"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(); + TensorBlock::new(values, &samples, &[valid_xyz_component()], &properties).unwrap() + } + + fn valid_position() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_position(), &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + // Empty systems slice, per-atom output + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[valid_xyz_component()], + &Labels::new(["position"], [[0]]), + ).unwrap(); + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &position, &[], None).unwrap(); + + // System with 0 atoms, per-atom output + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[valid_xyz_component()], + &Labels::new(["position"], [[0]]), + ).unwrap(); + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &position, &[system(0)], None).unwrap(); + } + + #[test] + fn selected_atoms() { + // Per-atom output with selected_atoms across multiple systems + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]), + &[valid_xyz_component()], + &Labels::new(["position"], [[0]]), + ).unwrap(); + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &position, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![4, 3, 1], vec![1.0; 12]).unwrap(), + &Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2], [1, 0]], + ), + &[valid_xyz_component()], + &Labels::new(["position"], [[0]]), + ).unwrap(); + + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &position, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let err = check(&request, &valid_position(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'position': expected one of [atom], got 'system'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let position = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'position': expected a single block, but found 0 blocks" + ); + + let position = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'position': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let position = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'position': expected a single block with key '_', but found key names [foo]" + ); + + let position = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'position': expected a single block with key value 0" + ); + } + + #[test] + fn wrong_property() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["wrong"], [[0]]), + ).unwrap(); + + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties for 'position': expected names [position], got [wrong]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["position"], [[1]]), + ).unwrap(); + + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties values for 'position': expected [[0]]" + ); + } + + #[test] + fn missing_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0; 3]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["position"], [[0]]), + ).unwrap(); + + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'position': expected 1 component(s), got 0" + ); + } + + #[test] + fn wrong_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["abc"], [[0], [1], [2]])], + &Labels::new(["position"], [[0]]), + ).unwrap(); + + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'position': expected names [xyz], got [abc]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["xyz"], [[1], [2], [3]])], + &Labels::new(["position"], [[0]]), + ).unwrap(); + + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components values for 'position': expected [[0], [1], [2]]" + ); + } + + #[test] + fn extra_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 3, 1], vec![1.0; 27]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[ + valid_xyz_component(), + Labels::new(["abc"], [[0], [1], [2]]), + ], + &Labels::new(["position"], [[0]]), + ).unwrap(); + + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'position': expected 1 component(s), got 2" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[valid_xyz_component()], + &Labels::new(["position"], [[0]]) + ).unwrap(); + + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'position': expected [system, atom], got [system]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["position"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![0.1, 0.2, 0.3]).unwrap(), + &Labels::new(["sample", "system", "atom"], [[0, 0, 0]]), + &[valid_xyz_component()], + &Labels::new(["position"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let position = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &position, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'position': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + // samples that are not in the selected_atoms + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["position"], [[0]]), + ).unwrap(); + let position = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &position, &[system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'position', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/quantity/quantities.rs b/metatomic-core/src/quantity/quantities.rs index cc4ca70fb..7c1ae1ac3 100644 --- a/metatomic-core/src/quantity/quantities.rs +++ b/metatomic-core/src/quantity/quantities.rs @@ -155,7 +155,7 @@ impl std::fmt::Display for QuantityName { } /// Different kind of samples a quantity can be associated with -#[derive(Debug, Clone, PartialEq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub enum SampleKind { /// The quantity is defined for each atom (e.g. atomic energy, charge, ...) Atom, @@ -198,14 +198,14 @@ impl std::fmt::Display for SampleKind { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { SampleKind::Atom => write!(f, "atom"), - SampleKind::AtomPair => write!(f, "atompair"), + SampleKind::AtomPair => write!(f, "atom_pair"), SampleKind::System => write!(f, "system"), } } } /// Different gradients that a quantity can have -#[derive(Debug, Clone, PartialEq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub enum Gradients { /// Gradients with respect to atomic positions Positions, @@ -365,8 +365,8 @@ mod tests { #[test] fn roundtrip_all_variants() { - for sample in [SampleKind::Atom, SampleKind::System, SampleKind::AtomPair] { - for grads in [ + for sample_kind in [SampleKind::Atom, SampleKind::System, SampleKind::AtomPair] { + for gradients in [ vec![], vec![Gradients::Positions], vec![Gradients::Strain], @@ -376,14 +376,14 @@ mod tests { name: QuantityName::new("test_ns::test".into()).unwrap(), unit: "unit".into(), description: Some("Hello".to_string()), - gradients: grads.clone(), - sample_kind: sample.clone(), + gradients: gradients.clone(), + sample_kind: sample_kind, }; let parsed = Quantity::try_from(&JsonValue::from(quantity.clone())).unwrap(); assert_eq!(parsed.name, quantity.name); assert_eq!(parsed.unit, quantity.unit); - assert_eq!(parsed.gradients, grads); - assert_eq!(parsed.sample_kind, sample); + assert_eq!(parsed.gradients, gradients); + assert_eq!(parsed.sample_kind, sample_kind); } } } diff --git a/metatomic-core/src/quantity/spin_multiplicity.rs b/metatomic-core/src/quantity/spin_multiplicity.rs new file mode 100644 index 000000000..c17125550 --- /dev/null +++ b/metatomic-core/src/quantity/spin_multiplicity.rs @@ -0,0 +1,302 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, SINGLE_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; + + +/// Check the layout of the "spin_multiplicity" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "spin_multiplicity"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::System])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[])?; + + let expected_properties = ExpectedLabels { + names: &["spin_multiplicity"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]" + }; + checks::it_should_have_expected_labels(&context, "properties", &block.properties(), expected_properties)?; + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("spin_multiplicity".into()).unwrap(), + unit: "dimensionless".into(), + description: None, + gradients: vec![], + sample_kind: SampleKind::System, + } + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new(["system"], [[0]]); + let properties = Labels::new(["spin_multiplicity"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(); + TensorBlock::new(values, &samples, &[], &properties).unwrap() + } + + fn valid_spin_multiplicity() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_spin_multiplicity(), &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + // Empty systems slice, per-system output + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 1], vec![]).unwrap(), + &Labels::new( + ["system"], + Array2::::from_shape_vec((0, 1), vec![]).unwrap(), + ), + &[], + &Labels::new(["spin_multiplicity"], [[0]]), + ).unwrap(); + let spin_multiplicity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&request, &spin_multiplicity, &[], None).unwrap(); + } + + #[test] + fn selected_atoms() { + // Per-system values with selected_atoms across multiple systems + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, 1], vec![5.0, 6.0]).unwrap(), + &Labels::new(["system"], [[0], [1]]), + &[], + &Labels::new(["spin_multiplicity"], [[0]]), + ).unwrap(); + let spin_multiplicity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &spin_multiplicity, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, 1], vec![1.0; 2]).unwrap(), + &Labels::new( + ["system"], + [[0], [1]], + ), + &[], + &Labels::new(["spin_multiplicity"], [[0]]), + ).unwrap(); + + let spin_multiplicity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &spin_multiplicity, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::Atom; + let err = check(&request, &valid_spin_multiplicity(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'spin_multiplicity': expected one of [system], got 'atom'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let spin_multiplicity = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &spin_multiplicity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'spin_multiplicity': expected a single block, but found 0 blocks" + ); + + let spin_multiplicity = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &spin_multiplicity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'spin_multiplicity': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let spin_multiplicity = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &spin_multiplicity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'spin_multiplicity': expected a single block with key '_', but found key names [foo]" + ); + + let spin_multiplicity = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &spin_multiplicity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'spin_multiplicity': expected a single block with key value 0" + ); + } + + #[test] + fn wrong_property() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["wrong"], [[0]]), + ).unwrap(); + + let spin_multiplicity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &spin_multiplicity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties for 'spin_multiplicity': expected names [spin_multiplicity], got [wrong]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["spin_multiplicity"], [[1]]), + ).unwrap(); + + let spin_multiplicity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &spin_multiplicity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties values for 'spin_multiplicity': expected [[0]]" + ); + } + + #[test] + fn has_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0; 3]).unwrap(), + &Labels::new(["system"], [[0]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &Labels::new(["spin_multiplicity"], [[0]]), + ).unwrap(); + + let spin_multiplicity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &spin_multiplicity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: components for 'spin_multiplicity' should be empty" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0]]), + &[], + &Labels::new(["spin_multiplicity"], [[0]]) + ).unwrap(); + + let spin_multiplicity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &spin_multiplicity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'spin_multiplicity': expected [system], got [system, atom]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 1], vec![1.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[], + &Labels::new(["spin_multiplicity"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![0.1, 0.2, 0.3]).unwrap(), + &Labels::new(["sample"], [[0]]), + &[Labels::new(["xyz"], [[0], [1], [2]])], + &Labels::new(["spin_multiplicity"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let spin_multiplicity = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &spin_multiplicity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'spin_multiplicity': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![2, 1], vec![3.0, 4.0]).unwrap(), + // systems that are not in the selected_atoms + &Labels::new(["system"], [[0], [1]]), + &[], + &Labels::new(["spin_multiplicity"], [[0]]), + ).unwrap(); + let spin_multiplicity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &spin_multiplicity, &[system(3), system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'spin_multiplicity', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/quantity/velocity.rs b/metatomic-core/src/quantity/velocity.rs new file mode 100644 index 000000000..c1a8936ee --- /dev/null +++ b/metatomic-core/src/quantity/velocity.rs @@ -0,0 +1,376 @@ +use metatensor::{Labels, TensorMap}; + +use super::Quantity; +use super::checks::{self, ExpectedLabels, SINGLE_LABELS_REFERENCE, XYZ_LABELS_REFERENCE}; + +use crate::{Error, SampleKind, System}; + + + +/// Check the layout of the "velocity" quantity. +pub(super) fn check( + request: &Quantity, + value: &TensorMap, + systems: &[System], + selected_atoms: Option<&Labels> +) -> Result<(), Error> { + assert!(!request.name.is_custom() && request.name.base() == "velocity"); + + let context = format!("'{}'", request.name.full()); + checks::it_should_have_valid_sample_kind(&context, request.sample_kind, &[SampleKind::Atom])?; + + checks::it_should_have_a_single_block(&context, value)?; + let block = value.block_by_id(0); + + checks::it_should_have_valid_samples(&context, request.sample_kind, block, systems, selected_atoms)?; + checks::it_should_have_expected_components(&context, block, &[ + ExpectedLabels { + names: &["xyz"], + values: &XYZ_LABELS_REFERENCE, + values_message: "[[0], [1], [2]]", + }, + ])?; + + let expected_properties = ExpectedLabels { + names: &["velocity"], + values: &SINGLE_LABELS_REFERENCE, + values_message: "[[0]]" + }; + checks::it_should_have_expected_labels(&context, "properties", &block.properties(), expected_properties)?; + checks::it_should_have_expected_gradients(&context, request, block, &[])?; + + return Ok(()); +} + +#[cfg(test)] +mod tests { + use metatensor::{Labels, TensorBlock, TensorMap}; + use ndarray::{Array1, Array2, ArrayD}; + use dlpk::DLPackTensor; + + use crate::{Quantity, QuantityName, SampleKind, System}; + + use super::check; + + fn system(n_atoms: usize) -> System { + let types: DLPackTensor = Array1::::from_vec(vec![1; n_atoms]).try_into().unwrap(); + let positions: DLPackTensor = Array2::::from_shape_vec((n_atoms, 3), vec![0.0; n_atoms * 3]).unwrap().try_into().unwrap(); + let cell: DLPackTensor = Array2::::from_shape_vec( + (3, 3), + vec![10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0], + ).unwrap().try_into().unwrap(); + + let pbc: DLPackTensor = Array1::::from_vec(vec![true, true, true]).try_into().unwrap(); + System::new("Angstrom".into(), types, positions, cell, pbc).unwrap() + } + + fn valid_request() -> Quantity { + Quantity { + name: QuantityName::new("velocity".into()).unwrap(), + unit: "Angstrom/ps".into(), + description: None, + gradients: vec![], + sample_kind: SampleKind::Atom, + } + } + + fn valid_xyz_component() -> Labels { + Labels::new(["xyz"], [[0], [1], [2]]) + } + + fn valid_block() -> TensorBlock { + let samples = Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2]], + ); + let properties = Labels::new(["velocity"], [[0]]); + let values = ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(); + TensorBlock::new(values, &samples, &[valid_xyz_component()], &properties).unwrap() + } + + fn valid_velocity() -> TensorMap { + let keys = Labels::new(["_"], [[0]]); + TensorMap::new(keys, vec![valid_block()]).unwrap() + } + + #[test] + fn ok() { + check(&valid_request(), &valid_velocity(), &[system(3)], None).unwrap(); + } + + #[test] + fn empty_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[valid_xyz_component()], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &velocity, &[], None).unwrap(); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![0, 3, 1], vec![]).unwrap(), + &Labels::new( + ["system", "atom"], + Array2::::from_shape_vec((0, 2), vec![]).unwrap(), + ), + &[valid_xyz_component()], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &velocity, &[system(0)], None).unwrap(); + } + + #[test] + fn selected_atoms() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]); + let systems = [system(3), system(1)]; + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [1, 0]]), + &[valid_xyz_component()], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &velocity, &systems, Some(&selected_atoms)).unwrap(); + } + + #[test] + fn multiple_systems() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![4, 3, 1], vec![1.0; 12]).unwrap(), + &Labels::new( + ["system", "atom"], + [[0, 0], [0, 1], [0, 2], [1, 0]], + ), + &[valid_xyz_component()], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + check(&valid_request(), &velocity, &[system(3), system(1)], None).unwrap(); + } + + #[test] + fn invalid_sample_kind() { + let mut request = valid_request(); + request.sample_kind = SampleKind::System; + let err = check(&request, &valid_velocity(), &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample_kind for 'velocity': expected one of [atom], got 'system'" + ); + } + + #[test] + fn wrong_number_of_blocks() { + let velocity = TensorMap::new(Labels::empty(vec!["_"]), vec![]).unwrap(); + + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'velocity': expected a single block, but found 0 blocks" + ); + + let velocity = TensorMap::new( + Labels::new(["_"], [[0], [1]]), + vec![valid_block(), valid_block()] + ).unwrap(); + + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'velocity': expected a single block, but found 2 blocks" + ); + } + + #[test] + fn wrong_key() { + let velocity = TensorMap::new(Labels::new(["foo"], [[0]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'velocity': expected a single block with key '_', but found key names [foo]" + ); + + let velocity = TensorMap::new(Labels::new(["_"], [[1]]), vec![valid_block()]).unwrap(); + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid 'velocity': expected a single block with key value 0" + ); + } + + #[test] + fn wrong_property() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["wrong"], [[0]]), + ).unwrap(); + + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties for 'velocity': expected names [velocity], got [wrong]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["velocity"], [[1]]), + ).unwrap(); + + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid properties values for 'velocity': expected [[0]]" + ); + } + + #[test] + fn missing_components() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 1], vec![1.0; 3]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'velocity': expected 1 component(s), got 0" + ); + } + + #[test] + fn wrong_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["abc"], [[0], [1], [2]])], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'velocity': expected names [xyz], got [abc]" + ); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[Labels::new(["xyz"], [[1], [2], [3]])], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components values for 'velocity': expected [[0], [1], [2]]" + ); + } + + #[test] + fn extra_component() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 3, 1], vec![1.0; 27]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[ + valid_xyz_component(), + Labels::new(["abc"], [[0], [1], [2]]), + ], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid components for 'velocity': expected 1 component(s), got 2" + ); + } + + #[test] + fn wrong_sample_names() { + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![1.0, 2.0, 3.0]).unwrap(), + &Labels::new(["system"], [[0]]), + &[valid_xyz_component()], + &Labels::new(["velocity"], [[0]]) + ).unwrap(); + + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid sample names for 'velocity': expected [system, atom], got [system]" + ); + } + + #[test] + fn gradients() { + let mut block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + + let gradient = TensorBlock::new( + ArrayD::::from_shape_vec(vec![1, 3, 1], vec![0.1, 0.2, 0.3]).unwrap(), + &Labels::new(["sample", "system", "atom"], [[0, 0, 0]]), + &[valid_xyz_component()], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + + block.add_gradient("positions", gradient).unwrap(); + + let velocity = TensorMap::new( + Labels::new(["_"], [[0]]), + vec![block] + ).unwrap(); + + let err = check(&valid_request(), &velocity, &[system(3)], None).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid gradients for 'velocity': expected no gradients, but found gradients with respect to [positions]" + ); + } + + #[test] + fn selected_atoms_error() { + let selected_atoms = Labels::new(["system", "atom"], [[0, 0], [0, 1]]); + + let block = TensorBlock::new( + ArrayD::::from_shape_vec(vec![3, 3, 1], vec![1.0; 9]).unwrap(), + // samples that are not in the selected_atoms + &Labels::new(["system", "atom"], [[0, 0], [0, 1], [0, 2]]), + &[valid_xyz_component()], + &Labels::new(["velocity"], [[0]]), + ).unwrap(); + let velocity = TensorMap::new(Labels::new(["_"], [[0]]), vec![block]).unwrap(); + let err = check(&valid_request(), &velocity, &[system(3)], Some(&selected_atoms)).unwrap_err(); + assert_eq!( + err.to_string(), + "invalid parameter: invalid samples for 'velocity', they do not match the `systems` and `selected_atoms`" + ); + } +} diff --git a/metatomic-core/src/system.rs b/metatomic-core/src/system.rs index a189e449d..493592436 100644 --- a/metatomic-core/src/system.rs +++ b/metatomic-core/src/system.rs @@ -6,7 +6,8 @@ use dlpk::{DLPackTensor, DLPackTensorRef}; use metatensor::{TensorBlock, TensorMap}; use crate::kernels::ReferenceValue; -use crate::{Error, PairListOptions, QuantityName}; +use crate::quantity::check_quantity; +use crate::{Error, Gradients, PairListOptions, Quantity, QuantityName, SampleKind}; /// Names that can never be used as custom data in a system static INVALID_DATA_NAMES: LazyLock> = LazyLock::new(|| { @@ -49,6 +50,20 @@ pub struct System { unsafe impl Send for System {} unsafe impl Sync for System {} +impl std::fmt::Debug for System { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("System") + .field("length_unit", &self.length_unit) + .field("types", &self.types) + .field("positions", &self.positions) + .field("cell", &self.cell) + .field("pbc", &self.pbc) + .field("pairs", &self.pairs.keys().collect::>()) + .field("custom_data", &self.custom_data.keys().collect::>()) + .finish() + } +} + impl System { /// Create a `System` from raw DLPack tensors pub fn new( @@ -224,43 +239,25 @@ impl System { ))); } - // validate the quantity name - let name = QuantityName::new(name)?; - - if !override_ && self.custom_data.contains_key(name.full()) { - return Err(Error::InvalidParameter(format!( - "custom data '{}' is already present in this system", - name - ))); - } - - if data.keys().count() == 0 { + if data.keys().is_empty() { return Err(Error::InvalidParameter(format!( "custom data '{}' has no blocks", name ))); } - // TODO: add TensorMap::device/dtype and use them here - let block = data.block_by_id(0); - let values = block.values(); - let data_device = values.device()?; - if data_device != self.device() { - return Err(Error::InvalidParameter(format!( - "device ({}:{}) of the custom data '{}' does not match this system device ({}:{})", - data_device.device_type, data_device.device_id, name, - self.device().device_type, self.device().device_id, - ))); - } + // validate the quantity + let name = QuantityName::new(name)?; + let quantity = quantity_for_data(name, &data)?; + check_quantity(&quantity, &data, std::slice::from_ref(self), None)?; - let values_dtype = values.dtype()?; - if values_dtype != self.dtype() { + if !override_ && self.custom_data.contains_key(quantity.name.full()) { return Err(Error::InvalidParameter(format!( - "dtype of custom data '{}' does not match this system dtype", - name, + "custom data '{}' is already present in this system", + quantity.name ))); } - self.custom_data.insert(name.full().to_string(), data); + self.custom_data.insert(quantity.name.full().to_string(), data); return Ok(()); } @@ -274,7 +271,7 @@ impl System { } return self.custom_data.get(name).ok_or_else(|| Error::InvalidParameter(format!( - "no data for '{}' found in this system", name + "no custom data for '{}' found in this system", name ))); } @@ -284,17 +281,83 @@ impl System { } /// The device used for all tensors in this system - fn device(&self) -> DLDevice { + pub fn device(&self) -> DLDevice { self.types.device() } /// The data type used for the `positions` and `cell` tensors in this /// system, as well as any pair lists and custom data added to this system. - fn dtype(&self) -> DLDataType { + pub fn dtype(&self) -> DLDataType { self.positions.dtype() } } +/// Guess the `SampleKind` corresponding to the provided `TensorMap`. +/// +/// If `allow_unknown` is `true`, this will return `SampleKind::System` when +/// unable to determine the sample kind. Otherwise, it will return an error. +fn sample_kind_from_sample_names(data: &TensorMap, allow_unknown: bool) -> Result { + assert!(!data.keys().is_empty()); + + let first_block = data.block_by_id(0); + let samples = first_block.samples(); + let sample_names = samples.names(); + + if sample_names == ["system"] { + Ok(SampleKind::System) + } else if sample_names == ["system", "atom"] { + Ok(SampleKind::Atom) + } else if sample_names == ["system", "first_atom", "second_atom", "cell_shift_a", "cell_shift_b", "cell_shift_c"] { + Ok(SampleKind::AtomPair) + } else if allow_unknown { + Ok(SampleKind::System) + } else { + Err(Error::InvalidParameter(format!( + "data has unknown sample names: [{}]", + sample_names.join(", ") + ))) + } +} + +/// Guess the `Quantity` corresponding to the provided custom data name and +/// `TensorMap`. +fn quantity_for_data(name: QuantityName, data: &TensorMap) -> Result { + assert!(!data.keys().is_empty()); + + if name.is_custom() { + return Ok(Quantity { + name: name, + unit: String::new(), + description: None, + gradients: vec![], + sample_kind: sample_kind_from_sample_names(data, true)?, + }); + } + + let mut gradients = Vec::new(); + let first_block = data.block_by_id(0); + for parameter in first_block.gradient_list() { + if parameter == "positions" { + gradients.push(Gradients::Positions); + } else if parameter == "cell" { + gradients.push(Gradients::Strain); + } else { + return Err(Error::InvalidParameter(format!( + "data '{}' has an unknown gradient '{}'", + name, parameter + ))); + } + } + + return Ok(Quantity { + name: name, + unit: data.get_info("unit").unwrap_or("").into(), + description: None, + gradients: gradients, + sample_kind: sample_kind_from_sample_names(data, false)?, + }); +} + fn validate_system_tensors( types: &DLPackTensor, positions: &DLPackTensor, @@ -352,9 +415,11 @@ fn validate_system_tensors( } if cell.dtype() != positions.dtype() { - return Err(Error::InvalidParameter( - "`cell` must have the same dtype as `positions`".into() - )); + return Err(Error::InvalidParameter(format!( + "`cell` must have the same dtype as `positions`, got {} and {}", + cell.dtype(), + positions.dtype() + ))); } let pbc_shape = pbc.shape(); @@ -481,14 +546,6 @@ mod tests { TensorMap::new(keys, vec![block]).unwrap() } - fn assert_error(result: Result, expected: &str) { - let error = match result { - Ok(_) => panic!("expected error"), - Err(error) => error, - }; - assert_eq!(error.to_string(), expected); - } - pub(crate) fn test_system() -> System { let mut system = System::new( "Angstrom".into(), @@ -549,73 +606,57 @@ mod tests { let cell = cell_tensor(0.0, "f32"); let pbc = pbc_tensor(&[true, true, true]); - assert_error( - System::new(length_unit.clone(), bad_types, positions, cell, pbc), - "invalid parameter: `types` must be a tensor of 32-bit integers", - ); + let err = System::new(length_unit.clone(), bad_types, positions, cell, pbc).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: `types` must be a tensor of 32-bit integers"); let bad_types: DLPackTensor = Array2::::from_shape_vec((2, 2), vec![1, 2, 3, 4]).unwrap().try_into().unwrap(); let positions = positions_tensor(2, "f32"); let cell = cell_tensor(0.0, "f32"); let pbc = pbc_tensor(&[true, true, true]); - assert_error( - System::new(length_unit.clone(), bad_types, positions, cell, pbc), - "invalid parameter: `types` must be a (n_atoms,) tensor, got a tensor with shape [2, 2]", - ); + let err = System::new(length_unit.clone(), bad_types, positions, cell, pbc).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: `types` must be a (n_atoms,) tensor, got a tensor with shape [2, 2]"); let types = type_tensor(&[1]); let bad_positions: DLPackTensor = Array2::::from_shape_vec((1, 3), vec![1, 2, 3]).unwrap().try_into().unwrap(); let cell = cell_tensor(0.0, "f32"); let pbc = pbc_tensor(&[true, true, true]); - assert_error( - System::new(length_unit.clone(), types, bad_positions, cell, pbc), - "invalid parameter: `positions` must be a tensor of 32 or 64-bit floating point data", - ); + let err = System::new(length_unit.clone(), types, bad_positions, cell, pbc).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: `positions` must be a tensor of 32 or 64-bit floating point data"); let types = type_tensor(&[1, 6]); let bad_positions = Array2::::from_shape_vec((2, 2), vec![0.0; 4]).unwrap().try_into().unwrap(); let cell = cell_tensor(0.0, "f32"); let pbc = pbc_tensor(&[true, true, true]); - assert_error( - System::new("Angstrom".into(), types, bad_positions, cell, pbc), - "invalid parameter: `positions` must be a (n_atoms x 3) tensor, got a tensor with shape [2, 2]", - ); + let err = System::new("Angstrom".into(), types, bad_positions, cell, pbc).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: `positions` must be a (n_atoms x 3) tensor, got a tensor with shape [2, 2]"); let types = type_tensor(&[1, 6]); let positions = positions_tensor(2, "f32"); let bad_cell = Array2::::from_shape_vec((2, 3), vec![0.0; 6]).unwrap().try_into().unwrap(); let pbc = pbc_tensor(&[true, true, true]); - assert_error( - System::new(length_unit.clone(), types, positions, bad_cell, pbc), - "invalid parameter: `cell` must be a (3 x 3) tensor, got a tensor with shape [2, 3]", - ); + let err = System::new(length_unit.clone(), types, positions, bad_cell, pbc).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: `cell` must be a (3 x 3) tensor, got a tensor with shape [2, 3]"); let types = type_tensor(&[1, 6]); let positions = positions_tensor(2, "f32"); let cell = cell_tensor(0.0, "f64"); let pbc = pbc_tensor(&[true, true, true]); - assert_error( - System::new(length_unit.clone(), types, positions, cell, pbc), - "invalid parameter: `cell` must have the same dtype as `positions`", - ); + let err = System::new(length_unit.clone(), types, positions, cell, pbc).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: `cell` must have the same dtype as `positions`, got f64 and f32"); let bad_pbc_dtype: DLPackTensor = Array1::::from_vec(vec![1, 0, 1]).try_into().unwrap(); let types = type_tensor(&[1, 6]); let positions = positions_tensor(2, "f32"); let cell = cell_tensor(0.0, "f32"); - assert_error( - System::new(length_unit.clone(), types, positions, cell, bad_pbc_dtype), - "invalid parameter: `pbc` must be a tensor of booleans", - ); + let err = System::new(length_unit.clone(), types, positions, cell, bad_pbc_dtype).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: `pbc` must be a tensor of booleans"); let types = type_tensor(&[1, 6]); let positions = positions_tensor(2, "f32"); let cell = cell_tensor(0.0, "f32"); let bad_pbc = pbc_tensor(&[true, true]); - assert_error( - System::new(length_unit, types, positions, cell, bad_pbc), - "invalid parameter: `pbc` must contain 3 entries, got a tensor with shape [2]", - ); + let err = System::new(length_unit, types, positions, cell, bad_pbc).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: `pbc` must contain 3 entries, got a tensor with shape [2]"); } #[test] @@ -651,10 +692,8 @@ mod tests { let positions = positions_tensor(1, "f32"); let cell = cell_tensor(10.0, "f32"); let pbc = pbc_tensor(&[true, false, true]); - assert_error( - System::new(length_unit.clone(), types, positions, cell, pbc), - "invalid parameter: invalid cell: for non-periodic dimensions, the corresponding cell vector must be zero, but cell[1] contains non-zero values", - ); + let err = System::new(length_unit.clone(), types, positions, cell, pbc).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: invalid cell: for non-periodic dimensions, the corresponding cell vector must be zero, but cell[1] contains non-zero values"); } #[test] @@ -669,6 +708,7 @@ mod tests { let options = PairListOptions { cutoff: 3.5, full_list: true, strict: false, requestors: vec![] }; let pairs = valid_pair_block("f32"); + let pairs_ptr = pairs.as_ptr(); system.add_pairs(options.clone(), pairs).unwrap(); assert_eq!(system.known_pairs().len(), 1); assert_eq!(system.get_pairs(&options).unwrap().properties().names(), ["distance"]); @@ -679,9 +719,9 @@ mod tests { strict: false, requestors: vec!["test-requestor".into()], }; - // TODO: check that this is the exact same block once we can get the - // pointer to check for id. - assert!(system.get_pairs(&options_with_requestor).is_some()); + + let pairs_from_system = system.get_pairs(&options_with_requestor).unwrap(); + assert_eq!(pairs_from_system.as_ptr(), pairs_ptr); system.add_pairs( PairListOptions { cutoff: 5.0, full_list: false, strict: true, requestors: vec![] }, @@ -706,10 +746,8 @@ mod tests { assert_eq!(system.known_custom_data(), vec!["test::my_data"]); assert_eq!(system.get_custom_data("test::my_data").unwrap().keys().names(), ["key"]); - assert_error( - system.add_custom_data("test::my_data", valid_custom_data("f32"), false), - "invalid parameter: custom data 'test::my_data' is already present in this system", - ); + let err = system.add_custom_data("test::my_data", valid_custom_data("f32"), false).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: custom data 'test::my_data' is already present in this system"); let replacement = valid_custom_data("f32"); system.add_custom_data("test::my_data", replacement, true).unwrap(); @@ -722,20 +760,27 @@ mod tests { cell_tensor(10.0, "f32"), pbc_tensor(&[true, true, true]), ).unwrap(); - system.add_custom_data("test::a", valid_custom_data("f32"), false).unwrap(); - system.add_custom_data("test::b", valid_custom_data("f32"), false).unwrap(); + + let test_data_a = valid_custom_data("f32"); + let test_data_a_ptr = test_data_a.as_ptr(); + system.add_custom_data("test::a", test_data_a, false).unwrap(); + + let test_data_b = valid_custom_data("f32"); + let test_data_b_ptr = test_data_b.as_ptr(); + system.add_custom_data("test::b", test_data_b, false).unwrap(); + let mut names = system.known_custom_data(); names.sort_unstable(); assert_eq!(names, vec!["test::a", "test::b"]); - // TODO: check we get back the same pointer - assert!(system.get_custom_data("test::a").is_ok()); - assert!(system.get_custom_data("test::b").is_ok()); + let data_a = system.get_custom_data("test::a").unwrap(); + assert_eq!(data_a.as_ptr(), test_data_a_ptr); - assert_error( - system.get_custom_data("no_such_data"), - "invalid parameter: no data for 'no_such_data' found in this system", - ); + let data_b = system.get_custom_data("test::b").unwrap(); + assert_eq!(data_b.as_ptr(), test_data_b_ptr); + + let err = system.get_custom_data("no_such_data").unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: no custom data for 'no_such_data' found in this system"); } #[test] @@ -749,28 +794,20 @@ mod tests { ).unwrap(); for name in ["types", "type", "Positions", "position", "CELL", "neighbors", "neighbor", "pair", "pairs", "Types", "POSITIONS", "Cell", "Neighbors"] { let data = valid_custom_data("f32"); - assert_error( - system.add_custom_data(name.to_string(), data, false), - &format!("invalid parameter: custom data can not be named '{}'", name), - ); + let err = system.add_custom_data(name.to_string(), data, false).unwrap_err(); + assert_eq!(err.to_string(), format!("invalid parameter: custom data can not be named '{}'", name)); } - assert_error( - system.add_custom_data("my_data", valid_custom_data("f32"), false), - "invalid parameter: 'my_data' is not a standard quantity name; custom quantity names must use '::'", - ); + let err = system.add_custom_data("my_data", valid_custom_data("f32"), false).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: 'my_data' is not a standard quantity name; custom quantity names must use '::'"); let keys = Labels::empty(vec!["key"]); let empty = TensorMap::new(keys, vec![]).unwrap(); - assert_error( - system.add_custom_data("test::empty", empty, false), - "invalid parameter: custom data 'test::empty' has no blocks", - ); + let err = system.add_custom_data("test::empty", empty, false).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: custom data 'test::empty' has no blocks"); let dtype_mismatch = valid_custom_data("f64"); - assert_error( - system.add_custom_data("test::dtype", dtype_mismatch, false), - "invalid parameter: dtype of custom data 'test::dtype' does not match this system dtype", - ); + let err = system.add_custom_data("test::dtype", dtype_mismatch, false).unwrap_err(); + assert_eq!(err.to_string(), "invalid parameter: invalid dtype for quantity 'test::dtype': expected f32, got f64"); } } diff --git a/metatomic-core/tests/system.cpp b/metatomic-core/tests/system.cpp index 638b31154..37a4a66db 100644 --- a/metatomic-core/tests/system.cpp +++ b/metatomic-core/tests/system.cpp @@ -237,7 +237,7 @@ TEST_CASE("system") { CHECK(system == nullptr); mta_last_error(&message, nullptr, nullptr); - CHECK(std::string(message) == "invalid parameter: `cell` must have the same dtype as `positions`"); + CHECK(std::string(message) == "invalid parameter: `cell` must have the same dtype as `positions`, got i32 and f32"); // wrong dtype for pbc (float instead of bool) status = mta_system_create( From ec1067132678898d0f8b68ad7d136438a063eed0 Mon Sep 17 00:00:00 2001 From: Guillaume Fraux Date: Fri, 31 Jul 2026 14:33:25 +0200 Subject: [PATCH 43/43] Update to metatensor-core v0.2.4 This includes SimpleDataArray which simplifies the tests --- .github/workflows/build-wheels.yml | 6 +- .github/workflows/rust-tests.yml | 1 + .github/workflows/torch-tests.yml | 1 + metatomic-core/CMakeLists.txt | 2 +- metatomic-core/Cargo.toml | 2 +- metatomic-core/tests/cxx/system.cpp | 10 +-- metatomic-core/tests/system.cpp | 92 ++-------------------------- metatomic-core/tests/utils/mod.rs | 2 +- python/metatomic_core/pyproject.toml | 2 +- 9 files changed, 17 insertions(+), 101 deletions(-) diff --git a/.github/workflows/build-wheels.yml b/.github/workflows/build-wheels.yml index 16ed75626..ca77610d0 100644 --- a/.github/workflows/build-wheels.yml +++ b/.github/workflows/build-wheels.yml @@ -365,9 +365,9 @@ jobs: - name: setup libmetatensor run: | - curl --location -O https://github.com/metatensor/metatensor/releases/download/metatensor-core-v0.2.3/metatensor-core-cxx-0.2.3.tar.gz - tar xf metatensor-core-cxx-0.2.3.tar.gz - cmake -B build-metatensor -S metatensor-core-cxx-0.2.3 \ + curl --location -O https://github.com/metatensor/metatensor/releases/download/metatensor-core-v0.2.4/metatensor-core-cxx-0.2.4.tar.gz + tar xf metatensor-core-cxx-0.2.4.tar.gz + cmake -B build-metatensor -S metatensor-core-cxx-0.2.4 \ -DMETATENSOR_INSTALL_BOTH_STATIC_SHARED=OFF \ -DCMAKE_INSTALL_PREFIX=$CMAKE_PREFIX_PATH \ -DCMAKE_BUILD_TYPE=Debug diff --git a/.github/workflows/rust-tests.yml b/.github/workflows/rust-tests.yml index 36eb59bff..5951f68bc 100644 --- a/.github/workflows/rust-tests.yml +++ b/.github/workflows/rust-tests.yml @@ -107,6 +107,7 @@ jobs: - name: install valgrind if: matrix.do-valgrind run: | + sudo apt-get update sudo apt-get install -y valgrind - name: Setup sccache diff --git a/.github/workflows/torch-tests.yml b/.github/workflows/torch-tests.yml index d9dd09dc3..38bc2a8f9 100644 --- a/.github/workflows/torch-tests.yml +++ b/.github/workflows/torch-tests.yml @@ -73,6 +73,7 @@ jobs: - name: install valgrind if: matrix.do-valgrind run: | + sudo apt-get update sudo apt-get install -y valgrind - name: Setup sccache diff --git a/metatomic-core/CMakeLists.txt b/metatomic-core/CMakeLists.txt index e7dd396a7..0eb97f20f 100644 --- a/metatomic-core/CMakeLists.txt +++ b/metatomic-core/CMakeLists.txt @@ -106,7 +106,7 @@ function(check_compatible_versions _actual_ _requested_) endfunction() -set(REQUIRED_METATENSOR_VERSION "0.2.0") +set(REQUIRED_METATENSOR_VERSION "0.2.4") # Either metatensor is built as part of the same CMake project, or we try to # find the corresponding CMake package if (TARGET metatensor) diff --git a/metatomic-core/Cargo.toml b/metatomic-core/Cargo.toml index db3482f68..071e97b34 100644 --- a/metatomic-core/Cargo.toml +++ b/metatomic-core/Cargo.toml @@ -14,7 +14,7 @@ name = "metatomic" bench = false [dependencies] -metatensor = { version = "0.5.0" } +metatensor = { version = "0.5.1" } dlpk = { version = "0.4", features = ["ndarray"]} json = "0.12" libloading = "0.9" diff --git a/metatomic-core/tests/cxx/system.cpp b/metatomic-core/tests/cxx/system.cpp index d21715ae2..79532d16a 100644 --- a/metatomic-core/tests/cxx/system.cpp +++ b/metatomic-core/tests/cxx/system.cpp @@ -71,20 +71,14 @@ static metatomic::DLPackTensor cell_tensor() { } static metatomic::DLPackTensor pbc_tensor() { - // `SimpleDataArray` does not compile (`std::vector` has no - // `data()` method), so we use `uint8_t` and patch the dtype code to - // `kDLBool`. - auto array = std::make_unique>( + auto array = std::make_unique>( std::vector{3}, std::vector{1, 0, 1} ); auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); DLDevice cpu = {kDLCPU, 0}; DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; - auto* tensor = mts.as_dlpack(cpu, nullptr, version); - tensor->dl_tensor.dtype.code = DLDataTypeCode::kDLBool; - - return metatomic::DLPackTensor(tensor); + return metatomic::DLPackTensor(mts.as_dlpack(cpu, nullptr, version)); } static metatomic::System test_system(size_t n_atoms = 4) { diff --git a/metatomic-core/tests/system.cpp b/metatomic-core/tests/system.cpp index 37a4a66db..2686e6ce7 100644 --- a/metatomic-core/tests/system.cpp +++ b/metatomic-core/tests/system.cpp @@ -72,24 +72,18 @@ template static DLManagedTensorVersioned* pbc_tensor() { return mts.as_dlpack(cpu, nullptr, version); } - -/// SimpleDataArray doesn't compile (std::vector has no data() -/// method). We use SimpleDataArray and patch the dtype code -/// from kDLUInt to kDLBool. +/// `SimpleDataArray` stores data as `uint8_t` internally, so the data +/// vector must use `uint8_t` as well. template <> DLManagedTensorVersioned* pbc_tensor() { std::vector pbc_data = {1, 0, 1}; - auto array = std::make_unique>( + auto array = std::make_unique>( std::vector{3}, std::move(pbc_data) ); auto mts = metatensor::DataArrayBase::to_mts_array(std::move(array)); DLDevice cpu = {kDLCPU, 0}; DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; - auto* tensor = mts.as_dlpack(cpu, nullptr, version); - - tensor->dl_tensor.dtype.code = DLDataTypeCode::kDLBool; - - return tensor; + return mts.as_dlpack(cpu, nullptr, version); } static mts_block_t* pair_block() { @@ -716,56 +710,6 @@ static void check_full_system_data(const mta_system_t* system) { CHECK(retrieved != nullptr); } -/// `DataArrayBase` storing boolean data as `uint8_t` (since -/// `std::vector` has no `data()` method, `SimpleDataArray` can not -/// be used). This class reports its dtype as `kDLBool` so that the metatensor -/// serialization code correctly handles it. -class BoolDataArray: public metatensor::SimpleDataArray { -public: - using SimpleDataArray::SimpleDataArray; - - DLDataType dtype() const override { - DLDataType dtype; - dtype.code = DLDataTypeCode::kDLBool; - dtype.bits = 8; - dtype.lanes = 1; - return dtype; - } - - DLManagedTensorVersioned* as_dlpack( - DLDevice device, - const int64_t* stream, - DLPackVersion max_version - ) override { - auto* managed = SimpleDataArray::as_dlpack(device, stream, max_version); - managed->dl_tensor.dtype.code = DLDataTypeCode::kDLBool; - return managed; - } - - std::unique_ptr copy(DLDevice device) const override { - if (device.device_type != kDLCPU) { - throw metatensor::Error("BoolDataArray only supports copying to CPU"); - } - return std::unique_ptr(new BoolDataArray(*this)); - } - - std::unique_ptr create( - std::vector shape, - metatensor::MtsArray fill_value - ) const override { - DLDevice cpu_device = {kDLCPU, 0}; - DLPackVersion version = {DLPACK_MAJOR_VERSION, DLPACK_MINOR_VERSION}; - auto fill_dlpack = fill_value.as_dlpack_array(cpu_device, nullptr, version); - - if (!fill_dlpack.shape().empty()) { - throw metatensor::Error("`fill_value` must be a single scalar"); - } - - auto scalar = fill_dlpack.data()[0]; - return std::unique_ptr(new BoolDataArray(std::move(shape), scalar)); - } -}; - /// `mts_realloc_buffer_t` callback backed by a `std::vector`. static uint8_t* vector_realloc(void* user_data, uint8_t* /*ptr*/, uintptr_t new_size) { auto* buffer = static_cast*>(user_data); @@ -773,30 +717,6 @@ static uint8_t* vector_realloc(void* user_data, uint8_t* /*ptr*/, uintptr_t new_ return buffer->data(); } -/// `mts_create_array_callback_t` that delegates to -/// `metatensor::details::default_create_array`, but handles `kDLBool` by -/// creating a `BoolDataArray` (`SimpleDataArray` does not compile since -/// `std::vector` has no `data()` method). Can be removed once -/// https://github.com/metatensor/metatensor/pull/1164 is released. -static mts_status_t create_array_with_bool( - const uintptr_t* shape_ptr, - uintptr_t shape_count, - DLDataType dtype, - mts_array_t* array -) { - if (dtype.code == kDLBool && dtype.bits == 8 && dtype.lanes == 1) { - auto shape = std::vector(); - for (uintptr_t i = 0; i < shape_count; i++) { - shape.push_back(shape_ptr[i]); - } - auto cxx_array = std::make_unique(shape); - *array = metatensor::DataArrayBase::to_mts_array(std::move(cxx_array)).release(); - return MTS_SUCCESS; - } - - return metatensor::details::default_create_array(shape_ptr, shape_count, dtype, array); -} - TEST_CASE("system serialization") { SECTION("save and load to a file") { auto* system = full_test_system(); @@ -808,7 +728,7 @@ TEST_CASE("system serialization") { mta_system_t* loaded = nullptr; auto status = mta_load( path.c_str(), - create_array_with_bool, + metatensor::details::default_create_array, &loaded ); CHECK(status == MTA_SUCCESS); @@ -837,7 +757,7 @@ TEST_CASE("system serialization") { mta_system_t* loaded = nullptr; status = mta_load_buffer( buffer.data(), buffer.size(), - create_array_with_bool, + metatensor::details::default_create_array, &loaded ); CHECK(status == MTA_SUCCESS); diff --git a/metatomic-core/tests/utils/mod.rs b/metatomic-core/tests/utils/mod.rs index 7a22f5e67..ff2ae89ff 100644 --- a/metatomic-core/tests/utils/mod.rs +++ b/metatomic-core/tests/utils/mod.rs @@ -245,7 +245,7 @@ pub fn setup_torch_pip(python: &Path) -> PathBuf { /// Install metatensor in a Python virtualenv with pip, and return the /// CMAKE_PREFIX_PATH for the installed libmetatensor. pub fn setup_metatensor_pip(python: &Path) -> PathBuf { - pip_install(python, &["metatensor-core >=0.2.2,<0.3"], PipInstallOptions::default()); + pip_install(python, &["metatensor-core >=0.2.4,<0.3"], PipInstallOptions::default()); let mut cmd = Command::new(python); cmd.arg("-c"); diff --git a/python/metatomic_core/pyproject.toml b/python/metatomic_core/pyproject.toml index 62e501ba9..b2320ca33 100644 --- a/python/metatomic_core/pyproject.toml +++ b/python/metatomic_core/pyproject.toml @@ -36,7 +36,7 @@ requires = [ "setuptools >=77", "packaging >=26", "cmake", - "metatensor-core >=0.2.2,<0.3", + "metatensor-core >=0.2.4,<0.3", ] build-backend = "setuptools.build_meta"