Sublinear-expander sparse attention / expander_attention.py
Unverified
1import math, time, itertools, json, random
2import numpy as np
3
4
5def random_regular(n, d, seed=0):
6 if d < 0 or d >= n: raise ValueError("degree must satisfy 0 <= d < n")
7 rng=np.random.default_rng(seed)
8 if (n*d) % 2: raise ValueError("n*d must be even")
9 # A randomly relabeled circulant is simple and exactly d-regular.
10 # For odd d, add a perfect matching (requiring even n).
11 if d % 2 and n % 2: raise ValueError("odd degree requires even n")
12 adj=[set() for _ in range(n)]
13 half=d//2
14 for i in range(n):
15 for z in range(1,half+1):
16 adj[i].add((i+z)%n); adj[i].add((i-z)%n)
17 if d % 2:
18 perm=rng.permutation(n)
19 for a,b in perm.reshape(-1,2):
20 adj[int(a)].add(int(b)); adj[int(b)].add(int(a))
21 perm=rng.permutation(n); inv=np.empty(n,dtype=int)
22 for i,x in enumerate(perm): inv[x]=i
23 out=[set() for _ in range(n)]
24 for i in range(n):
25 out[int(perm[i])]={int(perm[j]) for j in adj[i]}
26 return [sorted(x) for x in out]
27
28def ring(n, d):
29 adj=[set() for _ in range(n)]
30 for i in range(n):
31 for z in range(1,d//2+1): adj[i].update(((i-z)%n,(i+z)%n))
32 return [sorted(x) for x in adj]
33
34def boundary(adj, U):
35 us=set(U); out=set()
36 for i in U:
37 out.update(j for j in adj[i] if j not in us)
38 return len(out)
39
40def exact_stats(adj, k, eps):
41 n=len(adj); rows=[]
42 for s in range(k,n//2+1):
43 vals=[]
44 for U in itertools.combinations(range(n),s):
45 b=boundary(adj,U); bound=eps*s/(math.log(3*s/k)**2)
46 vals.append((b/s,b,bound))
47 a=np.array(vals)
48 rows.append((s,float(a[:,0].min()),float(a[:,0].mean()),float(a[:,0].max()),float(a[:,2].max())))
49 return rows
50
51def sampled_stats(adj,k,eps,seed=1,samples=1000):
52 rng=np.random.default_rng(seed); n=len(adj); rows=[]
53 for s in [k, min(2*k,n//2), min(4*k,n//2), n//2]:
54 vals=[]
55 for _ in range(samples):
56 U=rng.choice(n,s,replace=False); b=boundary(adj,U)
57 vals.append((b/s, eps/(math.log(3*s/k)**2)))
58 rows.append((s,float(np.min([x[0] for x in vals])),float(np.mean([x[0] for x in vals])),float(np.mean([x[1] for x in vals]))))
59 return rows
60
61def bfs_depth(adj, seedset, target):
62 seen=set(seedset); frontier=set(seedset); depth=0; history=[len(seen)]
63 while len(seen)<target and frontier and depth<100:
64 nxt=set()
65 for i in frontier: nxt.update(adj[i])
66 nxt-=seen; seen |= nxt; frontier=nxt; depth+=1; history.append(len(seen))
67 return depth,history
68
69def torch_benchmark(n=512,h=64,d=16,bs=8,iters=30):
70 import torch
71 device='cuda' if torch.cuda.is_available() else 'cpu'
72 def run(kind):
73 q=torch.randn(bs,n,h,device=device); k=torch.randn_like(q); v=torch.randn_like(q)
74 if kind=='dense':
75 fn=lambda: torch.softmax(torch.matmul(q,k.transpose(1,2))/math.sqrt(h),-1).matmul(v)
76 else:
77 adj=ring(n,d) if kind=='local' else random_regular(n,d,9)
78 ii=np.repeat(np.arange(n),d); jj=np.array([j for row in adj for j in row])
79 ti=torch.tensor(ii,device=device); tj=torch.tensor(jj,device=device)
80 def fn():
81 scores=(q[:,ti,:]*k[:,tj,:]).sum(-1)/math.sqrt(h)
82 scores=scores.view(bs,n,d); att=torch.softmax(scores,-1)
83 return (att.unsqueeze(-1)*v[:,tj,:].view(bs,n,d,h)).sum(2)
84 for _ in range(5): fn()
85 if device=='cuda': torch.cuda.synchronize()
86 t=time.perf_counter()
87 for _ in range(iters): fn()
88 if device=='cuda': torch.cuda.synchronize()
89 return (time.perf_counter()-t)/iters*1000
90 try: return device,{x:run(x) for x in ['dense','local','expander']}
91 except Exception as e:
92 return 'cpu-fallback',{'error':str(e)}
93
94def expansion_constant(adj, k):
95 n=len(adj); vals=[]
96 for size in range(k,n//2+1):
97 for U in itertools.combinations(range(n),size):
98 vals.append((boundary(adj,U)/size)*math.log(3*size/k)**2)
99 return min(vals)
100
101def max_exact_violation(adj, k, eps):
102 n=len(adj); worst=0.0
103 for size in range(k,n//2+1):
104 bound=eps/(math.log(3*size/k)**2)
105 for U in itertools.combinations(range(n),size):
106 worst=max(worst, bound-boundary(adj,U)/size)
107 return max(0.0,worst)
108
109def predicted_steps(eps, k, target):
110 x=float(k); steps=0
111 while x < target and steps < 10000:
112 x *= 1 + eps/(math.log(3*x/k)**2)
113 steps += 1
114 return steps, x
115
116def bfs_to_target(adj, seedset, target):
117 seen=set(seedset); frontier=set(seedset); depth=0
118 while len(seen)<target and frontier and depth<10000:
119 nxt=set()
120 for i in frontier: nxt.update(adj[i])
121 nxt-=seen; seen |= nxt; frontier=nxt; depth+=1
122 return depth, len(seen)
123
124def main():
125 np.random.seed(0); random.seed(0)
126 n=16; k=2
127 expander=random_regular(n,4,3); ring_graph=ring(n,4)
128 # The largest epsilon for which the definition holds is the minimum
129 # of observed boundary_ratio * log^2(3|U|/k).
130 c_exp=expansion_constant(expander,k); c_ring=expansion_constant(ring_graph,k)
131 eps=0.8*c_exp
132 exact=exact_stats(expander,k,eps)
133 max_violation=max_exact_violation(expander,k,eps)
134 # Prediction 1: every exact subset obeys the bound at eps <= epsilon_max.
135 epsilon_scaling=[]
136 for frac in [0.25,0.5,0.75,1.0,1.1]:
137 e=frac*c_exp; rows=exact_stats(expander,k,e)
138 viol=max_exact_violation(expander,k,e)
139 epsilon_scaling.append({'fraction_of_max':frac,'epsilon':e,'max_violation':viol})
140 # Prediction 2: better-mixed expander has stronger normalized expansion
141 # than a same-degree local ring.
142 family=[('ring',ring_graph),('random_regular',expander)]
143 family_stats=[]
144 for name,g in family:
145 rows=sampled_stats(g,k,eps,samples=600)
146 family_stats.append({'graph':name,'epsilon_max':expansion_constant(g,k),
147 'sampled_rows':rows})
148 # Prediction 3: increasing epsilon in the theoretical recurrence decreases
149 # the predicted number of layers to reach n/2; compare with BFS.
150 growth=[]
151 for d in [2,4,6,8]:
152 gg=random_regular(64,d,100+d); cm=min(x[1] * math.log(3*x[0]/2)**2 for x in sampled_stats(gg,2,1.0,samples=1200))
153 vals=[]
154 for frac in [0.25,0.5,0.75,1.0]:
155 e=frac*cm; pred,_=predicted_steps(e,2,32)
156 obs,reached=bfs_to_target(gg,[0,1],32)
157 vals.append({'fraction_of_max':frac,'predicted_layers_to_32':pred,
158 'observed_layers_to_32':obs,'reached':reached})
159 growth.append({'d':d,'epsilon_max':cm,'sweep':vals})
160 bench=torch_benchmark()
161 result={'epsilon_max_expander':c_exp,'epsilon_max_ring':c_ring,
162 'epsilon_used':eps,'exact_expansion_rows':exact,
163 'max_bound_minus_observed':max_violation,
164 'epsilon_scaling':epsilon_scaling,'graph_family':family_stats,
165 'growth_sweep':growth,'attention_ms':bench[1],'device':bench[0]}
166 print(json.dumps(result,indent=2))
167
168if __name__=='__main__': main()