File size: 3,421 Bytes
81ae663
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
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