File size: 5,442 Bytes
33d751b | 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 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 | # MiniMind - a minimal GPT-style decoder, built from scratch.
# Load with: model, cfg = load_model("config.json", "model.pt")
import json
import torch
import torch.nn as nn
from torch.nn import functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, n_embd, n_head, block_size, dropout):
super().__init__()
assert n_embd % n_head == 0
self.head_size = n_embd // n_head
self.n_head = n_head
self.key = nn.Linear(n_embd, n_embd, bias=False)
self.query = nn.Linear(n_embd, n_embd, bias=False)
self.value = nn.Linear(n_embd, n_embd, bias=False)
self.proj = nn.Linear(n_embd, n_embd)
self.dropout = nn.Dropout(dropout)
self.register_buffer("tril", torch.tril(torch.ones(block_size, block_size)))
def forward(self, x):
B, T, C = x.shape
k = self.key(x).view(B, T, self.n_head, self.head_size).transpose(1, 2)
q = self.query(x).view(B, T, self.n_head, self.head_size).transpose(1, 2)
v = self.value(x).view(B, T, self.n_head, self.head_size).transpose(1, 2)
att = (q @ k.transpose(-2, -1)) * (self.head_size ** -0.5)
att = att.masked_fill(self.tril[:T, :T] == 0, float("-inf"))
att = F.softmax(att, dim=-1)
att = self.dropout(att)
y = (att @ v).transpose(1, 2).contiguous().view(B, T, C)
return self.proj(y)
class FeedForward(nn.Module):
def __init__(self, n_embd, dropout):
super().__init__()
self.net = nn.Sequential(
nn.Linear(n_embd, 4 * n_embd),
nn.GELU(),
nn.Linear(4 * n_embd, n_embd),
nn.Dropout(dropout),
)
def forward(self, x):
return self.net(x)
class Block(nn.Module):
def __init__(self, n_embd, n_head, block_size, dropout):
super().__init__()
self.ln1 = nn.LayerNorm(n_embd)
self.attn = MultiHeadAttention(n_embd, n_head, block_size, dropout)
self.ln2 = nn.LayerNorm(n_embd)
self.ffwd = FeedForward(n_embd, dropout)
def forward(self, x):
x = x + self.attn(self.ln1(x))
x = x + self.ffwd(self.ln2(x))
return x
class MiniMind(nn.Module):
def __init__(self, vocab_size, n_embd, n_layer, n_head, block_size, dropout=0.0):
super().__init__()
self.block_size = block_size
self.token_emb = nn.Embedding(vocab_size, n_embd)
self.pos_emb = nn.Embedding(block_size, n_embd)
self.blocks = nn.Sequential(
*[Block(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)
def forward(self, idx, targets=None):
B, T = idx.shape
assert T <= self.block_size, f"input length {T} > block_size {self.block_size}"
tok_emb = self.token_emb(idx)
pos_emb = self.pos_emb(torch.arange(T, device=idx.device))
x = self.blocks(tok_emb + pos_emb)
x = self.ln_f(x)
logits = self.lm_head(x)
loss = None
if targets is not None:
Bt, Tt, C = logits.shape
loss = F.cross_entropy(logits.view(Bt * Tt, C), targets.view(Bt * Tt),
ignore_index=-100)
return logits, loss
@torch.no_grad()
def generate(self, idx, max_new_tokens, temperature=1.0, top_k=None,
top_p=None, repetition_penalty=1.1, eos_id=None, stop_on_eos=True):
generated = idx.clone()
for _ in range(max_new_tokens):
idx_cond = generated[:, -self.block_size:]
logits, _ = self(idx_cond)
logits = logits[:, -1, :] / temperature
if repetition_penalty != 1.0:
for b in range(generated.size(0)):
for t in set(generated[b].tolist()):
logits[b, t] = logits[b, t] / repetition_penalty
if top_k is not None:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = float("-inf")
if top_p is not None:
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
sorted_indices_to_remove = cumulative_probs > top_p
sorted_indices_to_remove[:, 1:] = sorted_indices_to_remove[:, :-1].clone()
sorted_indices_to_remove[:, 0] = False
indices_to_remove = sorted_indices_to_remove.scatter(
1, sorted_indices, sorted_indices_to_remove)
logits[indices_to_remove] = float("-inf")
probs = F.softmax(logits, dim=-1)
nxt = torch.multinomial(probs, num_samples=1)
generated = torch.cat([generated, nxt], dim=1)
if stop_on_eos and eos_id is not None and nxt.item() == eos_id:
break
return generated
def load_model(config_path="config.json", weights_path="model.pt", device="cpu"):
cfg = json.load(open(config_path))
model = MiniMind(
cfg["vocab_size"], cfg["n_embd"], cfg["n_layer"],
cfg["n_head"], cfg["block_size"], cfg.get("dropout", 0.0),
)
ck = torch.load(weights_path, map_location=device, weights_only=False)
model.load_state_dict(ck["model_state_dict"])
return model.to(device).eval(), cfg
|