Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
107 changes: 92 additions & 15 deletions fiddle/_src/codegen/auto_config/experimental_top_level_api_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"},
Expand Down
16 changes: 10 additions & 6 deletions fiddle/_src/codegen/auto_config/ir_to_cst.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}"
Expand Down
4 changes: 2 additions & 2 deletions fiddle/_src/codegen/auto_config/ir_to_cst_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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)

Expand Down
42 changes: 34 additions & 8 deletions fiddle/_src/codegen/auto_config/make_symbolic_references.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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,
Expand All @@ -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:
Expand Down Expand Up @@ -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],
Expand Down
14 changes: 13 additions & 1 deletion fiddle/_src/codegen/auto_config/test_fixtures.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
Loading