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
133 changes: 120 additions & 13 deletions meshed/scrap/cached_dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,27 @@ def __setitem__(self, key, value):
NoSuchKey = type("NoSuchKey", (), {})


def _is_same_value(a, b):
"""Whether ``a`` and ``b`` are the same value, never raising.

Identity counts as sameness (so a ``nan`` object equals itself). Values whose
``==`` doesn't give a plain truth value (e.g. numpy arrays) are only considered
the same if they are identical.

>>> _is_same_value(1, 1), _is_same_value(1, 2)
(True, False)
>>> nan = float('nan')
>>> _is_same_value(nan, nan), _is_same_value(nan, float('nan'))
(True, False)
"""
if a is b:
return True
try:
return bool(a == b)
except Exception:
return False


# TODO: Cache validation and invalidation
# TODO: Continue constructing uppward towards lazyprop-using class (instances are
# varnodes)
Expand Down Expand Up @@ -227,7 +248,11 @@ def __init__(self, dag, cache=True, name=None):
"This type of cache is not implemented (must resolve to a Mapping): "
f"{cache=}"
)
self._cache = ChainMap(self.defaults, self.cache)
else:
self.cache = cache
# Cached values (computed outputs and inputs the dag was called with) take
# precedence over the dag's defaults.
self._cache = ChainMap(self.cache, self.defaults)

@property
def __name__(self):
Expand All @@ -244,37 +269,110 @@ def func_node_id(self, k):
# TODO: Consider having args and kwargs instead of just input_kwargs.
# or making it (k, /, *args, **kwargs)
def __call__(self, k, /, **input_kwargs):
# print(f"Calling ({k=},{input_kwargs=})\t{self.cache=}")
input_kwargs = dict(input_kwargs)
if intersection := (input_kwargs.keys() & self.cache.keys()):
self._validate_inputs_against_cache(input_kwargs)
keys_before = set(self.cache)
try:
output = self._compute(k, input_kwargs)
except BaseException:
# Roll back the outputs computed during this failed call, so the cache
# stays consistent with its (uncached) inputs and the call can be retried.
try:
for key in set(self.cache) - keys_before:
del self.cache[key]
except Exception: # pragma: no cover - e.g. a cache without __delitem__
pass # never mask the original error with a rollback error
raise
# Only persist inputs once they led to a successful computation, so that a
# failed (e.g. mistaken) call doesn't pin values in the cache.
self._cache_inputs(input_kwargs)
return output

def _validate_inputs_against_cache(self, input_kwargs):
"""Raise a ``ValueError`` if ``input_kwargs`` contradicts the cache.

An input contradicts the cache if:

- it is already cached with a different value, or
- it is not cached, but cached outputs downstream of it were computed using
its default (or it has no default), and the given value differs from it, or
- it is not cached and not a root, but the cache (and the dag's defaults)
already determine everything it would be computed from.

Note that inputs are only validated against the *cache*: values given in the
same call are not checked against each other (``c('h', f=100, a=1)`` is
accepted even if ``f`` wouldn't be computed as ``100`` from ``a=1``).

Values are compared with ``_is_same_value``: for values without a plain
``==`` truth value (e.g. numpy arrays), only the very same object counts as
the same value.
"""
conflicts = {
name
for name in input_kwargs.keys() & self.cache.keys()
if not _is_same_value(input_kwargs[name], self.cache[name])
}
if conflicts:
# TODO: Can give the user a more informative/correct message, since the
# user has more options than just the root nodes: They some combination of
# intermediates would also satisfy requirements.
raise ValueError(
f"input_kwargs can't contain any keys that are already in cache! "
f"These names were in both: {intersection}"
f"input_kwargs can't contain keys that are already in cache with a "
f"different value! These names were in both: {conflicts}"
)
for name in input_kwargs.keys() - self.cache.keys():
if name not in self.var_nodes:
continue
cached_downstream = descendants(self.dag.graph_ids, [name]) & set(
self.cache
)
if cached_downstream and not _is_same_value(
input_kwargs[name], self.defaults.get(name, NoSuchKey)
):
raise ValueError(
f"The value given for {name!r} contradicts the cache: "
f"{sorted(cached_downstream)} were already computed without it."
)
if name not in self.roots and self._is_determined_by_cache(name):
cached_upstream = descendants(self.reversed_graph, [name]) & set(
self.cache
)
raise ValueError(
f"{name!r} is already determined by the cache "
f"({sorted(cached_upstream)}), so it can't be given as an input."
)

def _is_determined_by_cache(self, k):
"""Whether ``k``'s value is already fixed by the cache and the dag's defaults
(that is, whether ``self(k)``, with no inputs, would give a value)."""
if k in self.cache or k in self.defaults:
return True
func_node_id = self.func_node_id(k)
if func_node_id is None: # a root node with no value in sight
return False
func_node = self.func_node_of_id[func_node_id]
return all(
self._is_determined_by_cache(src) for src in func_node.bind.values()
)

def _compute(self, k, input_kwargs):
_cache = ChainMap(input_kwargs, self._cache)
if k in _cache:
return _cache[k]
input_kwargs = dict(input_kwargs)
func_node_id = self.func_node_id(k)
# print(f"{func_node_id=}")
if func_node_id:
if (output := self.cache.get(func_node_id)) is not None:
return output
else:
func_node = self.func_node_of_id[func_node_id]
input_sources = {
src: self(src, **input_kwargs) for src in func_node.bind.values()
src: self._compute(src, input_kwargs)
for src in func_node.bind.values()
}
# inputs = dict(input_sources, **input_kwargs) #
# TODO: do we need to include **self.defaults in the middle?
inputs = ChainMap(_cache, input_sources)
# print(f"Computing {func_node_id}: ", end=" ")
output = func_node.call_on_scope(inputs, write_output_into_scope=False)
self.cache[func_node_id] = output
# print(f"result -> {output}")
return output
else: # k is a root node
assert k in self.roots, f"Was expecting this to be a root node: {k}"
Expand All @@ -287,6 +385,15 @@ def __call__(self, k, /, **input_kwargs):
f"argument: '{k}'"
)

def _cache_inputs(self, input_kwargs):
"""Persist the (explicitly given) values of the dag's var nodes in the cache,
so later calls can reuse them (see https://github.com/i2mint/meshed/issues/34).
Keys that are not var nodes of the dag, or that are already cached (and were
validated to hold the same value), are not written."""
for name, value in input_kwargs.items():
if name in self.var_nodes and name not in self.cache:
self.cache[name] = value

def _call(self, k, /, **kwargs):
return self(k, **kwargs)

Expand Down Expand Up @@ -359,9 +466,9 @@ def g(a, y=2):
dag = DAG([f, g])

c = CachedDag(dag)
c("g", a=1)
assert c("g", a=1) == 2
assert c.cache == {"g": 2, "a": 1}
assert c("f" == 2)
assert c("f") == 2


def add(a, b=1):
Expand Down
190 changes: 190 additions & 0 deletions meshed/tests/test_cached_dag.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,190 @@
"""Tests for ``meshed.scrap.cached_dag.CachedDag`` (see i2mint/meshed#34)."""

import pytest

from meshed import DAG
from meshed.scrap.cached_dag import CachedDag, cached_dag_test


def f(a, x=1):
return a + x


def g(a, y=2):
return a * y


def _cached_dag(**kwargs):
return CachedDag(DAG([f, g]), **kwargs)


def test_cached_dag_test_function():
cached_dag_test()


def test_inputs_are_cached_and_reused():
c = _cached_dag()
assert c("g", a=1) == 2
assert c.cache == {"g": 2, "a": 1}
assert c("f") == 2 # uses the cached ``a``
assert c.cache == {"g": 2, "a": 1, "f": 2}


def test_repeating_an_input_with_the_same_value_is_allowed():
c = _cached_dag()
c("g", a=1)
assert c("f", a=1) == 2


def test_conflicting_input_raises():
c = _cached_dag()
c("g", a=1)
with pytest.raises(ValueError):
c("f", a=10)


def test_cached_input_takes_precedence_over_default():
c = _cached_dag()
assert c("g", a=1, y=5) == 5
assert c.cache["y"] == 5
assert c("y") == 5 # not the dag's default (2)


def test_non_var_node_inputs_are_not_cached():
c = _cached_dag()
c("g", a=1, not_a_node=3)
assert "not_a_node" not in c.cache


def test_failed_call_does_not_cache_inputs():
c = _cached_dag()
with pytest.raises(TypeError):
c("g", y=5) # missing ``a``
assert c.cache == {}


def test_custom_mapping_cache():
cache = {}
c = _cached_dag(cache=cache)
assert c("g", a=1) == 2
assert cache == {"g": 2, "a": 1}


def k2(y, a):
return y - a


def test_failed_call_does_not_cache_inputs_regardless_of_arg_order():
c = CachedDag(DAG([k2]))
with pytest.raises(TypeError):
c("k2", y=5) # missing ``a``, which comes *after* ``y``
assert c.cache == {}


def h(f, g):
return f + g


def test_multi_level_dag_with_values_without_plain_equality():
class Arr:
"""Stand-in for a numpy array: ``==`` has an ambiguous truth value."""

def __init__(self, v):
self.v = v

def __add__(self, other):
return Arr(self.v + (other.v if isinstance(other, Arr) else other))

__radd__ = __add__

def __mul__(self, other):
return Arr(self.v * other)

def __eq__(self, other):
class Ambiguous:
def __bool__(self):
raise ValueError("ambiguous")

return Ambiguous()

__ne__ = __eq__

c = CachedDag(DAG([f, g, h]))
a = Arr(1)
assert c("h", a=a).v == 4 # (1 + 1) + (1 * 2)
assert c.cache["a"] is a
assert c("h", a=a).v == 4 # same object again: allowed


def test_nan_input():
nan = float("nan")
c = _cached_dag()
out = c("f", a=nan)
assert out != out # nan
c("g", a=nan) # the same nan object is accepted again


def test_input_contradicting_cached_outputs_raises():
c = _cached_dag()
c("f", a=1) # computed with the default x=1
assert c("f", a=1, x=1) == 2 # same as the default used: fine
with pytest.raises(ValueError):
c("f", a=1, x=5) # f was computed with x=1


def test_intermediate_node_as_input():
c = CachedDag(DAG([f, g, h]))
assert c("h", f=100, a=1) == 102
assert c.cache["f"] == 100
with pytest.raises(ValueError):
c("h", f=5)


def k(f, b):
return f + b


def test_failed_call_can_be_retried():
c = CachedDag(DAG([f, k]))
with pytest.raises(TypeError):
c("k", a=1) # missing ``b`` (after ``f`` was computed)
assert c.cache == {} # partial results were rolled back
assert c("k", a=1, b=3) == 5
assert c.cache == {"f": 2, "k": 5, "a": 1, "b": 3}


def test_exception_in_function_can_be_retried():
calls = []

def flaky(a):
calls.append(a)
if len(calls) == 1:
raise RuntimeError("transient")
return a

def total(f, flaky):
return f + flaky

c = CachedDag(DAG([f, flaky, total]))
with pytest.raises(RuntimeError):
c("total", a=1)
assert c("total", a=1) == 3


def test_input_determined_by_cached_upstream_values_raises():
c = CachedDag(DAG([f, g, h]))
c("f", a=1) # caches a=1, which determines g (and therefore h)
with pytest.raises(ValueError):
c("h", g=99)


def test_input_not_determined_by_cache_is_allowed():
def u(t):
return t + 1

def v(u, s):
return u + s

c = CachedDag(DAG([u, v]))
c("u", t=1) # caches t, u -- but ``v`` also needs ``s``, so it isn't determined
assert c("v", v=10) == 10
Loading