ProCreations's picture
Audit group Steiner extension exactly
8005f3e
Raw
History Blame Contribute Delete
6.7 kB
#!/usr/bin/env python3
"""Exact finite group-Steiner/forest audit for the learning-augmented protocol."""
from __future__ import annotations
import heapq
import json
import numpy as np
def connected_nodes(edges: list[tuple[int, int, float]], chosen: set[int], root: int = 0) -> set[int]:
graph: dict[int, list[int]] = {}
for index in chosen:
left, right, _ = edges[index]
graph.setdefault(left, []).append(right)
graph.setdefault(right, []).append(left)
seen = {root}
stack = [root]
while stack:
node = stack.pop()
for neighbor in graph.get(node, []):
if neighbor not in seen:
seen.add(neighbor)
stack.append(neighbor)
return seen
def exact_group_steiner(
edges: list[tuple[int, int, float]], groups: tuple[frozenset[int], ...]
) -> tuple[float, set[int]]:
best_cost = float("inf")
best: set[int] = set()
for mask in range(1 << len(edges)):
chosen = {index for index in range(len(edges)) if mask & (1 << index)}
cost = sum(edges[index][2] for index in chosen)
if cost >= best_cost:
continue
reached = connected_nodes(edges, chosen)
if all(reached & set(group) for group in groups):
best_cost = cost
best = chosen
if not best:
raise RuntimeError("group instance has no feasible tree")
return best_cost, best
def shortest_path_to_group(
edges: list[tuple[int, int, float]], reached: set[int], group: frozenset[int]
) -> list[int]:
adjacency: dict[int, list[tuple[int, int, float]]] = {}
for index, (left, right, weight) in enumerate(edges):
adjacency.setdefault(left, []).append((right, index, weight))
adjacency.setdefault(right, []).append((left, index, weight))
queue: list[tuple[float, int]] = [(0.0, node) for node in reached]
distance = {node: 0.0 for node in reached}
previous: dict[int, tuple[int, int]] = {}
while queue:
cost, node = heapq.heappop(queue)
if cost != distance[node]:
continue
if node in group:
path: list[int] = []
current = node
while current not in reached:
parent, edge_index = previous[current]
path.append(edge_index)
current = parent
return path
for neighbor, edge_index, weight in adjacency.get(node, []):
proposal = cost + weight
if proposal < distance.get(neighbor, float("inf")):
distance[neighbor] = proposal
previous[neighbor] = (node, edge_index)
heapq.heappush(queue, (proposal, neighbor))
raise RuntimeError("group is unreachable")
def online_group_connector(
edges: list[tuple[int, int, float]],
groups: tuple[frozenset[int], ...],
order: tuple[int, ...],
preload: set[int],
) -> float:
chosen = set(preload)
for group_index in order:
reached = connected_nodes(edges, chosen)
if not reached & set(groups[group_index]):
chosen.update(shortest_path_to_group(edges, reached, groups[group_index]))
return sum(edges[index][2] for index in chosen)
def graph_family(family: int) -> tuple[list[tuple[int, int, float]], tuple[frozenset[int], ...]]:
rng = np.random.default_rng(20260729 + family)
nodes = 7
edges: list[tuple[int, int, float]] = [(0, node, 1.0 + 0.05 * node) for node in (1, 2)]
for left in range(1, nodes):
for right in range(left + 1, nodes):
if len(edges) >= 11:
break
if rng.random() < 0.48:
edges.append((left, right, float(rng.integers(2, 9) / 10.0)))
if len(edges) >= 11:
break
for node in range(3, nodes):
if not any(node in edge[:2] for edge in edges):
edges.append((1, node, 0.7 + 0.03 * node))
groups = (
frozenset((1, 3, 4)),
frozenset((2, 4, 5)),
frozenset((3, 5, 6)),
frozenset((1, 5, 6)),
)
return edges, groups
def main() -> None:
rows: list[dict[str, object]] = []
for family in range(12):
edges, groups = graph_family(family)
optimum, optimal_edges = exact_group_steiner(edges, groups)
for corruption in (0, 1, 2, 3):
predicted = set(optimal_edges)
available = [index for index in range(len(edges)) if index not in predicted]
for index in range(corruption):
if predicted:
predicted.remove(sorted(predicted)[index % len(predicted)])
predicted.add(available[(family + index) % len(available)])
eta = len(optimal_edges.symmetric_difference(predicted))
order = tuple((family + shift) % len(groups) for shift in range(len(groups)))
baseline = online_group_connector(edges, groups, order, set())
follow = online_group_connector(edges, groups, order, predicted)
combined = min(baseline, follow)
rows.append({
"family": family,
"nodes": 7,
"edges": len(edges),
"groups": len(groups),
"corruption": corruption,
"eta_symmetric_difference": eta,
"exact_group_steiner_optimum": optimum,
"baseline_cost": baseline,
"prediction_following_cost": follow,
"combined_cost": combined,
"combined_ratio": combined / optimum,
})
perfect = [row for row in rows if row["corruption"] == 0]
summary = {
"cells": len(rows),
"graph_families": 12,
"nodes": 7,
"groups_per_instance": 4,
"exact_edge_subset_optima": 12,
"corruption_levels": [0, 1, 2, 3],
"perfect_prediction_ratio_is_one": all(abs(row["combined_ratio"] - 1.0) < 1e-12 for row in perfect),
"combiner_never_worse_than_baseline": all(row["combined_cost"] <= row["baseline_cost"] + 1e-12 for row in rows),
"all_groups_covered": all(row["combined_cost"] >= row["exact_group_steiner_optimum"] - 1e-12 for row in rows),
"max_combined_ratio": max(row["combined_ratio"] for row in rows),
"eta_values": sorted({row["eta_symmetric_difference"] for row in rows}),
}
print(json.dumps({"schema": "exact-group-steiner-audit-v1", "summary": summary}, indent=2, sort_keys=True))
if not all(summary[key] for key in ("perfect_prediction_ratio_is_one", "combiner_never_worse_than_baseline", "all_groups_covered")):
raise SystemExit("group-Steiner audit gate failed")
if __name__ == "__main__":
main()