Download tests/test_math.py from ChatterjeeLab/DooABLe: direct link, hf CLI and curl.
- Browser
- Download file 3.42 kB
-
https://huggingface.co/ChatterjeeLab/DooABLe/resolve/main/tests/test_math.py
- Command line
-
hf download hf://ChatterjeeLab/DooABLe/tests/test_math.py
-
curl -L -o test_math.py https://huggingface.co/ChatterjeeLab/DooABLe/resolve/main/tests/test_math.py
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 | |
| 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 | |