Connectivity-Preserving Wedge Token Pooling / run_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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()