Sublinear-expander sparse attention / expander_attention.py

Unverified

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