File size: 4,652 Bytes
29f25be | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 | """Small pre-norm decoder-only GPT, with tied embedding/output weights."""
import math
from dataclasses import asdict, dataclass
import torch
from torch import nn
from torch.nn import functional as F
@dataclass(frozen=True)
class GPTConfig:
vocab_size: int = 16384
context_length: int = 128
d_model: int = 256
n_heads: int = 4
n_layers: int = 4
d_ff: int = 1024
dropout: float = 0.0
def __post_init__(self):
if min(self.vocab_size, self.context_length, self.d_model, self.n_heads, self.n_layers, self.d_ff) < 1:
raise ValueError("Model dimensions must be positive.")
if self.d_model % self.n_heads or not 0 <= self.dropout < 1:
raise ValueError("Invalid attention heads or dropout.")
class CausalAttention(nn.Module):
def __init__(self, config):
super().__init__()
self.heads = config.n_heads
self.dropout = config.dropout
self.qkv = nn.Linear(config.d_model, 3 * config.d_model)
self.projection = nn.Linear(config.d_model, config.d_model)
def forward(self, hidden):
batch, length, width = hidden.shape
qkv = self.qkv(hidden).view(batch, length, 3, self.heads, width // self.heads)
query, key, value = qkv.permute(2, 0, 3, 1, 4).unbind(0)
# Right padding is always AFTER valid tokens. Causality prevents a valid
# query from seeing PAD, so no dense per-sample mask is necessary.
attended = F.scaled_dot_product_attention(query, key, value, is_causal=True,
dropout_p=self.dropout if self.training else 0.0)
return self.projection(attended.transpose(1, 2).reshape(batch, length, width))
class DecoderBlock(nn.Module):
def __init__(self, config):
super().__init__()
self.attention_norm = nn.LayerNorm(config.d_model)
self.attention = CausalAttention(config)
self.ffn_norm = nn.LayerNorm(config.d_model)
self.ffn = nn.Sequential(nn.Linear(config.d_model, config.d_ff),
nn.GELU(), nn.Linear(config.d_ff, config.d_model))
self.dropout = nn.Dropout(config.dropout)
def forward(self, hidden):
hidden = hidden + self.dropout(self.attention(self.attention_norm(hidden)))
return hidden + self.dropout(self.ffn(self.ffn_norm(hidden)))
class TinyGPT(nn.Module):
def __init__(self, config=GPTConfig()):
super().__init__()
self.config = config
self.token_embedding = nn.Embedding(config.vocab_size, config.d_model)
self.position_embedding = nn.Embedding(config.context_length, config.d_model)
self.dropout = nn.Dropout(config.dropout)
self.blocks = nn.ModuleList(DecoderBlock(config) for _ in range(config.n_layers))
self.final_norm = nn.LayerNorm(config.d_model)
self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
self.apply(self._initialize)
self.lm_head.weight = self.token_embedding.weight
for block in self.blocks:
for projection in (block.attention.projection, block.ffn[2]):
nn.init.normal_(projection.weight, std=0.02 / math.sqrt(2 * config.n_layers))
@staticmethod
def _initialize(module):
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
def forward(self, input_ids, labels=None):
if input_ids.ndim != 2 or not 1 <= input_ids.shape[1] <= self.config.context_length:
raise ValueError("Expected [batch, time] input within model context.")
positions = torch.arange(input_ids.shape[1], device=input_ids.device)
hidden = self.dropout(self.token_embedding(input_ids) + self.position_embedding(positions))
for block in self.blocks:
hidden = block(hidden)
hidden = self.final_norm(hidden)
if labels is None:
return self.lm_head(hidden)
if labels.shape != input_ids.shape:
raise ValueError("Labels must match inputs; labels are already shifted by DataLoader.")
valid = labels != -100
# Avoid allocating [B,T,V] logits for PAD positions; supervision is unchanged.
logits = self.lm_head(hidden[valid])
return {"loss_sum": F.cross_entropy(logits, labels[valid], reduction="sum"),
"token_count": valid.sum()}
def parameter_count(self):
return sum(parameter.numel() for parameter in self.parameters())
def configuration(self):
return asdict(self.config)
|