diff --git a/fiddle/_src/codegen/auto_config/experimental_top_level_api_test.py b/fiddle/_src/codegen/auto_config/experimental_top_level_api_test.py index f0c6ee33..42adfefc 100644 --- a/fiddle/_src/codegen/auto_config/experimental_top_level_api_test.py +++ b/fiddle/_src/codegen/auto_config/experimental_top_level_api_test.py @@ -13,7 +13,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -import dataclasses import importlib import os import random @@ -151,22 +150,100 @@ def test_fuzz(self, kwargs): generated_config = module.config_fixture.as_buildable() self.assertDagEqual(config, generated_config) - def test_config_contains_tags_wo_default(self): - @dataclasses.dataclass - class Foo: - a: int = 1 + @parameterized.named_parameters( + {"testcase_name": "config", "buildable_cls": fdl.Config}, + {"testcase_name": "partial", "buildable_cls": fdl.Partial}, + ) + def test_config_contains_tags_wo_value( + self, buildable_cls: type[fdl.Buildable] + ): + config = buildable_cls(test_fixtures.bar, x=1) + fdl.add_tag(config, "y", test_fixtures.ATag) - config = fdl.Config(Foo) - fdl.set_tags(config, "a", [test_fixtures.ATag]) + code = experimental_top_level_api.auto_config_codegen(config) + self.assertIn( + "auto_config.with_tags(fdl.NO_VALUE, test_fixtures.ATag)", code + ) - with self.assertRaisesRegex( - ValueError, - ( - "assigning a value to the field first or removing field tags from " - "your config" - ), - ): - experimental_top_level_api.auto_config_codegen(config) + module = self._load_code_as_module(code) + generated_config = module.config_fixture.as_buildable() + self.assertDagEqual(config, generated_config) + self.assertEqual(fdl.get_tags(generated_config, "y"), {test_fixtures.ATag}) + self.assertNotIn("y", generated_config.__arguments__) + + @parameterized.named_parameters( + { + "testcase_name": "multiple_annotated_tags", + "config": fdl.Partial(test_fixtures.multi_annotated_bar, x=1), + }, + { + "testcase_name": "positional_only_unset", + "config": fdl.Partial( + test_fixtures.positional_only_annotated_bar, y=2 + ), + }, + { + "testcase_name": "positional_only_set", + "config": fdl.Partial( + test_fixtures.positional_only_annotated_bar, 5, y=2 + ), + }, + ) + def test_partial_with_annotated_tags_round_trips(self, config: fdl.Partial): + code = experimental_top_level_api.auto_config_codegen(config) + module = self._load_code_as_module(code) + generated_config = module.config_fixture.as_buildable() + self.assertDagEqual(config, generated_config) + + def test_partial_with_annotated_tag_unset_arg_unchanged(self): + config = fdl.Partial(test_fixtures.annotated_bar, x=1) + code = experimental_top_level_api.auto_config_codegen(config) + self.assertNotIn("NO_VALUE", code) + module = self._load_code_as_module(code) + self.assertDagEqual(config, module.config_fixture.as_buildable()) + # Called as plain Python, the fixture should use `y`'s default. + self.assertEqual(module.config_fixture()(), (1, 3)) + + @parameterized.product( + buildable_cls=(fdl.Config, fdl.Partial), + kwargs=({"x": 1}, {"x": 1, "y": 2}), + ) + def test_user_tag_on_annotated_arg_round_trips( + self, buildable_cls: type[fdl.Buildable], kwargs: dict[str, int] + ): + config = buildable_cls(test_fixtures.annotated_bar, **kwargs) + fdl.add_tag(config, "y", test_fixtures.BTag) + code = experimental_top_level_api.auto_config_codegen(config) + module = self._load_code_as_module(code) + self.assertDagEqual(config, module.config_fixture.as_buildable()) + + @parameterized.product( + buildable_cls=(fdl.Config, fdl.Partial), + kwargs=({"x": 1}, {"x": 1, "y": 2}), + clear_all_field_tags=(False, True), + ) + def test_removed_annotated_tag_is_reapplied( + self, + buildable_cls: type[fdl.Buildable], + kwargs: dict[str, int], + clear_all_field_tags: bool, + ): + config = buildable_cls(test_fixtures.annotated_bar, **kwargs) + if clear_all_field_tags: + config = tagging.materialize_tags(config, clear_field_tags=True) + else: + fdl.remove_tag(config, "y", test_fixtures.ATag) + self.assertEqual(fdl.get_tags(config, "y"), set()) + + code = experimental_top_level_api.auto_config_codegen(config) + module = self._load_code_as_module(code) + generated_config = module.config_fixture.as_buildable() + # Constructing the Buildable re-applies the annotation tag, which is the + # closest equivalent auto_config code can express. + self.assertDagEqual( + buildable_cls(test_fixtures.annotated_bar, **kwargs), generated_config + ) + self.assertEqual(fdl.get_tags(generated_config, "y"), {test_fixtures.ATag}) @parameterized.named_parameters( {"testcase_name": "basic", "api": "highlevel"}, diff --git a/fiddle/_src/codegen/auto_config/ir_to_cst.py b/fiddle/_src/codegen/auto_config/ir_to_cst.py index 22388a01..8fe73607 100644 --- a/fiddle/_src/codegen/auto_config/ir_to_cst.py +++ b/fiddle/_src/codegen/auto_config/ir_to_cst.py @@ -171,13 +171,17 @@ def _prepare_args_helper( elif isinstance(value, code_ir.WithTagsCall): attr = daglish.Attr("item_to_tag") item_to_tag = state.call(value.item_to_tag, attr) - call_args = [cst.Arg(item_to_tag)] - sorted_tags = sorted([tag for tag in value.tag_symbol_expressions]) - for tag in sorted_tags: - tag_name = cst.parse_expression(tag) - call_args.append(cst.Arg(tag_name)) + tag_exprs = [ + cst.parse_expression(tag) + for tag in sorted(value.tag_symbol_expressions) + ] + # `with_tags` accepts either a single tag or a collection of tags. + if len(tag_exprs) == 1: + tags_arg = tag_exprs[0] + else: + tags_arg = cst.List([cst.Element(tag) for tag in tag_exprs]) with_tags = cst.parse_expression("auto_config.with_tags") - return cst.Call(with_tags, args=call_args) + return cst.Call(with_tags, args=[cst.Arg(item_to_tag), cst.Arg(tags_arg)]) elif state.is_traversable(value): raise NotImplementedError( f"Expression generation is not implemented for {value!r}" diff --git a/fiddle/_src/codegen/auto_config/ir_to_cst_test.py b/fiddle/_src/codegen/auto_config/ir_to_cst_test.py index 6b48b130..b5b777d0 100644 --- a/fiddle/_src/codegen/auto_config/ir_to_cst_test.py +++ b/fiddle/_src/codegen/auto_config/ir_to_cst_test.py @@ -156,7 +156,7 @@ def test_tags_generation_config(self): def config_fixture(): return test_fixtures.bar(x=auto_config.with_tags(1, test_fixtures.ATag), - y=auto_config.with_tags(2, test_fixtures.ATag, test_fixtures.BTag)) + y=auto_config.with_tags(2, [test_fixtures.ATag, test_fixtures.BTag])) """ self.assertEqual(code.split(), expected.split(), msg=code) @@ -175,7 +175,7 @@ def test_tags_generation_partial(self): def config_fixture(): return functools.partial(test_fixtures.bar, x=auto_config.with_tags(1, test_fixtures.ATag), - y=auto_config.with_tags(2, test_fixtures.ATag, test_fixtures.BTag)) + y=auto_config.with_tags(2, [test_fixtures.ATag, test_fixtures.BTag])) """ self.assertEqual(code.split(), expected.split(), msg=code) diff --git a/fiddle/_src/codegen/auto_config/make_symbolic_references.py b/fiddle/_src/codegen/auto_config/make_symbolic_references.py index b057d71b..a2276642 100644 --- a/fiddle/_src/codegen/auto_config/make_symbolic_references.py +++ b/fiddle/_src/codegen/auto_config/make_symbolic_references.py @@ -24,6 +24,7 @@ from fiddle import daglish from fiddle._src import config as config_lib from fiddle._src import partial +from fiddle._src import tag_type from fiddle._src.codegen.auto_config import code_ir from fiddle._src.codegen.auto_config import import_manager_wrapper @@ -86,6 +87,31 @@ def replace_callables_and_configs_with_symbols( fn_name = None + def _no_value_expr() -> code_ir.CodegenNode: + """Returns an expression for `fdl.NO_VALUE`, e.g. for tagged unset args.""" + fiddle_module = task.import_manager.add_by_name("fiddle") + return code_ir.AttributeExpression( + code_ir.ModuleReference(code_ir.Name(fiddle_module)), "NO_VALUE" + ) + + def _unset_tagged_args_needing_no_value( + buildable: config_lib.Buildable, + ) -> list[str | int]: + """Returns names of unset tagged args that must be emitted as `NO_VALUE`.""" + annotated_tags: dict[str | int, set[tag_type.TagType]] = { + name: set(tags) + for name, tags in tag_type.find_tags_from_annotations( + config_lib.get_callable(buildable) + ).items() + } + return [ + name + for name, tags in buildable.__argument_tags__.items() + if tags + and name not in buildable.__arguments__ + and not tags <= annotated_tags.get(name, set()) + ] + def _handle_partial( value: partial.Partial, state: daglish.State, @@ -110,6 +136,9 @@ def _handle_partial( else: regular_args[name] = state.call(arg_value, daglish.Attr(name)) # pyrefly: ignore[bad-argument-type] + for name in _unset_tagged_args_needing_no_value(value): + regular_args[name] = _no_value_expr() + for dict_of_args in (arg_factory_args, regular_args): for arg in dict_of_args: if arg in all_tags: @@ -161,17 +190,14 @@ def traverse(value, state: daglish.State): ) if isinstance(value, config_lib.Config): all_tags = value.__argument_tags__ + no_value_args = _unset_tagged_args_needing_no_value(value) value = state.map_children(value) + for arg in no_value_args: + value.__arguments__[arg] = _no_value_expr() for arg, arg_tags in all_tags.items(): + if not arg_tags or arg not in value.__arguments__: + continue tag_expr = [task.import_manager.add(tag) for tag in arg_tags] - if arg not in value.__arguments__: - raise ValueError( - f"Tagged field '{arg}' of {value!r} is not found in its" - f" arguments: {value.__arguments__}. This is likely because the" - " tagged field doesn't yet have a value. Consider assigning a" - " value to the field first or removing field tags from your" - " config, for example using `fdl.clear_tags`." - ) value.__arguments__[arg] = code_ir.WithTagsCall( tag_symbol_expressions=tag_expr, item_to_tag=value.__arguments__[arg], diff --git a/fiddle/_src/codegen/auto_config/test_fixtures.py b/fiddle/_src/codegen/auto_config/test_fixtures.py index 7610ec4c..57aeba4c 100644 --- a/fiddle/_src/codegen/auto_config/test_fixtures.py +++ b/fiddle/_src/codegen/auto_config/test_fixtures.py @@ -17,7 +17,7 @@ import dataclasses import functools -from typing import Callable +from typing import Annotated, Callable import fiddle as fdl from fiddle import arg_factory @@ -43,6 +43,18 @@ class BTag(fdl.Tag): """Sample tag to test code generation of tags.""" +def annotated_bar(x, y: Annotated[int, ATag] = 3): + return (x, y) + + +def multi_annotated_bar(x, y: Annotated[int, ATag, BTag] = 3): + return (x, y) + + +def positional_only_annotated_bar(x: Annotated[int, ATag] = 3, /, y=4): + return (x, y) + + @dataclasses.dataclass class SharedType: x: int