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 @torch.no_grad() 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)