code-sprout-30m / model.py
Raymond1938-4's picture
Release experimental CodeSprout 30M checkpoint
61790a1 verified
Raw History Blame Contribute Delete
4.2 kB
"""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<cut,float('-inf'));nxt=torch.multinomial(F.softmax(scores,dim=-1),1)
ids=torch.cat((ids,nxt),dim=1)
if torch.all(nxt==eos):break
if i<steps-1:logits,_,cache=self(nxt,past=cache,use_cache=True)
return ids