diff --git a/meshed/scrap/cached_dag.py b/meshed/scrap/cached_dag.py index f6d4d74b..73a5cb99 100644 --- a/meshed/scrap/cached_dag.py +++ b/meshed/scrap/cached_dag.py @@ -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) @@ -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): @@ -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}" @@ -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) @@ -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): diff --git a/meshed/tests/test_cached_dag.py b/meshed/tests/test_cached_dag.py new file mode 100644 index 00000000..33e5d2eb --- /dev/null +++ b/meshed/tests/test_cached_dag.py @@ -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