import json, math from pathlib import Path import numpy as np from scipy.optimize import linear_sum_assignment SEED = 1234 def assignment_sample(cost, tau, rng): noisy = cost + tau * rng.gumbel(size=cost.shape) rows, cols = linear_sum_assignment(noisy) return list(zip(rows.tolist(), cols.tolist())) def frequency_estimate(cost, tau, K, rng): counts = np.zeros_like(cost, dtype=float) for _ in range(K): for i, j in assignment_sample(cost, tau, rng): counts[i, j] += 1.0 return counts / K def kruskal_tree(score): """Maximum-score spanning tree on the query/key bipartite graph.""" nq, nk = score.shape edges = [(float(score[i, j]), i, nq + j) for i in range(nq) for j in range(nk)] edges.sort(key=lambda x: (-x[0], x[1], x[2])) parent = list(range(nq + nk)) def find(x): while parent[x] != x: parent[x] = parent[parent[x]] x = parent[x] return x chosen = [] for s, u, v in edges: a, b = find(u), find(v) if a != b: parent[a] = b chosen.append((u, v - nq)) if len(chosen) == nq + nk - 1: break tree = np.zeros_like(score, dtype=bool) for i, j in chosen: tree[i, j] = True return tree def assignment_tree(cost, tau, K, rng): freq = frequency_estimate(cost, tau, K, rng) tree = kruskal_tree(freq) # Frequencies remain the edge bias; only tree edges are legal supports. return freq * tree, tree, freq def sparse_attention(q, k, v, tree_freq, alpha=1.0, eps=1e-6): logits = q @ k.T / math.sqrt(q.shape[1]) support = tree_freq > 0 # A tree can contain a query leaf with no incident positive-frequency edge # only in a degenerate sample; repair it without changing normal cases. for i in range(len(q)): if not support[i].any(): support[i, np.argmax(tree_freq[i])] = True biased = logits + alpha * np.log(eps + tree_freq) biased[~support] = -np.inf z = biased - np.max(biased, axis=1, keepdims=True) p = np.exp(z); p /= p.sum(axis=1, keepdims=True) return p @ v, p def dense_attention(q, k, v): logits = q @ k.T / math.sqrt(q.shape[1]) z = logits - logits.max(axis=1, keepdims=True) p = np.exp(z); p /= p.sum(axis=1, keepdims=True) return p @ v, p def main(): n, d, tau = 15, 8, 1.7 base = np.random.default_rng(SEED) q, k = base.normal(size=(n + 1, d)), base.normal(size=(n, d)) cost = ((q[:, None, :] - k[None, :, :]) ** 2).sum(axis=2) # Prediction 1: empirical assignment probabilities have Monte Carlo RMSE O(K^-1/2). ref = frequency_estimate(cost, tau, 100000, np.random.default_rng(999)) Ks = [4, 8, 16, 32, 64, 128, 256, 512] rmse, distinct, pred_distinct = [], [], [] positive = ref > 1e-4 p = ref[positive] for K in Ks: reps = 30 if K <= 128 else 15 vals, ds = [], [] for r in range(reps): f = frequency_estimate(cost, tau, K, np.random.default_rng(SEED + K*31 + r)) vals.append(np.sqrt(np.mean((f - ref) ** 2))) ds.append(np.count_nonzero(f)) rmse.append(float(np.mean(vals))); distinct.append(float(np.mean(ds))) pred_distinct.append(float(np.sum(1 - (1 - p) ** K))) slope = float(np.polyfit(np.log(Ks), np.log(rmse), 1)[0]) # Prediction 2: union occupancy follows the independent occupancy formula. # Prediction 3: alpha=1 increases probability-frequency alignment over alpha=0. f, tree, raw = assignment_tree(cost, tau, 128, np.random.default_rng(777)) v = base.normal(size=(n, 5)) _, p0 = sparse_attention(q, k, v, f, alpha=0.0) _, p1 = sparse_attention(q, k, v, f, alpha=1.0) support = tree corr0 = float(np.corrcoef(np.log(1e-6 + f[support]), p0[support])[0, 1]) corr1 = float(np.corrcoef(np.log(1e-6 + f[support]), p1[support])[0, 1]) # Transport sanity check: T=f/(n+1) always has exact query marginals; # key marginals are close only when the unmatched query is balanced. T = raw / (n + 1) row_err = float(np.max(np.abs(T.sum(1) - 1/(n+1)))) col_err = float(np.max(np.abs(T.sum(0) - 1/n))) # Tiny retrieval comparison, same inputs and fixed seeds. dense_err, sparse_err, tree_edges = [], [], [] for t in range(80): rr = np.random.default_rng(4000 + t) qq = rr.normal(size=(n + 1, d)); kk = qq[:n] + .8 * rr.normal(size=(n, d)) vv = rr.normal(size=(n, 3)) target = vv[np.argmin(((qq[:, None, :] - kk[None, :, :])**2).sum(2), axis=1) % n] cc = ((qq[:, None, :] - kk[None, :, :])**2).sum(2) yd, _ = dense_attention(qq, kk, vv) tf, tr, _ = assignment_tree(cc, tau, 16, np.random.default_rng(7000 + t)) ys, _ = sparse_attention(qq, kk, vv, tf, alpha=.5) dense_err.append(np.mean((yd - target)**2)); sparse_err.append(np.mean((ys - target)**2)) tree_edges.append(int(tr.sum())) out = { 'seed': SEED, 'shape': [n+1, n], 'tau': tau, 'prediction_1_rmse_slope_predicted': -0.5, 'prediction_1_rmse_slope_observed': slope, 'prediction_1_rmse': dict(zip(map(str, Ks), rmse)), 'prediction_2_union_edges': [{'K': K, 'observed': o, 'predicted': pr} for K,o,pr in zip(Ks,distinct,pred_distinct)], 'prediction_3_frequency_corr_alpha0': corr0, 'prediction_3_frequency_corr_alpha1': corr1, 'transport_max_row_marginal_error': row_err, 'transport_max_key_marginal_error': col_err, 'tree_edges_expected': n + 1 + n - 1, 'tree_edges_observed': int(tree.sum()), 'dense_retrieval_mse': float(np.mean(dense_err)), 'tree_sparse_retrieval_mse': float(np.mean(sparse_err)), 'tree_edges_K16_mean': float(np.mean(tree_edges)), 'raw_union_edges_K16_mean': float(np.mean(distinct[2:3])) } Path('results.json').write_text(json.dumps(out, indent=2)) print(json.dumps(out, indent=2)) if __name__ == '__main__': main()