import itertools, json, math, random, time import numpy as np import torch import torch.nn as nn import torch.nn.functional as F SEED=17 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) def intervals(L): # Half-open intervals [i,j], allowing empty insertions, in brace order. one=[(i,j) for i in range(L+1) for j in range(i,L+1)] return one def verify_math(L=6): ts=intervals(L) pairs=[(a,b) for a in ts for b in ts if a[1]<=b[0]] assert all(a[1]<=b[0] for a,b in pairs) assert len(set(pairs))==len(pairs) logits=torch.linspace(-1.0,1.0,len(pairs)) alpha=torch.softmax(logits,0) return {'num_ordered_pairs':len(pairs), 'expected_count':math.comb(L+4,4), 'alpha_sum':float(alpha.sum()), 'alpha_min':float(alpha.min()), 'all_nonoverlap':True} class Brace(nn.Module): def __init__(self,d,m=2,tau=0.7): super().__init__(); self.d=d; self.tau=tau self.experts=nn.ModuleList([nn.Sequential(nn.Linear(d,2*d),nn.Tanh(),nn.Linear(2*d,d)) for _ in range(m)]) self.score=nn.ModuleList([nn.Linear(d,1) for _ in range(m)]) self.register_buffer('dummy',torch.zeros(1)) def forward(self,x, return_stats=False): B,L,D=x.shape; its=intervals(L); Q=len(its) # Span score is mean pooled input projected by each expert scorer. scores=[] for k in range(2): sk=torch.zeros(B,Q,device=x.device) for q,(a,b) in enumerate(its): pool=x[:,a:b].mean(1) if b>a else x.mean(1)*0 sk[:,q]=self.score[k](pool).squeeze(-1) scores.append(sk) # brace score is additive, with a hard ordered/noncrossing support. logits=scores[0][:,None,:]+scores[1][:,:,None] # q index dimensions are interval 1 and interval 2; keep only ordered. valid=torch.tensor([a[1]<=c[0] for a in its for c in its],device=x.device).view(Q,Q) logits=logits.masked_fill(~valid[None],-1e9).flatten(1) alpha=F.softmax(logits/self.tau,1) # Expert outputs for every interval, then splice tokenwise. outs=[] for k in range(2): z=torch.empty(B,Q,L,D,device=x.device) for q,(a,b) in enumerate(its): z[:,q]=x if b>a: z[:,q,a:b]=x[:,a:b]+self.experts[k](x[:,a:b]) outs.append(z) y=torch.zeros_like(x) for q,(a,b) in enumerate(its): for r,(c,d) in enumerate(its): w=alpha[:,q*Q+r].view(B,1,1) # ordered tuple is represented by applying each operation to its own span zz=outs[0][:,q]; zz=outs[1][:,r] # Reconstruct from original to avoid overlap ambiguity. z=x.clone() if b>a: z[:,a:b]=outs[0][:,q,a:b] if d>c: z[:,c:d]=outs[1][:,r,c:d] y += w*z if return_stats: chosen=alpha.argmax(1); q=chosen//Q; r=chosen%Q collisions=sum(not (its[int(q[i])][1]<=its[int(r[i])][0]) for i in range(B)) return y, {'collision_rate':collisions/B,'entropy':float((-alpha*torch.log(alpha+1e-9)).sum(1).mean())} return y class Independent(nn.Module): def __init__(self,d): super().__init__(); self.experts=nn.ModuleList([nn.Sequential(nn.Linear(d,2*d),nn.Tanh(),nn.Linear(2*d,d)) for _ in range(2)]) self.score=nn.ModuleList([nn.Linear(d,1) for _ in range(2)]) def forward(self,x,return_stats=False): B,L,D=x.shape; its=intervals(L); z=x.clone(); picks=[] for k in range(2): ss=[] for a,b in its: p=x[:,a:b].mean(1) if b>a else x.mean(1)*0 ss.append(self.score[k](p).squeeze(-1)) q=torch.stack(ss,1).argmax(1); picks.append(q) for i in range(B): a,b=its[int(q[i])] if b>a: z[i,a:b]=z[i,a:b]+self.experts[k](x[i,a:b]) if return_stats: bad=0 for i in range(B): a,b=its[int(picks[0][i])]; c,d=its[int(picks[1][i])] bad += not (b<=c or d<=a) return z, {'collision_rate':bad/B} return z class Task(nn.Module): def __init__(self,router,d): super().__init__(); self.router=router; self.head=nn.Sequential(nn.Linear(d,16),nn.ReLU(),nn.Linear(16,2)) def forward(self,x,stats=False): y=sliceout=self.router(x,stats) if stats: y,st=sliceout h=y.mean(1); out=self.head(h) return (out,st) if stats else out def data(n,L,d): x=torch.randn(n,L,d); # ordered task: first half and second half have different signed evidence s=x[:,:L//2,0].mean(1)-x[:,L//2:,0].mean(1) y=(s>0).long(); return x,y def train(kind,train,test,epochs=35): torch.manual_seed(SEED+ (0 if kind=='brace' else 1)); d=train[0].shape[-1] router=Brace(d) if kind=='brace' else Independent(d); model=Task(router,d) opt=torch.optim.Adam(model.parameters(),lr=3e-3); x,y=train t=time.time() for _ in range(epochs): opt.zero_grad(); loss=F.cross_entropy(model(x),y); loss.backward(); opt.step() with torch.no_grad(): pred=model(test[0]); acc=(pred.argmax(1)==test[1]).float().mean().item(); vl=F.cross_entropy(pred,test[1]).item() _,st=model(test[0],True) return {'test_loss':vl,'accuracy':acc,'collision_rate':st['collision_rate'],'seconds':time.time()-t,'params':sum(p.numel() for p in model.parameters())} if __name__=='__main__': report={'math':verify_math()}; trainset=data(384,6,4); testset=data(256,6,4) report['baseline_independent']=train('independent',trainset,testset) report['idea_brace']=train('brace',trainset,testset) print(json.dumps(report,indent=2))