Skip to content

Commit cca28bf

Browse files
authored
Cache class parsers to improve performance and reduce test suite runtime (#903)
1 parent c180980 commit cca28bf

4 files changed

Lines changed: 72 additions & 28 deletions

File tree

CHANGELOG.rst

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,8 @@ Changed
2020
- Docs now reference methods via the public ``ArgumentParser`` class instead of
2121
internal mixin classes (`#901
2222
<https://github.com/omni-us/jsonargparse/pull/901>`__).
23+
- Cache class parsers to improve performance and reduce test suite runtime
24+
(`#903 <https://github.com/omni-us/jsonargparse/pull/903>`__).
2325

2426
Deprecated
2527
^^^^^^^^^^

jsonargparse/_typehints.py

Lines changed: 63 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,7 @@
7777
from ._paths import Path, PathError, change_to_path_dir
7878
from ._required import clear_required
7979
from ._subcommands import find_action, find_parent_action, parse_kwargs
80+
from ._type_checking import ArgumentParser
8081
from ._util import (
8182
NestedArg,
8283
NoneType,
@@ -213,6 +214,61 @@ def get_parse_optional_num_return() -> int:
213214
parse_optional_num_return = get_parse_optional_num_return()
214215

215216

217+
def freeze(value):
218+
if isinstance(value, dict):
219+
return tuple(sorted(((k, freeze(v)) for k, v in value.items()), key=lambda item: repr(item[0])))
220+
if isinstance(value, set):
221+
return tuple(sorted((freeze(v) for v in value), key=repr))
222+
if isinstance(value, (list, tuple)):
223+
return tuple(freeze(v) for v in value)
224+
return value
225+
226+
227+
_cached_class_parsers: dict[tuple, ArgumentParser] = {}
228+
229+
230+
def cached_get_class_parser(*, val_class, sub_add_kwargs, skip_args, parent_parser, nested_links):
231+
if isinstance(val_class, str):
232+
val_class = import_object(val_class)
233+
parser_class = type(parent_parser)
234+
cache_key = (
235+
val_class,
236+
parser_class,
237+
parent_parser.parser_mode,
238+
freeze(sub_add_kwargs),
239+
freeze(skip_args),
240+
freeze(nested_links),
241+
)
242+
if cache_key in _cached_class_parsers:
243+
parser = _cached_class_parsers[cache_key]
244+
parser.logger = parent_parser.logger
245+
return parser
246+
247+
kwargs = dict(sub_add_kwargs) if sub_add_kwargs else {}
248+
if skip_args:
249+
kwargs.setdefault("skip", set()).update(skip_args)
250+
251+
parser = parser_class(exit_on_error=False, logger=parent_parser.logger, parser_mode=parent_parser.parser_mode)
252+
remove_actions(parser, (ActionConfigFile, _ActionPrintConfig))
253+
if inspect.isclass(val_class) or inspect.isclass(get_typehint_origin(val_class)):
254+
parser.add_class_arguments(val_class, **kwargs)
255+
else:
256+
kwargs = {k: v for k, v in kwargs.items() if k != "instantiate"}
257+
parser.add_function_arguments(val_class, **kwargs)
258+
259+
if "linked_targets" in kwargs:
260+
for key in kwargs["linked_targets"]:
261+
clear_required(parser, key)
262+
263+
for link_kwargs in nested_links:
264+
parser.link_arguments(**link_kwargs)
265+
266+
parser._inner_parser = True
267+
268+
_cached_class_parsers[cache_key] = parser
269+
return parser
270+
271+
216272
class ActionTypeHint(Action):
217273
"""Action to parse a type hint."""
218274

@@ -633,33 +689,13 @@ def instantiate_classes(self, value):
633689

634690
@staticmethod
635691
def get_class_parser(val_class, sub_add_kwargs=None, skip_args=None):
636-
if isinstance(val_class, str):
637-
val_class = import_object(val_class)
638-
kwargs = dict(sub_add_kwargs) if sub_add_kwargs else {}
639-
if skip_args:
640-
kwargs.setdefault("skip", set()).update(skip_args)
641-
parser = parent_parser.get()
642-
from ._core import ArgumentParser
643-
644-
assert isinstance(parser, ArgumentParser)
645-
parser = type(parser)(exit_on_error=False, logger=parser.logger, parser_mode=parser.parser_mode)
646-
remove_actions(parser, (ActionConfigFile, _ActionPrintConfig))
647-
if inspect.isclass(val_class) or inspect.isclass(get_typehint_origin(val_class)):
648-
parser.add_class_arguments(val_class, **kwargs)
649-
else:
650-
kwargs = {k: v for k, v in kwargs.items() if k != "instantiate"}
651-
parser.add_function_arguments(val_class, **kwargs)
652-
653-
if "linked_targets" in kwargs:
654-
for key in kwargs["linked_targets"]:
655-
clear_required(parser, key)
656-
657-
for link_kwargs in nested_links.get():
658-
parser.link_arguments(**link_kwargs)
659-
660-
parser._inner_parser = True
661-
662-
return parser
692+
return cached_get_class_parser(
693+
val_class=val_class,
694+
sub_add_kwargs=sub_add_kwargs,
695+
skip_args=skip_args,
696+
parent_parser=parent_parser.get(),
697+
nested_links=nested_links.get(),
698+
)
663699

664700
def extra_help(self):
665701
extra = ""

jsonargparse_tests/conftest.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -168,10 +168,14 @@ def parsing_settings_patch():
168168

169169
@pytest.fixture
170170
def subclass_behavior(monkeypatch) -> Iterator[None]:
171+
from jsonargparse._typehints import _cached_class_parsers
172+
171173
monkeypatch.setattr("jsonargparse._common.subclasses_enabled_types", set())
172174
monkeypatch.setattr("jsonargparse._common.subclasses_disabled_types", set())
175+
_cached_class_parsers.clear()
173176
with patch.dict("jsonargparse._common.subclasses_disabled_selectors"):
174177
yield
178+
_cached_class_parsers.clear()
175179

176180

177181
@pytest.fixture

jsonargparse_tests/test_subclasses.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@
2323
lazy_instance,
2424
)
2525
from jsonargparse._instantiation import _global_class_instantiators
26-
from jsonargparse._typehints import implements_protocol, is_instance_or_supports_protocol
26+
from jsonargparse._typehints import _cached_class_parsers, implements_protocol, is_instance_or_supports_protocol
2727
from jsonargparse.typing import final
2828
from jsonargparse_tests.conftest import (
2929
capture_logs,
@@ -1039,6 +1039,7 @@ def __init__(self, cal: Union[Calendar, bool] = lazy_instance(OverrideMixed, par
10391039

10401040

10411041
def test_subclass_discard_init_args_mixed_type(parser, logger):
1042+
_cached_class_parsers.clear()
10421043
parser.logger = logger
10431044
parser.add_class_arguments(OverrideMixedMain, "main")
10441045
with capture_logs(logger) as logs:
@@ -1081,6 +1082,7 @@ def __init__(self, *args, param: str = "0", **kwargs):
10811082

10821083

10831084
def test_subclass_discard_init_args_with_default_config_files(parser, tmp_cwd, logger):
1085+
_cached_class_parsers.clear()
10841086
config = {
10851087
"class_path": f"{__name__}.OverrideDefaultConfig",
10861088
"init_args": {"firstweekday": 2, "param": "1"},

0 commit comments

Comments
 (0)