| |
| |
| |
| import math, torch |
| from torch import nn |
| import torch.nn.functional as F |
|
|
| G=4; KCODE=256 |
|
|
| def diag_scan(g,u,C=16): |
| B,T,D=g.shape; pad=(C-T%C)%C |
| if pad: g=F.pad(g,(0,0,0,pad),value=1.0); u=F.pad(u,(0,0,0,pad),value=0.0) |
| Tp=T+pad; NC=Tp//C; g=g.view(B,NC,C,D); u=u.view(B,NC,C,D) |
| s=torch.zeros(B,NC,D,device=g.device,dtype=g.dtype); P=torch.ones_like(s); loc=[]; Ps=[] |
| for t in range(C): |
| s=g[:,:,t,:]*s+u[:,:,t,:]; loc.append(s); P=P*g[:,:,t,:]; Ps.append(P) |
| s_local=torch.stack(loc,2); Pcum=torch.stack(Ps,2); Gc=Pcum[:,:,-1,:]; sl=s_local[:,:,-1,:] |
| carry=torch.zeros(B,D,device=g.device,dtype=g.dtype); car=[] |
| for c in range(NC): |
| car.append(carry); carry=Gc[:,c,:]*carry+sl[:,c,:] |
| car=torch.stack(car,1); return (s_local+Pcum*car[:,:,None,:]).view(B,Tp,D)[:,:T,:] |
|
|
| class FastBlock(nn.Module): |
| def __init__(self,d,ff): |
| super().__init__(); self.n1=nn.LayerNorm(d); self.proj=nn.Linear(d,4*d,bias=False); self.o=nn.Linear(d,d,bias=False) |
| self.n2=nn.LayerNorm(d); self.up=nn.Linear(d,2*ff,bias=False); self.down=nn.Linear(ff,d,bias=False); self.decay=nn.Parameter(torch.full((d,),2.0)) |
| def forward(self,x,C=16): |
| k,v,r,g=self.proj(self.n1(x)).chunk(4,-1); g=torch.sigmoid(g+self.decay); u=(1-g)*torch.tanh(k) |
| s=diag_scan(g,u,C); x=x+self.o(torch.sigmoid(r)*s*torch.sigmoid(v)) |
| a,b=self.up(self.n2(x)).chunk(2,-1); return x+self.down(F.silu(b)*a) |
|
|
| class STEDecode(torch.autograd.Function): |
| """fwd: decoded codon rows. bwd: pass grad straight through to shadow rows (STE).""" |
| @staticmethod |
| def forward(ctx,shadow_rows,decoded_rows): |
| return decoded_rows |
| @staticmethod |
| def backward(ctx,gout): |
| return gout,gout |
|
|
| class DnaChat(nn.Module): |
| def __init__(self,vocab=32000,d=512,layers=16,ff=1536,a=2048,b=2048,chunk=16): |
| super().__init__(); self.vocab=vocab; self.d=d; self.layers=layers; self.ff=ff; self.a=a; self.b=b; self.chunk=chunk |
| self.ncod=d//G; self.nrows=a*b |
| self.embed=nn.Embedding(vocab,d); nn.init.normal_(self.embed.weight,std=.02) |
| self.blocks=nn.ModuleList([FastBlock(d,ff) for _ in range(layers)]); self.norm=nn.LayerNorm(d) |
| self.ra=nn.Linear(d,a,bias=False); self.rb=nn.Linear(d,b,bias=False) |
| self.ca=nn.Parameter(torch.randn(a,d)*.02); self.cb=nn.Parameter(torch.randn(b,d)*.02) |
| self.head2=nn.Linear(d,vocab,bias=False); nn.init.normal_(self.head2.weight,std=.02) |
| |
| |
| self.codebook=nn.Parameter(torch.randn(4,KCODE,G)*.05) |
| self.register_buffer('codons',torch.randint(0,KCODE,(self.nrows,self.ncod),dtype=torch.uint8)) |
| frames=torch.arange(self.ncod)%4 |
| self.register_buffer('frames',frames) |
| |
| self.expand=nn.Linear(d,2*d,bias=False); self.contract=nn.Linear(2*d,d,bias=False) |
| @torch.no_grad() |
| def mutate_rows(self,idx,new_rows,chunk=65536): |
| """v17 mutation, CHUNKED to bound cdist memory (fixes OOM at large cadences).""" |
| for s0 in range(0,idx.size(0),chunk): |
| sl=slice(s0,min(s0+chunk,idx.size(0))) |
| x=new_rows[sl].view(-1,self.ncod,G) |
| out=torch.empty(x.size(0),self.ncod,dtype=torch.uint8,device=x.device) |
| for fr in range(4): |
| mask=(self.frames==fr) |
| xf=x[:,mask,:].reshape(-1,G) |
| d2=torch.cdist(xf.bfloat16(),self.codebook[fr].bfloat16()) |
| out[:,mask]=d2.argmin(-1).view(x.size(0),-1).to(torch.uint8) |
| self.codons[idx[sl]]=out |
| def decode_rows(self,idx): |
| cod=self.codons[idx].long() |
| fr=self.frames.unsqueeze(0).expand_as(cod) |
| dec=self.codebook[fr,cod].view(-1,self.d).to(getattr(self,'_hdtype',self.codebook.dtype)) |
| leaf=dec.detach().clone().requires_grad_(True) |
| self._leaf=leaf; self._leaf_idx=idx |
| return dec + (leaf - leaf.detach()) |
| def route(self,h): |
| self._hdtype=h.dtype |
| xa,xb=self.ra(h),self.rb(h); va,ia=xa.topk(2,-1); vb,ib=xb.topk(2,-1) |
| s0=va[:,0]+vb[:,0]; sa=va[:,1]+vb[:,0]; sb=va[:,0]+vb[:,1]; use=sa>=sb |
| k0=ia[:,0]*self.b+ib[:,0]; k1=torch.where(use,ia[:,1]*self.b+ib[:,0],ia[:,0]*self.b+ib[:,1]); s1=torch.where(use,sa,sb) |
| alpha=torch.softmax(torch.stack((s0,s1),1),1).unsqueeze(-1) |
| rows=torch.stack((k0,k1),1) |
| dec=self.decode_rows(rows.reshape(-1)).view(-1,2,self.d) |
| ext=(dec*alpha.to(dec.dtype)).sum(1) |
| ext=self.contract(F.silu(self.expand(ext)).view(ext.size(0),2,self.d).flatten(1)) |
| code=(self.ca[ia[:,0]]+self.cb[ib[:,0]]).to(xa.dtype) |
| pa,pb=xa.softmax(-1).mean(0),xb.softmax(-1).mean(0); balance=.5*(self.a*pa.square().sum()+self.b*pb.square().sum()) |
| return ext+code,balance,rows |
| def features(self,ids): |
| B,T=ids.shape; x=self.embed(ids) |
| for blk in self.blocks: x=blk(x,self.chunk) |
| mem,bal,rows=self.route(x.reshape(B*T,self.d)); return self.norm(x+mem.reshape(B,T,self.d)),bal,rows |
| def config(self): return {k:getattr(self,k) for k in ('vocab','d','layers','ff','a','b')} |
|
|
| def counts(m): |
| ctrl=sum(p.numel() for n,p in m.named_parameters()) |
| tbl_bytes=m.nrows*m.ncod |
| return ctrl,tbl_bytes |
|
|