diff --git a/packages/mace-core/README.md b/packages/mace-core/README.md index 0b2cb598b..4f36ee2a7 100644 --- a/packages/mace-core/README.md +++ b/packages/mace-core/README.md @@ -5,4 +5,17 @@ observables, the kernel-backend protocol, Clebsch-Gordan, neighbours and the data spec. It imports no torch, no jax and no e3nn, so both implementation packages can depend on it without depending on each other. -Distribution `mace-core`, import name `mace_core`. Scaffold only for now. +Distribution `mace-core`, import name `mace_core`. + +What is here so far: + +- `mace_core.config` — `BaseConfig`, the pydantic base every v1 config schema + derives from: one TOML/YAML/JSON file, validated once; unknown keys are + errors, a validated config is frozen, and its resolved export loads back to + the same config. `config.cli` holds the `--a.b value` override grammar for a + command line to compose with it. +- `mace_core.metadata` — `ModelMetadata`, the versioned record every trained + model carries (config as written and as resolved, provenance, a summary per + data source, per head the sources it consumed, DOI, citations and notes), + with a JSON round trip, `ConfigRecord.from_config()` to embed a config in + its fixed-point form, and `format_citations()`. diff --git a/packages/mace-core/pyproject.toml b/packages/mace-core/pyproject.toml index 7caf4c664..60075d184 100644 --- a/packages/mace-core/pyproject.toml +++ b/packages/mace-core/pyproject.toml @@ -17,7 +17,15 @@ classifiers = [ "Programming Language :: Python :: 3.13", "Operating System :: OS Independent", ] -dependencies = [] +dependencies = [ + # The suite passes at this floor (2.7.4 on Python 3.10) and at current + # releases. + "pydantic>=2.7", + "pyyaml>=6.0", + "tomli>=2.0; python_version < '3.11'", + # `typing.Self` is 3.11+; the floor is 3.10. + "typing-extensions>=4.4", +] [project.urls] Homepage = "https://github.com/ACEsuit/mace" diff --git a/packages/mace-core/src/mace_core/__init__.py b/packages/mace-core/src/mace_core/__init__.py index 4122cd7d8..c6920d9ae 100644 --- a/packages/mace-core/src/mace_core/__init__.py +++ b/packages/mace-core/src/mace_core/__init__.py @@ -1,11 +1,22 @@ """Framework-agnostic contract and pure math for MACE v1. -Scaffold only. The public surface arrives with the tickets that build it. +Re-exports the public surface of the lightweight submodules. Heavy ones (the +Clebsch-Gordan basis, neighbours) are imported from their own modules. """ from importlib.metadata import PackageNotFoundError, version -__all__ = ["__version__"] +from mace_core.config import BaseConfig, ConfigError, ConfigSection +from mace_core.metadata import ModelMetadata, format_citations + +__all__ = [ + "BaseConfig", + "ConfigError", + "ConfigSection", + "ModelMetadata", + "__version__", + "format_citations", +] #: Version of the installed `mace-core` distribution. Read from installed metadata #: rather than hardcoded, so it cannot drift from what pip resolved. diff --git a/packages/mace-core/src/mace_core/config/__init__.py b/packages/mace-core/src/mace_core/config/__init__.py new file mode 100644 index 000000000..d80e5fd27 --- /dev/null +++ b/packages/mace-core/src/mace_core/config/__init__.py @@ -0,0 +1,23 @@ +"""Configuration schemas for MACE v1. + +The base machinery lives in `base` and the command-line override grammar in +`cli`; the training schema sections (model, data, training, ...) arrive with +their own tickets and are re-exported from here. +""" + +from mace_core.config.base import ( + BaseConfig, + ConfigError, + ConfigSection, + read_config_file, +) +from mace_core.config.cli import apply_overrides, parse_overrides + +__all__ = [ + "BaseConfig", + "ConfigError", + "ConfigSection", + "apply_overrides", + "parse_overrides", + "read_config_file", +] diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py new file mode 100644 index 000000000..3fef486ac --- /dev/null +++ b/packages/mace-core/src/mace_core/config/base.py @@ -0,0 +1,188 @@ +"""The config base: one file, validated once by pydantic, exported as a dict. + +Everything the schema rejects, unknown keys included, is pydantic's +`ValidationError`; `ConfigError` is only for a file that cannot be read or +parsed. Nothing here knows about a command line: that is `cli`. +""" + +import json +import os +import sys +from collections.abc import Iterator, Mapping +from collections.abc import Set as AbstractSet +from pathlib import Path +from typing import Any, get_args, get_origin + +import yaml +from pydantic import BaseModel, ConfigDict +from typing_extensions import Self + +if sys.version_info >= (3, 11): + import tomllib +else: + import tomli as tomllib + +__all__ = ["BaseConfig", "ConfigError", "ConfigSection", "read_config_file"] + + +class ConfigError(ValueError): + """A config file, or a command-line override (`cli`), that cannot be read or + parsed. Schema errors are pydantic's.""" + + +def read_config_file(path: str | os.PathLike[str]) -> dict[str, Any]: + """Parse one TOML, YAML or JSON config file, chosen by its (case-folded) + extension, into a dict; an empty or comment-only file is `{}`. Raises + `ConfigError` for an unknown extension, an unreadable file, a parse failure + or a top level that is not a mapping.""" + path = Path(path) + parsers = { + ".toml": tomllib.loads, + ".yaml": yaml.safe_load, + ".yml": yaml.safe_load, + ".json": json.loads, + } + parse = parsers.get(path.suffix.lower()) + if parse is None: + raise ConfigError( + f"cannot read config file {path}: unknown extension {path.suffix!r}; " + "use .toml, .yaml, .yml or .json" + ) + try: + text = path.read_text(encoding="utf-8") + except (OSError, UnicodeDecodeError) as error: + raise ConfigError(f"cannot read config file {path}: {error}") from error + try: + document = parse(text) + except (ValueError, yaml.YAMLError) as error: # toml, json: ValueError + raise ConfigError(f"cannot parse config file {path}: {error}") from error + if document is None: # YAML reads an empty or comment-only file as None + return {} + if not isinstance(document, dict): + raise ConfigError( + f"config file {path} must be a mapping of keys to values at the top " + f"level, not {type(document).__name__}" + ) + return document + + +class ConfigSection(BaseModel): + """A node of a config tree: unknown keys are errors, and its attributes + cannot be reassigned (a list or dict it holds is not protected). A change + is a new validation, `from_dict` on an edited dict. + + Subclasses declare fields only. A subclass that would undo either rule, or + hold what does not load back from the export (a set, an alias, an excluded + or computed field, a plain `BaseModel`), is a `TypeError` when it is + defined, or on `load` if a forward reference left it incomplete. Declare a + forward-referenced class at module level: pydantic cannot resolve one + local to a function. + """ + + # inf/nan as the JSON constants Infinity and NaN, not pydantic's default + # null, which would turn a value into a different one. + model_config = ConfigDict(extra="forbid", frozen=True, ser_json_inf_nan="constants") + + @classmethod + def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: + super().__pydantic_init_subclass__(**kwargs) + extra = cls.model_config.get("extra") + if extra != "forbid": + raise TypeError( + f"{cls.__name__} sets extra={extra!r}; a section keeps " + "extra='forbid' so an unknown key stays an error" + ) + if not cls.model_config.get("frozen"): + raise TypeError( + f"{cls.__name__} sets frozen=False; a section stays frozen " + "so a validated config is not changed behind the validation" + ) + # A forward reference leaves the field types unknown for now; + # `_check_schema` checks the section on load instead. + if cls.__pydantic_complete__: + _check_section_fields(cls) + + +def _types_in(annotation: Any) -> Iterator[type]: + """Every class an annotation mentions, containers included.""" + outer = get_origin(annotation) or annotation + if isinstance(outer, type): + yield outer + for argument in get_args(annotation): + yield from _types_in(argument) + + +def _check_section_fields(section: type[ConfigSection]) -> None: + """Refuse, with a `TypeError`, what pydantic allows on a field but a config + section cannot have: a field that would not load back from the export, or + a child that would ignore unknown keys.""" + if section.model_computed_fields: + name = next(iter(section.model_computed_fields)) + raise TypeError( + f"{section.__name__}.{name} is a computed field; it would not load back" + ) + for name, field in section.model_fields.items(): + where = f"{section.__name__}.{name}" + if field.alias or field.validation_alias or field.serialization_alias: + raise TypeError(f"{where} has an alias; a config key is its field name") + if field.exclude: + raise TypeError(f"{where} is excluded from dumps; it would not load back") + for held in _types_in(field.annotation): + # Any set type, `set`, `frozenset` or an abstract one, dumps in an + # order that varies between runs, so the export would not be stable. + if issubclass(held, AbstractSet): + raise TypeError(f"{where} is typed as a set; order varies. Use a list") + if issubclass(held, BaseModel) and not issubclass(held, ConfigSection): + raise TypeError( + f"{where} holds {held.__name__}, which is not a ConfigSection; " + "unknown keys under it would be dropped" + ) + + +def _check_schema(root: type[ConfigSection]) -> None: + """Run `_check_section_fields` on the root and every section under it. + + A section is checked when it is defined, unless a forward reference left + its field types unknown then. Such a section is resolved and checked here, + on load. The others are checked a second time, which costs less than + tracking which ones were skipped.""" + to_visit, seen = [root], set() + while to_visit: + section = to_visit.pop() + if section in seen: + continue + seen.add(section) + if not section.__pydantic_complete__: + section.model_rebuild() # resolves the forward reference or raises + _check_section_fields(section) + for field in section.model_fields.values(): + for held in _types_in(field.annotation): + if issubclass(held, ConfigSection): + to_visit.append(held) + + +class BaseConfig(ConfigSection): + """Base class for the root of a config schema: subclass it and declare the + fields, then `load` a file or `from_dict` a parsed one.""" + + @classmethod + def load(cls, config_file: str | os.PathLike[str]) -> Self: + """Build the config from one TOML, YAML or JSON file.""" + return cls.from_dict(read_config_file(config_file)) + + @classmethod + def from_dict(cls, document: Mapping[str, Any]) -> Self: + """Build the config from an already parsed file, such as one a command + line applied its overrides to.""" + _check_schema(cls) + return cls.model_validate(document) + + def to_resolved_dict(self) -> dict[str, Any]: + """A JSON-compatible dict that loads back to this config, defaults + filled in.""" + return self.model_dump(mode="json") + + def to_user_dict(self) -> dict[str, Any]: + """A JSON-compatible dict that loads back to this config, holding only + the fields that were set.""" + return self.model_dump(mode="json", exclude_unset=True) diff --git a/packages/mace-core/src/mace_core/config/cli.py b/packages/mace-core/src/mace_core/config/cli.py new file mode 100644 index 000000000..da3f08aeb --- /dev/null +++ b/packages/mace-core/src/mace_core/config/cli.py @@ -0,0 +1,104 @@ +"""The command-line half of a config: `--a.b value` tokens, and a parsed file +edited before validation. + +`parse_overrides(tokens)` turns command-line tokens into a mapping of dotted +path to value; `apply_overrides(document, overrides)` writes such a mapping +into a copy of a parsed config file. A command line then builds its config as + + Config.from_dict(apply_overrides(read_config_file(path), parse_overrides(rest))) + +and every error the schema reports is pydantic's, at the dotted location the +override named. Nothing in `base` knows about this module; a CLI that exposes +explicit flags instead can hand `apply_overrides` a mapping it built itself. +""" + +import json +from collections.abc import Iterable, Mapping +from typing import Any + +from mace_core.config.base import ConfigError + +__all__ = ["apply_overrides", "parse_overrides"] + + +def parse_overrides(tokens: Iterable[str]) -> dict[str, Any]: + """Turn command-line tokens into a dict of dotted path to value. + + Both `--a.b value` and `--a.b=value` give `{"a.b": value}`. A value that + is `null` or starts with `[` or `{` is parsed as JSON; any other stays a + string for pydantic to coerce. A repeated path keeps its last value. A + token that fits none of this is a `ConfigError`.""" + if isinstance(tokens, str): + raise TypeError(f"tokens is a string, {tokens!r}; pass a list of tokens") + argv = list(tokens) + overrides: dict[str, Any] = {} + position = 0 + while position < len(argv): + token = argv[position] + position += 1 + dotted_path, has_inline_value, value = token[2:].partition("=") + if not token.startswith("--") or not dotted_path: + raise ConfigError( + f"unknown config option '{token}'; options are --key.path value" + ) + if not has_inline_value: + if position == len(argv): + raise ConfigError(f"override {token} is missing its value") + value = argv[position] + position += 1 + parsed: Any = value + if value == "null" or value[:1] in ("[", "{"): + try: + parsed = json.loads(value) + except ValueError as error: + raise ConfigError( + f"override {token} is not valid JSON: {error}" + ) from error + overrides.pop(dotted_path, None) # a repeated path applies where it is last + overrides[dotted_path] = parsed + return overrides + + +def apply_overrides( + document: Mapping[str, Any], overrides: Mapping[str, Any] +) -> dict[str, Any]: + """Return a copy of a parsed config file with the overrides written in. + + Each value goes to its dotted path, in order. A mapping merges into a + mapping already there, so `--model '{"depth": 3}'` keeps the file's other + `model` keys; any other value replaces what was there.""" + copy = _copy_tree(document) + for dotted_path, value in overrides.items(): + _set_at_path(copy, dotted_path.split("."), value) + return copy + + +def _set_at_path(mapping: dict[str, Any], keys: list[str], value: Any) -> None: + """Walk down the keys, creating dicts on the way (a parent that is not a + dict is replaced by one), and set or merge the value at the last key.""" + *parent_keys, last_key = keys + for key in parent_keys: + if not isinstance(mapping.get(key), dict): + mapping[key] = {} + mapping = mapping[key] + _set_or_merge(mapping, last_key, value) + + +def _set_or_merge(mapping: dict[str, Any], key: str, value: Any) -> None: + """Set `mapping[key]`. A dict merges, key by key, into a dict already + there; anything else replaces what was there.""" + if isinstance(value, Mapping) and isinstance(mapping.get(key), dict): + for inner_key, item in value.items(): + _set_or_merge(mapping[key], inner_key, item) + else: + mapping[key] = _copy_tree(value) + + +def _copy_tree(value: Any) -> Any: + """Copy the dicts and lists of a document, scalars as they are. Unlike + `copy.deepcopy`, two keys sharing one object get two copies.""" + if isinstance(value, Mapping): + return {key: _copy_tree(item) for key, item in value.items()} + if isinstance(value, list): + return [_copy_tree(item) for item in value] + return value diff --git a/packages/mace-core/src/mace_core/metadata.py b/packages/mace-core/src/mace_core/metadata.py new file mode 100644 index 000000000..fbd96f3b1 --- /dev/null +++ b/packages/mace-core/src/mace_core/metadata.py @@ -0,0 +1,180 @@ +"""`ModelMetadata`: the record every trained v1 model carries about how it was +made, stored as JSON beside the weights.""" + +from __future__ import annotations + +from collections.abc import Iterable +from typing import Any, Final + +from pydantic import BaseModel, ConfigDict, Field, model_validator + +from mace_core.config import BaseConfig + +__all__ = [ + "SCHEMA_VERSION", + "Citation", + "ConfigRecord", + "DataSourceSummary", + "DataSummary", + "HeadSummary", + "MetadataSchemaError", + "ModelMetadata", + "Provenance", + "format_citations", +] + +#: Bump when a field is added, removed or changes meaning. +SCHEMA_VERSION: Final = 1 + + +class MetadataSchemaError(ValueError): + """Metadata that would not read back from its JSON as the same record.""" + + +class _Record(BaseModel): + """Common ground: unknown keys are errors, so a typo cannot be stored.""" + + # inf/nan are written as JSON constants (Infinity, NaN) rather than + # pydantic's default null, which would turn a value into a different one. + model_config = ConfigDict(extra="forbid", ser_json_inf_nan="constants") + + +class ConfigRecord(_Record): + """Record of the training config, as written and as resolved.""" + + user: dict[str, Any] = Field(default_factory=dict) + resolved: dict[str, Any] = Field(default_factory=dict) + + @classmethod + def from_config(cls, config: BaseConfig) -> ConfigRecord: + return cls(user=config.to_user_dict(), resolved=config.to_resolved_dict()) + + +class Provenance(_Record): + """Which code produced the model.""" + + #: Version per distribution used, `{"mace-core": "1.0.2", "mace-torch": "1.1.0"}` + versions: dict[str, str] + #: Full hash of the commit the code was run from; None when not in a checkout. + git_commit: str | None = None + + +class DataSourceSummary(_Record): + """Summary of one data source.""" + + #: The data source's name in the config. + name: str + num_configurations: int | None = None + num_atoms: int | None = None + #: Atomic numbers of every element present. + elements: list[int] = Field(default_factory=list) + #: The reference quantities provided, each prefixed by the method that + #: produced it: `pbe_energy`, `pbe_forces`, `r2scan_energy`. + reference_keys: list[str] = Field(default_factory=list) + + +class DataSummary(_Record): + """One summary per data source, each source once.""" + + sources: list[DataSourceSummary] = Field(default_factory=list) + + +class HeadSummary(_Record): + """What one head was fitted on.""" + + #: Names from `DataSummary.sources`. + sources: list[str] = Field(default_factory=list) + + +class Citation(_Record): + """One work users of the model are asked to cite.""" + + title: str + authors: list[str] = Field(default_factory=list) + venue: str | None = None + year: int | None = None + doi: str | None = None + url: str | None = None + + +class ModelMetadata(_Record): + """How a model was made: its config, the code, the data and what to cite.""" + + schema_version: int = SCHEMA_VERSION + config: ConfigRecord + provenance: Provenance + data: DataSummary = Field(default_factory=DataSummary) + #: Keyed by head name, as in the config; a single-head model has one entry. + heads: dict[str, HeadSummary] = Field(default_factory=dict) + #: DOI of the model itself, not of the papers describing it. + doi: str | None = None + citations: list[Citation] = Field(default_factory=list) + notes: str = "" + + @model_validator(mode="before") + @classmethod + def _reads_only_its_own_schema_version(cls, data: Any) -> Any: + # Before the fields are read: a newer record may hold fields this code + # does not know, and the version is the error worth reporting. + if isinstance(data, dict): + version = data.get("schema_version", SCHEMA_VERSION) + if version != SCHEMA_VERSION: + raise ValueError( + f"model metadata has schema_version {version!r}, but this " + f"mace-core reads schema_version {SCHEMA_VERSION}; a newer " + "record needs an upgrade of mace-core" + ) + return data + + @model_validator(mode="after") + def _heads_name_known_sources(self) -> ModelMetadata: + names = [source.name for source in self.data.sources] + if len(set(names)) != len(names): + raise ValueError(f"data.sources names a source twice: {sorted(names)}") + for head, summary in self.heads.items(): + for name in summary.sources: + if name not in names: + raise ValueError( + f"heads.{head}.sources names {name!r}, which is not in " + f"data.sources; add its summary or drop the name" + ) + return self + + def to_json(self, indent: int | None = 2) -> str: + """Serialise to JSON, checking that it reads back as the same record.""" + text = self.model_dump_json(indent=indent) + if self.from_json(text) != self: + raise MetadataSchemaError( + "model metadata does not survive a JSON round trip: a field holds " + "a value JSON reads back differently, such as a tuple (a list " + "on the way back) or NaN (never equal to itself)" + ) + return text + + @classmethod + def from_json(cls, text: str) -> ModelMetadata: + """Read a record written by `to_json`.""" + return cls.model_validate_json(text) + + +def format_citations(citations: Iterable[Citation]) -> str: + """Render citations as a numbered block, one line each; unset fields are + left out.""" + lines = [] + for number, citation in enumerate(citations, start=1): + parts = [] + if citation.authors: + parts.append(", ".join(citation.authors)) + parts.append(citation.title) + if citation.venue and citation.year: + parts.append(f"{citation.venue} ({citation.year})") + elif citation.venue: + parts.append(citation.venue) + elif citation.year: + parts.append(str(citation.year)) + if citation.doi: + parts.append(f"https://doi.org/{citation.doi}") + elif citation.url: + parts.append(citation.url) + lines.append(f"[{number}] " + ". ".join(parts)) + return "\n".join(lines) diff --git a/packages/mace-core/tests/mace_core_demo.py b/packages/mace-core/tests/mace_core_demo.py new file mode 100644 index 000000000..faa6c294e --- /dev/null +++ b/packages/mace-core/tests/mace_core_demo.py @@ -0,0 +1,95 @@ +"""The demo schema the config tests share: two levels of nesting, a list, an +optional, a free dict, a kinds field; one config as a dict; and the helpers +that write it out in each format and read a `ValidationError`'s locations.""" + +import json +from typing import Annotated, Any, Literal + +import yaml +from mace_core.config import BaseConfig, ConfigSection +from pydantic import Field + + +class RadialSection(ConfigSection): + num_bessel: int = 8 + cutoff: float = 5.0 + + +class ModelSection(ConfigSection): + num_interactions: int = 2 + hidden_irreps: str = "128x0e + 128x1o" + radial: RadialSection = RadialSection() + + +class DataSection(ConfigSection): + train_file: str | None = None + valid_fraction: float = 0.1 + energy_key: str = "REF_energy" + heads: list[str] = Field(default_factory=lambda: ["default"]) + + +class StageTwoSection(ConfigSection): + start_epoch: int = 100 + energy_weight: float = 1000.0 + + +class Weighted(ConfigSection): + kind: Literal["weighted"] = "weighted" + stress_weight: float = 0.0 + + +class Huber(ConfigSection): + kind: Literal["huber"] = "huber" + delta: float = 0.01 + + +class DemoConfig(BaseConfig): + name: str = "mace" + seed: int = 123 + default_dtype: str = "float64" + model: ModelSection = ModelSection() + data: DataSection = DataSection() + #: A section left at its defaults unless a file or the CLI writes into it. + stage_two: StageTwoSection = StageTwoSection() + loss: Annotated[Weighted | Huber, Field(discriminator="kind")] = Weighted() + extra: dict[str, Any] = Field(default_factory=dict) + + +#: One config, as a dict. Each format test writes it out and loads it back. +FILE_VALUES = { + "name": "water", + "seed": 7, + "model": {"num_interactions": 4, "radial": {"cutoff": 4.5}}, + "data": {"train_file": "train.xyz", "heads": ["pbe", "r2scan"]}, +} + + +def to_toml(values, prefix=""): + """Enough TOML for a None-free config: scalars and lists share JSON's + literal syntax, nested dicts become `[a.b]` tables after the scalars.""" + lines = [ + f"{k} = {json.dumps(v)}" for k, v in values.items() if not isinstance(v, dict) + ] + for key, value in values.items(): + if isinstance(value, dict): + lines += [f"\n[{prefix}{key}]", to_toml(value, f"{prefix}{key}.")] + return "\n".join(lines) + + +def dump(values, extension): + if extension == ".toml": + return to_toml(values) + if extension == ".json": + return json.dumps(values) + return yaml.safe_dump(values) + + +def write_config(tmp_path, extension=".json", values=FILE_VALUES, name="config"): + path = tmp_path / f"{name}{extension}" + path.write_text(dump(values, extension), encoding="utf-8") + return path + + +def error_locations(excinfo): + """The dotted location of every error a `ValidationError` carries.""" + return [".".join(map(str, error["loc"])) for error in excinfo.value.errors()] diff --git a/packages/mace-core/tests/test_mace_core_config.py b/packages/mace-core/tests/test_mace_core_config.py new file mode 100644 index 000000000..420b44562 --- /dev/null +++ b/packages/mace-core/tests/test_mace_core_config.py @@ -0,0 +1,452 @@ +"""`BaseConfig`: file formats, `from_dict`, unknown keys, and the two +exports' fixed point.""" + +import inspect +import json +import math +import re +import subprocess +import sys +from collections.abc import Set as AbstractSet +from typing import Annotated + +import pytest +import yaml +from mace_core.config import BaseConfig, ConfigError, ConfigSection, read_config_file +from mace_core_demo import ( + FILE_VALUES, + DemoConfig, + RadialSection, + StageTwoSection, + dump, + error_locations, + write_config, +) +from pydantic import BaseModel, ConfigDict, Field, ValidationError, computed_field + +#: A warning the test did not ask for is a failure. +pytestmark = pytest.mark.filterwarnings("error") + + +# --------------------------------------------------------------------------- +# File loading + + +@pytest.mark.parametrize("extension", [".toml", ".yaml", ".yml", ".json"]) +def test_same_config_loads_identically_from_every_format(tmp_path, extension): + config = DemoConfig.load(write_config(tmp_path, extension)) + assert config == DemoConfig.model_validate(FILE_VALUES) + # The file set two fields at depth two; the sibling kept its default. + assert config.model.radial.cutoff == 4.5 + assert config.model.radial.num_bessel == 8 + + +def test_read_config_file_returns_the_parsed_dict(tmp_path): + assert read_config_file(write_config(tmp_path, ".toml")) == FILE_VALUES + + +def test_extension_is_matched_in_any_case(tmp_path): + path = tmp_path / "CONFIG.YAML" + path.write_text("seed: 5\n", encoding="utf-8") + assert DemoConfig.load(path).seed == 5 + + +def test_unknown_extension_is_an_error(tmp_path): + path = tmp_path / "config.ini" + path.write_text("seed = 1", encoding="utf-8") + with pytest.raises(ConfigError, match=r"unknown extension '\.ini'"): + DemoConfig.load(path) + + +def test_empty_file_is_all_defaults(tmp_path): + path = tmp_path / "empty.yaml" + path.write_text("", encoding="utf-8") + assert DemoConfig.load(path) == DemoConfig() + + +def test_comment_only_file_is_all_defaults(tmp_path): + path = tmp_path / "comments.yaml" + path.write_text("# nothing set yet\n", encoding="utf-8") + assert DemoConfig.load(path) == DemoConfig() + + +def test_file_must_be_a_mapping_at_the_top(tmp_path): + path = tmp_path / "list.json" + path.write_text("[1, 2]", encoding="utf-8") + with pytest.raises(ConfigError, match="mapping of keys to values at the top"): + DemoConfig.load(path) + + +def test_missing_file_is_a_config_error(tmp_path): + with pytest.raises(ConfigError, match=r"cannot read config file .*nope\.yaml"): + DemoConfig.load(tmp_path / "nope.yaml") + + +def test_unreadable_file_is_a_config_error(tmp_path): + path = tmp_path / "latin.yaml" + path.write_bytes(b"name: caf\xe9\n") + with pytest.raises(ConfigError, match=r"cannot read config file .*latin\.yaml"): + DemoConfig.load(path) + + +@pytest.mark.parametrize( + ("extension", "text"), + [(".toml", "seed = \n"), (".yaml", "seed: [1\n"), (".json", "{")], +) +def test_malformed_file_is_a_config_error(tmp_path, extension, text): + path = tmp_path / f"broken{extension}" + path.write_text(text, encoding="utf-8") + with pytest.raises(ConfigError, match=r"cannot parse config file .*broken"): + DemoConfig.load(path) + + +# --------------------------------------------------------------------------- +# Precedence: defaults < the file. Nothing else feeds a config. + + +def test_file_overrides_defaults(tmp_path): + config = DemoConfig.load(write_config(tmp_path, ".yaml")) + assert config.model.num_interactions == 4 # from the file + assert config.default_dtype == "float64" # untouched default + + +def test_values_are_handed_to_pydantic_as_given(tmp_path): + values = {"seed": "9", "data": {"train_file": None, "heads": ["a", "b"]}} + config = DemoConfig.load(write_config(tmp_path, ".json", values)) + assert config.seed == 9 # pydantic's lax coercion, not the loader's + assert config.data.train_file is None + assert config.data.heads == ["a", "b"] + + +def test_value_of_the_wrong_type_is_a_validation_error(tmp_path): + with pytest.raises(ValidationError, match="seed"): + DemoConfig.load(write_config(tmp_path, ".yaml", {"seed": "seven"})) + + +# --------------------------------------------------------------------------- +# `from_dict` is `load` for a caller that edits the parsed file first (a +# command line writing its flags); it validates the same way. + + +def test_from_dict_validates_a_parsed_document(tmp_path): + document = read_config_file(write_config(tmp_path, ".toml")) + document["seed"] = 9 + document.setdefault("stage_two", {})["start_epoch"] = 50 + config = DemoConfig.from_dict(document) + assert config.seed == 9 + assert config.model.num_interactions == 4 # the file's values survive + assert config.stage_two == StageTwoSection(start_epoch=50) + assert DemoConfig.load(write_config(tmp_path, ".toml")) == DemoConfig.from_dict( + read_config_file(write_config(tmp_path, ".toml")) + ) + + +def test_from_dict_does_not_write_into_the_document(): + document = {"model": {"radial": {"cutoff": 6.0}}} + DemoConfig.from_dict(document) + assert document == {"model": {"radial": {"cutoff": 6.0}}} + + +# --------------------------------------------------------------------------- +# Unknown keys are pydantic's error, at their dotted location; every error of +# one load is reported together. + + +def test_unknown_top_level_key_in_file(tmp_path): + path = tmp_path / "typo.yaml" + path.write_text("sead: 1\n", encoding="utf-8") + with pytest.raises(ValidationError) as excinfo: + DemoConfig.load(path) + assert error_locations(excinfo) == ["sead"] + assert excinfo.value.errors()[0]["type"] == "extra_forbidden" + + +def test_unknown_nested_key_in_file_names_the_dotted_path(tmp_path): + path = tmp_path / "typo.json" + path.write_text(json.dumps({"model": {"radial": {"cutof": 4.0}}}), encoding="utf-8") + with pytest.raises(ValidationError) as excinfo: + DemoConfig.load(path) + assert error_locations(excinfo) == ["model.radial.cutof"] + assert "model.radial.cutof" in str(excinfo.value) + + +def test_every_error_is_reported_at_once(tmp_path): + path = tmp_path / "typos.yaml" + path.write_text( + "sead: 1\nseed: seven\nmodel:\n num_interaction: 3\n", encoding="utf-8" + ) + with pytest.raises(ValidationError) as excinfo: + DemoConfig.load(path) + assert set(error_locations(excinfo)) == {"sead", "seed", "model.num_interaction"} + + +def test_unknown_key_in_a_document_is_reported_at_its_path(): + documents = { + "nmae": {"nmae": 1}, + "model.num_interaction": {"model": {"num_interaction": 1}}, + "stage_two.start": {"stage_two": {"start": 1}}, + } + for dotted_path, document in documents.items(): + with pytest.raises(ValidationError) as excinfo: + DemoConfig.from_dict(document) + assert error_locations(excinfo) == [dotted_path] + + +def test_every_bad_list_item_is_reported(tmp_path): + class Layers(BaseConfig): + layers: list[RadialSection] = Field(default_factory=list) + + path = write_config(tmp_path, ".json", {"layers": [{"cutof": 1}, {"nb": 2}]}) + with pytest.raises(ValidationError) as excinfo: + Layers.load(path) + assert error_locations(excinfo) == ["layers.0.cutof", "layers.1.nb"] + + +def test_direct_construction_rejects_unknown_keys_too(): + with pytest.raises(ValidationError, match="extra_forbidden"): + DemoConfig(model={"num_interaction": 3}) + + +# --------------------------------------------------------------------------- +# Resolved export + + +def test_resolved_dict_has_every_default_in_declaration_order(tmp_path): + # The file lists keys in the reverse of the schema's order. + path = tmp_path / "reversed.yaml" + path.write_text("seed: 1\nname: x\n", encoding="utf-8") + resolved = DemoConfig.load(path).to_resolved_dict() + assert list(resolved) == [ + "name", + "seed", + "default_dtype", + "model", + "data", + "stage_two", + "loss", + "extra", + ] + assert resolved["stage_two"] == {"start_epoch": 100, "energy_weight": 1000.0} + assert list(resolved["model"]) == ["num_interactions", "hidden_irreps", "radial"] + assert resolved["model"]["radial"] == {"num_bessel": 8, "cutoff": 5.0} + assert resolved["data"]["train_file"] is None + + +def assert_fixed_point(tmp_path, first, extension): + written = tmp_path / f"resolved{extension}" + written.write_text(dump(first, extension), encoding="utf-8") + second = DemoConfig.load(written).to_resolved_dict() + assert second == first + assert json.dumps(second) == json.dumps(first) # order included + + +@pytest.mark.parametrize("extension", [".yaml", ".json"]) +def test_file_to_resolved_to_file_to_resolved_is_a_fixed_point(tmp_path, extension): + values = {**FILE_VALUES, "data": {"train_file": None, "heads": ["pbe"]}} + first = DemoConfig.load(write_config(tmp_path, ".yaml", values)).to_resolved_dict() + assert first["data"]["train_file"] is None # a None is part of what has to survive + assert_fixed_point(tmp_path, first, extension) + + +def test_fixed_point_holds_through_toml_when_nothing_is_none(tmp_path): + # TOML has no null, so the file sets the optional file name; the resolved + # dict then goes through all three formats. + values = {**FILE_VALUES, "stage_two": {"start_epoch": 50}} + first = DemoConfig.load(write_config(tmp_path, ".yaml", values)).to_resolved_dict() + assert "null" not in json.dumps(first) + for extension in (".toml", ".yaml", ".json"): + assert_fixed_point(tmp_path, first, extension) + + +def test_exports_hold_json_native_values_under_a_free_field(): + # Without the JSON mode of the dump the tuple would stay a tuple. + config = DemoConfig.from_dict({"extra": {"window": (1, 2)}}) + assert config.to_resolved_dict()["extra"] == {"window": [1, 2]} + assert config.to_user_dict() == {"extra": {"window": [1, 2]}} + + +def test_inf_survives_the_exports_in_every_format(tmp_path): + # pydantic's JSON mode writes it as null by default under a free field + # (`extra`), which would put a different value into the model metadata. + # Each format spells it its own way: JSON Infinity, YAML .inf, TOML inf. + values = { + **FILE_VALUES, # the file sets the optional file name, so no null + "model": {"radial": {"cutoff": math.inf}}, + "extra": {"limit": math.inf}, + } + config = DemoConfig.load(write_config(tmp_path, ".yaml", values)) + resolved = config.to_resolved_dict() + assert resolved["model"]["radial"]["cutoff"] == math.inf + assert config.to_user_dict()["extra"] == {"limit": math.inf} + text = json.dumps(resolved) + assert "null" not in text and "Infinity" in text + for extension, body in [ + (".json", text), + (".yaml", yaml.safe_dump(resolved)), + (".toml", "[model.radial]\ncutoff = inf\n"), + ]: + path = tmp_path / f"special{extension}" + path.write_text(body, encoding="utf-8") + second = DemoConfig.load(path).to_resolved_dict() + assert second["model"]["radial"]["cutoff"] == math.inf, extension + assert DemoConfig.load(tmp_path / "special.json").to_resolved_dict() == resolved + + +class LenientSection(BaseModel): + cutoff: float = 5.0 + + +def test_field_shapes_the_contract_cannot_keep_are_rejected_at_class_definition(): + # Each shape would break a guarantee: set order varies with the hash + # seed; aliases, excluded and computed fields do not validate back; a + # lenient section would swallow typos. + shapes = { + r"tags is typed as a set.*Use a list": ("tags", list[set[str]]), + r"names is typed as a set.*Use a list": ("names", AbstractSet[str]), + r"num has an alias": ("num", Annotated[int, Field(alias="n")]), + r"vnum has an alias": ("vnum", Annotated[int, Field(validation_alias="n")]), + r"snum has an alias": ("snum", Annotated[int, Field(serialization_alias="n")]), + r"hidden is excluded from dumps": ( + "hidden", + Annotated[int, Field(exclude=True)], + ), + r"radial holds LenientSection, which is not a ConfigSection": ( + "radial", + LenientSection | None, + ), + } + for message, (name, annotation) in shapes.items(): + with pytest.raises(TypeError, match=message): + type("Bad", (ConfigSection,), {"__annotations__": {name: annotation}}) + + with pytest.raises(TypeError, match=r"double is a computed field"): + + class Computed(ConfigSection): + seed: int = 1 + + @computed_field + def double(self) -> int: + return 2 * self.seed + + +def test_a_section_cannot_reopen_extra_or_thaw(): + with pytest.raises(TypeError, match=r"Loose sets extra='allow'; a section keeps"): + + class Loose(ConfigSection): + model_config = ConfigDict(extra="allow") + seed: int = 1 + + with pytest.raises(TypeError, match=r"Thawed sets frozen=False; a section stays"): + + class Thawed(ConfigSection): + model_config = ConfigDict(frozen=False) + seed: int = 1 + + +def test_a_validated_config_is_immutable(tmp_path): + # What validation produced is what the run uses: a change is a new + # validation of an edited dict, never an assignment behind it. + config = DemoConfig.load(write_config(tmp_path, ".yaml")) + with pytest.raises(ValidationError, match="frozen"): + config.seed = 8 # ty: ignore[invalid-assignment] # the point of the test + with pytest.raises(ValidationError, match="frozen"): + config.model.radial.cutoff = 6.0 # ty: ignore[invalid-assignment] + assert config.seed == 7 and config.model.radial.cutoff == 4.5 + changed = DemoConfig.from_dict({**config.to_resolved_dict(), "seed": 8}) + assert changed.seed == 8 and changed.model == config.model + + +# A class that names a class defined below it is incomplete at definition: +# pydantic keeps the name, so the check cannot see through the field. The +# first `load` resolves the name and checks the class then. + + +class Forward(ConfigSection): + later: "Later" = Field(default_factory=lambda: Later()) + + +class Later(ConfigSection): + x: int = 1 + + +class ForwardConfig(BaseConfig): + forward: Forward = Field(default_factory=Forward) + + +class Leaking(ConfigSection): + plain: "PlainLater" = Field(default_factory=lambda: PlainLater()) + + +class PlainLater(BaseModel): # not a ConfigSection: it would swallow a typo + a: int = 1 + + +class LeakingConfig(BaseConfig): + leaking: Leaking = Field(default_factory=Leaking) + + +def test_a_forward_reference_is_resolved_and_checked_on_load(tmp_path): + empty = write_config(tmp_path, ".yaml", {}) + assert ForwardConfig.load(empty).forward.later == Later() + config = ForwardConfig.from_dict({"forward": {"later": {"x": 2}}}) + assert config.forward.later == Later(x=2) + path = tmp_path / "typo.json" + path.write_text(json.dumps({"forward": {"later": {"x": 2, "typo": 1}}})) + with pytest.raises(ValidationError) as excinfo: + ForwardConfig.load(path) + assert error_locations(excinfo) == ["forward.later.typo"] + + +def test_a_lenient_section_behind_a_forward_reference_is_rejected(tmp_path): + with pytest.raises( + TypeError, match=r"Leaking.plain holds PlainLater, which is not a ConfigSection" + ): + LeakingConfig.load(write_config(tmp_path, ".yaml", {})) + + +def test_user_dict_holds_only_what_was_set(tmp_path): + config = DemoConfig.load(write_config(tmp_path, ".json")) + assert config.to_user_dict() == FILE_VALUES + assert config.to_user_dict() is not FILE_VALUES + + +# --------------------------------------------------------------------------- +# Nothing but the file feeds a config. + + +def test_environment_variables_are_ignored(tmp_path, monkeypatch): + monkeypatch.setenv("NAME", "from-the-environment") + monkeypatch.setenv("SEED", "99") + config = DemoConfig.load(write_config(tmp_path, ".yaml", {})) + assert config.name == "mace" + assert config.seed == 123 + + +@pytest.mark.parametrize("key", ["SEED", "Seed", "_env_file", "_cli_parse_args"]) +def test_root_keys_are_validated_like_any_section(tmp_path, key): + # Neither case variants nor BaseSettings-style private constructor + # options are special at the top level: unknown is unknown. + path = tmp_path / "root.json" + path.write_text(json.dumps({key: 1}), encoding="utf-8") + with pytest.raises(ValidationError) as excinfo: + DemoConfig.load(path) + assert error_locations(excinfo) == [key] + + +def test_config_module_imports_neither_torch_nor_jax(): + """In a fresh interpreter, so another test's imports cannot mask a leak.""" + code = ( + "import sys, mace_core.config; " + "leaked = {'torch', 'jax', 'e3nn'} & set(sys.modules); " + "assert not leaked, leaked" + ) + subprocess.run([sys.executable, "-c", code], check=True) + + +def test_the_base_does_not_import_the_command_line_module(): + # The package init re-exports both, so `sys.modules` cannot tell; the + # source can: `cli` imports `ConfigError` from `base`, never the reverse. + import mace_core.config.base as base_module + + source = inspect.getsource(base_module) + assert not re.search(r"^\s*(from|import) .*\bcli\b", source, re.MULTILINE) diff --git a/packages/mace-core/tests/test_mace_core_config_cli.py b/packages/mace-core/tests/test_mace_core_config_cli.py new file mode 100644 index 000000000..18a07f5d4 --- /dev/null +++ b/packages/mace-core/tests/test_mace_core_config_cli.py @@ -0,0 +1,195 @@ +"""`mace_core.config.cli`: the `--a.b value` grammar to a mapping of dotted +paths, and that mapping written into a parsed file before one validation. The +config base knows nothing of either; a command line composes them with +`read_config_file` and `from_dict`.""" + +import json + +import pytest +from mace_core.config import ( + ConfigError, + apply_overrides, + parse_overrides, + read_config_file, +) +from mace_core_demo import FILE_VALUES, DemoConfig, Huber, error_locations, write_config +from pydantic import ValidationError + +#: A warning the test did not ask for is a failure. +pytestmark = pytest.mark.filterwarnings("error") + +HEADS_JSON = '["a", "b"]' + + +def load(tmp_path, argv, values=FILE_VALUES, root=DemoConfig): + """What a command line does with its config file and the tokens after it.""" + document = read_config_file(write_config(tmp_path, ".json", values)) + return root.from_dict(apply_overrides(document, parse_overrides(argv))) + + +# --------------------------------------------------------------------------- +# parse_overrides: tokens to a mapping of dotted path to value. + + +def test_both_forms_and_the_value_types(): + argv = ["--seed=9", "--data.train_file", "null", "--data.heads", HEADS_JSON] + assert parse_overrides(argv) == { + "seed": "9", # a string: pydantic coerces it, the grammar does not + "data.train_file": None, + "data.heads": ["a", "b"], + } + assert parse_overrides(["--model", json.dumps({"num_interactions": 3})]) == { + "model": {"num_interactions": 3} + } + assert parse_overrides([]) == {} + + +def test_paths_keep_their_order_and_a_repeated_path_moves_to_where_it_is_last(): + overrides = parse_overrides(["--b", "1", "--a", "2", "--b", "3"]) + assert list(overrides.items()) == [("a", "2"), ("b", "3")] + # So a later write under a parent the same command line replaced survives. + argv = ["--m.r.c", "4", "--m.r", "null", "--m.r.c", "5"] + assert apply_overrides({}, parse_overrides(argv)) == {"m": {"r": {"c": "5"}}} + + +def test_an_empty_inline_value_does_not_hide_the_next_option(): + assert parse_overrides(["--name=", "--seed", "5"]) == {"name": "", "seed": "5"} + + +def test_a_value_starting_with_dashes_works_in_both_forms(): + assert parse_overrides(["--name=--odd"]) == {"name": "--odd"} + assert parse_overrides(["--name", "--odd"]) == {"name": "--odd"} + + +def test_a_path_missing_its_value_is_a_config_error(): + with pytest.raises(ConfigError, match="override --seed is missing its value"): + parse_overrides(["--seed"]) + + +def test_a_value_that_is_not_valid_json_is_a_config_error(): + with pytest.raises(ConfigError, match="override --model is not valid JSON"): + parse_overrides(["--model", "{oops"]) + + +@pytest.mark.parametrize("token", ["--", "--=5", "seed=5", "-s"]) +def test_a_token_that_is_not_a_dotted_option_is_a_config_error(token): + with pytest.raises(ConfigError, match=rf"unknown config option '{token}'"): + parse_overrides([token, "--seed", "5"]) + + +def test_tokens_given_as_one_string_are_refused(): + with pytest.raises(TypeError, match="tokens is a string"): + parse_overrides("--seed 5") + + +# --------------------------------------------------------------------------- +# apply_overrides: a dotted path is written into a copy of the parsed file. + + +def test_a_top_level_key_is_written(): + assert apply_overrides({}, {"seed": 9}) == {"seed": 9} + assert apply_overrides({"name": "x"}, {"seed": 9}) == {"name": "x", "seed": 9} + + +def test_a_nested_path_writes_into_the_section_the_file_set(): + document = {"model": {"num_interactions": 4, "radial": {"cutoff": 4.5}}} + assert apply_overrides(document, {"model.radial.cutoff": 6.0}) == { + "model": {"num_interactions": 4, "radial": {"cutoff": 6.0}} + } + + +def test_the_sections_a_path_passes_through_are_created(): + assert apply_overrides({}, {"model.radial.cutoff": 6.0, "stage_two.x": 1}) == { + "model": {"radial": {"cutoff": 6.0}}, + "stage_two": {"x": 1}, + } + + +@pytest.mark.parametrize("parent", [5, [1, 2], None], ids=["scalar", "list", "null"]) +def test_a_parent_that_is_not_a_mapping_is_replaced(parent): + assert apply_overrides({"a": parent}, {"a.b": 2}) == {"a": {"b": 2}} + + +def test_a_mapping_value_merges_into_a_mapping_but_a_list_replaces(): + document = {"by_name": {"pbe": {"cutoff": 4.0}}, "heads": ["a", "b"]} + merged = apply_overrides( + document, {"by_name": {"r2scan": {"cutoff": 6.0}}, "heads": ["c"]} + ) + assert merged == { + "by_name": {"pbe": {"cutoff": 4.0}, "r2scan": {"cutoff": 6.0}}, + "heads": ["c"], + } + # Deeper too, and a key of the JSON value is one key even with a dot in it. + assert apply_overrides( + {"a": {"b": {"c": 1, "d": 2}}}, {"a": {"b": {"c": 3}, "x.y": 4}} + ) == {"a": {"b": {"c": 3, "d": 2}, "x.y": 4}} + # A mapping over a scalar, or a scalar over a mapping, replaces. + assert apply_overrides({"a": 1}, {"a": {"b": 2}}) == {"a": {"b": 2}} + assert apply_overrides({"a": {"b": 2}}, {"a": 1}) == {"a": 1} + + +def test_paths_apply_in_order_on_top_of_the_file(): + # A dotted value followed by the whole section keeps both. + document = {"stage_two": {"energy_weight": 5.0}} + argv = ["--stage_two.start_epoch", "5", "--model.radial.cutoff", "4"] + argv += ["--model", json.dumps({"num_interactions": 3})] + assert apply_overrides(document, parse_overrides(argv)) == { + "stage_two": {"energy_weight": 5.0, "start_epoch": "5"}, + "model": {"radial": {"cutoff": "4"}, "num_interactions": 3}, + } + + +def test_neither_the_document_nor_the_overrides_are_written_into(): + document = {"extra": {"a": {"x": {}}}} + overrides = {"extra.a": {"x": {"y": 1}}, "extra.a.x.z": [1]} + copy = apply_overrides(document, overrides) + assert copy == {"extra": {"a": {"x": {"y": 1, "z": [1]}}}} + assert document == {"extra": {"a": {"x": {}}}} + assert overrides == {"extra.a": {"x": {"y": 1}}, "extra.a.x.z": [1]} + copy["extra"]["a"]["x"]["z"].append(2) # the value was copied too + assert overrides["extra.a.x.z"] == [1] + + +def test_a_yaml_anchor_does_not_share_an_override(tmp_path): + path = tmp_path / "anchors.yaml" + path.write_text("a: &empty {}\nb: *empty\n", encoding="utf-8") + document = read_config_file(path) + assert document["a"] is document["b"] # what the parser hands over + assert apply_overrides(document, {"a.x": 1}) == {"a": {"x": 1}, "b": {}} + + +# --------------------------------------------------------------------------- +# Composed with the base: the override beats the file, which beats the +# defaults, and every schema error is pydantic's at the path the override named. + + +def test_an_override_beats_the_file_which_beats_the_defaults(tmp_path): + config = load(tmp_path, ["--model.num_interactions", "3"]) + assert config.model.num_interactions == 3 # the override beats the file's 4 + assert config.model.radial.cutoff == 4.5 # the file's other values survive + assert config.model.radial.num_bessel == 8 # defaults fill the rest + assert load(tmp_path, []) == DemoConfig.from_dict(FILE_VALUES) + # A section the file left alone is created on the way and validated like + # any other; the value stays the string the grammar handed over. + assert load(tmp_path, ["--stage_two.start_epoch", "50"]).stage_two.start_epoch == 50 + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, ["--stage_two.start", "50"]) + assert error_locations(excinfo) == ["stage_two.start"] + + +def test_an_override_writes_into_a_kinds_field(tmp_path): + config = load(tmp_path, ["--loss.delta", "2"], {"loss": {"kind": "huber"}}) + assert config.loss == Huber(kind="huber", delta=2.0) + config = load(tmp_path, ["--loss", json.dumps({"kind": "huber", "delta": 2})], {}) + assert config.loss == Huber(kind="huber", delta=2.0) + assert config.to_user_dict() == {"loss": {"kind": "huber", "delta": 2.0}} + + +def test_an_override_that_changes_the_kind_leaves_the_old_settings_to_pydantic( + tmp_path, +): + # The file's `delta` stays in the dict and is an unknown key of the new kind. + huber = {"loss": {"kind": "huber", "delta": 0.5}} + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, ["--loss.kind", "weighted"], huber) + assert error_locations(excinfo) == ["loss.weighted.delta"] diff --git a/packages/mace-core/tests/test_mace_core_config_kinds.py b/packages/mace-core/tests/test_mace_core_config_kinds.py new file mode 100644 index 000000000..f23029bf5 --- /dev/null +++ b/packages/mace-core/tests/test_mace_core_config_kinds.py @@ -0,0 +1,264 @@ +"""A field of several kinds of section (a discriminated union on `kind`) is an +ordinary pydantic feature under the one-file contract: a config +writes it in the tagged form (`loss: {kind: huber, delta: 0.1}`), pydantic +picks the variant by the tag and reports every error under it +(`loss.huber.delta`), and both exports write the tag back so the dict loads to +the same config. Nothing in the loader knows about kinds.""" + +import json +from typing import Annotated, Literal + +import pytest +from mace_core.config import BaseConfig, ConfigSection +from pydantic import Field, ValidationError + +#: A warning the test did not ask for is a failure. +pytestmark = pytest.mark.filterwarnings("error") + +# --------------------------------------------------------------------------- +# The schema: a loss of three kinds, one of which holds a field of two kinds; +# the same choice with "none of them" as a fourth kind and the default, and +# inside the values of a dict and the items of a list. + + +class SubX(ConfigSection): + kind: Literal["x"] = "x" + a: float = 1.0 + + +class SubY(ConfigSection): + kind: Literal["y"] = "y" + b: float = 2.0 + + +class Weighted(ConfigSection): + kind: Literal["weighted"] = "weighted" + stress_weight: float = 0.0 + + +class Huber(ConfigSection): + kind: Literal["huber"] = "huber" + delta: float = 0.01 + sub: Annotated[SubX | SubY, Field(discriminator="kind")] = SubX() + + +class Universal(ConfigSection): + kind: Literal["universal"] = "universal" + huber_delta: float = 0.01 + + +Choice = Annotated[Weighted | Huber | Universal, Field(discriminator="kind")] + + +class NoChoice(ConfigSection): + """None of the other kinds is a kind too.""" + + kind: Literal["none"] = "none" + + +OptChoice = Annotated[ + Weighted | Huber | Universal | NoChoice, Field(discriminator="kind") +] + + +class HeadSection(ConfigSection): + loss: Choice = Weighted() + + +class LossConfig(BaseConfig): + energy_weight: float = 1.0 + choice: Choice = Weighted() + opt: OptChoice = NoChoice() + per_head: dict[str, HeadSection] = Field(default_factory=dict) + layers: list[HeadSection] = Field(default_factory=list) + + +class HuberRequired(ConfigSection): + kind: Literal["huber"] = "huber" + path: str + delta: float = 0.01 + + +class RequiredConfig(BaseConfig): + choice: Annotated[Weighted | HuberRequired, Field(discriminator="kind")] = ( + Weighted() + ) + + +def huber(delta=0.01, sub=SubX, **sub_fields): + return lambda c: ( + isinstance(c.choice, Huber) + and c.choice.delta == delta + and isinstance(c.choice.sub, sub) + and all(getattr(c.choice.sub, k) == v for k, v in sub_fields.items()) + ) + + +def load(tmp_path, root, file_values): + path = tmp_path / "config.json" + path.write_text(json.dumps(file_values)) + return root.load(path) + + +def error_locations(excinfo): + """The dotted location of every error a `ValidationError` carries.""" + return [".".join(map(str, error["loc"])) for error in excinfo.value.errors()] + + +# --------------------------------------------------------------------------- +# Loading: the tagged form, and pydantic's errors under the tag. + + +@pytest.mark.parametrize("extension", [".toml", ".yaml", ".json"]) +def test_kinds_load_from_every_file_format(tmp_path, extension): + text = { + ".toml": '[choice]\nkind = "huber"\ndelta = 0.5\n[choice.sub]\nkind = "y"\n' + 'b = 3\n[opt]\nkind = "universal"\n', + ".yaml": "choice:\n kind: huber\n delta: 0.5\n sub:\n kind: y\n" + " b: 3\nopt:\n kind: universal\n", + ".json": json.dumps( + { + "choice": {"kind": "huber", "delta": 0.5, "sub": {"kind": "y", "b": 3}}, + "opt": {"kind": "universal"}, + } + ), + }[extension] + path = tmp_path / f"config{extension}" + path.write_text(text) + config = LossConfig.load(path) + assert huber(0.5, SubY, b=3.0)(config) and isinstance(config.opt, Universal) + + +@pytest.mark.parametrize("value", [{"delta": 0.5}, {}], ids=["settings", "empty"]) +def test_a_kinds_field_without_its_tag_is_a_validation_error(tmp_path, value): + # Even though the field has a default: pydantic needs the tag to pick. + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, LossConfig, {"choice": value}) + assert error_locations(excinfo) == ["choice"] + assert excinfo.value.errors()[0]["type"] == "union_tag_not_found" + + +def test_a_wrong_tag_is_a_validation_error_listing_the_kinds(tmp_path): + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, LossConfig, {"choice": {"kind": "hubr"}}) + assert error_locations(excinfo) == ["choice"] + (error,) = excinfo.value.errors() + assert error["type"] == "union_tag_invalid" + assert all(kind in error["msg"] for kind in ("weighted", "huber", "universal")) + + +@pytest.mark.parametrize( + ("file_values", "location"), + [ + ( + {"choice": {"kind": "huber", "stress_weight": 1}}, + "choice.huber.stress_weight", + ), + ( + {"choice": {"kind": "huber", "sub": {"kind": "y", "a": 1}}}, + "choice.huber.sub.y.a", + ), + ( + {"per_head": {"h": {"loss": {"kind": "huber", "stress_weight": 1}}}}, + "per_head.h.loss.huber.stress_weight", + ), + ( + {"layers": [{"loss": {"kind": "huber", "stress_weight": 1}}]}, + "layers.0.loss.huber.stress_weight", + ), + ], + ids=["top", "nested", "dict value", "list item"], +) +def test_unknown_keys_under_a_kind_are_reported_under_the_tag( + tmp_path, file_values, location +): + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, LossConfig, file_values) + assert error_locations(excinfo) == [location] + assert excinfo.value.errors()[0]["type"] == "extra_forbidden" + + +def test_a_required_key_of_the_kind_is_reported_under_the_tag(tmp_path): + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, RequiredConfig, {"choice": {"kind": "huber"}}) + assert error_locations(excinfo) == ["choice.huber.path"] + assert excinfo.value.errors()[0]["type"] == "missing" + config = load(tmp_path, RequiredConfig, {"choice": {"kind": "huber", "path": "p"}}) + assert isinstance(config.choice, HuberRequired) and config.choice.path == "p" + + +# --------------------------------------------------------------------------- +# The dicts: the tag is written back, a fixed point, and only what was set. + + +@pytest.mark.parametrize( + "file_values", + [ + {}, + {"choice": {"kind": "huber"}}, + { + "choice": {"kind": "huber", "sub": {"kind": "y", "b": 7}}, + "opt": {"kind": "universal"}, + }, + { + "per_head": {"h": {"loss": {"kind": "huber"}}}, + "layers": [{"loss": {"kind": "universal"}}], + }, + ], + ids=["defaults", "huber", "nested", "collections"], +) +def test_resolved_dict_is_a_fixed_point_and_writes_the_tag(tmp_path, file_values): + config = load(tmp_path, LossConfig, file_values) + resolved = config.to_resolved_dict() + assert resolved["choice"]["kind"] == config.choice.kind + reloaded = load(tmp_path, LossConfig, resolved) + assert reloaded.to_resolved_dict() == resolved + assert reloaded == config + + +def test_resolved_dict_of_the_defaults(): + assert LossConfig().to_resolved_dict() == { + "energy_weight": 1.0, + "choice": {"kind": "weighted", "stress_weight": 0.0}, + "opt": {"kind": "none"}, + "per_head": {}, + "layers": [], + } + + +def test_user_dict_holds_the_tag_and_only_what_was_set(tmp_path): + config = load(tmp_path, LossConfig, {"choice": {"kind": "huber", "delta": 2.0}}) + assert config.to_user_dict() == {"choice": {"kind": "huber", "delta": 2.0}} + assert load(tmp_path, LossConfig, config.to_user_dict()) == config + # The documented caveat: built in code without its tag, a variant exports + # without `kind` and the dict does not load back; pass the tag, or use + # `to_resolved_dict`. + assert LossConfig(choice=Huber(delta=2)).to_user_dict() == { + "choice": {"delta": 2.0} + } + tagged = LossConfig(choice=Huber(kind="huber", delta=2)) + assert tagged.to_user_dict() == {"choice": {"kind": "huber", "delta": 2.0}} + assert LossConfig(choice=Huber(delta=2)).to_resolved_dict()["choice"]["kind"] == ( + "huber" + ) + + +def test_json_schema_is_produced_in_both_modes(): + for mode in ("validation", "serialization"): + schema = LossConfig.model_json_schema(mode=mode) + assert set(schema["properties"]) == set(LossConfig.model_fields), mode + + +def test_a_variant_class_as_a_plain_field_keeps_kind_as_a_key(tmp_path): + class Reuse(BaseConfig): + direct: Huber = Huber() + + resolved = { + "direct": {"kind": "huber", "delta": 0.01, "sub": {"kind": "x", "a": 1.0}} + } + assert Reuse().to_resolved_dict() == resolved + assert load(tmp_path, Reuse, resolved).to_resolved_dict() == resolved + config = load(tmp_path, Reuse, {"direct": {"kind": "huber", "sub": {"kind": "y"}}}) + assert config.direct == Huber(sub=SubY()) + with pytest.raises(ValidationError, match=r"direct\.kind"): + load(tmp_path, Reuse, {"direct": {"kind": "weighted"}}) diff --git a/packages/mace-core/tests/test_mace_core_metadata.py b/packages/mace-core/tests/test_mace_core_metadata.py new file mode 100644 index 000000000..a971f9bd8 --- /dev/null +++ b/packages/mace-core/tests/test_mace_core_metadata.py @@ -0,0 +1,191 @@ +"""`ModelMetadata`: JSON round trip, schema versioning, citation rendering.""" + +import json +import subprocess +import sys + +import pytest +from mace_core.config import BaseConfig, ConfigSection +from mace_core.metadata import ( + SCHEMA_VERSION, + Citation, + ConfigRecord, + DataSourceSummary, + DataSummary, + HeadSummary, + MetadataSchemaError, + ModelMetadata, + Provenance, + format_citations, +) +from pydantic import ValidationError + +MACE_PAPER = Citation( + title="MACE: Higher Order Equivariant Message Passing Neural Networks " + "for Fast and Accurate Force Fields", + authors=["I. Batatia", "D. P. Kovacs", "G. N. C. Simm", "C. Ortner", "G. Csanyi"], + venue="Advances in Neural Information Processing Systems", + year=2022, + url="https://arxiv.org/abs/2206.07697", +) + + +def full_record() -> ModelMetadata: + """Every field set, so the round trip is tested on all of them.""" + return ModelMetadata( + config=ConfigRecord( + user={"model": {"num_interactions": 3}}, + resolved={"name": "mace", "model": {"num_interactions": 3, "cutoff": 5.0}}, + ), + provenance=Provenance( + versions={"mace-core": "1.0.2", "mace-torch": "1.1.0"}, + git_commit="a" * 40, + ), + data=DataSummary( + sources=[ + DataSourceSummary( + name="water", + num_configurations=1200, + num_atoms=64_000, + elements=[1, 8], + reference_keys=["pbe_energy", "pbe_forces"], + ), + DataSourceSummary(name="ice", elements=[1, 8]), + ] + ), + heads={ + "pbe": HeadSummary(sources=["water", "ice"]), + "r2scan": HeadSummary(sources=["ice"]), + }, + doi="10.5281/zenodo.0000000", + citations=[MACE_PAPER, Citation(title="A dataset paper", doi="10.1000/xyz")], + notes="Trained for the round-trip test.", + ) + + +# --------------------------------------------------------------------------- +# Round trip and schema version + + +def test_json_round_trip_is_lossless(): + record = full_record() + assert ModelMetadata.from_json(record.to_json()) == record + + +def test_minimal_record_round_trips_too(): + record = ModelMetadata( + config=ConfigRecord(), provenance=Provenance(versions={"mace-core": "0.0.0"}) + ) + assert ModelMetadata.from_json(record.to_json()) == record + assert record.heads == {} + + +def test_heads_must_name_summarised_sources(): + record = full_record() + record.heads["pbe"].sources.append("vapour") + with pytest.raises(ValidationError, match=r"heads\.pbe\.sources names 'vapour'"): + ModelMetadata.model_validate(record.model_dump()) + record = full_record() + record.data.sources.append(DataSourceSummary(name="ice")) + with pytest.raises(ValidationError, match="names a source twice"): + ModelMetadata.model_validate(record.model_dump()) + + +def test_config_and_provenance_are_mandatory(): + with pytest.raises(ValidationError, match="config"): + ModelMetadata.model_validate({"provenance": {"versions": {"mace-core": "0"}}}) + + +def test_lossy_value_is_refused_rather_than_stored(): + record = full_record() + record.config.user["shape"] = (2, 3) # JSON brings it back as a list + with pytest.raises(MetadataSchemaError, match="does not survive a JSON round trip"): + record.to_json() + + +def test_schema_version_is_written(): + assert json.loads(full_record().to_json())["schema_version"] == SCHEMA_VERSION + + +@pytest.mark.parametrize( + "read", + [ModelMetadata.model_validate, lambda d: ModelMetadata.from_json(json.dumps(d))], + ids=["model_validate", "from_json"], +) +def test_another_schema_version_is_rejected_on_every_route(read): + document = json.loads(full_record().to_json()) + document["schema_version"] = SCHEMA_VERSION + 1 + document["a_field_of_the_next_version"] = 1 + with pytest.raises(ValidationError) as excinfo: + read(document) + assert excinfo.value.error_count() == 1 # the version, not the unknown field + message = str(excinfo.value) + assert f"schema_version {SCHEMA_VERSION + 1}" in message + assert f"reads schema_version {SCHEMA_VERSION}" in message + assert "upgrade" in message + + +def test_infinity_survives_and_nan_is_refused(): + # pydantic's default writes inf/nan as null, which would silently turn a + # config value into a different one. NaN is never equal to itself, so it + # cannot pass the round-trip check; a config value that is NaN is a bug + # upstream, not something to store. + record = full_record() + record.config.resolved["model"]["cutoff"] = float("inf") + back = ModelMetadata.from_json(record.to_json()) + assert back.config.resolved["model"]["cutoff"] == float("inf") + record.config.resolved["cutoff"] = float("nan") + with pytest.raises(MetadataSchemaError, match="does not survive"): + record.to_json() + + +def test_unknown_fields_are_rejected(): + with pytest.raises(ValidationError, match="extra_forbidden"): + ModelMetadata.model_validate( + {"config": {}, "provenance": {"versions": {"mace-core": "0"}}, "note": "x"} + ) + + +def test_config_record_is_built_from_a_config(): + class Section(ConfigSection): + cutoff: float = 5.0 + + class Config(BaseConfig): + seed: int = 1 + model: Section = Section() + + record = ConfigRecord.from_config(Config.model_validate({"model": {"cutoff": 4.0}})) + assert record.user == {"model": {"cutoff": 4.0}} + assert record.resolved == {"seed": 1, "model": {"cutoff": 4.0}} + # The embedded form is the fixed point: resolving it again changes nothing. + assert Config.model_validate(record.resolved).to_resolved_dict() == record.resolved + + +# --------------------------------------------------------------------------- +# Citations + + +def test_citations_render_to_a_numbered_block(): + block = format_citations(full_record().citations) + assert block.splitlines() == [ + "[1] I. Batatia, D. P. Kovacs, G. N. C. Simm, C. Ortner, G. Csanyi. " + "MACE: Higher Order Equivariant Message Passing Neural Networks for " + "Fast and Accurate Force Fields. " + "Advances in Neural Information Processing Systems (2022). " + "https://arxiv.org/abs/2206.07697", + "[2] A dataset paper. https://doi.org/10.1000/xyz", + ] + + +def test_no_citations_render_to_nothing(): + assert format_citations([]) == "" + + +def test_metadata_module_imports_neither_torch_nor_jax(): + """In a fresh interpreter, so another test's imports cannot mask a leak.""" + code = ( + "import sys, mace_core.metadata; " + "leaked = {'torch', 'jax', 'e3nn'} & set(sys.modules); " + "assert not leaked, leaked" + ) + subprocess.run([sys.executable, "-c", code], check=True)