Download src/dooable/exact.py from ChatterjeeLab/DooABLe: direct link, hf CLI and curl.
- Browser
- Download file 6.53 kB
-
https://huggingface.co/ChatterjeeLab/DooABLe/resolve/main/src/dooable/exact.py
- Command line
-
hf download hf://ChatterjeeLab/DooABLe/src/dooable/exact.py
-
curl -L -o exact.py https://huggingface.co/ChatterjeeLab/DooABLe/resolve/main/src/dooable/exact.py
6.53 kB
| """Sparse dynamic programs for endpoint-normalized path laws. | |
| No root-to-terminal path enumeration is needed. Runtime is linear in the | |
| materialized graph size. Full catalog expansion remains a separate cost. | |
| """ | |
| from dataclasses import dataclass | |
| import numpy as np | |
| from scipy.special import logsumexp | |
| from .graph import Graph | |
| class Solution: | |
| """Exact node log partitions, edge probabilities, and endpoint weights.""" | |
| log_prefix: np.ndarray | |
| backward: np.ndarray | |
| forward: np.ndarray | |
| target: dict[str, float] | |
| log_z: float | |
| temperature: float | |
| def normalize_rewards(graph, log_rewards): | |
| """Normalize finite log weights over canonical outcomes and return log Z.""" | |
| keys = set(graph.terminals) | |
| if set(log_rewards) != keys: | |
| raise ValueError("Reward keys must equal the terminal outcome set") | |
| vals = np.array([log_rewards[y] for y in keys], dtype=float) | |
| if not np.all(np.isfinite(vals)): | |
| raise ValueError("Log rewards must be finite") | |
| z = float(logsumexp(vals)) | |
| return {y: float(np.exp(log_rewards[y] - z)) for y in graph.terminals}, z | |
| def prefix_values(graph, temperature=1.0): | |
| """Compute one log partition per node in topological order.""" | |
| if not np.isfinite(temperature) or temperature <= 0: | |
| raise ValueError("Temperature must be positive") | |
| a = np.full(len(graph.nodes), -np.inf) | |
| a[graph.root_index] = 0.0 | |
| for v in graph.order: | |
| if v == graph.root_index: | |
| continue | |
| a[v] = logsumexp( | |
| [ | |
| a[graph.node_index[graph.edges[i].source]] | |
| - graph.edges[i].cost / temperature | |
| for i in graph.incoming[v] | |
| ] | |
| ) | |
| return a | |
| def backward_policy(graph, values, temperature): | |
| """Local normalization preserves a complete conditional route law.""" | |
| q = np.zeros(len(graph.edges)) | |
| for v in graph.order: | |
| ids = graph.incoming[v] | |
| if not ids: | |
| continue | |
| logits = np.array( | |
| [ | |
| values[graph.node_index[graph.edges[i].source]] | |
| - graph.edges[i].cost / temperature | |
| for i in ids | |
| ] | |
| ) | |
| q[ids] = np.exp(logits - logsumexp(logits)) | |
| return q | |
| def forward_from_backward(graph, q, target): | |
| """Propagate endpoint probability backward, then reverse each edge flow.""" | |
| logmass = np.full(len(graph.nodes), -np.inf) | |
| for y, v in graph.terminals.items(): | |
| logmass[v] = np.log(target[y]) | |
| for v in reversed(graph.order): | |
| for ei in graph.incoming[v]: | |
| u = graph.node_index[graph.edges[ei].source] | |
| if q[ei] > 0: | |
| logmass[u] = np.logaddexp(logmass[u], logmass[v] + np.log(q[ei])) | |
| p = np.zeros(len(graph.edges)) | |
| for u in graph.order: | |
| ids = graph.outgoing[u] | |
| if not ids: | |
| continue | |
| lp = np.array( | |
| [ | |
| np.log(q[i]) + logmass[graph.node_index[graph.edges[i].target]] | |
| for i in ids | |
| ] | |
| ) | |
| p[ids] = np.exp(lp - logsumexp(lp)) | |
| return p | |
| def solve(graph: Graph, log_rewards, temperature=1.0): | |
| """Construct the exact joint law on a finite stored executable graph.""" | |
| target, z = normalize_rewards(graph, log_rewards) | |
| a = prefix_values(graph, temperature) | |
| q = backward_policy(graph, a, temperature) | |
| return Solution( | |
| a, q, forward_from_backward(graph, q, target), target, z, temperature | |
| ) | |
| def endpoint_distribution(graph, forward): | |
| """Propagate a normalized edge policy and return exact terminal masses.""" | |
| mass = np.zeros(len(graph.nodes)) | |
| mass[graph.root_index] = 1.0 | |
| for u in graph.order: | |
| ids = graph.outgoing[u] | |
| if ids and not np.isclose(np.sum(forward[ids]), 1, atol=1e-6): | |
| raise ValueError("Invalid forward normalization") | |
| for ei in ids: | |
| mass[graph.node_index[graph.edges[ei].target]] += mass[u] * forward[ei] | |
| return {y: float(mass[v]) for y, v in graph.terminals.items()} | |
| def sample(graph, forward, n=100, seed=0): | |
| """Draw n complete paths, including edge IDs, action records, and costs.""" | |
| rng = np.random.default_rng(seed) | |
| result = [] | |
| for _ in range(n): | |
| u = graph.root_index | |
| path = [] | |
| cost = 0.0 | |
| lp = 0.0 | |
| while graph.outgoing[u]: | |
| ids = graph.outgoing[u] | |
| probs = np.asarray(forward[ids]) | |
| probs = probs / probs.sum() | |
| j = int(rng.choice(len(ids), p=probs)) | |
| ei = ids[j] | |
| e = graph.edges[ei] | |
| path.append(ei) | |
| cost += e.cost | |
| lp += np.log(probs[j]) | |
| u = graph.node_index[e.target] | |
| result.append( | |
| { | |
| "outcome": graph.nodes[u].outcome, | |
| "edge_ids": [graph.edges[i].id for i in path], | |
| "edge_indices": path, | |
| "cost": cost, | |
| "log_probability": float(lp), | |
| "actions": [graph.edges[i].action for i in path], | |
| } | |
| ) | |
| return result | |
| def expected_cost(graph, forward): | |
| """Compute mean additive execution cost by propagating node masses.""" | |
| mass = np.zeros(len(graph.nodes)) | |
| mass[graph.root_index] = 1.0 | |
| cost = 0.0 | |
| for u in graph.order: | |
| for ei in graph.outgoing[u]: | |
| e = graph.edges[ei] | |
| f = mass[u] * forward[ei] | |
| mass[graph.node_index[e.target]] += f | |
| cost += f * e.cost | |
| return float(cost) | |
| def uniform_policy(graph): | |
| """Return equal probabilities across each node's outgoing actions.""" | |
| p = np.zeros(len(graph.edges)) | |
| for ids in graph.outgoing: | |
| if ids: | |
| p[ids] = 1 / len(ids) | |
| return p | |
| def tilted_reference(graph, log_rewards, temperature=None): | |
| """Reference endpoint tilting; optional path-cost penalty is explicit.""" | |
| p = uniform_policy(graph) | |
| h = np.full(len(graph.nodes), -np.inf) | |
| for y, v in graph.terminals.items(): | |
| h[v] = log_rewards[y] | |
| for u in reversed(graph.order): | |
| ids = graph.outgoing[u] | |
| if ids: | |
| logits = np.array( | |
| [ | |
| np.log(p[i]) | |
| + h[graph.node_index[graph.edges[i].target]] | |
| - (graph.edges[i].cost / temperature if temperature else 0) | |
| for i in ids | |
| ] | |
| ) | |
| h[u] = logsumexp(logits) | |
| p[ids] = np.exp(logits - h[u]) | |
| return p | |