Connectivity-Preserving Wedge Token Pooling / wedge_pool.py

Mechanism confirmed, baseline not beaten

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