Files
2026-06-06 17:06:08 -05:00

266 lines
10 KiB
Python

"""A compact GPT-style dense decoder-only Transformer (~1B params with defaults).
Default shape is intended for a custom 128k vocabulary tokenizer:
- vocab_size=128000
- n_layer=16, n_embd=2048, n_head=16, n_kv_head=8, intermediate_size=5632
- tied input/output embeddings
This lands at about 1.02B trainable parameters.
"""
from __future__ import annotations
import json
import math
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
@dataclass
class GPTConfig:
vocab_size: int = 128_000
block_size: int = 4096
n_layer: int = 16
n_head: int = 16
n_kv_head: int = 8
n_embd: int = 2048
intermediate_size: int = 5632
dropout: float = 0.0
bias: bool = False
rope_theta: float = 10_000.0
norm_eps: float = 1e-5
tie_embeddings: bool = True
@classmethod
def from_json(cls, path: str | Path) -> "GPTConfig":
with open(path, "r", encoding="utf-8") as f:
return cls(**json.load(f))
def to_json(self, path: str | Path) -> None:
with open(path, "w", encoding="utf-8") as f:
json.dump(asdict(self), f, indent=2)
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: torch.Tensor) -> torch.Tensor:
dtype = x.dtype
x = x.float()
x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
return (self.weight * x).to(dtype)
def rotate_half(x: torch.Tensor) -> torch.Tensor:
x1, x2 = x[..., : x.shape[-1] // 2], x[..., x.shape[-1] // 2 :]
return torch.cat((-x2, x1), dim=-1)
def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
# x: [B, H, T, D], cos/sin: [1, 1, T, D]
return (x * cos) + (rotate_half(x) * sin)
class CausalSelfAttention(nn.Module):
def __init__(self, config: GPTConfig):
super().__init__()
assert config.n_embd % config.n_head == 0
assert config.n_head % config.n_kv_head == 0
self.n_head = config.n_head
self.n_kv_head = config.n_kv_head
self.head_dim = config.n_embd // config.n_head
self.kv_repeat = config.n_head // config.n_kv_head
self.dropout = config.dropout
self.q_proj = nn.Linear(config.n_embd, config.n_head * self.head_dim, bias=config.bias)
self.k_proj = nn.Linear(config.n_embd, config.n_kv_head * self.head_dim, bias=config.bias)
self.v_proj = nn.Linear(config.n_embd, config.n_kv_head * self.head_dim, bias=config.bias)
self.o_proj = nn.Linear(config.n_embd, config.n_embd, bias=config.bias)
self.resid_dropout = nn.Dropout(config.dropout)
inv_freq = 1.0 / (
config.rope_theta ** (torch.arange(0, self.head_dim, 2).float() / self.head_dim)
)
self.register_buffer("inv_freq", inv_freq, persistent=False)
def _rope_cache(self, seq_len: int, device: torch.device, dtype: torch.dtype):
t = torch.arange(seq_len, device=device, dtype=self.inv_freq.dtype)
freqs = torch.outer(t, self.inv_freq.to(device))
emb = torch.cat((freqs, freqs), dim=-1)
cos = emb.cos()[None, None, :, :].to(dtype)
sin = emb.sin()[None, None, :, :].to(dtype)
return cos, sin
def forward(self, x: torch.Tensor) -> torch.Tensor:
bsz, seq_len, embd = x.size()
q = self.q_proj(x).view(bsz, seq_len, self.n_head, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(bsz, seq_len, self.n_kv_head, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(bsz, seq_len, self.n_kv_head, self.head_dim).transpose(1, 2)
cos, sin = self._rope_cache(seq_len, x.device, q.dtype)
q = apply_rope(q, cos, sin)
k = apply_rope(k, cos, sin)
if self.kv_repeat != 1:
k = k.repeat_interleave(self.kv_repeat, dim=1)
v = v.repeat_interleave(self.kv_repeat, dim=1)
y = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=None,
dropout_p=self.dropout if self.training else 0.0,
is_causal=True,
)
y = y.transpose(1, 2).contiguous().view(bsz, seq_len, embd)
return self.resid_dropout(self.o_proj(y))
class SwiGLU(nn.Module):
def __init__(self, config: GPTConfig):
super().__init__()
self.gate_proj = nn.Linear(config.n_embd, config.intermediate_size, bias=config.bias)
self.up_proj = nn.Linear(config.n_embd, config.intermediate_size, bias=config.bias)
self.down_proj = nn.Linear(config.intermediate_size, config.n_embd, bias=config.bias)
self.dropout = nn.Dropout(config.dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = F.silu(self.gate_proj(x)) * self.up_proj(x)
return self.dropout(self.down_proj(x))
class Block(nn.Module):
def __init__(self, config: GPTConfig):
super().__init__()
self.ln_1 = RMSNorm(config.n_embd, eps=config.norm_eps)
self.attn = CausalSelfAttention(config)
self.ln_2 = RMSNorm(config.n_embd, eps=config.norm_eps)
self.mlp = SwiGLU(config)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attn(self.ln_1(x))
x = x + self.mlp(self.ln_2(x))
return x
class GPT(nn.Module):
def __init__(self, config: GPTConfig):
super().__init__()
self.config = config
self.transformer = nn.ModuleDict(
dict(
wte=nn.Embedding(config.vocab_size, config.n_embd),
drop=nn.Dropout(config.dropout),
h=nn.ModuleList([Block(config) for _ in range(config.n_layer)]),
ln_f=RMSNorm(config.n_embd, eps=config.norm_eps),
)
)
self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
if config.tie_embeddings:
self.lm_head.weight = self.transformer.wte.weight
self.apply(self._init_weights)
# Slightly scale residual projections as in GPT-2 for stability.
for name, param in self.named_parameters():
if name.endswith("o_proj.weight") or name.endswith("down_proj.weight"):
nn.init.normal_(param, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer))
def _init_weights(self, module: nn.Module) -> None:
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
def forward(
self, idx: torch.Tensor, targets: Optional[torch.Tensor] = None
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
bsz, seq_len = idx.shape
if seq_len > self.config.block_size:
raise ValueError(f"sequence length {seq_len} exceeds block_size {self.config.block_size}")
x = self.transformer.wte(idx)
x = self.transformer.drop(x)
for block in self.transformer.h:
x = block(x)
x = self.transformer.ln_f(x)
if targets is not None:
logits = self.lm_head(x)
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=-1)
else:
logits = self.lm_head(x[:, [-1], :])
loss = None
return logits, loss
@torch.no_grad()
def generate(
self,
idx: torch.Tensor,
max_new_tokens: int,
temperature: float = 1.0,
top_k: Optional[int] = None,
top_p: Optional[float] = None,
eos_token_id: Optional[int] = None,
) -> torch.Tensor:
for _ in range(max_new_tokens):
idx_cond = idx[:, -self.config.block_size :]
logits, _ = self(idx_cond)
logits = logits[:, -1, :]
if temperature <= 0:
next_id = torch.argmax(logits, dim=-1, keepdim=True)
else:
logits = logits / temperature
if top_k is not None and top_k > 0:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = -float("inf")
if top_p is not None and 0 < top_p < 1:
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
probs = torch.softmax(sorted_logits, dim=-1)
cumprobs = torch.cumsum(probs, dim=-1)
mask = cumprobs > top_p
mask[..., 1:] = mask[..., :-1].clone()
mask[..., 0] = False
sorted_logits[mask] = -float("inf")
logits = torch.full_like(logits, -float("inf"))
logits.scatter_(dim=-1, index=sorted_indices, src=sorted_logits)
probs = torch.softmax(logits, dim=-1)
next_id = torch.multinomial(probs, num_samples=1)
idx = torch.cat((idx, next_id), dim=1)
if eos_token_id is not None and torch.all(next_id.squeeze(-1) == eos_token_id):
break
return idx
def configure_optimizers(self, weight_decay: float, learning_rate: float, betas: tuple[float, float]):
decay_params = []
nodecay_params = []
for name, p in self.named_parameters():
if not p.requires_grad:
continue
if p.dim() >= 2:
decay_params.append(p)
else:
nodecay_params.append(p)
optim_groups = [
{"params": decay_params, "weight_decay": weight_decay},
{"params": nodecay_params, "weight_decay": 0.0},
]
return torch.optim.AdamW(optim_groups, lr=learning_rate, betas=betas, fused=torch.cuda.is_available())
def num_parameters(self, non_embedding: bool = False) -> int:
n = sum(p.numel() for p in self.parameters())
if non_embedding:
n -= self.transformer.wte.weight.numel()
return n