diff --git a/nemo_run/config.py b/nemo_run/config.py index a4babd9c..88ab7fc0 100644 --- a/nemo_run/config.py +++ b/nemo_run/config.py @@ -577,7 +577,7 @@ def _construct_args( for name, parameter in params.items(): arg = kwargs.get(name, None) - if arg: + if arg is not None: if dataclasses.is_dataclass(arg): final_args[name] = fdl.cast( Config, diff --git a/test/test_config.py b/test/test_config.py index cbf3054c..0cca952d 100644 --- a/test/test_config.py +++ b/test/test_config.py @@ -66,6 +66,10 @@ def train_manual( return optim +def with_falsy_values(count: int = 1, flag: bool = True, label: str = "x"): + return count, flag, label + + @dataclass class Data: name: str @@ -222,6 +226,13 @@ def test_default_optional(self, train_func): assert fn() == fdl.build(optimizer()) + def test_falsy_primitive_args_are_preserved(self): + """Falsy primitive values like 0, False, and '' must not be dropped.""" + partial = run.Partial(with_falsy_values, count=0, flag=False, label="") + fn = fdl.build(partial) + + assert fn() == (0, False, "") + def test_clone(self): partial = run.Partial(train, model=dummy_model(), optim=optimizer()) clone = partial.clone()