Connectivity-Preserving Wedge Token Pooling / run_experiment.py
Mechanism confirmed, baseline not beaten
1import json, time
2import numpy as np
3from wedge_pool import pool_wedge, random_pool, kmeans_pool, connected, sse
4
5
6def graph_with_two_regions(n_each=12):
7 # Two noisy cliques joined by one bridge: connected graph with a clear graph signal.
8 n = 2*n_each
9 adj = [set() for _ in range(n)]
10 for lo in (0, n_each):
11 hi = lo+n_each
12 for i in range(lo, hi):
13 for j in range(i+1, hi):
14 adj[i].add(j); adj[j].add(i)
15 adj[n_each-1].add(n_each); adj[n_each].add(n_each-1)
16 return [sorted(x) for x in adj]
17
18
19def recon_sse(X, pooled):
20 return float(((X-pooled)**2).sum())
21
22
23def main():
24 rng = np.random.default_rng(7)
25 adj = graph_with_two_regions(12)
26 n = len(adj); d = 4
27 # Smooth within-region signal, with a strong graph-aligned discontinuity.
28 X = np.zeros((n,d), float)
29 X[:12] = np.array([1., 0., .2, -.1]) + .10*rng.normal(size=(12,d))
30 X[12:] = np.array([-1., 0., .2, -.1]) + .10*rng.normal(size=(12,d))
31
32 # Direct numerical claim checks: mean beats arbitrary constants and children stay connected.
33 region = list(range(n))
34 global_cost, mu = sse(X, region)
35 alt_cost = float(((X-(mu+0.37))**2).sum())
36 out = {"mean_optimality_gap": alt_cost-global_cost, "connectivity": {}, "results": {}}
37 for M in (2, 4, 8):
38 w = pool_wedge(adj, X, M)
39 assert all(connected(adj, R) for R in w["regions"])
40 assert len(w["regions"]) == M
41 out["connectivity"][str(M)] = True
42 methods = {"wedge": w["pooled"]}
43 rm, rp, _ = random_pool(X, M, np.random.default_rng(100+M))
44 km, kp, _ = kmeans_pool(X, M, seed=100+M)
45 methods["random"] = rp; methods["kmeans"] = kp
46 vals = {}
47 for name, recon in methods.items():
48 t0=time.perf_counter();
49 for _ in range(30):
50 if name == "wedge": pool_wedge(adj, X, M)
51 elif name == "kmeans": kmeans_pool(X, M, seed=100+M)
52 else: random_pool(X, M, np.random.default_rng(100+M))
53 elapsed=(time.perf_counter()-t0)/30
54 vals[name] = {"reconstruction_sse": recon_sse(X,recon), "pool_seconds": elapsed,
55 "attention_quadratic_ratio": (M/n)**2}
56 out["results"][str(M)] = vals
57 print(json.dumps(out, indent=2))
58
59if __name__ == '__main__': main()