diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 6f21fd1e..763ffe3a 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -23,6 +23,11 @@ Added `__). - New ``ActionFail`` for arguments that should fail parsing with a given error message (`#759 `__). +- Experimental ``omegaconf+`` parser mode that supports variable interpolation + and resolving across configs and command line arguments. Depending on + community feedback, in v5.0.0 this new mode could replace the current + ``omegaconf`` mode, introducing a breaking change (`#765 + `__). Fixed ^^^^^ @@ -34,6 +39,8 @@ Fixed - Environment variable names not shown in help for positional arguments when ``default_env`` is true (`#763 `__). +- ``parse_object`` not parsing correctly configs (`#765 + `__). Changed ^^^^^^^ diff --git a/DOCUMENTATION.rst b/DOCUMENTATION.rst index 04ef8310..8b6c8f0d 100644 --- a/DOCUMENTATION.rst +++ b/DOCUMENTATION.rst @@ -2398,8 +2398,8 @@ instantiates :class:`Data` first, then use the ``num_classes`` attribute to instantiate :class:`Model`. -Variable interpolation -====================== +OmegaConf variable interpolation +================================ One of the possible reasons to add a parser mode (see :ref:`custom-loaders`) can be to have support for variable interpolation in yaml files. Any library could @@ -2463,10 +2463,22 @@ This yaml could be parsed as follows: .. note:: - The ``parser_mode='omegaconf'`` provides support for `OmegaConf's - `__ variable interpolation in a single - yaml file. It is not possible to do interpolation across multiple yaml files - or in an isolated individual command line argument. + The ``parser_mode="omegaconf"`` provides support for `OmegaConf's resolvers + `__ in a + single YAML file. It is not possible to do interpolation across multiple + YAML files or in an isolated individual command line argument. + +Experimental ``omegaconf+`` mode +-------------------------------- + +There is a new experimental ``omegaconf+`` parser mode that doesn't suffer from +the limitations of ``omegaconf`` mentioned above. Instead of applying OmegaConf +resolvers for each YAML config, the resolving is applied once at the end of +parsing. Because of this, in nested subconfigs, references to config nodes need +to be relative to work correctly. + +Depending on feedback from the community, this mode might become the default +``omegaconf`` mode in v5.0.0. .. _environment-variables: diff --git a/jsonargparse/_common.py b/jsonargparse/_common.py index 789550e0..c0c12178 100644 --- a/jsonargparse/_common.py +++ b/jsonargparse/_common.py @@ -63,6 +63,7 @@ def __call__(self, class_type: Type[ClassType], *args, **kwargs) -> ClassType: class_instantiators: ContextVar[Optional[InstantiatorsDictType]] = ContextVar("class_instantiators", default=None) nested_links: ContextVar[List[dict]] = ContextVar("nested_links", default=[]) applied_instantiation_links: ContextVar[Optional[set]] = ContextVar("applied_instantiation_links", default=None) +path_dump_preserve_relative: ContextVar[bool] = ContextVar("path_dump_preserve_relative", default=False) parser_context_vars = { @@ -74,6 +75,7 @@ def __call__(self, class_type: Type[ClassType], *args, **kwargs) -> ClassType: "class_instantiators": class_instantiators, "nested_links": nested_links, "applied_instantiation_links": applied_instantiation_links, + "path_dump_preserve_relative": path_dump_preserve_relative, } diff --git a/jsonargparse/_core.py b/jsonargparse/_core.py index fab10cc9..558a5e96 100644 --- a/jsonargparse/_core.py +++ b/jsonargparse/_core.py @@ -82,6 +82,7 @@ fsspec_support, import_fsspec, import_jsonnet, + omegaconf_apply, pyyaml_available, ) from ._parameter_resolvers import UnknownDefault @@ -378,9 +379,12 @@ def _parse_common( with parser_context(lenient_check=True): ActionTypeHint.add_sub_defaults(self, cfg) - _ActionPrintConfig.print_config_if_requested(self, cfg) - with parser_context(parent_parser=self): + if not lenient_check.get() and self.parser_mode == "omegaconf+": + cfg = omegaconf_apply(self, cfg) + + _ActionPrintConfig.print_config_if_requested(self, cfg) + try: ActionLink.apply_parsing_links(self, cfg) except Exception as ex: @@ -1401,6 +1405,13 @@ def _apply_actions( if isinstance(action, _ActionConfigLoad): config_keys.add(action_dest) keys.append(action_dest) + elif isinstance(action, ActionConfigFile): + if isinstance(value, str): + cfg.pop(action_dest) + preserve = Namespace({k: cfg[k] for k in keys[num:]}) + ActionConfigFile.apply_config(self, cfg, action_dest, value) + cfg.update(preserve) + continue elif getattr(action, "jsonnet_ext_vars", False): prev_cfg[action_dest] = value cfg[action_dest] = value @@ -1450,7 +1461,7 @@ def _check_value_key( value = action.check_type(value, self) elif hasattr(action, "_check_type"): with parser_context(parent_parser=self): - value = action._check_type_(value, cfg=cfg, append=append) # type: ignore[attr-defined] + value = action._check_type_(value, cfg=cfg, append=append, mode=self.parser_mode) # type: ignore[attr-defined] elif action.type is not None: try: if action.nargs in {None, "?"} or action.nargs == 0: @@ -1593,8 +1604,8 @@ def parser_mode(self) -> str: @parser_mode.setter def parser_mode(self, parser_mode: str): - if parser_mode == "omegaconf": - set_omegaconf_loader() + if parser_mode in {"omegaconf", "omegaconf+"}: + set_omegaconf_loader(parser_mode) if parser_mode not in loaders: raise ValueError(f"The only accepted values for parser_mode are {set(loaders)}.") if parser_mode == "jsonnet": diff --git a/jsonargparse/_loaders_dumpers.py b/jsonargparse/_loaders_dumpers.py index f3867446..0354bc7c 100644 --- a/jsonargparse/_loaders_dumpers.py +++ b/jsonargparse/_loaders_dumpers.py @@ -335,11 +335,12 @@ def set_dumper(format_name: str, dumper_fn: Callable[[Any], str]): dumpers[format_name] = dumper_fn -def set_omegaconf_loader(): - if omegaconf_support and "omegaconf" not in loaders: +def set_omegaconf_loader(mode="omegaconf"): + if omegaconf_support and mode not in loaders: from ._optionals import get_omegaconf_loader - set_loader("omegaconf", get_omegaconf_loader(), get_loader_exceptions("yaml")) + loader = yaml_load if mode == "omegaconf+" else get_omegaconf_loader() + set_loader(mode, loader, get_loader_exceptions("yaml")) set_loader("jsonnet", jsonnet_load, get_loader_exceptions("jsonnet")) diff --git a/jsonargparse/_optionals.py b/jsonargparse/_optionals.py index de7becc2..082da482 100644 --- a/jsonargparse/_optionals.py +++ b/jsonargparse/_optionals.py @@ -286,6 +286,24 @@ def omegaconf_load(value): return omegaconf_load +def omegaconf_apply(parser, cfg): + if "${" not in str(cfg): + return cfg + + with missing_package_raise("omegaconf", "omegaconf_apply"): + from omegaconf import OmegaConf + + from ._common import parser_context + + with parser_context(path_dump_preserve_relative=True): + cfg_dict = parser.dump( + cfg, format="json_compact", skip_validation=True, skip_none=False, skip_link_targets=False + ) + cfg_omegaconf = OmegaConf.create(cfg_dict) + cfg_dict = OmegaConf.to_container(cfg_omegaconf, resolve=True) + return parser._apply_actions(cfg_dict) + + annotated_alias = typing_extensions_import("_AnnotatedAlias") diff --git a/jsonargparse/_typehints.py b/jsonargparse/_typehints.py index 2eab9ae2..c9024ee9 100644 --- a/jsonargparse/_typehints.py +++ b/jsonargparse/_typehints.py @@ -54,6 +54,7 @@ get_unaliased_type, is_dataclass_like, is_subclass, + lenient_check, nested_links, parent_parser, parser_context, @@ -524,7 +525,7 @@ def __call__(self, *args, **kwargs): if "nargs" in kwargs and kwargs["nargs"] == 0: raise ValueError("ActionTypeHint does not allow nargs=0.") return ActionTypeHint(**kwargs) - cfg, val, opt_str = args[1:] + parser, cfg, val, opt_str = args if not (self.nargs == "?" and val is None): if isinstance(opt_str, str) and opt_str.startswith(f"--{self.dest}."): if opt_str.startswith(f"--{self.dest}.init_args."): @@ -533,7 +534,7 @@ def __call__(self, *args, **kwargs): sub_opt = opt_str[len(f"--{self.dest}.") :] val = NestedArg(key=sub_opt, val=val) append = opt_str == f"--{self.dest}+" - val = self._check_type_(val, append=append, cfg=cfg) + val = self._check_type_(val, append=append, cfg=cfg, mode=parser.parser_mode) if is_subclass_spec(val): prev_val = cfg.get(self.dest) if is_subclass_spec(prev_val) and "init_args" in prev_val: @@ -545,7 +546,7 @@ def __call__(self, *args, **kwargs): cfg.update(val, self.dest) return None - def _check_type(self, value, append=False, cfg=None): + def _check_type(self, value, append=False, cfg=None, mode=None): islist = _is_action_value_list(self) if not islist: value = [value] @@ -584,7 +585,14 @@ def _check_type(self, value, append=False, cfg=None): val = adapt_typehints(orig_val, self._typehint, default=self.default, **kwargs) ex = None except ValueError: - if self._enable_path and config_path is None and isinstance(orig_val, str): + if ( + lenient_check.get() + and mode == "omegaconf+" + and isinstance(orig_val, str) + and "${" in orig_val + ): + ex = None + elif self._enable_path and config_path is None and isinstance(orig_val, str): msg = f"\n- Expected a config path but {orig_val} either not accessible or invalid\n- " raise type(ex)(msg + str(ex)) from ex if ex: diff --git a/jsonargparse/_util.py b/jsonargparse/_util.py index d9adbf9b..2e4d3405 100644 --- a/jsonargparse/_util.py +++ b/jsonargparse/_util.py @@ -282,11 +282,13 @@ def get_typehint_origin(typehint): @contextmanager -def change_to_path_dir(path: Optional["Path"]) -> Iterator[Optional[str]]: +def change_to_path_dir(path: Optional[Union["Path", str]]) -> Iterator[Optional[str]]: """A context manager for running code in the directory of a path.""" path_dir = current_path_dir.get() chdir: Union[bool, str] = False if path is not None: + if isinstance(path, str): + path = Path(path, mode="d") if path._url_data and (path.is_url or path.is_fsspec): scheme = path._url_data.scheme path_dir = path._url_data.url_path diff --git a/jsonargparse/typing.py b/jsonargparse/typing.py index 81e2125f..1708c684 100644 --- a/jsonargparse/typing.py +++ b/jsonargparse/typing.py @@ -13,9 +13,9 @@ else: _TypeAlias = type -from ._common import is_final_class +from ._common import is_final_class, path_dump_preserve_relative from ._optionals import final, pydantic_support -from ._util import Path, get_import_path, get_private_kwargs, import_object +from ._util import Path, change_to_path_dir, get_import_path, get_private_kwargs, import_object __all__ = [ "final", @@ -219,6 +219,15 @@ def _is_path_type(value, type_class): return isinstance(value, Path) +def _serialize_path(path: Path): + if path_dump_preserve_relative.get() and path.relative != path.absolute: + return { + "relative": path._relative, + "cwd": path._cwd, + } + return str(path) + + def path_type(mode: str, docstring: Optional[str] = None, **kwargs) -> _TypeAlias: """Creates or returns an already registered path type class. @@ -249,10 +258,14 @@ class PathType(Path): _expression = name _mode = mode _skip_check = skip_check - _type = str + _type = _serialize_path def __init__(self, v, **k): - super().__init__(v, mode=self._mode, skip_check=self._skip_check, **k) + if isinstance(v, dict) and set(v) == {"cwd", "relative"}: + with change_to_path_dir(v["cwd"]): + super().__init__(v["relative"], mode=self._mode, skip_check=self._skip_check, **k) + else: + super().__init__(v, mode=self._mode, skip_check=self._skip_check, **k) restricted_type = type(name, (PathType,), {"__doc__": docstring}) add_type(restricted_type, register_key, type_check=_is_path_type) diff --git a/jsonargparse_tests/test_core.py b/jsonargparse_tests/test_core.py index ead53f2c..d272bd62 100644 --- a/jsonargparse_tests/test_core.py +++ b/jsonargparse_tests/test_core.py @@ -191,6 +191,18 @@ def test_parse_object_simple(parser): pytest.raises(ArgumentError, lambda: parser.parse_object({"undefined": True})) +def test_parse_object_config(parser): + parser.add_argument("--cfg", action="config") + parser.add_argument("--a", type=int) + parser.add_argument("--b", type=int) + path = Path("config.json") + path.write_text('{"a": 1, "b": 2}') + cfg = parser.parse_object({"b": 0, "cfg": str(path), "a": 3}) + popped_cfg = cfg.pop("cfg") + assert popped_cfg[0].relative == "config.json" + assert cfg == Namespace(a=3, b=2) + + def test_parse_object_nested(parser): parser.add_argument("--l1.l2.op", type=float) assert parser.parse_object({"l1": {"l2": {"op": 2.1}}}).l1.l2.op == 2.1 diff --git a/jsonargparse_tests/test_loaders_dumpers.py b/jsonargparse_tests/test_loaders_dumpers.py index 56c9e6b6..f9eb9b86 100644 --- a/jsonargparse_tests/test_loaders_dumpers.py +++ b/jsonargparse_tests/test_loaders_dumpers.py @@ -1,7 +1,6 @@ from __future__ import annotations import json -import os from dataclasses import dataclass from pathlib import Path from typing import List @@ -11,8 +10,8 @@ from jsonargparse import ArgumentParser, get_loader, set_dumper, set_loader from jsonargparse._common import parser_context -from jsonargparse._loaders_dumpers import load_value, loaders, yaml_dump -from jsonargparse._optionals import omegaconf_support, pyyaml_available, toml_dump_available, toml_load_available +from jsonargparse._loaders_dumpers import load_value +from jsonargparse._optionals import pyyaml_available, toml_dump_available, toml_load_available from jsonargparse_tests.conftest import get_parse_args_stdout, json_or_yaml_dump, json_or_yaml_load, skip_if_no_pyyaml if pyyaml_available: @@ -170,54 +169,3 @@ def test_toml_print_config(parser): parser.add_argument("--group.child2", type=List[float], default=[3.0, 4.5]) out = get_parse_args_stdout(parser, ["--print_config"]) assert out.strip() == toml_config.strip() - - -# omegaconf tests - - -@pytest.mark.skipif(not omegaconf_support, reason="omegaconf package is required") -def test_parser_mode_omegaconf_interpolation(): - parser = ArgumentParser(parser_mode="omegaconf") - parser.add_argument("--server.host", type=str) - parser.add_argument("--server.port", type=int) - parser.add_argument("--client.url", type=str) - parser.add_argument("--config", action="config") - - config = { - "server": { - "host": "localhost", - "port": 80, - }, - "client": { - "url": "http://${server.host}:${server.port}/", - }, - } - cfg = parser.parse_args([f"--config={yaml_dump(config)}"]) - assert cfg.client.url == "http://localhost:80/" - assert "url: http://localhost:80/" in parser.dump(cfg) - - -@pytest.mark.skipif(not omegaconf_support, reason="omegaconf package is required") -def test_parser_mode_omegaconf_interpolation_in_subcommands(parser, subparser): - subparser.add_argument("--config", action="config") - subparser.add_argument("--source", type=str) - subparser.add_argument("--target", type=str) - - parser.parser_mode = "omegaconf" - subcommands = parser.add_subcommands() - subcommands.add_subcommand("sub", subparser) - - config = { - "source": "hello", - "target": "${source}", - } - cfg = parser.parse_args(["sub", f"--config={yaml_dump(config)}"]) - assert cfg.sub.target == "hello" - - -@pytest.mark.skipif( - not (omegaconf_support and "JSONARGPARSE_OMEGACONF_FULL_TEST" in os.environ), - reason="only for omegaconf as the yaml loader", -) -def test_omegaconf_as_yaml_loader(): - assert loaders["yaml"] is loaders["omegaconf"] diff --git a/jsonargparse_tests/test_omegaconf.py b/jsonargparse_tests/test_omegaconf.py new file mode 100644 index 00000000..f61b6f34 --- /dev/null +++ b/jsonargparse_tests/test_omegaconf.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +import json +import math +import os +from dataclasses import dataclass +from pathlib import Path +from unittest.mock import patch + +import pytest + +from jsonargparse import ArgumentParser, Namespace +from jsonargparse._common import parser_context +from jsonargparse._loaders_dumpers import loaders, yaml_dump +from jsonargparse._optionals import omegaconf_support +from jsonargparse.typing import Path_fr +from jsonargparse_tests.conftest import get_parser_help + +if omegaconf_support: + from omegaconf import OmegaConf + +skip_if_omegaconf_unavailable = pytest.mark.skipif( + not omegaconf_support, + reason="omegaconf package is required", +) + + +@pytest.mark.skipif( + not (omegaconf_support and "JSONARGPARSE_OMEGACONF_FULL_TEST" in os.environ), + reason="only for omegaconf as the yaml loader", +) +def test_omegaconf_as_yaml_loader(): + assert loaders["yaml"] is loaders["omegaconf"] + + +@skip_if_omegaconf_unavailable +@pytest.mark.parametrize("mode", ["omegaconf", "omegaconf+"]) +def test_omegaconf_interpolation(mode): + parser = ArgumentParser(parser_mode=mode) + parser.add_argument("--server.host", type=str) + parser.add_argument("--server.port", type=int) + parser.add_argument("--client.url", type=str) + parser.add_argument("--config", action="config") + + config = { + "server": { + "host": "localhost", + "port": 80, + }, + "client": { + "url": "http://${server.host}:${server.port}/", + }, + } + cfg = parser.parse_args([f"--config={yaml_dump(config)}"]) + assert cfg.client.url == "http://localhost:80/" + assert "url: http://localhost:80/" in parser.dump(cfg) + + +@skip_if_omegaconf_unavailable +@pytest.mark.parametrize("mode", ["omegaconf", "omegaconf+"]) +def test_omegaconf_interpolation_in_subcommands(mode, parser, subparser): + subparser.add_argument("--config", action="config") + subparser.add_argument("--source", type=str) + subparser.add_argument("--target", type=str) + + parser.parser_mode = mode + subcommands = parser.add_subcommands() + subcommands.add_subcommand("sub", subparser) + + config = { + "source": "hello", + "target": "${source}" if mode == "omegaconf" else "${.source}", + } + cfg = parser.parse_args(["sub", f"--config={yaml_dump(config)}"]) + assert cfg.sub.target == "hello" + + +@dataclass +class Server: + host: str = "localhost" + port: int = 80 + + +@dataclass +class Client: + url: str = "http://example.com:8080" + + +@skip_if_omegaconf_unavailable +def test_omegaconf_global_interpolation(parser): + parser.parser_mode = "omegaconf+" + parser.add_class_arguments(Server, "server") + parser.add_class_arguments(Client, "client") + + config = {"url": "http://${server.host}:${..server.port}/"} + cfg = parser.parse_args([f"--client={yaml_dump(config)}"]) + assert cfg.client == Namespace(url="http://localhost:80/") + + cfg = parser.parse_args([f"--client={yaml_dump(config)}", "--server.port=9000"]) + assert cfg.client == Namespace(url="http://localhost:9000/") + + +@skip_if_omegaconf_unavailable +def test_omegaconf_global_resolver_config(parser): + OmegaConf.register_new_resolver("increment", lambda x: x + 1) + + parser.parser_mode = "omegaconf+" + parser.add_argument("--config", action="config") + parser.add_argument("--value", type=int, default=0) + parser.add_argument("--incremented", type=int, default=0) + + assert parser.parse_args([]) == Namespace(config=None, value=0, incremented=0) + + config = {"value": 1, "incremented": "${increment:${value}}"} + cfg = parser.parse_args([f"--config={yaml_dump(config)}", "--value=5"]) + assert cfg == Namespace(value=5, incremented=6) # currently config is lost + + OmegaConf.clear_resolver("increment") + + +@skip_if_omegaconf_unavailable +def test_omegaconf_global_resolver_argument(parser): + def const(expr: str): + allowed = {"pi": math.pi} + return eval(expr, {"__builtins__": None}, allowed) + + OmegaConf.register_new_resolver("const", const) + + parser.parser_mode = "omegaconf+" + parser.add_argument("--value", type=float) + cfg = parser.parse_args(["--value=${const:3*pi/4}"]) + assert cfg.value == 3 * math.pi / 4 + + OmegaConf.clear_resolver("const") + + +@skip_if_omegaconf_unavailable +@patch.dict(os.environ, {"X": "true"}) +def test_omegaconf_global_resolver_default(parser): + parser.parser_mode = "omegaconf+" + action = parser.add_argument("--env", type=bool, default="${oc.env:X}") + assert action.default == "${oc.env:X}" + + help_str = get_parser_help(parser) + assert "default: ${oc.env:X}" in help_str + + cfg = parser.parse_args([]) + assert cfg.env is True + + +@dataclass +class Nested: + path: Path_fr + + +@skip_if_omegaconf_unavailable +@patch.dict(os.environ, {"X": "Y"}) +def test_omegaconf_global_path_preserve_relative(parser, tmp_cwd): + import yaml + + parser.parser_mode = "omegaconf+" + parser.add_class_arguments(Nested, "nested") + parser.add_argument("--env") + + subdir = Path("sub") + subdir.mkdir() + (subdir / "file").touch() + nested = subdir / "nested.json" + nested.write_text(json.dumps({"path": "file"})) + + cfg = parser.parse_args([f"--nested={nested}", "--env=${oc.env:X}"]) + assert cfg.env == "Y" + assert cfg.nested.path.relative == "file" + assert cfg.nested.path.cwd == str(tmp_cwd / subdir) + + with parser_context(path_dump_preserve_relative=True): + dump = yaml.safe_load(parser.dump(cfg))["nested"]["path"] + assert dump == {"relative": "file", "cwd": str(tmp_cwd / subdir)}