DooABLe / tests /test_math.py
pranamanam's picture
Upload 309 files
81ae663 verified
Raw
History Blame Contribute Delete
3.42 kB
import numpy as np
import pytest
from dooable.graph import Graph, Node, Edge, toy_graph, grid_graph
from dooable.exact import (
solve,
endpoint_distribution,
sample,
expected_cost,
prefix_values,
backward_policy,
forward_from_backward,
)
def enumerate_paths(graph):
out = []
def visit(u, path, cost):
if not graph.outgoing[u]:
out.append((graph.nodes[u].outcome, path, cost))
return
for i in graph.outgoing[u]:
e = graph.edges[i]
visit(graph.node_index[e.target], path + [i], cost + e.cost)
visit(graph.root_index, [], 0.0)
return out
@pytest.mark.parametrize("multiplicity", [1, 2, 8, 16])
@pytest.mark.parametrize("temperature", [0.1, 1.0, 3.0])
def test_joint_law_against_route_enumeration(multiplicity, temperature):
g = toy_graph(multiplicity)
r = {"A": np.log(3), "B": 0.0}
s = solve(g, r, temperature)
paths = enumerate_paths(g)
for y, path, cost in paths:
partition = sum(np.exp(-c / temperature) for yy, _, c in paths if yy == y)
expected = s.target[y] * np.exp(-cost / temperature) / partition
assert np.prod(s.forward[path]) == pytest.approx(expected, abs=1e-11)
assert endpoint_distribution(g, s.forward) == pytest.approx({"A": 0.75, "B": 0.25})
def test_backward_error_changes_routes_and_preserves_endpoints():
g = toy_graph(8)
s = solve(g, {"A": 0.0, "B": 0.0}, 0.5)
v = prefix_values(g, 0.5)
v = v + np.random.default_rng(4).normal(size=len(v)) * 3
v[g.root_index] = 0
q = backward_policy(g, v, 0.5)
p = forward_from_backward(g, q, s.target)
assert endpoint_distribution(g, p) == pytest.approx(s.target)
assert expected_cost(g, p) != pytest.approx(expected_cost(g, s.forward))
def test_budgeted_cycles_are_acyclic_after_augmentation():
g = grid_graph(4, 5)
s = solve(g, {y: 0.0 for y in g.terminals})
for route in sample(g, s.forward, 200, 3):
assert sum(a.get("kind") == "move" for a in route["actions"]) <= 5
def test_cycle_rejected():
with pytest.raises(ValueError):
Graph(
[Node("root"), Node("a"), Node("b"), Node("t", "x")],
[
Edge("1", "root", "a"),
Edge("2", "a", "b"),
Edge("3", "b", "a"),
Edge("4", "b", "t"),
],
)
def test_route_temperature_reduces_expected_cost():
g = toy_graph(16)
r = {"A": 0.0, "B": 0.0}
assert expected_cost(g, solve(g, r, 0.05).forward) < expected_cost(
g, solve(g, r, 2.0).forward
)
def test_finite_residual_joint_bound():
g = toy_graph(5)
s = solve(g, {"A": 0.8, "B": -0.3}, 0.7)
v = s.log_prefix + 0.05 * np.sin(np.arange(len(g.nodes)))
v[g.root_index] = 0
q = backward_policy(g, v, 0.7)
p = forward_from_backward(g, q, s.target)
from scipy.special import logsumexp
eps = max(
abs(
v[j]
- logsumexp(
[
v[g.node_index[g.edges[i].source]] - g.edges[i].cost / 0.7
for i in ids
]
)
)
for j, ids in enumerate(g.incoming)
if ids
)
kl = 0.0
for _, path, _ in enumerate_paths(g):
pp = np.prod(p[path])
qq = np.prod(s.forward[path])
kl += pp * np.log(pp / qq)
assert kl <= 2 * 2 * eps + 1e-12