Skip to content
Open

core-2 #1733

Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
14 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 14 additions & 1 deletion packages/mace-core/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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()`.
10 changes: 9 additions & 1 deletion packages/mace-core/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
15 changes: 13 additions & 2 deletions packages/mace-core/src/mace_core/__init__.py
Original file line number Diff line number Diff line change
@@ -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.
Expand Down
23 changes: 23 additions & 0 deletions packages/mace-core/src/mace_core/config/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
]
188 changes: 188 additions & 0 deletions packages/mace-core/src/mace_core/config/base.py
Original file line number Diff line number Diff line change
@@ -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
Comment thread
steffen-wedig marked this conversation as resolved.
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)
104 changes: 104 additions & 0 deletions packages/mace-core/src/mace_core/config/cli.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading