Assignment Tree Attention / assignment_tree_attention.py
Mechanism confirmed, baseline not beaten
1import json, math
2from pathlib import Path
3import numpy as np
4from scipy.optimize import linear_sum_assignment
5
6SEED = 1234
7
8
9def assignment_sample(cost, tau, rng):
10 noisy = cost + tau * rng.gumbel(size=cost.shape)
11 rows, cols = linear_sum_assignment(noisy)
12 return list(zip(rows.tolist(), cols.tolist()))
13
14
15def frequency_estimate(cost, tau, K, rng):
16 counts = np.zeros_like(cost, dtype=float)
17 for _ in range(K):
18 for i, j in assignment_sample(cost, tau, rng):
19 counts[i, j] += 1.0
20 return counts / K
21
22
23def kruskal_tree(score):
24 """Maximum-score spanning tree on the query/key bipartite graph."""
25 nq, nk = score.shape
26 edges = [(float(score[i, j]), i, nq + j) for i in range(nq) for j in range(nk)]
27 edges.sort(key=lambda x: (-x[0], x[1], x[2]))
28 parent = list(range(nq + nk))
29 def find(x):
30 while parent[x] != x:
31 parent[x] = parent[parent[x]]
32 x = parent[x]
33 return x
34 chosen = []
35 for s, u, v in edges:
36 a, b = find(u), find(v)
37 if a != b:
38 parent[a] = b
39 chosen.append((u, v - nq))
40 if len(chosen) == nq + nk - 1:
41 break
42 tree = np.zeros_like(score, dtype=bool)
43 for i, j in chosen:
44 tree[i, j] = True
45 return tree
46
47
48def assignment_tree(cost, tau, K, rng):
49 freq = frequency_estimate(cost, tau, K, rng)
50 tree = kruskal_tree(freq)
51 # Frequencies remain the edge bias; only tree edges are legal supports.
52 return freq * tree, tree, freq
53
54
55def sparse_attention(q, k, v, tree_freq, alpha=1.0, eps=1e-6):
56 logits = q @ k.T / math.sqrt(q.shape[1])
57 support = tree_freq > 0
58 # A tree can contain a query leaf with no incident positive-frequency edge
59 # only in a degenerate sample; repair it without changing normal cases.
60 for i in range(len(q)):
61 if not support[i].any():
62 support[i, np.argmax(tree_freq[i])] = True
63 biased = logits + alpha * np.log(eps + tree_freq)
64 biased[~support] = -np.inf
65 z = biased - np.max(biased, axis=1, keepdims=True)
66 p = np.exp(z); p /= p.sum(axis=1, keepdims=True)
67 return p @ v, p
68
69
70def dense_attention(q, k, v):
71 logits = q @ k.T / math.sqrt(q.shape[1])
72 z = logits - logits.max(axis=1, keepdims=True)
73 p = np.exp(z); p /= p.sum(axis=1, keepdims=True)
74 return p @ v, p
75
76
77def main():
78 n, d, tau = 15, 8, 1.7
79 base = np.random.default_rng(SEED)
80 q, k = base.normal(size=(n + 1, d)), base.normal(size=(n, d))
81 cost = ((q[:, None, :] - k[None, :, :]) ** 2).sum(axis=2)
82
83 # Prediction 1: empirical assignment probabilities have Monte Carlo RMSE O(K^-1/2).
84 ref = frequency_estimate(cost, tau, 100000, np.random.default_rng(999))
85 Ks = [4, 8, 16, 32, 64, 128, 256, 512]
86 rmse, distinct, pred_distinct = [], [], []
87 positive = ref > 1e-4
88 p = ref[positive]
89 for K in Ks:
90 reps = 30 if K <= 128 else 15
91 vals, ds = [], []
92 for r in range(reps):
93 f = frequency_estimate(cost, tau, K, np.random.default_rng(SEED + K*31 + r))
94 vals.append(np.sqrt(np.mean((f - ref) ** 2)))
95 ds.append(np.count_nonzero(f))
96 rmse.append(float(np.mean(vals))); distinct.append(float(np.mean(ds)))
97 pred_distinct.append(float(np.sum(1 - (1 - p) ** K)))
98 slope = float(np.polyfit(np.log(Ks), np.log(rmse), 1)[0])
99
100 # Prediction 2: union occupancy follows the independent occupancy formula.
101 # Prediction 3: alpha=1 increases probability-frequency alignment over alpha=0.
102 f, tree, raw = assignment_tree(cost, tau, 128, np.random.default_rng(777))
103 v = base.normal(size=(n, 5))
104 _, p0 = sparse_attention(q, k, v, f, alpha=0.0)
105 _, p1 = sparse_attention(q, k, v, f, alpha=1.0)
106 support = tree
107 corr0 = float(np.corrcoef(np.log(1e-6 + f[support]), p0[support])[0, 1])
108 corr1 = float(np.corrcoef(np.log(1e-6 + f[support]), p1[support])[0, 1])
109
110 # Transport sanity check: T=f/(n+1) always has exact query marginals;
111 # key marginals are close only when the unmatched query is balanced.
112 T = raw / (n + 1)
113 row_err = float(np.max(np.abs(T.sum(1) - 1/(n+1))))
114 col_err = float(np.max(np.abs(T.sum(0) - 1/n)))
115
116 # Tiny retrieval comparison, same inputs and fixed seeds.
117 dense_err, sparse_err, tree_edges = [], [], []
118 for t in range(80):
119 rr = np.random.default_rng(4000 + t)
120 qq = rr.normal(size=(n + 1, d)); kk = qq[:n] + .8 * rr.normal(size=(n, d))
121 vv = rr.normal(size=(n, 3))
122 target = vv[np.argmin(((qq[:, None, :] - kk[None, :, :])**2).sum(2), axis=1) % n]
123 cc = ((qq[:, None, :] - kk[None, :, :])**2).sum(2)
124 yd, _ = dense_attention(qq, kk, vv)
125 tf, tr, _ = assignment_tree(cc, tau, 16, np.random.default_rng(7000 + t))
126 ys, _ = sparse_attention(qq, kk, vv, tf, alpha=.5)
127 dense_err.append(np.mean((yd - target)**2)); sparse_err.append(np.mean((ys - target)**2))
128 tree_edges.append(int(tr.sum()))
129
130 out = {
131 'seed': SEED, 'shape': [n+1, n], 'tau': tau,
132 'prediction_1_rmse_slope_predicted': -0.5, 'prediction_1_rmse_slope_observed': slope,
133 'prediction_1_rmse': dict(zip(map(str, Ks), rmse)),
134 'prediction_2_union_edges': [{'K': K, 'observed': o, 'predicted': pr} for K,o,pr in zip(Ks,distinct,pred_distinct)],
135 'prediction_3_frequency_corr_alpha0': corr0, 'prediction_3_frequency_corr_alpha1': corr1,
136 'transport_max_row_marginal_error': row_err, 'transport_max_key_marginal_error': col_err,
137 'tree_edges_expected': n + 1 + n - 1, 'tree_edges_observed': int(tree.sum()),
138 'dense_retrieval_mse': float(np.mean(dense_err)), 'tree_sparse_retrieval_mse': float(np.mean(sparse_err)),
139 'tree_edges_K16_mean': float(np.mean(tree_edges)), 'raw_union_edges_K16_mean': float(np.mean(distinct[2:3]))
140 }
141 Path('results.json').write_text(json.dumps(out, indent=2))
142 print(json.dumps(out, indent=2))
143
144if __name__ == '__main__':
145 main()