From acc78b5fe02c2d0821b208b223026a250c939111 Mon Sep 17 00:00:00 2001 From: Andrei Luzan Date: Tue, 5 Aug 2025 17:42:55 +0300 Subject: [PATCH 1/2] fetch defaults from function signature; resolve default config values --- configuronic/config.py | 7 ++++- tests/test_cli.py | 12 ++++++++ tests/test_config.py | 62 ++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 80 insertions(+), 1 deletion(-) diff --git a/configuronic/config.py b/configuronic/config.py index 38f2a45..9f596ed 100644 --- a/configuronic/config.py +++ b/configuronic/config.py @@ -190,6 +190,11 @@ def _get_creator_module() -> ModuleType | None: return module +def _get_default_argument_values(target: Callable) -> dict[str, Any]: + sig = inspect.signature(target) + return {name: param.default for name, param in sig.parameters.items() if param.default != inspect.Parameter.empty} + + class Config: def __init__(self, target, *args, **kwargs): """ @@ -231,7 +236,7 @@ def __init__(self, target, *args, **kwargs): assert callable(target), f'Target must be callable, got object of type {type(target)}.' self.target = target self.args = [_resolve_value(arg) for arg in args] # TODO: cover argument override with tests - self.kwargs = {} + self.kwargs = _get_default_argument_values(target) self._override_inplace(**kwargs) self._creator_module = _get_creator_module() diff --git a/tests/test_cli.py b/tests/test_cli.py index b6a5a79..794536d 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -37,3 +37,15 @@ def identity(a, b): cfn.cli(identity) out, err = capfd.readouterr() assert "This is a test function." in out + + +def test_cli_help_contains_function_default_values(capfd): + @cfn.config() + def identity(a=1, b=2): + print(a, b) + + with patch('sys.argv', ['script.py', '--help']): + cfn.cli(identity) + out, err = capfd.readouterr() + assert "a: 1" in out + assert "b: 2" in out diff --git a/tests/test_config.py b/tests/test_config.py index b6d5869..d7e053c 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -913,5 +913,67 @@ def func(x): cfn.Config(func, '.return1') +def test_config_config_default_values_are_resolved_properly(): + + @cfn.config() + def return1(): + return 1 + + @cfn.config() + def add1(a=return1): + return a + 1 + + assert add1.instantiate() == 2 + + +def test_config_config_multiple_default_values_are_resolved_properly(): + + @cfn.config() + def returnx(x=1): + return x + + return1 = returnx.override(x=1) + return2 = returnx.override(x=2) + + @cfn.config() + def add1(a, b=return1, c=return2): + return f"{a} + {b} + {c}" + + assert add1(a=0) == "0 + 1 + 2" + + +def test_config_default_function_kwargs_are_overriden_by_config_kwargs(): + @cfn.config() + def return1(): + return 1 + + @cfn.config() + def return2(): + return 2 + + @cfn.config(a=return1) + def add1(a=return2): + return a + 1 + + assert add1.instantiate() == 2 + + +def test_config_args_and_kwargs_kwargs_are_overriden_by_config_kwargs(): + @cfn.config() + def return1(): + return 1 + + @cfn.config() + def return2(): + return 2 + + def func(*args, b=return2): + return f"{args} + {b}" + + func = cfn.Config(func, 1, 2, b=return1) + + assert func.instantiate() == "(1, 2) + 1" + + if __name__ == '__main__': pytest.main() From cdaa6c1013325dc7715318d0ce9cb3438c746356 Mon Sep 17 00:00:00 2001 From: Andrei Luzan Date: Tue, 5 Aug 2025 17:55:33 +0300 Subject: [PATCH 2/2] change unused arg in README.md --- README.md | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/README.md b/README.md index f9b02d6..bb37f2a 100644 --- a/README.md +++ b/README.md @@ -78,13 +78,13 @@ In configuronic, configurations are **closures** - callables that store both the import configuronic as cfn # Create a configuration -@cfn.config(batch_size=32, lr=0.001) -def create_optimizer(batch_size: int, lr: float): - return torch.optim.Adam(lr=lr) +@cfn.config(lr=1e-3, weight_decay=0) +def create_optimizer(lr: float, weight_decay: float): + return torch.optim.Adam(lr=lr, weight_decay=weight_decay) # Override and create variants -fast_optimizer = create_optimizer.override(lr=0.01) -large_batch_optimizer = create_optimizer.override(batch_size=128) +fast_optimizer = create_optimizer.override(lr=1e-2) +regularaized_optimizer = create_optimizer.override(weight_decay=1e-3) # Instantiate when needed optimizer = fast_optimizer.instantiate()