achavan1211's picture
Upload instruction-tuned MiniMind (137.7M params, step 1000)
33d751b verified
Raw
History Blame Contribute Delete
5.44 kB
# 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