Download model.py from kaafivikrant/First5M: direct link, hf CLI and curl.
- Browser
- Download file 11.2 kB
-
https://huggingface.co/kaafivikrant/First5M/resolve/main/model.py
- Command line
-
hf download hf://kaafivikrant/First5M/model.py
-
curl -L -o model.py https://huggingface.co/kaafivikrant/First5M/resolve/main/model.py
11.2 kB
| """ | |
| model.py - Standalone Transformer LM for inference. | |
| A ~5M parameter decoder-only Transformer language model trained on OpenWebText. | |
| Built from scratch following "Attention Is All You Need" (Vaswani et al., 2017). | |
| Usage: | |
| import torch, tiktoken | |
| from model import ModelConfig, TransformerLM | |
| config = ModelConfig() | |
| model = TransformerLM(config) | |
| state_dict = torch.load("pytorch_model.pt", map_location="cpu", weights_only=True) | |
| model.load_state_dict(state_dict, strict=False) # strict=False: lm_head is weight-tied | |
| model.eval() | |
| enc = tiktoken.get_encoding("gpt2") | |
| ids = torch.tensor([enc.encode("Once upon a time")]) | |
| out = model.generate(ids, max_new_tokens=100, temperature=0.8, top_k=50) | |
| print(enc.decode(out[0].tolist())) | |
| """ | |
| import math | |
| from dataclasses import dataclass | |
| from typing import List, Optional, Tuple | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| HAS_SDPA = hasattr(F, "scaled_dot_product_attention") | |
| class ModelConfig: | |
| """Architecture hyperparameters.""" | |
| vocab_size: int = 50257 # GPT-2 BPE vocabulary size | |
| d_model: int = 256 # Hidden dimension | |
| n_heads: int = 4 # Number of attention heads | |
| n_layers: int = 6 # Number of Transformer blocks | |
| d_ff: int = 1024 # Feed-forward inner dimension (4 * d_model) | |
| max_seq_len: int = 256 # Maximum sequence length (context window) | |
| dropout: float = 0.1 # Dropout rate | |
| class SinusoidalPositionalEncoding(nn.Module): | |
| """Sinusoidal Positional Encoding (Section 3.5 of the original paper).""" | |
| def __init__(self, d_model: int, max_seq_len: int = 5000, dropout: float = 0.1): | |
| super().__init__() | |
| self.dropout = nn.Dropout(p=dropout) | |
| pe = torch.zeros(max_seq_len, d_model) | |
| position = torch.arange(0, max_seq_len, dtype=torch.float).unsqueeze(1) | |
| div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) | |
| pe[:, 0::2] = torch.sin(position * div_term) | |
| pe[:, 1::2] = torch.cos(position * div_term) | |
| self.register_buffer("pe", pe.unsqueeze(0)) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = x + self.pe[:, :x.size(1)] | |
| return self.dropout(x) | |
| class MultiHeadSelfAttention(nn.Module): | |
| """Multi-Head Self-Attention with causal masking and KV-cache support.""" | |
| def __init__(self, d_model: int, n_heads: int, max_seq_len: int = 512, dropout: float = 0.1): | |
| super().__init__() | |
| assert d_model % n_heads == 0 | |
| self.n_heads = n_heads | |
| self.d_k = d_model // n_heads | |
| self.dropout = dropout | |
| self.W_q = nn.Linear(d_model, d_model, bias=False) | |
| self.W_k = nn.Linear(d_model, d_model, bias=False) | |
| self.W_v = nn.Linear(d_model, d_model, bias=False) | |
| self.W_o = nn.Linear(d_model, d_model, bias=False) | |
| self.attn_dropout = nn.Dropout(dropout) | |
| self.resid_dropout = nn.Dropout(dropout) | |
| if not HAS_SDPA: | |
| self.register_buffer( | |
| "causal_mask", | |
| torch.tril(torch.ones(max_seq_len, max_seq_len)).view(1, 1, max_seq_len, max_seq_len), | |
| ) | |
| def forward(self, x: torch.Tensor, kv_cache: Optional[Tuple[torch.Tensor, torch.Tensor]] = None): | |
| B, T, C = x.shape | |
| q = self.W_q(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2) | |
| k = self.W_k(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2) | |
| v = self.W_v(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2) | |
| new_cache = None | |
| if kv_cache is not None: | |
| k_prev, v_prev = kv_cache | |
| k = torch.cat([k_prev, k], dim=2) | |
| v = torch.cat([v_prev, v], dim=2) | |
| new_cache = (k, v) | |
| if HAS_SDPA: | |
| out = F.scaled_dot_product_attention( | |
| q, k, v, | |
| is_causal=(kv_cache is None), | |
| dropout_p=self.dropout if self.training else 0.0, | |
| ) | |
| else: | |
| S = k.size(2) | |
| attn = (q @ k.transpose(-2, -1)) * (self.d_k ** -0.5) | |
| if kv_cache is None: | |
| attn = attn.masked_fill(self.causal_mask[:, :, :T, :T] == 0, float("-inf")) | |
| attn = self.attn_dropout(F.softmax(attn, dim=-1)) | |
| out = attn @ v | |
| out = out.transpose(1, 2).contiguous().view(B, T, C) | |
| return self.resid_dropout(self.W_o(out)), new_cache | |
| class FeedForward(nn.Module): | |
| """Position-wise Feed-Forward Network: expand -> GELU -> contract.""" | |
| def __init__(self, d_model: int, d_ff: int, dropout: float = 0.1): | |
| super().__init__() | |
| self.net = nn.Sequential( | |
| nn.Linear(d_model, d_ff), | |
| nn.GELU(), | |
| nn.Linear(d_ff, d_model), | |
| nn.Dropout(dropout), | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.net(x) | |
| class TransformerBlock(nn.Module): | |
| """Pre-LayerNorm Transformer decoder block (attention + FFN + residuals).""" | |
| def __init__(self, d_model: int, n_heads: int, d_ff: int, max_seq_len: int = 512, dropout: float = 0.1): | |
| super().__init__() | |
| self.ln1 = nn.LayerNorm(d_model) | |
| self.attn = MultiHeadSelfAttention(d_model, n_heads, max_seq_len, dropout) | |
| self.ln2 = nn.LayerNorm(d_model) | |
| self.ff = FeedForward(d_model, d_ff, dropout) | |
| def forward(self, x: torch.Tensor, kv_cache=None): | |
| attn_out, new_cache = self.attn(self.ln1(x), kv_cache=kv_cache) | |
| x = x + attn_out | |
| x = x + self.ff(self.ln2(x)) | |
| return x, new_cache | |
| class TransformerLM(nn.Module): | |
| """ | |
| Decoder-only Transformer Language Model. | |
| Features: | |
| - Pre-LayerNorm architecture (GPT-2 style) | |
| - Sinusoidal positional encoding | |
| - Weight tying between embedding and output head | |
| - KV-cache for efficient autoregressive generation | |
| - Repetition penalty for better generation quality | |
| """ | |
| def __init__(self, config: ModelConfig): | |
| super().__init__() | |
| self.config = config | |
| self.token_embedding = nn.Embedding(config.vocab_size, config.d_model) | |
| self.pos_encoding = SinusoidalPositionalEncoding( | |
| config.d_model, config.max_seq_len + 2048, config.dropout | |
| ) | |
| self.blocks = nn.ModuleList([ | |
| TransformerBlock(config.d_model, config.n_heads, config.d_ff, | |
| config.max_seq_len, config.dropout) | |
| for _ in range(config.n_layers) | |
| ]) | |
| self.ln_f = nn.LayerNorm(config.d_model) | |
| self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False) | |
| # Weight tying: embedding and output head share the same weights | |
| self.token_embedding.weight = self.lm_head.weight | |
| self.apply(self._init_weights) | |
| for pn, p in self.named_parameters(): | |
| if pn.endswith("W_o.weight") or pn.endswith("net.2.weight"): | |
| torch.nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layers)) | |
| def _init_weights(self, module): | |
| if isinstance(module, nn.Linear): | |
| torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) | |
| if module.bias is not None: | |
| torch.nn.init.zeros_(module.bias) | |
| elif isinstance(module, nn.Embedding): | |
| torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) | |
| def forward(self, idx, targets=None): | |
| x = self.pos_encoding(self.token_embedding(idx)) | |
| for block in self.blocks: | |
| x, _ = block(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, max_new_tokens, temperature=1.0, top_k=None, | |
| repetition_penalty=1.2): | |
| """ | |
| Autoregressive text generation with KV-cache and repetition penalty. | |
| Args: | |
| idx: Prompt token IDs, shape [batch, prompt_len] | |
| max_new_tokens: Number of tokens to generate | |
| temperature: Sampling temperature (default 1.0) | |
| top_k: Only sample from top K tokens (default None = all) | |
| repetition_penalty: Penalty for repeated tokens (1.0 = off, 1.2 = default) | |
| Returns: | |
| idx: Prompt + generated tokens, shape [batch, prompt_len + max_new_tokens] | |
| """ | |
| kv_caches: List[Optional[Tuple[torch.Tensor, torch.Tensor]]] = [None] * len(self.blocks) | |
| # Phase 1: Prefill — process entire prompt | |
| x = self.pos_encoding(self.token_embedding(idx)) | |
| for i, block in enumerate(self.blocks): | |
| x, kv_caches[i] = block(x) | |
| logits = self.lm_head(self.ln_f(x)) | |
| logits = logits[:, -1, :] | |
| if repetition_penalty != 1.0: | |
| for b in range(idx.size(0)): | |
| seen = idx[b].unique() | |
| for token_id in seen: | |
| if logits[b, token_id] > 0: | |
| logits[b, token_id] /= repetition_penalty | |
| else: | |
| logits[b, token_id] *= repetition_penalty | |
| logits = logits / temperature | |
| if top_k is not None: | |
| v, _ = torch.topk(logits, min(top_k, logits.size(-1))) | |
| logits[logits < v[:, [-1]]] = float("-inf") | |
| probs = F.softmax(logits, dim=-1) | |
| idx_next = torch.multinomial(probs, num_samples=1) | |
| idx = torch.cat((idx, idx_next), dim=1) | |
| # Phase 2: Decode — generate one token at a time with KV-cache | |
| for _ in range(max_new_tokens - 1): | |
| seq_pos = idx.size(1) - 1 | |
| x = self.token_embedding(idx_next) | |
| x = x + self.pos_encoding.pe[:, seq_pos:seq_pos + 1] | |
| for i, block in enumerate(self.blocks): | |
| x, kv_caches[i] = block(x, kv_cache=kv_caches[i]) | |
| logits = self.lm_head(self.ln_f(x)) | |
| logits = logits[:, -1, :] | |
| if repetition_penalty != 1.0: | |
| for b in range(idx.size(0)): | |
| seen = idx[b].unique() | |
| for token_id in seen: | |
| if logits[b, token_id] > 0: | |
| logits[b, token_id] /= repetition_penalty | |
| else: | |
| logits[b, token_id] *= repetition_penalty | |
| logits = logits / temperature | |
| if top_k is not None: | |
| v, _ = torch.topk(logits, min(top_k, logits.size(-1))) | |
| logits[logits < v[:, [-1]]] = float("-inf") | |
| probs = F.softmax(logits, dim=-1) | |
| idx_next = torch.multinomial(probs, num_samples=1) | |
| idx = torch.cat((idx, idx_next), dim=1) | |
| if idx.size(1) > self.config.max_seq_len: | |
| for i in range(len(kv_caches)): | |
| if kv_caches[i] is not None: | |
| k, v_tensor = kv_caches[i] | |
| kv_caches[i] = (k[:, :, -self.config.max_seq_len:, :], | |
| v_tensor[:, :, -self.config.max_seq_len:, :]) | |
| return idx | |