Critical-depth sparse attention / critical_depth_experiment.py

Mechanism failed

Raw ⬇ ZIP
 1import json, math
 2from pathlib import Path
 3import numpy as np
 4
 5SEED = 17
 6
 7def dyadic_rectangles(n):
 8    out=[]
 9    L=int(round(math.log2(n)))
10    for ah in range(L+1):
11        h=2**ah
12        for aw in range(L+1):
13            w=2**aw
14            for y in range(0,n,h):
15                for x in range(0,n,w): out.append((y,x,h,w))
16    return out
17
18def contains(a,b):
19    y,x,h,w=a; yy,xx,hh,ww=b
20    return y<=yy and x<=xx and y+h>=yy+hh and x+w>=xx+ww
21
22def antichain(rects):
23    kept=[]
24    for r in rects:
25        if not any(contains(k,r) or contains(r,k) for k in kept): kept.append(r)
26    return kept
27
28def halo(mask,threshold=.5):
29    n=mask.shape[0]; out=np.zeros_like(mask,dtype=bool)
30    for y,x,h,w in dyadic_rectangles(n):
31        if mask[y:y+h,x:x+w].mean()>threshold: out[y:y+h,x:x+w]=True
32    return out
33
34def ancestors(r, n):
35    y,x,h,w=r; hs=[]; vs=[]
36    for _ in range(int(round(math.log2(n)))+1):
37        hs.append((y,x,h,w)); nw=min(n,2*w); x=(x//nw)*nw; w=nw
38    y,x,h,w=r
39    for _ in range(int(round(math.log2(n)))+1):
40        vs.append((y,x,h,w)); nh=min(n,2*h); y=(y//nh)*nh; h=nh
41    return hs,vs
42
43def in_halo_rect(oh,r):
44    y,x,h,w=r
45    return bool(oh[y:y+h,x:x+w].all())
46
47def depth(r, oh, n):
48    ha,va=ancestors(r,n)
49    e1=sum(in_halo_rect(oh,a) for a in ha)-1
50    e2=sum(in_halo_rect(oh,a) for a in va)-1
51    return max(0,e1),max(0,e2)
52
53def build_H(mask,rects,s):
54    n=mask.shape[0]; oh=halo(mask); H=np.zeros((n,n),float); info=[]
55    for r in rects:
56        y,x,h,w=r; e1,e2=depth(r,oh,n)
57        alpha=(e1+1)**(-s)*(e2+1)**(-(1-s))
58        H[y:y+h,x:x+w]+=alpha
59        info.append((r,e1,e2,alpha))
60    return H,info
61
62def metrics(H,mask):
63    v=H[mask]
64    return {'mean':float(v.mean()),'p99':float(np.quantile(v,.99)),'max':float(v.max()),'exp_mean':float(np.exp(.25*v).mean())}
65
66def selected_rects(n, seed, count=28):
67    rs=dyadic_rectangles(n); rng=np.random.default_rng(seed)
68    weights=np.array([1/(r[2]*r[3])**.35 for r in rs]); weights/=weights.sum()
69    idx=rng.choice(len(rs),size=min(count*4,len(rs)),replace=False,p=weights)
70    return [rs[i] for i in idx]
71
72def router_demo(n=16):
73    mask=np.ones((n,n),bool); rs=selected_rects(n,101,32)
74    hb,ib=build_H(mask,rs,.5); ra=antichain(rs); ha,ia=build_H(mask,ra,.5)
75    return {'candidate_count':len(rs),'antichain_count':len(ra),'baseline':metrics(hb,mask),'idea':metrics(ha,mask)}
76
77def main():
78    n=16; mask=np.ones((n,n),bool); base=selected_rects(n,22,40)
79    sym=[]
80    for s in [.1,.2,.3,.5,.7,.8,.9]:
81        H1,_=build_H(mask,base,s); H2,_=build_H(mask,base,1-s)
82        sym.append({'s':s,'relative_symmetry_error':float(np.max(np.abs(H1-H2))/(np.max(H1)+1e-12)), 'mean':float(H1.mean())})
83    H,info=build_H(mask,antichain(base),.5); scale=[]
84    for lam in [.25,.5,1.,2.,4.]:
85        HH=lam*H; scale.append({'lambda':lam,'mean_H':float(HH.mean()),'log_exp_penalty':float(np.log(np.exp(.25*HH).mean()))})
86    directional=[]
87    for s in [.01,.1,.25,.5,.75,.9,.99]:
88        Hs,_=build_H(mask,antichain(base),s)
89        directional.append({'s':s,'p99':float(np.quantile(Hs[mask],.99)),'max':float(Hs.max())})
90    result={'seed':SEED,'grid':n,'symmetry_sweep':sym,'lambda_sweep':scale,'directional_sweep':directional,'router_demo':router_demo(n)}
91    Path('results.json').write_text(json.dumps(result,indent=2)); print(json.dumps(result,indent=2))
92
93if __name__=='__main__': main()