# 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