Skip to content

Commit b71090e

Browse files
authored
Allow multiple alf.define_config calls for nested confs (#1816)
1 parent ecbde3e commit b71090e

2 files changed

Lines changed: 17 additions & 3 deletions

File tree

alf/config_util.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,6 @@
2121
import os
2222
import pprint
2323
import runpy
24-
import shutil
2524

2625
__all__ = [
2726
'config',
@@ -994,6 +993,9 @@ def define_config(name, default_value):
994993
995994
Its value can be retrieved by ``get_config_value("_CONFIG._USER.{name}")``.
996995
996+
If the configurable has already been defined, subsequent define_config calls
997+
will the same name will be ignored and the already set default value will be returned.
998+
997999
Args:
9981000
name (str): name of the configurable value
9991001
default_value (Any): default value
@@ -1002,8 +1004,15 @@ def define_config(name, default_value):
10021004
"""
10031005
node = _Config()
10041006
node.set_default_value(default_value)
1005-
_add_to_conf_tree(['_CONFIG'], '_USER', name, node)
1006-
_DEFINED_CONFIGS.append('_CONFIG._USER.' + name)
1007+
try:
1008+
_add_to_conf_tree(['_CONFIG'], '_USER', name, node)
1009+
_DEFINED_CONFIGS.append('_CONFIG._USER.' + name)
1010+
except ValueError:
1011+
already_set_default_value = get_config_value("_CONFIG._USER." + name)
1012+
logging.warning(
1013+
f"Config {name} has already been configured. Provided value {default_value} will be ignored "
1014+
f"in favor of {already_set_default_value}")
1015+
10071016
return get_config_value("_CONFIG._USER." + name)
10081017

10091018

alf/config_util_test.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -342,6 +342,11 @@ def __init__(self, x, y, z=0):
342342
test_partial = partial(TestPositionalArgs, y=0)
343343
test_partial()
344344

345+
def test_define_config(self):
346+
val = alf.define_config("foobar", 3)
347+
new_val = alf.define_config("foobar", 4)
348+
self.assertEqual(val, new_val)
349+
345350

346351
if __name__ == '__main__':
347352
alf.test.main()

0 commit comments

Comments
 (0)