File size: 2,882 Bytes
1372fb6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
"""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))