Connectivity-Preserving Wedge Token Pooling / verify_math.py
Mechanism confirmed, baseline not beaten
1import numpy as np
2from wedge_pool import all_pairs_dist, wedge_split, connected, sse
3from collections import deque
4
5
6def connected_graph(adj):
7 seen={0}; q=deque([0])
8 while q:
9 v=q.popleft()
10 for z in adj[v]:
11 if z not in seen: seen.add(z); q.append(z)
12 return len(seen)==len(adj)
13
14
15def main():
16 rng=np.random.default_rng(123)
17 failures=[]; checks=0
18 for trial in range(500):
19 n=int(rng.integers(5,14))
20 adj=[set() for _ in range(n)]
21 # random connected backbone plus extra edges
22 for i in range(1,n):
23 j=int(rng.integers(i)); adj[i].add(j); adj[j].add(i)
24 for i in range(n):
25 for j in range(i):
26 if rng.random()<.18: adj[i].add(j); adj[j].add(i)
27 adj=[sorted(x) for x in adj]
28 R=list(range(n)); D=all_pairs_dist(adj,R)
29 for u in range(n):
30 for w in range(u+1,n):
31 A,B=wedge_split(adj,R,u,w,D)
32 if A and B:
33 checks+=1
34 if not connected(adj,A) or not connected(adj,B):
35 failures.append((trial,n,u,w,A,B)); break
36 if failures: break
37 if failures: break
38 # Independent mean identity: SSE around c = SSE around mean + |R| ||c-mu||^2.
39 X=rng.normal(size=(17,3)); R=list(range(17)); c=rng.normal(size=3)
40 cost,mu=sse(X,R); lhs=float(((X-c)**2).sum()); rhs=cost+len(R)*float(((c-mu)**2).sum())
41 print({'random_graph_split_checks':checks,'connectivity_failures':len(failures),
42 'first_failure':failures[:1], 'mean_identity_abs_error':abs(lhs-rhs),
43 'mean_cost':cost,'arbitrary_cost':lhs})
44
45if __name__=='__main__': main()