Noncrossing Brace Attention / brace_mvp.py

Mechanism failed

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