diff --git a/tinyml-modelmaker/tests/test_config_dict.py b/tinyml-modelmaker/tests/test_config_dict.py index 1372f544..2255d786 100644 --- a/tinyml-modelmaker/tests/test_config_dict.py +++ b/tinyml-modelmaker/tests/test_config_dict.py @@ -44,3 +44,9 @@ def test_none_input(self): cfg = ConfigDict(None) # Should create an empty config without error assert isinstance(cfg, ConfigDict) + + def test_constructor_args_deep_merge_nested_dict(self): + default = dict(training=dict(a=1, b=2)) + user = dict(training=dict(b=99)) + cfg = ConfigDict(default, user) + assert dict(cfg.training) == {"a": 1, "b": 99} diff --git a/tinyml-modelmaker/tinyml_modelmaker/utils/config_dict.py b/tinyml-modelmaker/tinyml_modelmaker/utils/config_dict.py index 953fa23d..d40e4cf2 100644 --- a/tinyml-modelmaker/tinyml_modelmaker/utils/config_dict.py +++ b/tinyml-modelmaker/tinyml_modelmaker/utils/config_dict.py @@ -57,7 +57,7 @@ def __init__(self, input=None, *args, **kwargs): # override the entries with args for value in args: if isinstance(value, (dict, ConfigDict)): - input_dict.update(value) + self._deep_merge(input_dict, value) # # # override the entries with kwargs @@ -95,6 +95,17 @@ def __setstate__(self, state): def _initialize(self): pass + @staticmethod + def _deep_merge(target, source): + for key, value in source.items(): + if key in target and isinstance(target[key], (dict, ConfigDict)) and isinstance(value, (dict, ConfigDict)): + ConfigDict._deep_merge(target[key], value) + else: + target[key] = value + # + # + return target + def _parse_include_files(self, include_files, include_base_path): input_dict = {} include_files = list(include_files)