File size: 5,553 Bytes
4c29028 | 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 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 | """Dense decoder-only GPT model used by train.py."""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
from config import ModelConfig
class RMSNorm(nn.Module):
def __init__(self, width: int, eps: float = 1e-5):
super().__init__()
self.weight = nn.Parameter(torch.ones(width))
self.eps = eps
def forward(self, x: torch.Tensor) -> torch.Tensor:
variance = x.float().square().mean(dim=-1, keepdim=True)
return (x * torch.rsqrt(variance + self.eps)).type_as(x) * self.weight
class Attention(nn.Module):
def __init__(self, cfg: ModelConfig):
super().__init__()
assert cfg.d_model % cfg.n_heads == 0
self.n_heads = cfg.n_heads
self.head_dim = cfg.d_model // cfg.n_heads
self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model, bias=False)
self.output = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
self.dropout = cfg.dropout
def forward(self, x: torch.Tensor) -> torch.Tensor:
batch, length, width = x.shape
q, k, v = self.qkv(x).chunk(3, dim=-1)
q = q.view(batch, length, self.n_heads, self.head_dim).transpose(1, 2)
k = k.view(batch, length, self.n_heads, self.head_dim).transpose(1, 2)
v = v.view(batch, length, self.n_heads, self.head_dim).transpose(1, 2)
output = F.scaled_dot_product_attention(
q, k, v, is_causal=True, dropout_p=self.dropout if self.training else 0.0
)
return self.output(output.transpose(1, 2).contiguous().view(batch, length, width))
class SwiGLU(nn.Module):
def __init__(self, cfg: ModelConfig):
super().__init__()
self.gate_up = nn.Linear(cfg.d_model, 2 * cfg.mlp_hidden, bias=False)
self.down = nn.Linear(cfg.mlp_hidden, cfg.d_model, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
gate, up = self.gate_up(x).chunk(2, dim=-1)
return self.down(F.silu(gate) * up)
class TransformerBlock(nn.Module):
def __init__(self, cfg: ModelConfig):
super().__init__()
self.attention_norm = RMSNorm(cfg.d_model)
self.attention = Attention(cfg)
self.mlp_norm = RMSNorm(cfg.d_model)
self.mlp = SwiGLU(cfg)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attention(self.attention_norm(x))
return x + self.mlp(self.mlp_norm(x))
class GPT(nn.Module):
"""A dense ~50M model. Tied embeddings keep vocabulary capacity efficient."""
def __init__(self, cfg: ModelConfig):
super().__init__()
self.cfg = cfg
self.token_embedding = nn.Embedding(cfg.vocab_size, cfg.d_model)
self.position_embedding = nn.Embedding(cfg.block_size, cfg.d_model)
self.blocks = nn.ModuleList([TransformerBlock(cfg) for _ in range(cfg.n_layers)])
self.final_norm = RMSNorm(cfg.d_model)
self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
self.lm_head.weight = self.token_embedding.weight
self.gradient_checkpointing = False
self.apply(self._init_weights)
@staticmethod
def _init_weights(module: nn.Module) -> None:
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
def forward(self, input_ids: torch.Tensor, targets: torch.Tensor | None = None):
_, length = input_ids.shape
if length > self.cfg.block_size:
raise ValueError(f"Sequence length {length} exceeds {self.cfg.block_size}.")
positions = torch.arange(length, device=input_ids.device)
x = self.token_embedding(input_ids) + self.position_embedding(positions)
checkpointing = self.training and self.gradient_checkpointing
for block in self.blocks:
x = torch.utils.checkpoint.checkpoint(block, x, use_reentrant=False) if checkpointing else block(x)
x = self.final_norm(x)
if targets is None:
return self.lm_head(x[:, [-1]]), None
# Do not allocate [batch, sequence, 32k] logits at once: this is VRAM-safe on 6GB.
hidden = x[:, :-1].contiguous().view(-1, self.cfg.d_model)
labels = targets[:, 1:].contiguous().view(-1)
loss_sum = hidden.new_zeros((), dtype=torch.float32)
for start in range(0, labels.numel(), 256):
loss_sum = loss_sum + F.cross_entropy(
self.lm_head(hidden[start : start + 256]).float(), labels[start : start + 256], reduction="sum"
)
return None, loss_sum / labels.numel()
@torch.inference_mode()
def generate(
self,
input_ids: torch.Tensor,
max_new_tokens: int,
temperature: float = 0.8,
top_k: int = 50,
eos_id: int | None = None,
):
for _ in range(max_new_tokens):
logits, _ = self(input_ids[:, -self.cfg.block_size :])
logits = logits[:, -1] / max(temperature, 1e-5)
if top_k > 0:
threshold = torch.topk(logits, min(top_k, logits.size(-1))).values[:, [-1]]
logits = logits.masked_fill(logits < threshold, float("-inf"))
next_token = torch.multinomial(F.softmax(logits, dim=-1), 1)
input_ids = torch.cat((input_ids, next_token), dim=1)
if eos_id is not None and (next_token == eos_id).all():
break
return input_ids
def parameter_count(model: nn.Module) -> int:
return sum(parameter.numel() for parameter in model.parameters())
|