import json, math import numpy as np SEED = 7 rng = np.random.default_rng(SEED) def affinity_graph(q, k, r): n, d = q.shape s = (q @ k.T + k @ q.T) / (2.0 * math.sqrt(d)) np.fill_diagonal(s, -np.inf) adj = np.zeros((n, n), dtype=bool) for i in range(n): adj[i, np.argpartition(s[i], -r)[-r:]] = True adj |= adj.T np.fill_diagonal(adj, False) return adj def planted_graph(n=128, groups=4, width=30, bridge=8, p=.20): """Community graph with a small globally-connected bridge set.""" a = np.zeros((n, n), dtype=bool) b = np.arange(bridge); rest = np.arange(bridge, bridge + groups * width) for g in range(groups): nodes = rest[g*width:(g+1)*width] z = rng.random((width, width)) < p z = np.triu(z, 1); a[np.ix_(nodes, nodes)] |= z | z.T # bridge tokens provide the intended global communication. a[np.ix_(b, nodes)] = True; a[np.ix_(nodes, b)] = True np.fill_diagonal(a, False) return a def components(adj, active): active = np.asarray(active, bool); seen = np.zeros(len(active), bool); out=[] for st in np.flatnonzero(active): if seen[st]: continue stack=[int(st)]; seen[st]=1; c=[] while stack: u=stack.pop(); c.append(u) for v in np.flatnonzero(adj[u] & active): if not seen[v]: seen[v]=1; stack.append(int(v)) out.append(np.array(c)) return out def triangle_nodes(adj, active): a = adj & active[:, None] & active[None, :]; tri=np.zeros(len(active)) for i in np.flatnonzero(active): ns=np.flatnonzero(a[i]); tri[i]=np.count_nonzero(a[np.ix_(ns,ns)])/2 return tri def separator_mask(adj, max_sep=2): n=len(adj); active=np.ones(n,bool); sep=[] for _ in range(max_sep): tri=triangle_nodes(adj,active); deg=(adj & active[:,None] & active[None,:]).sum(1) score=np.where(active,tri/(deg+1),-1); x=int(np.argmax(score)) if tri[x] < 1: break sep.append(x) closed=(np.arange(n)==x) | adj[x] active[closed]=False comps=components(adj,active); covered=np.zeros(n,bool) if sep: covered[sep]=True; covered |= adj[:,sep].any(1) mask=covered[:,None] | covered[None,:] for c in comps: mask[np.ix_(c,c)]=True np.fill_diagonal(mask,True) return mask, sep, comps, covered def local_mask(n, radius=2): ix=np.arange(n); return abs(ix[:,None]-ix[None,:])<=radius def attention(q,k,v,allow): z=q@k.T/math.sqrt(q.shape[1]); z=np.where(allow,z,-1e9) z-=z.max(1,keepdims=True); p=np.exp(z); p/=p.sum(1,keepdims=True) return p@v def evaluate_case(name, graph=None, r=4, trials=40): n,d,dv=128,32,16; vals={x:[] for x in ('dense','local','separator')}; costs={x:[] for x in vals}; ss=[]; cc=[]; cov=[] for _ in range(trials): q=rng.normal(size=(n,d)); k=rng.normal(size=(n,d)); v=rng.normal(size=(n,dv)) g=affinity_graph(q,k,r) if graph is None else graph sm,sep,comps,covered=separator_mask(g,2); dense=np.ones((n,n),bool); lm=local_mask(n) ref=attention(q,k,v,dense) for x,m in [('dense',dense),('local',lm),('separator',sm)]: vals[x].append(float(np.mean((attention(q,k,v,m)-ref)**2))); costs[x].append(int(m.sum())) ss.append(len(sep)); cc.append(len(comps)); cov.append(covered.mean()) 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))} def structural_check(): g=np.zeros((10,10),bool) for block in ([0,1,2],[5,6,7]): for i in block: for j in block: if i!=j:g[i,j]=1 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)) formula=covered[:,None]|covered[None,:] for x in c: formula[np.ix_(x,x)]=1 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))} if __name__=='__main__': result={'seed':SEED,'structural_check':structural_check(),'random':evaluate_case('random'),'planted':evaluate_case('planted',planted_graph())} with open('results.json','w') as f: json.dump(result,f,indent=2) print(json.dumps(result,indent=2))