Biclique-free hierarchical attention / experiment.py
Mechanism confirmed, baseline not beaten
1import json
2import math
3import random
4import time
5from collections import defaultdict
6from pathlib import Path
7
8import numpy as np
9
10
11def hierarchical_candidates(n, max_bucket=32, radius=0):
12 """Generate shared dyadic-bucket edges through inverted indices."""
13 buckets = defaultdict(list)
14 level = 1
15 while 2 ** level <= max_bucket:
16 width = 2 ** level
17 for x in range(n):
18 b = x // width
19 for bb in range(max(0, b - radius), b + radius + 1):
20 buckets[(level, bb)].append(x)
21 level += 1
22 edges = set()
23 for members in buckets.values():
24 for i in members:
25 for j in members:
26 edges.add((i, j))
27 return edges
28
29
30def local_edges(n, window=8):
31 return {(i, j) for i in range(n)
32 for j in range(max(0, i - window + 1), min(n, i + window))}
33
34
35def greedy_k22_free(edges, n):
36 """Keep edges iff the new edge cannot complete a K_{2,2}."""
37 nbr = [set() for _ in range(n)]
38 kept = []
39 for i, j in sorted(edges):
40 # If q already contains j and shares any prior key with i,
41 # (i,q) x (that key,j) plus (i,j) completes K_2,2.
42 creates = any(q != i and j in nbr[q] and nbr[i].intersection(nbr[q])
43 for q in range(n))
44 if not creates:
45 nbr[i].add(j)
46 kept.append((i, j))
47 return kept, nbr
48
49
50def max_common_keys(nbr):
51 best = 0
52 for i in range(len(nbr)):
53 for q in range(i):
54 best = max(best, len(nbr[i] & nbr[q]))
55 return best
56
57
58def brute_has_k22(edges, n):
59 e = set(edges)
60 for i in range(n):
61 for q in range(i):
62 for j in range(n):
63 for ell in range(j):
64 if (i, j) in e and (i, ell) in e and (q, j) in e and (q, ell) in e:
65 return True
66 return False
67
68
69def masked_attention(q, k, v, edges):
70 n, d = q.shape
71 out = np.zeros_like(v)
72 by_q = [[] for _ in range(n)]
73 for i, j in edges:
74 by_q[i].append(j)
75 for i, js in enumerate(by_q):
76 if not js:
77 js = [i]
78 logits = (k[js] @ q[i]) / math.sqrt(d)
79 logits -= logits.max()
80 weights = np.exp(logits)
81 weights /= weights.sum()
82 out[i] = weights @ v[js]
83 return out
84
85
86def timed_attention(q, k, v, edges, repeats=3):
87 start = time.perf_counter()
88 for _ in range(repeats):
89 masked_attention(q, k, v, edges)
90 return (time.perf_counter() - start) / repeats
91
92
93def run():
94 random.seed(7)
95 np.random.seed(7)
96 sizes = [32, 64, 128, 256, 512, 1024]
97 rows = []
98 for n in sizes:
99 raw = hierarchical_candidates(n, max_bucket=32)
100 repaired, nbr = greedy_k22_free(raw, n)
101 rows.append({
102 "n": n, "dense": n * n, "local": len(local_edges(n, 8)),
103 "hierarchical_raw": len(raw), "hierarchical_repaired": len(repaired),
104 "max_common_keys": max_common_keys(nbr),
105 "raw_over_n": len(raw) / n, "repaired_over_n": len(repaired) / n,
106 })
107
108 # Exhaustive detector validation on all edge subsets of tiny bipartite graphs.
109 exhaustive_ok = True
110 for mask in range(1 << 9):
111 e = [(i, j) for i in range(3) for j in range(3) if mask & (1 << (3 * i + j))]
112 nbr = [set(j for ii, j in e if ii == i) for i in range(3)]
113 exhaustive_ok &= (max_common_keys(nbr) >= 2) == brute_has_k22(e, 3)
114
115 n, d = 256, 16
116 rng = np.random.default_rng(11)
117 q, k, v = rng.normal(size=(n, d)), rng.normal(size=(n, d)), rng.normal(size=(n, d))
118 raw = hierarchical_candidates(n, max_bucket=32)
119 repaired, _ = greedy_k22_free(raw, n)
120 dense_edges = [(i, j) for i in range(n) for j in range(n)]
121 local = list(local_edges(n, 8))
122 dense_out = masked_attention(q, k, v, dense_edges)
123 sparse_out = masked_attention(q, k, v, repaired)
124 rel_error = np.linalg.norm(dense_out - sparse_out) / max(np.linalg.norm(dense_out), 1e-12)
125
126 x = np.log(np.array(sizes, dtype=float))
127 slopes = {}
128 for name in ["dense", "local", "hierarchical_raw", "hierarchical_repaired"]:
129 y = np.log(np.array([r[name] for r in rows], dtype=float))
130 slopes[name] = float(np.polyfit(x, y, 1)[0])
131 result = {
132 "rows": rows, "loglog_slopes": slopes,
133 "attention_n": n, "attention_dense_edges": n * n,
134 "attention_local_edges": len(local), "attention_sparse_edges": len(repaired),
135 "edge_reduction_vs_dense": (n * n) / len(repaired),
136 "edge_reduction_vs_local": len(local) / len(repaired),
137 "relative_dense_output_error": float(rel_error),
138 "k22_verified": all(r["max_common_keys"] <= 1 for r in rows),
139 "k22_detector_exhaustive_3x3": bool(exhaustive_ok),
140 "timing_seconds_per_call": {
141 "dense": timed_attention(q, k, v, dense_edges),
142 "local": timed_attention(q, k, v, local),
143 "hierarchical_repaired": timed_attention(q, k, v, repaired),
144 },
145 "settings": {"max_bucket": 32, "t": 2, "seed": 7},
146 }
147 Path("results.json").write_text(json.dumps(result, indent=2))
148 print(json.dumps(result, indent=2))
149
150
151if __name__ == "__main__":
152 run()