Inital commit
This commit is contained in:
@@ -0,0 +1,265 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user