| """Load and generate with Compactbot/tinystories-50m (custom GPT, 54.8M params). |
| Requires: torch, safetensors, tokenizers. GPU optional (CPU works, slower). |
| """ |
| import torch, torch.nn as nn, torch.nn.functional as F |
| from safetensors.torch import load_file |
| from tokenizers import Tokenizer |
|
|
| VOCAB,D,L,H,FFN,SEQ=8192,512,16,8,2048,512 |
|
|
| class RMSNorm(nn.Module): |
| def __init__(self,d): super().__init__(); self.w=nn.Parameter(torch.ones(d)) |
| def forward(self,x): return self.w*x*torch.rsqrt(x.float().pow(2).mean(-1,keepdim=True)+1e-6) |
|
|
| class Block(nn.Module): |
| def __init__(self,d,h): |
| super().__init__() |
| self.ln1=RMSNorm(d); self.ln2=RMSNorm(d) |
| self.qkv=nn.Linear(d,3*d,bias=False); self.proj=nn.Linear(d,d,bias=False) |
| self.fc1=nn.Linear(d,FFN,bias=False); self.fc2=nn.Linear(FFN,d,bias=False) |
| self.h,self.d=h,d |
| def forward(self,x): |
| B,T,Dd=x.shape; h=self.ln1(x) |
| qkv=self.qkv(h).view(B,T,3,self.h,Dd//self.h).transpose(2,1) |
| q,k,v=qkv[:,0].transpose(1,2),qkv[:,1].transpose(1,2),qkv[:,2].transpose(1,2) |
| att=F.scaled_dot_product_attention(q,k,v,is_causal=True) |
| x=x+self.proj(att.transpose(1,2).reshape(B,T,Dd)) |
| return x+self.fc2(F.gelu(self.fc1(self.ln2(x)))) |
|
|
| class GPT(nn.Module): |
| def __init__(self): |
| super().__init__() |
| self.tok=nn.Embedding(VOCAB,D); self.pos=nn.Embedding(SEQ,D) |
| self.blocks=nn.ModuleList([Block(D,H) for _ in range(L)]); self.ln_f=RMSNorm(D) |
| def forward(self,x,targets=None): |
| h=self.tok(x)+self.pos(torch.arange(x.shape[1],device=x.device)) |
| for b in self.blocks: h=b(h) |
| logits=self.ln_f(h)@self.tok.weight.t() |
| if targets is not None: return F.cross_entropy(logits.view(-1,VOCAB),targets.view(-1)) |
| return logits |
|
|
| def load(repo_dir=".",device="cuda" if torch.cuda.is_available() else "cpu"): |
| m=GPT().to(device); m.load_state_dict(load_file(f"{repo_dir}/model.safetensors"),strict=True) |
| m.eval() |
| return m,Tokenizer.from_file(f"{repo_dir}/tokenizer.json") |
|
|
| def generate(model,tok,prompt,max_new=120,seed=0,temperature=0.8,device="cuda"): |
| g=torch.Generator(device=device).manual_seed(seed) |
| ids=torch.tensor([tok.encode(prompt).ids],device=device) |
| if ids.shape[1]>SEQ-4: ids=ids[:,-(SEQ-4):] |
| with torch.no_grad(): |
| for _ in range(max_new): |
| with torch.autocast(device_type=device,dtype=torch.bfloat16): |
| logits=model(ids) |
| p=torch.softmax(logits[:,-1].float()/temperature,dim=-1) |
| nxt=torch.multinomial(p,1,generator=g); ids=torch.cat([ids,nxt],1) |
| if nxt.item()==1: break |
| return tok.decode(ids[0].tolist()) |
|
|
| if __name__=="__main__": |
| m,t=load(".") |
| print("params:",sum(p.numel() for p in m.parameters())) |
| print(generate(m,t,"Once upon a time, there was a little girl named Lily.",seed=0)) |
|
|