First5M / model.py
kaafivikrant's picture
Upload folder using huggingface_hub
1be11eb verified
Raw History Blame Contribute Delete
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")
@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