""" Full definition of a GPT Language Model, all of it in this single file. Includes: - cheap Q gate - Q-MLP gate (128 hidden) - Q-MLP gate (192 hidden) -- param-matched control for Q+A MLP - full-width Q gate - Q head-shared elementwise gate - full-width X gate - bottleneck X gate - paper-faithful X variants: - head-specific elementwise G1 - head-specific headwise G1 - head-shared elementwise G1 - Q+A dual-signal variants: - cheap Q+A gate (shared Linear 128->64 applied per head/token) - Q+A MLP gate (shared MLP 128->128->64 applied per head/token) - legacy Q+A head-shared elementwise gate (shared Linear 128->64; kept for compatibility) - bilinear diagonal Q+A gate (learned elementwise q*y interaction) - normed Q+A gate (RMSNorm on q and y before shared 128->64 mix) - lowrank Q+A gate (shared 128->16->64 bottleneck mix) - soft-QA gate (independent learned soft blends for q and y) - per-dimension soft-QA gate (independent 64-dim learned blend scales for q and y) """ import math import inspect from dataclasses import dataclass import torch import torch.nn as nn from torch.nn import functional as F class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x): rms = torch.sqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) return x / rms * self.weight class LayerNorm(nn.Module): def __init__(self, ndim, bias): super().__init__() self.weight = nn.Parameter(torch.ones(ndim)) self.bias = nn.Parameter(torch.zeros(ndim)) if bias else None def forward(self, input): return F.layer_norm(input, self.weight.shape, self.weight, self.bias, 1e-5) class CausalSelfAttention(nn.Module): def __init__(self, config): super().__init__() assert config.n_embd % config.n_head == 0 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.qv_variant = config.qv_variant hs = config.n_embd // config.n_head # -------------------------------------------------- # Q-conditioned gates # -------------------------------------------------- if self.qv_variant in ['dynamic', 'dynamic_swiglu']: self.qv_gate_proj = nn.Linear(hs, hs, bias=False) if self.qv_variant == 'dynamic_qconditioned_mlp128': self.q_gate_fc1 = nn.Linear(hs, 128, bias=False) self.q_gate_fc2 = nn.Linear(128, hs, bias=False) if self.qv_variant == 'dynamic_qconditioned_mlp192': self.q_gate_fc1_192 = nn.Linear(hs, 192, bias=False) self.q_gate_fc2_192 = nn.Linear(192, hs, bias=False) if self.qv_variant == 'dynamic_qconditioned_fullwidth': self.q_gate_proj_full = nn.Linear(config.n_embd, config.n_embd, bias=False) if self.qv_variant == 'dynamic_qconditioned_fullwidth_headspecific': # full-width Q source (768) with separate 768->64 map for each head self.q_gate_proj_full_headspecific = nn.Parameter(torch.empty(config.n_head, config.n_embd, hs)) nn.init.normal_(self.q_gate_proj_full_headspecific, mean=0.0, std=0.02) if self.qv_variant == 'dynamic_q_headshared_elementwise': self.q_gate_headshared = nn.Linear(hs, hs, bias=False) # -------------------------------------------------- # X-conditioned gates # -------------------------------------------------- if self.qv_variant == 'dynamic_xconditioned_g1': self.x_gate_proj = nn.Linear(config.n_embd, config.n_embd, bias=False) if self.qv_variant == 'dynamic_xconditioned_fullwidth_headspecific': # full-width X source (768) with separate 768->64 map for each head self.x_gate_proj_full_headspecific = nn.Parameter(torch.empty(config.n_head, config.n_embd, hs)) nn.init.normal_(self.x_gate_proj_full_headspecific, mean=0.0, std=0.02) if self.qv_variant == 'dynamic_xconditioned_bottleneck': self.x_gate_bottleneck = nn.Linear(config.n_embd, hs, bias=False) if self.qv_variant == 'dynamic_x_g1_headspecific_elementwise': self.x_gate_headspecific_elementwise = nn.Parameter(torch.empty(config.n_head, hs, hs)) nn.init.normal_(self.x_gate_headspecific_elementwise, mean=0.0, std=0.02) if self.qv_variant == 'dynamic_x_g1_headspecific_headwise': self.x_gate_headspecific_headwise = nn.Parameter(torch.empty(config.n_head, hs, 1)) nn.init.normal_(self.x_gate_headspecific_headwise, mean=0.0, std=0.02) if self.qv_variant == 'dynamic_x_g1_headshared_elementwise': self.x_gate_headshared_elementwise = nn.Linear(hs, hs, bias=False) # -------------------------------------------------- # Q+A dual-signal gates # -------------------------------------------------- # random gates have no learned parameters — nothing to init here # dot product gates also have no learned parameters — nothing to init here if self.qv_variant == 'dynamic_a_conditioned': # A-only: per-head Linear(64->64), conditioned on y only self.a_gate_proj = nn.Linear(hs, hs, bias=False) if self.qv_variant == 'dynamic_qa_conditioned': # cheap Q+A: shared Linear(128->64) applied to each head/token self.qa_gate_proj = nn.Linear(hs * 2, hs, bias=False) if self.qv_variant == 'dynamic_qa_conditioned_headspecific': # true head-specific Q+A: separate 128->64 matrix for each head self.qa_gate_proj_headspecific = nn.Parameter(torch.empty(config.n_head, hs * 2, hs)) nn.init.normal_(self.qa_gate_proj_headspecific, mean=0.0, std=0.02) if self.qv_variant == 'dynamic_qa_conditioned_mlp128': # main Q+A MLP: per-head MLP(128->128->64) self.qa_gate_fc1 = nn.Linear(hs * 2, 128, bias=False) self.qa_gate_fc2 = nn.Linear(128, hs, bias=False) if self.qv_variant == 'dynamic_qa_headshared_elementwise': # legacy compatibility variant: shared Linear(128->64) applied per head/token self.qa_gate_headshared = nn.Linear(hs * 2, hs, bias=False) if self.qv_variant == 'dynamic_qa_bilinear_diag': # learned elementwise interaction on q*y; 64 params per layer self.qa_bilinear_diag = nn.Parameter(torch.ones(hs)) if self.qv_variant == 'dynamic_qa_conditioned_normed': # RMSNorm q and y separately before shared Linear(128->64) self.qa_q_norm = RMSNorm(hs) self.qa_y_norm = RMSNorm(hs) self.qa_gate_proj_normed = nn.Linear(hs * 2, hs, bias=False) if self.qv_variant == 'dynamic_qa_conditioned_normed_yonly': # RMSNorm y only, raw q, before shared Linear(128->64) self.qa_y_norm_yonly = RMSNorm(hs) self.qa_gate_proj_normed_yonly = nn.Linear(hs * 2, hs, bias=False) if self.qv_variant == 'dynamic_qa_conditioned_normed_qonly': # RMSNorm q only, raw y, before shared Linear(128->64) self.qa_q_norm_qonly = RMSNorm(hs) self.qa_gate_proj_normed_qonly = nn.Linear(hs * 2, hs, bias=False) if self.qv_variant == 'dynamic_qa_conditioned_softq': # Learnable blend between RMSNorm(q) and raw q, with raw y self.qa_q_norm_soft = RMSNorm(hs) self.qa_gate_proj_softq = nn.Linear(hs * 2, hs, bias=False) self.qa_q_blend_alpha = nn.Parameter(torch.tensor(0.5)) if self.qv_variant == 'dynamic_qa_conditioned_softqa': # Independent learnable blends for RMSNorm(q)/raw q and RMSNorm(y)/raw y self.qa_q_norm_softqa = RMSNorm(hs) self.qa_y_norm_softqa = RMSNorm(hs) self.qa_gate_proj_softqa = nn.Linear(hs * 2, hs, bias=False) self.qa_q_blend_alpha = nn.Parameter(torch.tensor(0.5)) self.qa_y_blend_alpha = nn.Parameter(torch.tensor(0.5)) if self.qv_variant == 'dynamic_qa_conditioned_softqa_perlayer': # Per-layer independent learnable blends for RMSNorm(q)/raw q and RMSNorm(y)/raw y self.qa_q_norm_softqa_pl = RMSNorm(hs) self.qa_y_norm_softqa_pl = RMSNorm(hs) self.qa_gate_proj_softqa_pl = nn.Linear(hs * 2, hs, bias=False) self.qa_q_blend_alpha = nn.Parameter(torch.tensor(0.0)) self.qa_y_blend_alpha = nn.Parameter(torch.tensor(0.0)) if self.qv_variant == 'dynamic_qa_conditioned_softqa_perdim': # Per-dimension independent learnable blends for RMSNorm(q)/raw q and RMSNorm(y)/raw y self.qa_q_norm_softqa_pd = RMSNorm(hs) self.qa_y_norm_softqa_pd = RMSNorm(hs) self.qa_gate_proj_softqa_pd = nn.Linear(hs * 2, hs, bias=False) self.qa_q_dim_scale = nn.Parameter(torch.zeros(hs)) self.qa_y_dim_scale = nn.Parameter(torch.zeros(hs)) if self.qv_variant == 'dynamic_qa_conditioned_softqa_perdim_informed': # Per-dimension independent learnable blends with informed initialization: # 1.0459 = sigmoid^-1(0.74), matching the Phase 17 converged scalar value. self.qa_q_norm_softqa_pd = RMSNorm(hs) self.qa_y_norm_softqa_pd = RMSNorm(hs) self.qa_gate_proj_softqa_pd = nn.Linear(hs * 2, hs, bias=False) self.qa_q_dim_scale = nn.Parameter(torch.full((hs,), 1.0459)) self.qa_y_dim_scale = nn.Parameter(torch.full((hs,), 1.0459)) if self.qv_variant == 'dynamic_qa_conditioned_softqa_perdim_free': # Per-dimension independent learnable blends with tiny random symmetry breaking: # start near sigmoid(scale)=0.5, but each dimension differs slightly. self.qa_q_norm_softqa_pd = RMSNorm(hs) self.qa_y_norm_softqa_pd = RMSNorm(hs) self.qa_gate_proj_softqa_pd = nn.Linear(hs * 2, hs, bias=False) self.qa_q_dim_scale = nn.Parameter(torch.empty(hs)) self.qa_y_dim_scale = nn.Parameter(torch.empty(hs)) nn.init.normal_(self.qa_q_dim_scale, mean=0.0, std=0.01) nn.init.normal_(self.qa_y_dim_scale, mean=0.0, std=0.01) if self.qv_variant == 'dynamic_qa_conditioned_lowrank16': # shared low-rank Q+A mixer: 128->16->64 per head/token self.qa_gate_lowrank_fc1 = nn.Linear(hs * 2, 16, bias=False) self.qa_gate_lowrank_fc2 = nn.Linear(16, hs, bias=False) # -------------------------------------------------- # Static / normalization controls # -------------------------------------------------- if self.qv_variant == 'post_rmsnorm_y': self.post_attn_norm = RMSNorm(hs) elif self.qv_variant == 'static_gate': self.static_gate_param = nn.Parameter(torch.zeros(config.n_embd)) elif self.qv_variant == 'static_gate_prehead': self.static_gate_prehead_param = nn.Parameter(torch.zeros(config.n_head, hs)) self.q_norm = None self.v_norm = None self.qv_modulator = None variant = getattr(config, 'qv_variant', 'none') if variant in ['qvnorm']: self.q_norm = RMSNorm(hs) if variant in ['vnorm', 'qvnorm']: self.v_norm = RMSNorm(hs) 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 self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention') if not self.flash: self.register_buffer( "bias", torch.tril(torch.ones(config.block_size, config.block_size)).view( 1, 1, config.block_size, config.block_size ), ) def forward(self, x, x_prenorm=None): B, T, C = x.size() hs = C // self.n_head q, k, v = self.c_attn(x).split(self.n_embd, dim=2) k = k.view(B, T, self.n_head, hs).transpose(1, 2) q = q.view(B, T, self.n_head, hs).transpose(1, 2) v = v.view(B, T, self.n_head, hs).transpose(1, 2) if self.q_norm is not None: q = self.q_norm(q) if self.v_norm is not None: v = self.v_norm(v) if self.flash: y = torch.nn.functional.scaled_dot_product_attention( q, k, v, is_causal=True, dropout_p=self.dropout if self.training else 0 ) else: att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1))) att = att.masked_fill(self.bias[:, :, :T, :T] == 0, float('-inf')) att = F.softmax(att, dim=-1) att = self.attn_dropout(att) y = att @ v # -------------------------------------------------- # Gate application — y shape: (B, n_head, T, hs) # -------------------------------------------------- # Q-conditioned gates if self.qv_variant in ['dynamic', 'dynamic_swiglu']: gate_logit = self.qv_gate_proj(q) gate = torch.sigmoid(gate_logit) if self.qv_variant == 'dynamic' else gate_logit * torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qconditioned_mlp128': gate_hidden = F.silu(self.q_gate_fc1(q)) gate_logit = self.q_gate_fc2(gate_hidden) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qconditioned_mlp192': gate_hidden = F.silu(self.q_gate_fc1_192(q)) gate_logit = self.q_gate_fc2_192(gate_hidden) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qconditioned_fullwidth': q_full = q.transpose(1, 2).contiguous().view(B, T, C) gate_logit = self.q_gate_proj_full(q_full) gate = torch.sigmoid(gate_logit) gate = gate.view(B, T, self.n_head, hs).transpose(1, 2) y = y * gate elif self.qv_variant == 'dynamic_qconditioned_fullwidth_headspecific': q_full = q.transpose(1, 2).contiguous().view(B, T, C) gate_logit = torch.einsum('btc,hce->bhte', q_full, self.q_gate_proj_full_headspecific) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_q_headshared_elementwise': gate_logit = self.q_gate_headshared(q) gate = torch.sigmoid(gate_logit) y = y * gate # Static / normalization controls elif self.qv_variant == 'post_rmsnorm_y': y = self.post_attn_norm(y) elif self.qv_variant == 'static_gate': gate = torch.sigmoid(self.static_gate_param) y = y.transpose(1, 2).contiguous().view(B, T, C) y = y * gate[None, None, :] elif self.qv_variant == 'static_gate_prehead': gate = torch.sigmoid(self.static_gate_prehead_param) y = y * gate[None, :, None, :] # X-conditioned gates elif self.qv_variant == 'dynamic_xconditioned_g1': gate_logit = self.x_gate_proj(x_prenorm) gate = torch.sigmoid(gate_logit) gate = gate.view(B, T, self.n_head, hs).transpose(1, 2) y = y * gate elif self.qv_variant == 'dynamic_xconditioned_fullwidth_headspecific': gate_logit = torch.einsum('btc,hce->bhte', x_prenorm, self.x_gate_proj_full_headspecific) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_xconditioned_bottleneck': z = self.x_gate_bottleneck(x_prenorm) gate = torch.sigmoid(z) gate = gate.unsqueeze(1) y = y * gate elif self.qv_variant == 'dynamic_x_g1_headspecific_elementwise': xh = x_prenorm.view(B, T, self.n_head, hs).transpose(1, 2) gate_logit = torch.einsum('bhtd,hde->bhte', xh, self.x_gate_headspecific_elementwise) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_x_g1_headspecific_headwise': xh = x_prenorm.view(B, T, self.n_head, hs).transpose(1, 2) gate_logit = torch.einsum('bhtd,hde->bhte', xh, self.x_gate_headspecific_headwise) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_x_g1_headshared_elementwise': xh = x_prenorm.view(B, T, self.n_head, hs).transpose(1, 2) gate_logit = self.x_gate_headshared_elementwise(xh) gate = torch.sigmoid(gate_logit) y = y * gate # Random / ablation gates — no learned params elif self.qv_variant == 'dynamic_random_gate': # Uniform random [0,1] gate — tests if location matters at all gate = torch.rand_like(y) y = y * gate elif self.qv_variant == 'dynamic_ones_gate': # Always 1.0 — pure identity, sanity check (should == baseline) pass # y unchanged elif self.qv_variant == 'dynamic_random_normal': # Normal(0.5, 0.2) gate clamped to [0,1] — different noise shape gate = torch.randn_like(y) * 0.2 + 0.5 gate = gate.clamp(0.0, 1.0) y = y * gate elif self.qv_variant == 'dynamic_bernoulli_gate': # Binary gate: randomly zeros 50% of head dims — spicy dropout variant gate = torch.bernoulli(torch.full_like(y, 0.5)) # Scale by 2.0 to preserve expected value (like standard dropout) y = y * gate * 2.0 # A-only gate elif self.qv_variant == 'dynamic_a_conditioned': gate_logit = self.a_gate_proj(y) gate = torch.sigmoid(gate_logit) y = y * gate # Q+A dual-signal gates elif self.qv_variant == 'dynamic_qa_conditioned': # cheap Q+A: cat(q, y) -> shared Linear(128->64) -> sigmoid gate_in = torch.cat([q, y], dim=-1) # (B, n_head, T, 128) gate_logit = self.qa_gate_proj(gate_in) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_conditioned_headspecific': # true head-specific Q+A: per-head (128->64) -> sigmoid gate_in = torch.cat([q, y], dim=-1) # (B, n_head, T, 128) gate_logit = torch.einsum('bhtd,hde->bhte', gate_in, self.qa_gate_proj_headspecific) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_conditioned_mlp128': # main Q+A MLP: cat(q, y) -> Linear(128->128) -> SiLU -> Linear(128->64) -> sigmoid gate_in = torch.cat([q, y], dim=-1) # (B, n_head, T, 128) gate_hidden = F.silu(self.qa_gate_fc1(gate_in)) gate_logit = self.qa_gate_fc2(gate_hidden) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_headshared_elementwise': # legacy compatibility variant: cat(q, y) -> shared Linear(128->64) -> sigmoid gate_in = torch.cat([q, y], dim=-1) # (B, n_head, T, 128) gate_logit = self.qa_gate_headshared(gate_in) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_bilinear_diag': # learned diagonal bilinear interaction: sigmoid((q * y * w) / sqrt(d_head)) gate_logit = (q * y * self.qa_bilinear_diag[None, None, None, :]) / math.sqrt(hs) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_conditioned_normed': qn = self.qa_q_norm(q) yn = self.qa_y_norm(y) gate_in = torch.cat([qn, yn], dim=-1) gate_logit = self.qa_gate_proj_normed(gate_in) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_conditioned_normed_yonly': yn = self.qa_y_norm_yonly(y) gate_in = torch.cat([q, yn], dim=-1) gate_logit = self.qa_gate_proj_normed_yonly(gate_in) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_conditioned_normed_qonly': qn = self.qa_q_norm_qonly(q) gate_in = torch.cat([qn, y], dim=-1) gate_logit = self.qa_gate_proj_normed_qonly(gate_in) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_conditioned_softq': qn = self.qa_q_norm_soft(q) alpha = torch.sigmoid(self.qa_q_blend_alpha) q_blend = alpha * qn + (1.0 - alpha) * q gate_in = torch.cat([q_blend, y], dim=-1) gate_logit = self.qa_gate_proj_softq(gate_in) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_conditioned_softqa': qn = self.qa_q_norm_softqa(q) yn = self.qa_y_norm_softqa(y) alpha_q = torch.sigmoid(self.qa_q_blend_alpha) alpha_y = torch.sigmoid(self.qa_y_blend_alpha) q_blend = alpha_q * qn + (1.0 - alpha_q) * q y_blend = alpha_y * yn + (1.0 - alpha_y) * y gate_in = torch.cat([q_blend, y_blend], dim=-1) gate_logit = self.qa_gate_proj_softqa(gate_in) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_conditioned_softqa_perlayer': qn = self.qa_q_norm_softqa_pl(q) yn = self.qa_y_norm_softqa_pl(y) alpha_q = torch.sigmoid(self.qa_q_blend_alpha) alpha_y = torch.sigmoid(self.qa_y_blend_alpha) q_blend = alpha_q * qn + (1.0 - alpha_q) * q y_blend = alpha_y * yn + (1.0 - alpha_y) * y gate_in = torch.cat([q_blend, y_blend], dim=-1) gate_logit = self.qa_gate_proj_softqa_pl(gate_in) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant in ['dynamic_qa_conditioned_softqa_perdim', 'dynamic_qa_conditioned_softqa_perdim_informed', 'dynamic_qa_conditioned_softqa_perdim_free']: qn = self.qa_q_norm_softqa_pd(q) yn = self.qa_y_norm_softqa_pd(y) scale_q = torch.sigmoid(self.qa_q_dim_scale).to(dtype=q.dtype, device=q.device) scale_y = torch.sigmoid(self.qa_y_dim_scale).to(dtype=y.dtype, device=y.device) q_blend = scale_q * qn + (1.0 - scale_q) * q y_blend = scale_y * yn + (1.0 - scale_y) * y gate_in = torch.cat([q_blend, y_blend], dim=-1) gate_logit = self.qa_gate_proj_softqa_pd(gate_in) gate = torch.sigmoid(gate_logit) y = y * gate elif self.qv_variant == 'dynamic_qa_conditioned_lowrank16': gate_in = torch.cat([q, y], dim=-1) gate_hidden = F.silu(self.qa_gate_lowrank_fc1(gate_in)) gate_logit = self.qa_gate_lowrank_fc2(gate_hidden) gate = torch.sigmoid(gate_logit) y = y * gate # Zero-parameter dot product gates — raw Q·y geometry elif self.qv_variant == 'dynamic_dot_scalar': # scalar gate: sigmoid(sum(q*y) / sqrt(d_head)) # one scalar per head per token, zero params hs = q.shape[-1] dot = (q * y).sum(dim=-1, keepdim=True) # (B, n_head, T, 1) gate = torch.sigmoid(dot / math.sqrt(hs)) # scalar broadcast y = y * gate elif self.qv_variant == 'dynamic_dot_elementwise': # elementwise gate: sigmoid(q*y / sqrt(d_head)) # same shape as cheap_qa gate but zero params hs = q.shape[-1] gate = torch.sigmoid((q * y) / math.sqrt(hs)) # (B, n_head, T, 64) y = y * gate if self.qv_variant != 'static_gate': y = y.transpose(1, 2).contiguous().view(B, T, C) y = self.resid_dropout(self.c_proj(y)) return y class MLP(nn.Module): def __init__(self, config): 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, x): x = self.c_fc(x) x = self.gelu(x) x = self.c_proj(x) x = self.dropout(x) return x class Block(nn.Module): def __init__(self, config): super().__init__() self.ln_1 = LayerNorm(config.n_embd, bias=config.bias) self.attn = CausalSelfAttention(config) self.ln_2 = LayerNorm(config.n_embd, bias=config.bias) self.mlp = MLP(config) def forward(self, x): x_norm = self.ln_1(x) attn_out = self.attn(x_norm, x_norm) x = x + attn_out x = x + self.mlp(self.ln_2(x)) return x @dataclass class GPTConfig: block_size: int = 1024 vocab_size: int = 50304 n_layer: int = 12 n_head: int = 12 n_embd: int = 768 dropout: float = 0.0 bias: bool = True qv_variant: str = 'none' class GPT(nn.Module): def __init__(self, config): super().__init__() assert config.vocab_size is not None assert config.block_size is not None self.config = config self.transformer = nn.ModuleDict(dict( 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 = 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 pn, p in self.named_parameters(): if pn.endswith('c_proj.weight'): torch.nn.init.normal_(p, mean=0.0, std=0.02/math.sqrt(2 * config.n_layer)) print("number of parameters: %.2fM" % (self.get_num_params()/1e6,)) def get_num_params(self, non_embedding=True): n_params = sum(p.numel() for p in self.parameters()) if non_embedding: n_params -= self.transformer.wpe.weight.numel() return n_params def _init_weights(self, module): 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, idx, targets=None): device = idx.device b, t = idx.size() assert t <= self.config.block_size pos = torch.arange(0, t, dtype=torch.long, device=device) tok_emb = self.transformer.wte(idx) pos_emb = self.transformer.wpe(pos) x = self.transformer.drop(tok_emb + pos_emb) 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 def crop_block_size(self, block_size): assert block_size <= self.config.block_size self.config.block_size = block_size self.transformer.wpe.weight = nn.Parameter(self.transformer.wpe.weight[:block_size]) for block in self.transformer.h: if hasattr(block.attn, 'bias'): block.attn.bias = block.attn.bias[:, :, :block_size, :block_size] @classmethod def from_pretrained(cls, model_type, override_args=None): assert model_type in {'gpt2', 'gpt2-medium', 'gpt2-large', 'gpt2-xl'} override_args = override_args or {} assert all(k == 'dropout' for k in override_args) from transformers import GPT2LMHeadModel print("loading weights from pretrained gpt: %s" % model_type) config_args = { 'gpt2': dict(n_layer=12, n_head=12, n_embd=768), 'gpt2-medium': dict(n_layer=24, n_head=16, n_embd=1024), 'gpt2-large': dict(n_layer=36, n_head=20, n_embd=1280), 'gpt2-xl': dict(n_layer=48, n_head=25, n_embd=1600), }[model_type] config_args['vocab_size'] = 50257 config_args['block_size'] = 1024 config_args['bias'] = True if 'dropout' in override_args: config_args['dropout'] = override_args['dropout'] config = GPTConfig(**config_args) model = GPT(config) sd = model.state_dict() sd_keys = [k for k in sd.keys() if not k.endswith('.attn.bias')] model_hf = GPT2LMHeadModel.from_pretrained(model_type) sd_hf = model_hf.state_dict() sd_keys_hf = [k for k in sd_hf.keys() if not k.endswith('.attn.masked_bias') and not k.endswith('.attn.bias')] transposed = ['attn.c_attn.weight', 'attn.c_proj.weight', 'mlp.c_fc.weight', 'mlp.c_proj.weight'] assert len(sd_keys_hf) == len(sd_keys) for k in sd_keys_hf: if any(k.endswith(w) for w in transposed): assert sd_hf[k].shape[::-1] == sd[k].shape with torch.no_grad(): sd[k].copy_(sd_hf[k].t()) else: assert sd_hf[k].shape == sd[k].shape with torch.no_grad(): sd[k].copy_(sd_hf[k]) return model def configure_optimizers(self, weight_decay, learning_rate, betas, device_type): param_dict = {pn: p for pn, p in self.named_parameters() if p.requires_grad} decay_params = [p for n, p in param_dict.items() if p.dim() >= 2] nodecay_params = [p for n, p in param_dict.items() if p.dim() < 2] optim_groups = [ {'params': decay_params, 'weight_decay': weight_decay}, {'params': nodecay_params, 'weight_decay': 0.0} ] num_decay_params = sum(p.numel() for p in decay_params) num_nodecay_params = sum(p.numel() for p in nodecay_params) print(f"num decayed parameter tensors: {len(decay_params)}, with {num_decay_params:,} parameters") print(f"num non-decayed parameter tensors: {len(nodecay_params)}, with {num_nodecay_params:,} parameters") fused_available = 'fused' in inspect.signature(torch.optim.AdamW).parameters use_fused = fused_available and device_type == 'cuda' optimizer = torch.optim.AdamW(optim_groups, lr=learning_rate, betas=betas, fused=use_fused) print(f"using fused AdamW: {use_fused}") return optimizer def estimate_mfu(self, fwdbwd_per_iter, dt): N = self.get_num_params() cfg = self.config L, H, Q, T = cfg.n_layer, cfg.n_head, cfg.n_embd // cfg.n_head, cfg.block_size flops_per_token = 6 * N + 12 * L * H * Q * T flops_per_fwdbwd = flops_per_token * T flops_per_iter = flops_per_fwdbwd * fwdbwd_per_iter flops_achieved = flops_per_iter * (1.0 / dt) flops_promised = 312e12 mfu = flops_achieved / flops_promised return mfu @torch.no_grad() def generate(self, idx, max_new_tokens, temperature=1.0, top_k=None): for _ in range(max_new_tokens): idx_cond = idx if idx.size(1) <= self.config.block_size else idx[:, -self.config.block_size:] logits, _ = self(idx_cond) logits = logits[:, -1, :] / temperature if top_k is not None: v, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < v[:, [-1]]] = -float('Inf') probs = F.softmax(logits, dim=-1) idx_next = torch.multinomial(probs, num_samples=1) idx = torch.cat((idx, idx_next), dim=1) return idx