"""Small decoder-only code model; tied embeddings and causal SDPA attention.""" import json from dataclasses import dataclass,asdict import torch from torch import nn from torch.nn import functional as F @dataclass class Config: vocab_size:int=8192 context:int=512 width:int=512 layers:int=8 heads:int=8 dropout:float=.1 class Attention(nn.Module): def __init__(self,c): super().__init__();self.heads=c.heads;self.dropout=c.dropout self.qkv=nn.Linear(c.width,c.width*3);self.out=nn.Linear(c.width,c.width) def forward(self,x,cache=None,use_cache=False): b,t,d=x.shape;q,k,v=self.qkv(x).chunk(3,dim=-1) q,k,v=[z.view(b,t,self.heads,d//self.heads).transpose(1,2) for z in (q,k,v)] mask=None if cache is not None: offset=cache[0].shape[2];k=torch.cat((cache[0],k),dim=2);v=torch.cat((cache[1],v),dim=2) mask=torch.arange(k.shape[2],device=x.device)[None,:]<=torch.arange(offset,offset+t,device=x.device)[:,None] a=F.scaled_dot_product_attention(q,k,v,attn_mask=mask,is_causal=cache is None,dropout_p=self.dropout if self.training else 0.) result=self.out(a.transpose(1,2).contiguous().view(b,t,d)) return (result,(k,v)) if use_cache else result class Block(nn.Module): def __init__(self,c): super().__init__();self.ln1=nn.LayerNorm(c.width);self.attn=Attention(c);self.ln2=nn.LayerNorm(c.width) self.mlp=nn.Sequential(nn.Linear(c.width,c.width*4),nn.GELU(),nn.Linear(c.width*4,c.width),nn.Dropout(c.dropout)) def forward(self,x,cache=None,use_cache=False): if use_cache: a,cache=self.attn(self.ln1(x),cache,use_cache=True);return self.finish(x+a),cache return self.finish(x+self.attn(self.ln1(x))) def finish(self,x):return x+self.mlp(self.ln2(x)) class CodeSprout(nn.Module): def __init__(self,c=Config()): super().__init__();self.config=c;self.tokens=nn.Embedding(c.vocab_size,c.width);self.positions=nn.Embedding(c.context,c.width) self.blocks=nn.ModuleList([Block(c) for _ in range(c.layers)]);self.ln=nn.LayerNorm(c.width);self.lm_head=nn.Linear(c.width,c.vocab_size,bias=False);self.lm_head.weight=self.tokens.weight self.apply(self.init) def init(self,m): if isinstance(m,(nn.Linear,nn.Embedding)): nn.init.normal_(m.weight,std=.02) if isinstance(m,nn.Linear) and m.bias is not None:nn.init.zeros_(m.bias) def forward(self,ids,targets=None,past=None,use_cache=False): offset=past[0][0].shape[2] if past else 0 if not torch.jit.is_tracing() and ids.shape[1]+offset>self.config.context:raise ValueError('Prompt exceeds context limit') x=self.tokens(ids)+self.positions(torch.arange(offset,offset+ids.shape[1],device=ids.device));caches=[] for i,block in enumerate(self.blocks): if use_cache:x,c=block(x,past[i] if past else None,use_cache=True);caches.append(c) else:x=block(x) logits=self.lm_head(self.ln(x));loss=F.cross_entropy(logits.reshape(-1,logits.size(-1)),targets.reshape(-1)) if targets is not None else None return (logits,loss,caches) if use_cache else (logits,loss) @torch.inference_mode() def generate(self,ids,steps=120,temperature=.7,top_k=30,eos=0): self.eval() steps=min(steps,self.config.context-1) limit=max(1,self.config.context-steps) prefix=ids if ids.shape[1]<=limit else torch.cat((ids[:,:1],ids[:,-(limit-1):]),dim=1) if limit>1 else ids[:,:1] logits,_,cache=self(prefix,use_cache=True) for i in range(steps): scores=logits[:,-1,:]/max(.05,temperature) if temperature<=0:nxt=scores.argmax(dim=-1,keepdim=True) else: k=min(top_k,scores.shape[-1]);cut=torch.topk(scores,k).values[:,-1,None];scores=scores.masked_fill(scores