Download model.py from Raymond1938-4/code-sprout-30m: direct link, hf CLI and curl.
- Browser
- Download file 4.2 kB
-
https://huggingface.co/Raymond1938-4/code-sprout-30m/resolve/main/model.py
- Command line
-
hf download hf://Raymond1938-4/code-sprout-30m/model.py
-
curl -L -o model.py https://huggingface.co/Raymond1938-4/code-sprout-30m/resolve/main/model.py
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 | |
| 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) | |
| 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 | |