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