Neighborhood-separator attention / experiment.py
Failed on benchmark
1import json, math
2import numpy as np
3
4SEED = 7
5rng = np.random.default_rng(SEED)
6
7
8def affinity_graph(q, k, r):
9 n, d = q.shape
10 s = (q @ k.T + k @ q.T) / (2.0 * math.sqrt(d))
11 np.fill_diagonal(s, -np.inf)
12 adj = np.zeros((n, n), dtype=bool)
13 for i in range(n): adj[i, np.argpartition(s[i], -r)[-r:]] = True
14 adj |= adj.T
15 np.fill_diagonal(adj, False)
16 return adj
17
18
19def planted_graph(n=128, groups=4, width=30, bridge=8, p=.20):
20 """Community graph with a small globally-connected bridge set."""
21 a = np.zeros((n, n), dtype=bool)
22 b = np.arange(bridge); rest = np.arange(bridge, bridge + groups * width)
23 for g in range(groups):
24 nodes = rest[g*width:(g+1)*width]
25 z = rng.random((width, width)) < p
26 z = np.triu(z, 1); a[np.ix_(nodes, nodes)] |= z | z.T
27 # bridge tokens provide the intended global communication.
28 a[np.ix_(b, nodes)] = True; a[np.ix_(nodes, b)] = True
29 np.fill_diagonal(a, False)
30 return a
31
32
33def components(adj, active):
34 active = np.asarray(active, bool); seen = np.zeros(len(active), bool); out=[]
35 for st in np.flatnonzero(active):
36 if seen[st]: continue
37 stack=[int(st)]; seen[st]=1; c=[]
38 while stack:
39 u=stack.pop(); c.append(u)
40 for v in np.flatnonzero(adj[u] & active):
41 if not seen[v]: seen[v]=1; stack.append(int(v))
42 out.append(np.array(c))
43 return out
44
45
46def triangle_nodes(adj, active):
47 a = adj & active[:, None] & active[None, :]; tri=np.zeros(len(active))
48 for i in np.flatnonzero(active):
49 ns=np.flatnonzero(a[i]); tri[i]=np.count_nonzero(a[np.ix_(ns,ns)])/2
50 return tri
51
52
53def separator_mask(adj, max_sep=2):
54 n=len(adj); active=np.ones(n,bool); sep=[]
55 for _ in range(max_sep):
56 tri=triangle_nodes(adj,active); deg=(adj & active[:,None] & active[None,:]).sum(1)
57 score=np.where(active,tri/(deg+1),-1); x=int(np.argmax(score))
58 if tri[x] < 1: break
59 sep.append(x)
60 closed=(np.arange(n)==x) | adj[x]
61 active[closed]=False
62 comps=components(adj,active); covered=np.zeros(n,bool)
63 if sep:
64 covered[sep]=True; covered |= adj[:,sep].any(1)
65 mask=covered[:,None] | covered[None,:]
66 for c in comps: mask[np.ix_(c,c)]=True
67 np.fill_diagonal(mask,True)
68 return mask, sep, comps, covered
69
70
71def local_mask(n, radius=2):
72 ix=np.arange(n); return abs(ix[:,None]-ix[None,:])<=radius
73
74
75def attention(q,k,v,allow):
76 z=q@k.T/math.sqrt(q.shape[1]); z=np.where(allow,z,-1e9)
77 z-=z.max(1,keepdims=True); p=np.exp(z); p/=p.sum(1,keepdims=True)
78 return p@v
79
80
81def evaluate_case(name, graph=None, r=4, trials=40):
82 n,d,dv=128,32,16; vals={x:[] for x in ('dense','local','separator')}; costs={x:[] for x in vals}; ss=[]; cc=[]; cov=[]
83 for _ in range(trials):
84 q=rng.normal(size=(n,d)); k=rng.normal(size=(n,d)); v=rng.normal(size=(n,dv))
85 g=affinity_graph(q,k,r) if graph is None else graph
86 sm,sep,comps,covered=separator_mask(g,2); dense=np.ones((n,n),bool); lm=local_mask(n)
87 ref=attention(q,k,v,dense)
88 for x,m in [('dense',dense),('local',lm),('separator',sm)]:
89 vals[x].append(float(np.mean((attention(q,k,v,m)-ref)**2))); costs[x].append(int(m.sum()))
90 ss.append(len(sep)); cc.append(len(comps)); cov.append(covered.mean())
91 return {'case':name,'metrics':{x:{'mse_to_dense_mean':float(np.mean(vals[x])), 'mse_std':float(np.std(vals[x])), 'allowed_pairs_mean':float(np.mean(costs[x])), 'relative_attention_cost':float(np.mean(costs[x])/(n*n)), 'speedup':float(n*n/np.mean(costs[x]))} for x in vals},'separator_size_mean':float(np.mean(ss)),'components_mean':float(np.mean(cc)),'covered_fraction_mean':float(np.mean(cov))}
92
93
94def structural_check():
95 g=np.zeros((10,10),bool)
96 for block in ([0,1,2],[5,6,7]):
97 for i in block:
98 for j in block:
99 if i!=j:g[i,j]=1
100 m,s,c,covered=separator_mask(g,2); anti=all(not g[np.ix_(c[i],c[j])].any() for i in range(len(c)) for j in range(i))
101 formula=covered[:,None]|covered[None,:]
102 for x in c: formula[np.ix_(x,x)]=1
103 return {'separator_size':len(s),'components':[len(x) for x in c],'anti_adjacent_components':anti,'mask_formula_exact':bool(np.array_equal(m,formula))}
104
105if __name__=='__main__':
106 result={'seed':SEED,'structural_check':structural_check(),'random':evaluate_case('random'),'planted':evaluate_case('planted',planted_graph())}
107 with open('results.json','w') as f: json.dump(result,f,indent=2)
108 print(json.dumps(result,indent=2))