Download model.py from mohar07/PunyPunk: direct link, hf CLI and curl.
- Browser
- Download file 3.07 kB
-
https://huggingface.co/mohar07/PunyPunk/resolve/main/model.py
- Command line
-
hf download hf://mohar07/PunyPunk/model.py
-
curl -L -o model.py https://huggingface.co/mohar07/PunyPunk/resolve/main/model.py
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 | |
| 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) | |