Connectivity-Preserving Wedge Token Pooling / wedge_pool.py
Mechanism confirmed, baseline not beaten
1"""Connectivity-preserving wedge token pooling for small unweighted graphs."""
2from collections import deque
3import numpy as np
4
5
6def all_pairs_dist(adj, nodes=None):
7 """Exact BFS distances restricted to nodes; unreachable distances are inf."""
8 n = len(adj)
9 allowed = np.ones(n, dtype=bool) if nodes is None else np.array([i in set(nodes) for i in range(n)])
10 out = np.full((n, n), np.inf)
11 for s in np.flatnonzero(allowed):
12 out[s, s] = 0
13 q = deque([int(s)])
14 while q:
15 v = q.popleft()
16 for z in adj[v]:
17 if allowed[z] and np.isinf(out[s, z]):
18 out[s, z] = out[s, v] + 1
19 q.append(z)
20 return out
21
22
23def connected(adj, region):
24 region = set(region)
25 if not region: return False
26 seen, q = {next(iter(region))}, deque([next(iter(region))])
27 while q:
28 v = q.popleft()
29 for z in adj[v]:
30 if z in region and z not in seen:
31 seen.add(z); q.append(z)
32 return len(seen) == len(region)
33
34
35def sse(X, region):
36 ids = np.asarray(sorted(region), dtype=int)
37 mu = X[ids].mean(axis=0)
38 return float(((X[ids] - mu) ** 2).sum()), mu
39
40
41def wedge_split(adj, region, u, w, distances=None):
42 region = sorted(region)
43 D = all_pairs_dist(adj, region) if distances is None else distances
44 A = [v for v in region if D[u, v] <= D[w, v]]
45 B = [v for v in region if D[u, v] > D[w, v]]
46 return A, B
47
48
49def pooled_adjacency(adj, assignment, M, mode="sum"):
50 """Aggregate original undirected edges between final regions."""
51 out = np.zeros((M, M), dtype=float)
52 for v, nbrs in enumerate(adj):
53 for z in nbrs:
54 if z > v:
55 i, j = int(assignment[v]), int(assignment[z])
56 if i != j:
57 out[i, j] += 1; out[j, i] += 1
58 if mode == "mean":
59 sizes = np.bincount(assignment, minlength=M).astype(float)
60 denom = sizes[:, None] * sizes[None, :]
61 out = np.divide(out, denom, out=np.zeros_like(out), where=denom > 0)
62 elif mode != "sum":
63 raise ValueError("mode must be 'sum' or 'mean'")
64 return out
65
66
67def pool_wedge(adj, X, M, seed_pairs=None):
68 """Greedy maximum SSE-gain wedge splits. Returns tokens and partition metadata."""
69 n = len(adj)
70 if M < 1 or M > n: raise ValueError("M must be in [1,n]")
71 regions = [list(range(n))]
72 history = []
73 pairs = [(u, w) for u in range(n) for w in range(u + 1, n)] if seed_pairs is None else seed_pairs
74 while len(regions) < M:
75 best = None
76 for ri, R in enumerate(regions):
77 if len(R) < 2: continue
78 D = all_pairs_dist(adj, R)
79 base, _ = sse(X, R)
80 for u, w in pairs:
81 if u not in R or w not in R: continue
82 A, B = wedge_split(adj, R, u, w, D)
83 if not A or not B or not connected(adj, A) or not connected(adj, B): continue
84 gain = base - sse(X, A)[0] - sse(X, B)[0]
85 key = (gain, -min(A), -min(B), -u, -w)
86 if best is None or key > best[0]: best = (key, ri, u, w, A, B, gain)
87 if best is None or best[6] <= 1e-12: break
88 _, ri, u, w, A, B, gain = best
89 old = regions.pop(ri)
90 regions.extend([A, B])
91 history.append({"parent": sorted(old), "children": [sorted(A), sorted(B)], "seeds": [u, w], "gain": gain})
92 means = np.array([X[R].mean(axis=0) for R in regions])
93 assignment = np.empty(n, dtype=int)
94 for i, R in enumerate(regions): assignment[R] = i
95 pooled = means[assignment]
96 sizes = np.array([len(R) for R in regions])
97 return {"tokens": means, "pooled": pooled, "regions": regions, "assignment": assignment,
98 "sizes": sizes, "pooled_adjacency": pooled_adjacency(adj, assignment, len(regions)),
99 "history": history, "stopped": len(regions) < M}
100
101
102def random_pool(X, M, rng):
103 n = len(X); labels = np.repeat(np.arange(M), np.ceil(n/M))[:n]; rng.shuffle(labels)
104 means = np.array([X[labels == i].mean(0) for i in range(M)])
105 return means, means[labels], labels
106
107
108def kmeans_pool(X, M, seed=0):
109 from sklearn.cluster import KMeans
110 labels = KMeans(M, n_init=10, random_state=seed).fit_predict(X)
111 means = np.array([X[labels == i].mean(0) for i in range(M)])
112 return means, means[labels], labels