Team Ai
Apppublic

ujwaliyengar/Decoder_Only_Model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py221 linesDownload Raw Back to root
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)