"""Connectivity-preserving wedge token pooling for small unweighted graphs.""" from collections import deque import numpy as np def all_pairs_dist(adj, nodes=None): """Exact BFS distances restricted to nodes; unreachable distances are inf.""" n = len(adj) allowed = np.ones(n, dtype=bool) if nodes is None else np.array([i in set(nodes) for i in range(n)]) out = np.full((n, n), np.inf) for s in np.flatnonzero(allowed): out[s, s] = 0 q = deque([int(s)]) while q: v = q.popleft() for z in adj[v]: if allowed[z] and np.isinf(out[s, z]): out[s, z] = out[s, v] + 1 q.append(z) return out def connected(adj, region): region = set(region) if not region: return False seen, q = {next(iter(region))}, deque([next(iter(region))]) while q: v = q.popleft() for z in adj[v]: if z in region and z not in seen: seen.add(z); q.append(z) return len(seen) == len(region) def sse(X, region): ids = np.asarray(sorted(region), dtype=int) mu = X[ids].mean(axis=0) return float(((X[ids] - mu) ** 2).sum()), mu def wedge_split(adj, region, u, w, distances=None): region = sorted(region) D = all_pairs_dist(adj, region) if distances is None else distances A = [v for v in region if D[u, v] <= D[w, v]] B = [v for v in region if D[u, v] > D[w, v]] return A, B def pooled_adjacency(adj, assignment, M, mode="sum"): """Aggregate original undirected edges between final regions.""" out = np.zeros((M, M), dtype=float) for v, nbrs in enumerate(adj): for z in nbrs: if z > v: i, j = int(assignment[v]), int(assignment[z]) if i != j: out[i, j] += 1; out[j, i] += 1 if mode == "mean": sizes = np.bincount(assignment, minlength=M).astype(float) denom = sizes[:, None] * sizes[None, :] out = np.divide(out, denom, out=np.zeros_like(out), where=denom > 0) elif mode != "sum": raise ValueError("mode must be 'sum' or 'mean'") return out def pool_wedge(adj, X, M, seed_pairs=None): """Greedy maximum SSE-gain wedge splits. Returns tokens and partition metadata.""" n = len(adj) if M < 1 or M > n: raise ValueError("M must be in [1,n]") regions = [list(range(n))] history = [] pairs = [(u, w) for u in range(n) for w in range(u + 1, n)] if seed_pairs is None else seed_pairs while len(regions) < M: best = None for ri, R in enumerate(regions): if len(R) < 2: continue D = all_pairs_dist(adj, R) base, _ = sse(X, R) for u, w in pairs: if u not in R or w not in R: continue A, B = wedge_split(adj, R, u, w, D) if not A or not B or not connected(adj, A) or not connected(adj, B): continue gain = base - sse(X, A)[0] - sse(X, B)[0] key = (gain, -min(A), -min(B), -u, -w) if best is None or key > best[0]: best = (key, ri, u, w, A, B, gain) if best is None or best[6] <= 1e-12: break _, ri, u, w, A, B, gain = best old = regions.pop(ri) regions.extend([A, B]) history.append({"parent": sorted(old), "children": [sorted(A), sorted(B)], "seeds": [u, w], "gain": gain}) means = np.array([X[R].mean(axis=0) for R in regions]) assignment = np.empty(n, dtype=int) for i, R in enumerate(regions): assignment[R] = i pooled = means[assignment] sizes = np.array([len(R) for R in regions]) return {"tokens": means, "pooled": pooled, "regions": regions, "assignment": assignment, "sizes": sizes, "pooled_adjacency": pooled_adjacency(adj, assignment, len(regions)), "history": history, "stopped": len(regions) < M} def random_pool(X, M, rng): n = len(X); labels = np.repeat(np.arange(M), np.ceil(n/M))[:n]; rng.shuffle(labels) means = np.array([X[labels == i].mean(0) for i in range(M)]) return means, means[labels], labels def kmeans_pool(X, M, seed=0): from sklearn.cluster import KMeans labels = KMeans(M, n_init=10, random_state=seed).fit_predict(X) means = np.array([X[labels == i].mean(0) for i in range(M)]) return means, means[labels], labels