| |
| |
| 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 |
|
|