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
11 changes: 10 additions & 1 deletion meshed/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
170 changes: 160 additions & 10 deletions meshed/dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 "<string>" for a filename
is_generated = frame.f_code.co_filename == "<string>"
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)
Expand Down Expand Up @@ -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))
Expand All @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
#
Expand Down Expand Up @@ -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

Expand All @@ -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.
Expand All @@ -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).
Expand All @@ -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
Expand Down Expand Up @@ -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!
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down
Loading
Loading