PunyPunk / model.py
mohar07's picture
train slm: 17.78M params
fa974cd verified
Raw History Blame Contribute Delete
3.07 kB
import os
import json
from huggingface_hub import snapshot_download
import torch
import torch.nn as nn
import torch.nn.functional as F
class DenseTransformerBlock(nn.Module):
def __init__(self, n_embd, n_head, block_size, dropout):
super().__init__()
self.ln1 = nn.LayerNorm(n_embd)
self.attn = nn.MultiheadAttention(n_embd, n_head, dropout=dropout, batch_first=True)
self.ln2 = nn.LayerNorm(n_embd)
self.mlp = nn.Sequential(
nn.Linear(n_embd, 4 * n_embd),
nn.GELU(),
nn.Linear(4 * n_embd, n_embd),
nn.Dropout(dropout),
)
self.register_buffer("mask", torch.triu(torch.full((block_size, block_size), float("-inf")), diagonal=1))
def forward(self, x):
t = x.shape[1]
a, _ = self.attn(self.ln1(x), self.ln1(x), self.ln1(x), attn_mask=self.mask[:t, :t])
x = x + a
x = x + self.mlp(self.ln2(x))
return x
class SLM(nn.Module):
def __init__(self, vocab_size, n_embd=384, n_head=8, n_layer=10, block_size=64, dropout=0.1):
super().__init__()
self.token_emb = nn.Embedding(vocab_size, n_embd)
self.pos_emb = nn.Embedding(block_size, n_embd)
self.blocks = nn.Sequential(*[DenseTransformerBlock(n_embd, n_head, block_size, dropout) for _ in range(n_layer)])
self.ln_f = nn.LayerNorm(n_embd)
self.lm_head = nn.Linear(n_embd, vocab_size, bias=False)
self.token_emb.weight = self.lm_head.weight
self.block_size = block_size
def forward(self, idx, targets=None):
b, t = idx.shape
x = self.token_emb(idx) + self.pos_emb(torch.arange(t, device=idx.device))
x = self.blocks(x)
logits = self.lm_head(self.ln_f(x))
loss = None
if targets is not None:
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))
return logits, loss
@torch.no_grad()
def generate(self, idx, new_tokens, temperature=1.0):
for _ in range(new_tokens):
idx_cond = idx[:, -self.block_size :]
logits, _ = self.forward(idx_cond)
logits = logits[:, -1, :] / temperature
probs = F.softmax(logits, dim=-1)
idx = torch.cat([idx, torch.multinomial(probs, 1)], dim=1)
return idx
def load_model(repo_or_dir):
ckpt_dir = repo_or_dir if os.path.isdir(repo_or_dir) else snapshot_download(repo_or_dir)
cfg = json.load(open(os.path.join(ckpt_dir, "config.json")))
vocab = cfg["vocab"]
arch = cfg["arch"]
model = SLM(len(vocab), **arch)
model.load_state_dict(torch.load(os.path.join(ckpt_dir, "pytorch_model.bin"), map_location="cpu"))
return model, vocab
def sample(model, vocab, prompt, new_tokens=200, temperature=1.0):
stoi = {c: i for i, c in enumerate(vocab)}
itos = {i: c for c, i in stoi.items()}
idx = torch.tensor([[stoi[c] for c in prompt]], dtype=torch.long)
out = model.generate(idx, new_tokens, temperature=temperature)[0].tolist()
return "".join(itos[i] for i in out)