Biclique-free hierarchical attention / experiment.py

Mechanism confirmed, baseline not beaten

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