trtd56's picture
Release phase2 staging
58e0fa2 verified
Raw
History Blame Contribute Delete
6.03 kB
from __future__ import annotations
import math
from dataclasses import asdict, dataclass
import torch
import torch.nn as nn
from torch.nn import functional as F
@dataclass
class GPTConfig:
block_size: int = 1024
vocab_size: int = 8192
n_layer: int = 12
n_head: int = 12
n_embd: int = 768
dropout: float = 0.0
bias: bool = False
class CausalSelfAttention(nn.Module):
def __init__(self, config: GPTConfig) -> None:
super().__init__()
if config.n_embd % config.n_head:
raise ValueError("embedding dimension must be divisible by number of heads")
self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd, bias=config.bias)
self.c_proj = nn.Linear(config.n_embd, config.n_embd, bias=config.bias)
self.attn_dropout = nn.Dropout(config.dropout)
self.resid_dropout = nn.Dropout(config.dropout)
self.n_head = config.n_head
self.n_embd = config.n_embd
self.dropout = config.dropout
def forward(self, value: torch.Tensor) -> torch.Tensor:
batch, time, channels = value.size()
query, key, val = self.c_attn(value).split(self.n_embd, dim=2)
head_size = channels // self.n_head
query = query.view(batch, time, self.n_head, head_size).transpose(1, 2)
key = key.view(batch, time, self.n_head, head_size).transpose(1, 2)
val = val.view(batch, time, self.n_head, head_size).transpose(1, 2)
attended = F.scaled_dot_product_attention(
query,
key,
val,
attn_mask=None,
dropout_p=self.dropout if self.training else 0,
is_causal=True,
)
attended = attended.transpose(1, 2).contiguous().view(batch, time, channels)
return self.resid_dropout(self.c_proj(attended))
class MLP(nn.Module):
def __init__(self, config: GPTConfig) -> None:
super().__init__()
self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd, bias=config.bias)
self.gelu = nn.GELU()
self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd, bias=config.bias)
self.dropout = nn.Dropout(config.dropout)
def forward(self, value: torch.Tensor) -> torch.Tensor:
return self.dropout(self.c_proj(self.gelu(self.c_fc(value))))
class Block(nn.Module):
def __init__(self, config: GPTConfig) -> None:
super().__init__()
self.ln_1 = nn.LayerNorm(config.n_embd, bias=config.bias)
self.attn = CausalSelfAttention(config)
self.ln_2 = nn.LayerNorm(config.n_embd, bias=config.bias)
self.mlp = MLP(config)
def forward(self, value: torch.Tensor) -> torch.Tensor:
value = value + self.attn(self.ln_1(value))
return value + self.mlp(self.ln_2(value))
class GPT(nn.Module):
def __init__(self, config: GPTConfig) -> None:
super().__init__()
self.config = config
self.transformer = nn.ModuleDict(
{
"wte": nn.Embedding(config.vocab_size, config.n_embd),
"wpe": nn.Embedding(config.block_size, config.n_embd),
"drop": nn.Dropout(config.dropout),
"h": nn.ModuleList([Block(config) for _ in range(config.n_layer)]),
"ln_f": nn.LayerNorm(config.n_embd, bias=config.bias),
}
)
self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
self.transformer.wte.weight = self.lm_head.weight
self.apply(self._init_weights)
for name, parameter in self.named_parameters():
if name.endswith("c_proj.weight"):
torch.nn.init.normal_(
parameter, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layer)
)
@staticmethod
def _init_weights(module: nn.Module) -> None:
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, index: torch.Tensor, targets: torch.Tensor | None = None
) -> tuple[torch.Tensor, torch.Tensor | None]:
_, time = index.shape
if time > self.config.block_size:
raise ValueError("sequence exceeds model block size")
positions = torch.arange(0, time, dtype=torch.long, device=index.device)
value = self.transformer.drop(self.transformer.wte(index) + self.transformer.wpe(positions))
for block in self.transformer.h:
value = block(value)
value = self.transformer.ln_f(value)
logits = self.lm_head(value)
loss = (
F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=-1)
if targets is not None
else None
)
return logits, loss
@torch.no_grad()
def generate(
self,
index: torch.Tensor,
max_new_tokens: int,
temperature: float = 0.8,
top_k: int | None = 200,
) -> torch.Tensor:
for _ in range(max_new_tokens):
cropped = index[:, -self.config.block_size :]
logits, _ = self(cropped)
logits = logits[:, -1, :] / temperature
if top_k is not None:
values, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < values[:, [-1]]] = -float("Inf")
probabilities = F.softmax(logits, dim=-1)
index = torch.cat((index, torch.multinomial(probabilities, num_samples=1)), dim=1)
return index
def parameter_count(self, non_embedding: bool = False) -> int:
count = sum(parameter.numel() for parameter in self.parameters())
if non_embedding:
count -= self.transformer.wpe.weight.numel()
return count
def config_dict(self) -> dict[str, object]:
return asdict(self.config)