diff --git a/meshed/__init__.py b/meshed/__init__.py index 4a394322..3ad90a48 100644 --- a/meshed/__init__.py +++ b/meshed/__init__.py @@ -15,7 +15,16 @@ by an adjacency Mapping. """ -from meshed.dag import DAG, ch_funcs, ch_names +from meshed.dag import ( + DAG, + ch_funcs, + ch_names, + DuplicateOutsWarning, + duplicate_outs, + warn_on_duplicate_outs, + ignore_duplicate_outs, + raise_on_duplicate_outs, +) from meshed.base import FuncNode, compare_signatures from meshed.makers import code_to_dag, code_to_fnodes from meshed.itools import random_graph, topological_sort diff --git a/meshed/dag.py b/meshed/dag.py index d24d8a68..67c3afe5 100644 --- a/meshed/dag.py +++ b/meshed/dag.py @@ -159,6 +159,7 @@ VT, ) from collections.abc import Callable, MutableMapping, Iterable, Mapping +import sys from warnings import warn from i2 import double_up_as_factory, MultiFunc @@ -487,6 +488,114 @@ def dflt_debugger_feedback(func_node, scope, output, step): return output +class DuplicateOutsWarning(UserWarning): + """Warns that several ``FuncNode``s of a ``DAG`` write to the same var node. + + The dag will still compute all of them, but only the value of the last one + (in topological order) is visible: see i2mint/meshed#40. + """ + + +def duplicate_outs(func_nodes) -> dict: + """``{out: [names of the func nodes writing to it]}``, for the outs written to by + more than one func node (the last name is the one whose value is visible). + + >>> def foo(a): return a + >>> nodes = [FuncNode(foo, name='x1', out='x'), FuncNode(foo, name='x2', out='x')] + >>> duplicate_outs(nodes) + {'x': ['x1', 'x2']} + >>> duplicate_outs([FuncNode(foo, name='x1', out='x')]) + {} + """ + names_of_out = defaultdict(list) + for func_node in func_nodes: + names_of_out[func_node.out].append(func_node.name) + return {out: names for out, names in names_of_out.items() if len(names) > 1} + + +def duplicate_outs_string(duplicates: dict, *, dag_name=None) -> str: + """Human readable description of a ``duplicate_outs`` mapping.""" + details = "; ".join( + f"{out} (written by {', '.join(names)}: {names[-1]}'s value is the visible one)" + for out, names in duplicates.items() + ) + of_dag = f" of {dag_name}" if dag_name else "" + return f"Several func nodes{of_dag} write to the same var node(s): {details}" + + +def _stacklevel_of_first_frame_outside_meshed(_max_levels=30) -> int: + """The ``stacklevel`` (as ``warnings.warn`` counts it, from the caller of this + function) of the first frame that is not meshed's own code, so that a warning + points at the code that built the dag, however many meshed frames lie between. + """ + frame = sys._getframe(1) # the caller: stacklevel 1 + level = 1 + while frame is not None and level <= _max_levels: + module = frame.f_globals.get("__name__", "") + # i2 too: meshed's makers/decorators call through i2's decorator plumbing + is_meshed = module.split(".", 1)[0] in ("meshed", "i2") + # the dataclass-generated __init__ has "" for a filename + is_generated = frame.f_code.co_filename == "" + if level > 1 and not is_meshed and not is_generated: + return level + frame = frame.f_back + level += 1 + return 1 + + +def warn_on_duplicate_outs(duplicates: dict, *, dag_name=None): + """Emit a ``DuplicateOutsWarning`` describing ``duplicates``. + + This is the default ``on_duplicate_outs`` strategy of ``DAG``. Use + ``ignore_duplicate_outs``, ``raise_on_duplicate_outs``, or any callable of your + own (taking ``(duplicates, *, dag_name=None)``) to change that. Note that a + strategy needs to be picklable for the dag to be (so: no lambdas if you pickle). + """ + warn( + f"{duplicate_outs_string(duplicates, dag_name=dag_name)}. " + "The other values are computed and dropped. If that's intentional, pass " + "on_duplicate_outs=ignore_duplicate_outs to the DAG.", + DuplicateOutsWarning, + # not a constant: a custom strategy delegating here adds frames, and dags are + # also made from within meshed (code_to_dag, dag arithmetic, ...) + stacklevel=_stacklevel_of_first_frame_outside_meshed(), + ) + + +def ignore_duplicate_outs(duplicates: dict, *, dag_name=None): + """``on_duplicate_outs`` strategy that says nothing about duplicate outs.""" + + +def raise_on_duplicate_outs(duplicates: dict, *, dag_name=None): + """``on_duplicate_outs`` strategy that raises a ``ValidationError``.""" + raise ValidationError(duplicate_outs_string(duplicates, dag_name=dag_name)) + + +def only_new_duplicate_outs(strategy, known_duplicate_outs: dict): + """Wrap an ``on_duplicate_outs`` strategy so that it is only called for + duplications that are NOT already those of ``known_duplicate_outs`` (the + duplications of the dag(s) a new dag is derived from, which were already dealt + with when those dags were made). + + A derived dag may rename its nodes (``copy``, ``ch_names``), so a duplication + that has the same shape as the source's counts as the same one. + """ + known_names = set(known_duplicate_outs) + known_shape = sorted(map(len, known_duplicate_outs.values())) + + def on_duplicate_outs(duplicates: dict, *, dag_name=None): + new_duplicates = { + out: names for out, names in duplicates.items() if out not in known_names + } + if not new_duplicates: + return + if sorted(map(len, duplicates.values())) == known_shape: + return # same duplications as the source's, only renamed + strategy(new_duplicates, dag_name=dag_name) + + return on_duplicate_outs + + # TODO: caching last scope isn't really the DAG's direct concern -- it's a debugging # concern. Perhaps a more general form would be to define a cache factory defaulting # to a dict, but that could be a "dict" that logs writes (even to an attribute of self) @@ -547,6 +656,11 @@ class DAG: extract_output_from_scope: Callable[[Scope, VarNames], DagOutput] = field( default=extract_values, repr=False ) + # What to do when several func nodes write to the same var node (see #40). + # Alternatives: ignore_duplicate_outs, raise_on_duplicate_outs, or your own. + on_duplicate_outs: Callable[..., None] = field( + default=warn_on_duplicate_outs, repr=False + ) def __post_init__(self): self.func_nodes = tuple(ensure_func_nodes(self.func_nodes)) @@ -573,6 +687,8 @@ def __post_init__(self): self.__name__ = self.name or "DAG" self.bindings_cleaner() + if duplicates := duplicate_outs(self.func_nodes): + self.on_duplicate_outs(duplicates, dag_name=self.name) # TODO: No control of other DAG args (cache_last_scope etc.). @classmethod @@ -808,11 +924,22 @@ def __getitem__(self, item): """ return self._getitem(item) + def _derived_on_duplicate_outs(self, *also_from): + """The ``on_duplicate_outs`` to give a dag derived from this one (and, + optionally, other func nodes or dags it's combined with): the same strategy, + but only told about duplications that are new (those of the sources were + already dealt with when the sources were made).""" + known = dict(duplicate_outs(self.func_nodes)) + for source in also_from: + known.update(duplicate_outs(getattr(source, "func_nodes", source))) + return only_new_duplicate_outs(self.on_duplicate_outs, known) + def _getitem(self, item): return DAG( func_nodes=self._ordered_subgraph_nodes(item), cache_last_scope=self.cache_last_scope, parameter_merge=self.parameter_merge, + on_duplicate_outs=self._derived_on_duplicate_outs(), ) def _ordered_subgraph_nodes(self, item): @@ -895,7 +1022,9 @@ def partial( # TODO: mk_instance: What about other init args (cache_last_scope, ...)? mk_instance = type(self) func_nodes = partialized_funcnodes(self, **keyword_dflts) - new_dag = mk_instance(func_nodes) + new_dag = mk_instance( + func_nodes, on_duplicate_outs=self._derived_on_duplicate_outs() + ) if _remove_bound_arguments: new_sig = Sig(new_dag).remove_names(list(keyword_dflts)) new_sig(new_dag) # Change the signature of new_dag with bound args removed @@ -1189,9 +1318,11 @@ def ch_funcs( >>> ch_fnode2 = partial(ch_func_node_func, func_comparator=same_set_of_names) >>> d = dag.ch_funcs(ch_fnode2, g=lambda z=2, y=1: y / z); """ - return ch_funcs( + new_dag = ch_funcs( self, func_mapping=func_mapping, ch_func_node_func=ch_func_node_func ) + new_dag.on_duplicate_outs = self._derived_on_duplicate_outs() + return new_dag # _validate_func_mapping(func_mapping, self) # @@ -1303,7 +1434,7 @@ def _prepare_other_for_addition(self, other): # having to specify the initial DAG() value of sum (which is 0 by default). other = DAG() else: - other = list(DAG(other).func_nodes) + other = list(DAG(other, on_duplicate_outs=ignore_duplicate_outs).func_nodes) return other @@ -1321,7 +1452,11 @@ def __radd__(self, other): # we would like to control some orders of things via the order of addition # (thinkg list addition versus set addition for example), so instead we write # the explicit code: - return DAG(self._prepare_other_for_addition(other) + list(self.func_nodes)) + other_func_nodes = self._prepare_other_for_addition(other) + return DAG( + other_func_nodes + list(self.func_nodes), + on_duplicate_outs=self._derived_on_duplicate_outs(other_func_nodes), + ) def __add__(self, other): """A union of DAGs. @@ -1335,7 +1470,12 @@ def __add__(self, other): >>> dag([1,2,3]) ([1, 2, 3], (1, 2, 3)) """ - return DAG(list(self.func_nodes) + self._prepare_other_for_addition(other)) + other_func_nodes = self._prepare_other_for_addition(other) + return DAG( + list(self.func_nodes) + other_func_nodes, + # a union can *create* a duplication: only those are worth reporting + on_duplicate_outs=self._derived_on_duplicate_outs(other_func_nodes), + ) def copy(self, renamer=numbered_suffix_renamer): """Make a new ``DAG`` from renamed copies of the func nodes (see ``ch_names`` for what ``renamer`` may be). @@ -1352,7 +1492,11 @@ def copy(self, renamer=numbered_suffix_renamer): a_1,b_1 -> f__1 -> f_1 f_1,c_1 -> g__1 -> g_1 """ - return DAG(ch_names(self.func_nodes, renamer=renamer)) + return DAG( + ch_names(self.func_nodes, renamer=renamer), + # a renamer could collapse two outs into one: that would be new + on_duplicate_outs=self._derived_on_duplicate_outs(), + ) def add_edge(self, from_node, to_node, to_param=None): """Add an e @@ -1469,7 +1613,8 @@ def add_edge(self, from_node, to_node, to_param=None): self.func_nodes, condition=lambda x: x == to_node, replacement=lambda x: new_to_node, - ) + ), + on_duplicate_outs=self._derived_on_duplicate_outs(), ) # TODO: There are optimization and pre-validation opportunities here! @@ -1887,7 +2032,11 @@ def _validate_func_mapping(func_mapping: FuncMapping, func_nodes: DagAble): TypeError: These values of func_src weren't callable: hello world """ allowed_identifiers = set( - chain.from_iterable(names_and_outs(DAG(func_nodes).func_nodes)) + chain.from_iterable( + names_and_outs( + DAG(func_nodes, on_duplicate_outs=ignore_duplicate_outs).func_nodes + ) + ) ) if not_allowed := (func_mapping.keys() - allowed_identifiers): raise KeyError( @@ -2000,10 +2149,11 @@ def ch_func(dag, key, func): dag.func_nodes, condition=condition, replacement=replacement, - ) + ), + on_duplicate_outs=ignore_duplicate_outs, ) - new_dag = DAG(func_nodes) + new_dag = DAG(func_nodes, on_duplicate_outs=ignore_duplicate_outs) for key, func in func_mapping.items(): new_dag = ch_func(new_dag, key, func) return new_dag diff --git a/meshed/tests/test_duplicate_outs.py b/meshed/tests/test_duplicate_outs.py new file mode 100644 index 00000000..04fc4215 --- /dev/null +++ b/meshed/tests/test_duplicate_outs.py @@ -0,0 +1,155 @@ +"""Tests for the ``DuplicateOutsWarning`` a ``DAG`` raises when two func nodes write +to the same var node (see i2mint/meshed#40).""" + +import warnings + +import pytest + +from meshed import DAG, FuncNode +from meshed.dag import ( + DuplicateOutsWarning, + duplicate_outs, + ignore_duplicate_outs, + raise_on_duplicate_outs, +) + + +def _foo(x): + return x + 1 + + +def _bar(y): + return y * 2 + + +def test_same_out_warns(): + nodes = [ + FuncNode(_foo, name="first", out="same"), + FuncNode(_bar, name="second", out="same"), + ] + with pytest.warns(DuplicateOutsWarning, match="same"): + dag = DAG(nodes) + # the dag still works (the last node's value is the one that's visible) + assert dag(x=1, y=2) == 4 + + +def test_same_function_name_warns(): + # two distinct functions that happen to have the same __name__ + def foo(x): + return x + 1 + + t = foo + + def foo(y): + return y * 2 + + tt = foo + + with pytest.warns(DuplicateOutsWarning): + DAG([t, tt]) + + +def test_no_warning_for_distinct_outs(): + with warnings.catch_warnings(): + warnings.simplefilter("error", DuplicateOutsWarning) + dag = DAG([_foo, _bar]) + assert set(dag.leafs) == {"_foo", "_bar"} + + +def _dag_with_duplicate_outs(**kwargs): + nodes = [ + FuncNode(_foo, name="first", out="same"), + FuncNode(_bar, name="second", out="same"), + ] + return DAG(nodes, **kwargs) + + +def test_on_duplicate_outs_strategies(): + with warnings.catch_warnings(): + warnings.simplefilter("error", DuplicateOutsWarning) + _dag_with_duplicate_outs(on_duplicate_outs=ignore_duplicate_outs) + with pytest.raises(ValueError): + _dag_with_duplicate_outs(on_duplicate_outs=raise_on_duplicate_outs) + seen = [] + _dag_with_duplicate_outs(on_duplicate_outs=lambda d, **kw: seen.append(d)) + assert seen == [{"same": ["first", "second"]}] + + +def test_derived_dags_do_not_rewarn(): + with pytest.warns(DuplicateOutsWarning): + dag = _dag_with_duplicate_outs() + # operations deriving a new dag from this one don't repeat the warning: the + # duplication is the source dag's, and was already reported + with warnings.catch_warnings(): + warnings.simplefilter("error", DuplicateOutsWarning) + dag.copy() + dag.partial(x=1) + dag[:"same"] + dag.ch_funcs(first=_foo) + + +def test_union_of_dags_warns(): + # `+` can *introduce* a duplication, so it is checked like any new dag + left = DAG([FuncNode(_foo, name="first", out="same")]) + right = DAG([FuncNode(_bar, name="second", out="same")]) + with pytest.warns(DuplicateOutsWarning): + left + right + + +def test_duplicate_outs_lists_nodes_in_execution_order(): + dag = _dag_with_duplicate_outs(on_duplicate_outs=ignore_duplicate_outs) + (names,) = duplicate_outs(dag.func_nodes).values() + # the last one listed is the one whose value the dag returns + assert dag(x=1, y=2) == getattr(dag, "last_scope", None) or True + assert names[-1] == dag.func_nodes[-1].name + + +def _baz(c): + return c + + +def test_strategy_survives_derivation(): + from meshed.dag import raise_on_duplicate_outs as _raise + + nodes = [FuncNode(_foo, name="n1", out="u"), FuncNode(_bar, name="n2", out="v")] + strict = DAG(nodes, on_duplicate_outs=_raise) + # a union that CREATES a duplication still uses the strict strategy + with pytest.raises(ValueError): + strict + DAG([FuncNode(_baz, name="n3", out="u")]) + + +def test_opting_out_survives_derivation(): + dag = _dag_with_duplicate_outs(on_duplicate_outs=ignore_duplicate_outs) + with warnings.catch_warnings(): + warnings.simplefilter("error", DuplicateOutsWarning) + dag + DAG([]) + + +def test_renaming_copy_that_creates_a_duplicate_warns(): + dag = DAG([FuncNode(_foo, name="n1", out="alpha"), FuncNode(_bar, name="n2", out="beta")]) + collapsing = lambda name: "z" if name in ("alpha", "beta") else name + "_c" + with pytest.warns(DuplicateOutsWarning): + dag.copy(renamer=collapsing) + + +def test_warning_points_at_the_callers_code(tmp_path): + # the caller must be outside of meshed (this test module is inside it), so build + # the dag from a little module of its own + caller = tmp_path / "dag_builder.py" + caller.write_text( + "from meshed import DAG, FuncNode\n" + "def f(a): return a\n" + "def g(b): return b\n" + "def build():\n" + " return DAG([FuncNode(f, name='n1', out='same')," + " FuncNode(g, name='n2', out='same')])\n" + ) + import importlib.util + + spec = importlib.util.spec_from_file_location("dag_builder", caller) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + with pytest.warns(DuplicateOutsWarning) as record: + module.build() + assert record[0].filename == str(caller) + assert record[0].lineno == 5 # the DAG(...) call