ujwaliyengar/Decoder_Only_Model
0
1# Solving for residual std scaling issue2import os3import math4import time5import torch6import torch.nn as nn7from torch.nn import functional as F8from dataclasses import dataclass9 10 11class CausalSelfAttention(nn.Module):12 def __init__(self, config):13 super().__init__()14 assert config.n_embd % config.n_head == 015 # Key, query, value projections for all heads, but in a batch16 self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd)17 # Output projection18 self.c_proj = nn.Linear(config.n_embd, config.n_embd)19 self.c_proj.NANGPT_SCALE_INIT = 120 self.n_head = config.n_head21 self.n_embd = config.n_embd22 self.register_buffer(23 "bias",24 torch.tril(torch.ones(config.block_size, config.block_size)).view(25 1, 1, config.block_size, config.block_size26 )27 )28 29 def forward(self, x):30 B, T, C = x.size()31 qkv = self.c_attn(x)32 q, k, v = qkv.split(self.n_embd, dim=2)33 k = k.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)34 q = q.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)35 v = v.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)36 37 att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1)))38 att = att.masked_fill(self.bias[:, :, :T, :T] == 0, float('-inf'))39 att = F.softmax(att, dim=-1)40 y = att @ v41 42 y = y.transpose(1, 2).contiguous().view(B, T, C)43 y = self.c_proj(y)44 return y45 46 47class MLP(nn.Module):48 def __init__(self, config):49 super().__init__()50 self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd)51 self.gelu = nn.GELU(approximate='tanh')52 self.dropout = nn.Dropout(0.1)53 self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd)54 self.c_proj.NANOGPT_SCALE_INIT = 155 56 def forward(self, x):57 x = self.c_fc(x)58 x = self.gelu(x)59 x = self.dropout(x)60 x = self.c_proj(x)61 return x62 63 64class Block(nn.Module):65 def __init__(self, config):66 super().__init__()67 self.ln_1 = nn.LayerNorm(config.n_embd)68 self.attn = CausalSelfAttention(config)69 self.ln_2 = nn.LayerNorm(config.n_embd)70 self.mlp = MLP(config)71 72 def forward(self, x):73 x = x + self.attn(self.ln_1(x))74 x = x + self.mlp(self.ln_2(x))75 return x76 77 78@dataclass79class GPTConfig:80 block_size: int = 102481 vocab_size: int = 5025782 n_layer: int = 2483 n_head: int = 1684 n_embd: int = 102485 86 87class GPT(nn.Module):88 def __init__(self, config):89 super().__init__()90 self.config = config91 self.transformer = nn.ModuleDict(92 dict(93 wte=nn.Embedding(config.vocab_size, config.n_embd),94 wpe=nn.Embedding(config.block_size, config.n_embd),95 h=nn.ModuleList([Block(config) for _ in range(config.n_layer)]),96 ln_f=nn.LayerNorm(config.n_embd),97 )98 )99 self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)100 self.transformer.wte.weight = self.lm_head.weight101 self.apply(self._init_weights)102 103 def _init_weights(self, module):104 if isinstance(module, nn.Linear):105 std = 0.02106 if hasattr(module, 'NANGPT_SCALE_INIT'):107 std *= (2 * self.config.n_layer) ** -0.5108 torch.nn.init.normal_(module.weight, mean=0.0, std=std)109 if module.bias is not None:110 torch.nn.init.zeros_(module.bias)111 elif isinstance(module, nn.Embedding):112 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)113 114 def forward(self, idx, targets=None):115 B, T = idx.size()116 assert T <= self.config.block_size, f"Cannot forward sequence of length {T}, block size is only {self.config.block_size}"117 pos = torch.arange(0, T, dtype=torch.long, device=idx.device)118 pos_emb = self.transformer.wpe(pos)119 tok_emb = self.transformer.wte(idx)120 x = tok_emb + pos_emb121 for block in self.transformer.h:122 x = block(x)123 x = self.transformer.ln_f(x)124 logits = self.lm_head(x)125 loss = None126 if targets is not None:127 loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))128 return logits, loss129 130 131# Define device132device = 'cpu'133if torch.cuda.is_available():134 device = 'cuda'135elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():136 device = "mps"137print(f"Using device: {device}")138 139# Seed for reproducibility140torch.manual_seed(1337)141if torch.cuda.is_available():142 torch.cuda.manual_seed(1337)143 144# Tokenizer setup145import tiktoken146enc = tiktoken.get_encoding('gpt2')147 148# DataLoaderLite149class DataLoaderLite:150 def __init__(self, B, T):151 self.B = B152 self.T = T153 with open('input.txt', 'r') as f:154 text = f.read()155 tokens = enc.encode(text)156 self.tokens = torch.tensor(tokens)157 print(f"Loaded {len(self.tokens)} tokens")158 print(f"1 epoch = {len(self.tokens) // (B * T)} batches")159 self.current_position = 0160 161 def next_batch(self):162 B, T = self.B, self.T163 buf = self.tokens[self.current_position: self.current_position + B * T + 1]164 x = (buf[:-1]).view(B, T)165 y = (buf[1:]).view(B, T)166 self.current_position += B * T167 if self.current_position + (B * T + 1) > len(self.tokens):168 self.current_position = 0169 return x, y170 171 172# Initialize the model and training setup173model = GPT(GPTConfig())174model.to(device)175 176num_return_sequences = 5177train_loader = DataLoaderLite(B=num_return_sequences, T=32)178 179from torch.optim.lr_scheduler import CosineAnnealingLR180 181optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)182scheduler = CosineAnnealingLR(optimizer, T_max=50)183 184# Training loop185for i in range(20):186 x, y = train_loader.next_batch()187 x, y = x.to(device), y.to(device)188 optimizer.zero_grad()189 logits, loss = model(x, y)190 loss.backward()191 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)192 optimizer.step()193 scheduler.step()194 print(f"Step {i}, Loss: {loss.item()}")195 if loss.item() < 0.1:196 print("Target loss reached. Stopping training.")197 break198 199# Generation loop200torch.manual_seed(42)201torch.cuda.manual_seed(42)202 203# Define max_length here204max_length = 50205 206x = torch.randint(0, model.config.vocab_size, (num_return_sequences, 1), device=device)207 208while x.size(1) < max_length:209 with torch.no_grad():210 logits = model(x)[0]211 logits = logits[:, -1, :]212 probs = F.softmax(logits, dim=-1)213 topk_probs, topk_indices = torch.topk(probs, 50, dim=-1)214 ix = torch.multinomial(topk_probs, 1)215 xcol = torch.gather(topk_indices, -1, ix)216 x = torch.cat((x, xcol), dim=1)217 218for i in range(num_return_sequences):219 tokens = x[i, :max_length].tolist()220 decoded = enc.decode(tokens)221 print(">", decoded)