import json import math import random import time from collections import defaultdict from pathlib import Path import numpy as np def hierarchical_candidates(n, max_bucket=32, radius=0): """Generate shared dyadic-bucket edges through inverted indices.""" buckets = defaultdict(list) level = 1 while 2 ** level <= max_bucket: width = 2 ** level for x in range(n): b = x // width for bb in range(max(0, b - radius), b + radius + 1): buckets[(level, bb)].append(x) level += 1 edges = set() for members in buckets.values(): for i in members: for j in members: edges.add((i, j)) return edges def local_edges(n, window=8): return {(i, j) for i in range(n) for j in range(max(0, i - window + 1), min(n, i + window))} def greedy_k22_free(edges, n): """Keep edges iff the new edge cannot complete a K_{2,2}.""" nbr = [set() for _ in range(n)] kept = [] for i, j in sorted(edges): # If q already contains j and shares any prior key with i, # (i,q) x (that key,j) plus (i,j) completes K_2,2. creates = any(q != i and j in nbr[q] and nbr[i].intersection(nbr[q]) for q in range(n)) if not creates: nbr[i].add(j) kept.append((i, j)) return kept, nbr def max_common_keys(nbr): best = 0 for i in range(len(nbr)): for q in range(i): best = max(best, len(nbr[i] & nbr[q])) return best def brute_has_k22(edges, n): e = set(edges) for i in range(n): for q in range(i): for j in range(n): for ell in range(j): if (i, j) in e and (i, ell) in e and (q, j) in e and (q, ell) in e: return True return False def masked_attention(q, k, v, edges): n, d = q.shape out = np.zeros_like(v) by_q = [[] for _ in range(n)] for i, j in edges: by_q[i].append(j) for i, js in enumerate(by_q): if not js: js = [i] logits = (k[js] @ q[i]) / math.sqrt(d) logits -= logits.max() weights = np.exp(logits) weights /= weights.sum() out[i] = weights @ v[js] return out def timed_attention(q, k, v, edges, repeats=3): start = time.perf_counter() for _ in range(repeats): masked_attention(q, k, v, edges) return (time.perf_counter() - start) / repeats def run(): random.seed(7) np.random.seed(7) sizes = [32, 64, 128, 256, 512, 1024] rows = [] for n in sizes: raw = hierarchical_candidates(n, max_bucket=32) repaired, nbr = greedy_k22_free(raw, n) rows.append({ "n": n, "dense": n * n, "local": len(local_edges(n, 8)), "hierarchical_raw": len(raw), "hierarchical_repaired": len(repaired), "max_common_keys": max_common_keys(nbr), "raw_over_n": len(raw) / n, "repaired_over_n": len(repaired) / n, }) # Exhaustive detector validation on all edge subsets of tiny bipartite graphs. exhaustive_ok = True for mask in range(1 << 9): e = [(i, j) for i in range(3) for j in range(3) if mask & (1 << (3 * i + j))] nbr = [set(j for ii, j in e if ii == i) for i in range(3)] exhaustive_ok &= (max_common_keys(nbr) >= 2) == brute_has_k22(e, 3) n, d = 256, 16 rng = np.random.default_rng(11) q, k, v = rng.normal(size=(n, d)), rng.normal(size=(n, d)), rng.normal(size=(n, d)) raw = hierarchical_candidates(n, max_bucket=32) repaired, _ = greedy_k22_free(raw, n) dense_edges = [(i, j) for i in range(n) for j in range(n)] local = list(local_edges(n, 8)) dense_out = masked_attention(q, k, v, dense_edges) sparse_out = masked_attention(q, k, v, repaired) rel_error = np.linalg.norm(dense_out - sparse_out) / max(np.linalg.norm(dense_out), 1e-12) x = np.log(np.array(sizes, dtype=float)) slopes = {} for name in ["dense", "local", "hierarchical_raw", "hierarchical_repaired"]: y = np.log(np.array([r[name] for r in rows], dtype=float)) slopes[name] = float(np.polyfit(x, y, 1)[0]) result = { "rows": rows, "loglog_slopes": slopes, "attention_n": n, "attention_dense_edges": n * n, "attention_local_edges": len(local), "attention_sparse_edges": len(repaired), "edge_reduction_vs_dense": (n * n) / len(repaired), "edge_reduction_vs_local": len(local) / len(repaired), "relative_dense_output_error": float(rel_error), "k22_verified": all(r["max_common_keys"] <= 1 for r in rows), "k22_detector_exhaustive_3x3": bool(exhaustive_ok), "timing_seconds_per_call": { "dense": timed_attention(q, k, v, dense_edges), "local": timed_attention(q, k, v, local), "hierarchical_repaired": timed_attention(q, k, v, repaired), }, "settings": {"max_bucket": 32, "t": 2, "seed": 7}, } Path("results.json").write_text(json.dumps(result, indent=2)) print(json.dumps(result, indent=2)) if __name__ == "__main__": run()