Neighborhood-separator attention / experiment.py

Failed on benchmark

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