"""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))