Noncrossing Brace Attention / brace_mvp.py
Mechanism failed
1import itertools, json, math, random, time
2import numpy as np
3import torch
4import torch.nn as nn
5import torch.nn.functional as F
6
7SEED=17
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9torch.set_num_threads(4)
10
11def intervals(L):
12 # Half-open intervals [i,j], allowing empty insertions, in brace order.
13 one=[(i,j) for i in range(L+1) for j in range(i,L+1)]
14 return one
15
16def verify_math(L=6):
17 ts=intervals(L)
18 pairs=[(a,b) for a in ts for b in ts if a[1]<=b[0]]
19 assert all(a[1]<=b[0] for a,b in pairs)
20 assert len(set(pairs))==len(pairs)
21 logits=torch.linspace(-1.0,1.0,len(pairs))
22 alpha=torch.softmax(logits,0)
23 return {'num_ordered_pairs':len(pairs), 'expected_count':math.comb(L+4,4),
24 'alpha_sum':float(alpha.sum()), 'alpha_min':float(alpha.min()),
25 'all_nonoverlap':True}
26
27class Brace(nn.Module):
28 def __init__(self,d,m=2,tau=0.7):
29 super().__init__(); self.d=d; self.tau=tau
30 self.experts=nn.ModuleList([nn.Sequential(nn.Linear(d,2*d),nn.Tanh(),nn.Linear(2*d,d)) for _ in range(m)])
31 self.score=nn.ModuleList([nn.Linear(d,1) for _ in range(m)])
32 self.register_buffer('dummy',torch.zeros(1))
33 def forward(self,x, return_stats=False):
34 B,L,D=x.shape; its=intervals(L); Q=len(its)
35 # Span score is mean pooled input projected by each expert scorer.
36 scores=[]
37 for k in range(2):
38 sk=torch.zeros(B,Q,device=x.device)
39 for q,(a,b) in enumerate(its):
40 pool=x[:,a:b].mean(1) if b>a else x.mean(1)*0
41 sk[:,q]=self.score[k](pool).squeeze(-1)
42 scores.append(sk)
43 # brace score is additive, with a hard ordered/noncrossing support.
44 logits=scores[0][:,None,:]+scores[1][:,:,None]
45 # q index dimensions are interval 1 and interval 2; keep only ordered.
46 valid=torch.tensor([a[1]<=c[0] for a in its for c in its],device=x.device).view(Q,Q)
47 logits=logits.masked_fill(~valid[None],-1e9).flatten(1)
48 alpha=F.softmax(logits/self.tau,1)
49 # Expert outputs for every interval, then splice tokenwise.
50 outs=[]
51 for k in range(2):
52 z=torch.empty(B,Q,L,D,device=x.device)
53 for q,(a,b) in enumerate(its):
54 z[:,q]=x
55 if b>a: z[:,q,a:b]=x[:,a:b]+self.experts[k](x[:,a:b])
56 outs.append(z)
57 y=torch.zeros_like(x)
58 for q,(a,b) in enumerate(its):
59 for r,(c,d) in enumerate(its):
60 w=alpha[:,q*Q+r].view(B,1,1)
61 # ordered tuple is represented by applying each operation to its own span
62 zz=outs[0][:,q]; zz=outs[1][:,r]
63 # Reconstruct from original to avoid overlap ambiguity.
64 z=x.clone()
65 if b>a: z[:,a:b]=outs[0][:,q,a:b]
66 if d>c: z[:,c:d]=outs[1][:,r,c:d]
67 y += w*z
68 if return_stats:
69 chosen=alpha.argmax(1); q=chosen//Q; r=chosen%Q
70 collisions=sum(not (its[int(q[i])][1]<=its[int(r[i])][0]) for i in range(B))
71 return y, {'collision_rate':collisions/B,'entropy':float((-alpha*torch.log(alpha+1e-9)).sum(1).mean())}
72 return y
73
74class Independent(nn.Module):
75 def __init__(self,d):
76 super().__init__(); self.experts=nn.ModuleList([nn.Sequential(nn.Linear(d,2*d),nn.Tanh(),nn.Linear(2*d,d)) for _ in range(2)])
77 self.score=nn.ModuleList([nn.Linear(d,1) for _ in range(2)])
78 def forward(self,x,return_stats=False):
79 B,L,D=x.shape; its=intervals(L); z=x.clone(); picks=[]
80 for k in range(2):
81 ss=[]
82 for a,b in its:
83 p=x[:,a:b].mean(1) if b>a else x.mean(1)*0
84 ss.append(self.score[k](p).squeeze(-1))
85 q=torch.stack(ss,1).argmax(1); picks.append(q)
86 for i in range(B):
87 a,b=its[int(q[i])]
88 if b>a: z[i,a:b]=z[i,a:b]+self.experts[k](x[i,a:b])
89 if return_stats:
90 bad=0
91 for i in range(B):
92 a,b=its[int(picks[0][i])]; c,d=its[int(picks[1][i])]
93 bad += not (b<=c or d<=a)
94 return z, {'collision_rate':bad/B}
95 return z
96
97class Task(nn.Module):
98 def __init__(self,router,d):
99 super().__init__(); self.router=router; self.head=nn.Sequential(nn.Linear(d,16),nn.ReLU(),nn.Linear(16,2))
100 def forward(self,x,stats=False):
101 y=sliceout=self.router(x,stats)
102 if stats: y,st=sliceout
103 h=y.mean(1); out=self.head(h)
104 return (out,st) if stats else out
105
106def data(n,L,d):
107 x=torch.randn(n,L,d); # ordered task: first half and second half have different signed evidence
108 s=x[:,:L//2,0].mean(1)-x[:,L//2:,0].mean(1)
109 y=(s>0).long(); return x,y
110
111def train(kind,train,test,epochs=35):
112 torch.manual_seed(SEED+ (0 if kind=='brace' else 1)); d=train[0].shape[-1]
113 router=Brace(d) if kind=='brace' else Independent(d); model=Task(router,d)
114 opt=torch.optim.Adam(model.parameters(),lr=3e-3); x,y=train
115 t=time.time()
116 for _ in range(epochs):
117 opt.zero_grad(); loss=F.cross_entropy(model(x),y); loss.backward(); opt.step()
118 with torch.no_grad():
119 pred=model(test[0]); acc=(pred.argmax(1)==test[1]).float().mean().item(); vl=F.cross_entropy(pred,test[1]).item()
120 _,st=model(test[0],True)
121 return {'test_loss':vl,'accuracy':acc,'collision_rate':st['collision_rate'],'seconds':time.time()-t,'params':sum(p.numel() for p in model.parameters())}
122
123if __name__=='__main__':
124 report={'math':verify_math()}; trainset=data(384,6,4); testset=data(256,6,4)
125 report['baseline_independent']=train('independent',trainset,testset)
126 report['idea_brace']=train('brace',trainset,testset)
127 print(json.dumps(report,indent=2))