Critical-depth sparse attention / critical_depth_experiment.py
Mechanism failed
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()