import json, time import numpy as np from wedge_pool import pool_wedge, random_pool, kmeans_pool, connected, sse def graph_with_two_regions(n_each=12): # Two noisy cliques joined by one bridge: connected graph with a clear graph signal. n = 2*n_each adj = [set() for _ in range(n)] for lo in (0, n_each): hi = lo+n_each for i in range(lo, hi): for j in range(i+1, hi): adj[i].add(j); adj[j].add(i) adj[n_each-1].add(n_each); adj[n_each].add(n_each-1) return [sorted(x) for x in adj] def recon_sse(X, pooled): return float(((X-pooled)**2).sum()) def main(): rng = np.random.default_rng(7) adj = graph_with_two_regions(12) n = len(adj); d = 4 # Smooth within-region signal, with a strong graph-aligned discontinuity. X = np.zeros((n,d), float) X[:12] = np.array([1., 0., .2, -.1]) + .10*rng.normal(size=(12,d)) X[12:] = np.array([-1., 0., .2, -.1]) + .10*rng.normal(size=(12,d)) # Direct numerical claim checks: mean beats arbitrary constants and children stay connected. region = list(range(n)) global_cost, mu = sse(X, region) alt_cost = float(((X-(mu+0.37))**2).sum()) out = {"mean_optimality_gap": alt_cost-global_cost, "connectivity": {}, "results": {}} for M in (2, 4, 8): w = pool_wedge(adj, X, M) assert all(connected(adj, R) for R in w["regions"]) assert len(w["regions"]) == M out["connectivity"][str(M)] = True methods = {"wedge": w["pooled"]} rm, rp, _ = random_pool(X, M, np.random.default_rng(100+M)) km, kp, _ = kmeans_pool(X, M, seed=100+M) methods["random"] = rp; methods["kmeans"] = kp vals = {} for name, recon in methods.items(): t0=time.perf_counter(); for _ in range(30): if name == "wedge": pool_wedge(adj, X, M) elif name == "kmeans": kmeans_pool(X, M, seed=100+M) else: random_pool(X, M, np.random.default_rng(100+M)) elapsed=(time.perf_counter()-t0)/30 vals[name] = {"reconstruction_sse": recon_sse(X,recon), "pool_seconds": elapsed, "attention_quadratic_ratio": (M/n)**2} out["results"][str(M)] = vals print(json.dumps(out, indent=2)) if __name__ == '__main__': main()