Assignment Tree Attention / assignment_tree_attention.py

Mechanism confirmed, baseline not beaten

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