""" train_gpt_simple.py This file descends from the [NanoGPT speedrun](https://github.com/KellerJordan/modded-nanogpt). It was prepared as a simplified version of the speedrun for use in neural net optimization research. """ import os import sys with open(sys.argv[0]) as f: code = f.read() # read the code of this file ASAP, for logging import uuid import time from pathlib import Path import torch from torch import Tensor, nn from torch.optim import AdamW import torch.nn.functional as F import torch.distributed as dist ######################################## # Dataloader # ######################################## def _load_data_shard(file: Path): header = torch.from_file(str(file), False, 256, dtype=torch.int32) # header is 256 int32 assert header[0] == 20240520, "magic number mismatch in the data .bin file" assert header[1] == 1, "unsupported version" num_tokens = int(header[2]) # number of tokens (claimed) with file.open("rb", buffering=0) as f: tokens = torch.empty(num_tokens, dtype=torch.uint16, pin_memory=True) f.seek(256 * 4) nbytes = f.readinto(tokens.numpy()) # avoid bytes->array copy assert nbytes == 2 * num_tokens, "number of tokens read does not match header" return tokens def distributed_data_generator(filename_pattern: str, batch_size: int, seq_len=1024): world_size = dist.get_world_size() rank = dist.get_rank() files = sorted(Path.cwd().glob(filename_pattern)) assert batch_size % world_size == 0 local_batch_size = batch_size // world_size file_iter = iter(files) tokens, pos = _load_data_shard(next(file_iter)), 0 while True: if pos + batch_size + 1 >= len(tokens): tokens, pos = _load_data_shard(next(file_iter)), 0 buf = tokens[pos + rank * local_batch_size:][:local_batch_size + 1] inputs = buf[:-1].to(device="cuda", dtype=torch.int32, non_blocking=True) targets = buf[1:].to(device="cuda", dtype=torch.int64, non_blocking=True) pos += batch_size yield inputs.view(-1, seq_len), targets.view(-1, seq_len) ######################################## # Architecture # ######################################## def norm(x: Tensor): return F.rms_norm(x, (x.size(-1),)) class Linear(nn.Linear): def __init__(self, in_features, out_features): super().__init__(in_features, out_features, bias=True) def forward(self, x): return F.linear(x, self.weight.type_as(x), self.bias.type_as(x)) class RMSNorm(nn.Module): def __init__(self, dim): super().__init__() self.gains = nn.Parameter(torch.ones(dim)) def forward(self, x): return (norm(x.float()) * self.gains).type_as(x) class Rotary(nn.Module): def __init__(self, dim: int): super().__init__() # half-truncate RoPE (w/ base freq tuning) angular_freq = (1 / 1024) ** torch.linspace(0, 1, steps=dim//4, dtype=torch.float32) self.register_buffer("angular_freq", torch.cat([angular_freq, angular_freq.new_zeros(dim//4)])) def forward(self, x_BTHD: Tensor): pos = torch.arange(x_BTHD.size(1), dtype=torch.float32, device=x_BTHD.device) theta = torch.outer(pos, self.angular_freq)[None, :, None, :] cos, sin = theta.cos(), theta.sin() x1, x2 = x_BTHD.to(dtype=torch.float32).chunk(2, dim=-1) y1 = x1 * cos + x2 * sin y2 = x1 * (-sin) + x2 * cos return torch.cat((y1, y2), 3).type_as(x_BTHD) class CausalSelfAttention(nn.Module): def __init__(self, dim: int, head_dim=128): super().__init__() self.num_heads = dim // head_dim self.head_dim = head_dim hdim = self.num_heads * self.head_dim self.q = Linear(dim, hdim) self.k = Linear(dim, hdim) self.v = Linear(dim, hdim) self.proj = Linear(hdim, dim) self.rotary = Rotary(head_dim) def forward(self, x: Tensor): B, T = x.size(0), x.size(1) q = self.q(x).view(B, T, self.num_heads, self.head_dim) k = self.k(x).view(B, T, self.num_heads, self.head_dim) v = self.v(x).view(B, T, self.num_heads, self.head_dim) q, k = norm(q), norm(k) q, k = self.rotary(q), self.rotary(k) y = F.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), scale=0.12, is_causal=True).transpose(1, 2) y = y.contiguous().view(B, T, self.num_heads * self.head_dim) y = self.proj(y) return y class MLP(nn.Module): def __init__(self, dim: int): super().__init__() hdim = 4 * dim self.fc = Linear(dim, hdim) self.proj = Linear(hdim, dim) def forward(self, x: Tensor): x = self.fc(x) x = F.relu(x).square() x = self.proj(x) return x class Block(nn.Module): def __init__(self, dim: int): super().__init__() self.attn = CausalSelfAttention(dim) self.mlp = MLP(dim) self.norm1 = RMSNorm(dim) self.norm2 = RMSNorm(dim) def forward(self, x: Tensor): x = x + self.attn(self.norm1(x)) x = x + self.mlp(self.norm2(x)) return x class GPT(nn.Module): def __init__(self, vocab_size: int, num_layers: int, model_dim: int): super().__init__() self.embed = nn.Embedding(vocab_size, model_dim).bfloat16() self.blocks = nn.ModuleList([Block(model_dim) for _ in range(num_layers)]) self.proj = Linear(model_dim, vocab_size) self.norm1 = RMSNorm(model_dim) self.norm2 = RMSNorm(model_dim) def forward(self, inputs: Tensor, targets: Tensor): x = self.norm1(self.embed(inputs)) for block in self.blocks: x = block(x) logits = self.proj(self.norm2(x)).float() logits = 15 * logits * (logits.square() + 15**2).rsqrt() return F.cross_entropy(logits.view(targets.numel(), -1), targets.view(-1), reduction="sum") ######################################## # Optimizer # ######################################## def zeropower_via_newtonschulz5(G: Tensor) -> Tensor: assert G.ndim >= 2 X = G.bfloat16() if G.size(-2) > G.size(-1): X = X.mT # Ensure spectral norm is at most 1 X = X / (X.norm(dim=(-2, -1), keepdim=True) + 1e-7) # Perform the NS iterations, not optimizing for wallclock speed a, b, c = 2, -1.5, 0.5 for _ in range(12): A = X @ X.mT B = b * A + c * A @ A X = a * X + B @ X if G.size(-2) > G.size(-1): X = X.mT return X @torch.compile def muon_update(grad, momentum, beta=0.95, nesterov=True): momentum.lerp_(grad, 1 - beta) update = grad.lerp_(momentum, beta) if nesterov else momentum update = zeropower_via_newtonschulz5(update) update *= max(1, grad.size(-2) / grad.size(-1))**0.5 return update class Muon(torch.optim.Optimizer): def __init__(self, params, lr=0.02, weight_decay=0, momentum=0.95): defaults = dict(lr=lr, weight_decay=weight_decay, momentum=momentum) assert isinstance(params, list) and len(params) >= 1 and isinstance(params[0], torch.nn.Parameter) params = sorted(params, key=lambda x: x.size(), reverse=True) super().__init__(params, defaults) @torch.no_grad() def step(self): world_size = dist.get_world_size() rank = dist.get_rank() for group in self.param_groups: params = group["params"] params_pad = params + [torch.empty_like(params[-1])] * (world_size - len(params) % world_size) for base_i in range(len(params))[::world_size]: if base_i + rank < len(params): p = params[base_i + rank] state = self.state[p] if len(state) == 0: state["momentum_buffer"] = torch.zeros_like(p) update = muon_update(p.grad, state["momentum_buffer"], beta=group["momentum"]) p.mul_(1 - group["lr"] * group["weight_decay"]) p.add_(update, alpha=-group["lr"]) dist.all_gather(params_pad[base_i:base_i + world_size], params_pad[base_i + rank]) ######################################## # Setup # ######################################## # torchrun sets these env variables device = torch.device("cuda", int(os.environ["LOCAL_RANK"])) torch.cuda.set_device(device) dist.init_process_group(backend="nccl", device_id=device) dist.barrier() # this code can be run equivalently with 1, 2, 4, or 8 gpus. assert 8 % dist.get_world_size() == 0 # logging setup if dist.get_rank() == 0: os.makedirs("logs", exist_ok=True) logfile = f"logs/{uuid.uuid4()}.txt" print(logfile) def print0(s, console=False, log=True): if dist.get_rank() == 0: with open(logfile, "a") as f: if console: print(s) if log: print(s, file=f) # we begin by logging this file itself print0(code) print0("="*100) print0(f"Running PyTorch {torch.version.__version__} compiled for CUDA {torch.version.cuda}") print0("="*100) val_tokens = 20 * 524288 batch_size = 8 * 64 * 1024 mbs = 64 train_loader = distributed_data_generator("data/fineweb10B/fineweb_train_*.bin", batch_size) val_inputs, val_targets = next(distributed_data_generator("data/fineweb10B/fineweb_val_*.bin", val_tokens)) model = GPT(vocab_size=50304, num_layers=12, model_dim=768).cuda() model.compile(dynamic=False) ######################################## # Init & Optim Hyperparams # ######################################## # we want to minimize this while still reaching 3.28 val loss train_steps = 5625 # AdamW replacement for the block matrix parameters that Muon optimizes in the baseline. # The auxiliary AdamW parameter groups below are intentionally left unchanged. block_adamw_lr = 0.0015 block_adamw_weight_decay = 0.10 block_adamw_warmup_steps = 250 block_adamw_betas = (0.9, 0.95) # initialize model parameters for name, p in model.named_parameters(): if "proj" in name: p.data.zero_() dist.broadcast(p.detach(), 0) # create the optimizer(s) optimizer1 = AdamW([dict(params=[model.embed.weight], lr=0.3), dict(params=[model.proj.weight], lr=1/320), dict(params=[p for p in model.parameters() if p.ndim < 2], lr=0.01)], betas=(0.8, 0.95), eps=1e-10, weight_decay=0, fused=True) optimizer2 = AdamW([p for p in model.blocks.parameters() if p.ndim >= 2], lr=block_adamw_lr, betas=block_adamw_betas, eps=1e-10, weight_decay=block_adamw_weight_decay, fused=True) optimizers = [optimizer1, optimizer2] assert set(p for opt in optimizers for group in opt.param_groups for p in group["params"]) == set(model.parameters()) for opt in optimizers: for group in opt.param_groups: group["initial_lr"] = group["lr"] group["warmup_steps"] = block_adamw_warmup_steps # learning rate schedule: stable then decay def set_hparams(step, cooldown_frac=0.7): progress = step / train_steps assert 0 <= progress < 1 if progress < 1 - cooldown_frac: eta = 1.0 else: eta = (1 - progress) / cooldown_frac for opt in optimizers: for group in opt.param_groups: warmup_steps = group["warmup_steps"] warmup = min(1.0, (step + 1) / warmup_steps) if warmup_steps > 0 else 1.0 group["lr"] = group["initial_lr"] * warmup * eta ######################################## # Training and Validation # ######################################## training_time = 0 # start the clock dist.barrier() t0 = time.perf_counter() for step in range(train_steps + 1): # --------------- VALIDATION SECTION ----------------- if step == train_steps or step % 125 == 0: # stop the clock dist.barrier() training_time += time.perf_counter() - t0 model.eval() val_loss = 0 with torch.no_grad(): assert len(val_inputs) % mbs == 0 for i in range(len(val_inputs) // mbs): val_loss += model(val_inputs[i*mbs:(i+1)*mbs], val_targets[i*mbs:(i+1)*mbs]) dist.all_reduce(val_loss, op=dist.ReduceOp.SUM) val_loss /= val_tokens print0(f"step:{step}/{train_steps} val_loss:{val_loss:.5f} train_time:{training_time:.3f}s" + f" step_avg:{1000*training_time/max(step, 1):.2f}ms", console=True) model.train() # start the clock again dist.barrier() t0 = time.perf_counter() if step == train_steps: break # --------------- TRAINING SECTION ----------------- inputs, targets = next(train_loader) # accumulate across microbatches in case we are running with fewer than 8 gpus assert len(inputs) % mbs == 0 for i in range(len(inputs) // mbs): model(inputs[i*mbs:(i+1)*mbs], targets[i*mbs:(i+1)*mbs]).backward() for name, p in model.named_parameters(): assert p.grad is not None, name dist.all_reduce(p.grad, op=dist.ReduceOp.SUM) # set optimization hyperparameters and take a step set_hparams(step) for opt in optimizers: opt.step() model.zero_grad(set_to_none=True) approx_training_time = training_time + (time.perf_counter() - t0) print0(f"step:{step+1}/{train_steps} train_time:{approx_training_time:.3f}s" + f" step_avg:{1000*approx_training_time/(step + 1):.2f}ms", console=True, log=False) dist.destroy_process_group() ==================================================================================================== Running PyTorch 2.10.0+cu128 compiled for CUDA 12.8 ==================================================================================================== step:0/5625 val_loss:10.82584 train_time:0.000s step_avg:0.45ms step:125/5625 val_loss:6.18993 train_time:20.672s step_avg:165.37ms step:250/5625 val_loss:5.07445 train_time:39.141s step_avg:156.57ms step:375/5625 val_loss:4.41418 train_time:57.606s step_avg:153.62ms step:500/5625 val_loss:4.15174 train_time:76.051s step_avg:152.10ms step:625/5625 val_loss:3.99071 train_time:94.504s step_avg:151.21ms step:750/5625 val_loss:3.89882 train_time:112.943s step_avg:150.59ms step:875/5625 val_loss:3.83496 train_time:131.392s step_avg:150.16ms step:1000/5625 val_loss:3.77288 train_time:149.846s step_avg:149.85ms step:1125/5625 val_loss:3.74011 train_time:168.290s step_avg:149.59ms step:1250/5625 val_loss:3.69335 train_time:186.735s step_avg:149.39ms step:1375/5625 val_loss:3.67990 train_time:205.192s step_avg:149.23ms step:1500/5625 val_loss:3.63656 train_time:223.647s step_avg:149.10ms step:1625/5625 val_loss:3.61797 train_time:242.103s step_avg:148.99ms step:1750/5625 val_loss:3.59600 train_time:260.562s step_avg:148.89ms step:1875/5625 val_loss:3.57059 train_time:279.022s step_avg:148.81ms step:2000/5625 val_loss:3.55221 train_time:297.487s step_avg:148.74ms step:2125/5625 val_loss:3.53439 train_time:315.952s step_avg:148.68ms step:2250/5625 val_loss:3.51779 train_time:334.423s step_avg:148.63ms step:2375/5625 val_loss:3.50229 train_time:352.885s step_avg:148.58ms step:2500/5625 val_loss:3.48864 train_time:371.351s step_avg:148.54ms step:2625/5625 val_loss:3.47456 train_time:389.811s step_avg:148.50ms step:2750/5625 val_loss:3.46283 train_time:408.268s step_avg:148.46ms step:2875/5625 val_loss:3.45253 train_time:426.803s step_avg:148.45ms step:3000/5625 val_loss:3.44107 train_time:445.259s step_avg:148.42ms step:3125/5625 val_loss:3.42847 train_time:463.713s step_avg:148.39ms step:3250/5625 val_loss:3.41711 train_time:482.167s step_avg:148.36ms step:3375/5625 val_loss:3.40741 train_time:500.627s step_avg:148.33ms step:3500/5625 val_loss:3.39651 train_time:519.077s step_avg:148.31ms step:3625/5625 val_loss:3.38970 train_time:537.536s step_avg:148.29ms step:3750/5625 val_loss:3.38059 train_time:555.999s step_avg:148.27ms step:3875/5625 val_loss:3.37081 train_time:574.467s step_avg:148.25ms step:4000/5625 val_loss:3.36230 train_time:592.935s step_avg:148.23ms step:4125/5625 val_loss:3.35442 train_time:611.412s step_avg:148.22ms step:4250/5625 val_loss:3.34739 train_time:629.891s step_avg:148.21ms step:4375/5625 val_loss:3.33899 train_time:648.377s step_avg:148.20ms step:4500/5625 val_loss:3.33182 train_time:666.861s step_avg:148.19ms step:4625/5625 val_loss:3.32404 train_time:685.343s step_avg:148.18ms step:4750/5625 val_loss:3.31602 train_time:703.832s step_avg:148.18ms step:4875/5625 val_loss:3.30929 train_time:722.341s step_avg:148.17ms step:5000/5625 val_loss:3.30259 train_time:740.836s step_avg:148.17ms step:5125/5625 val_loss:3.29661 train_time:759.310s step_avg:148.16ms step:5250/5625 val_loss:3.29050 train_time:777.788s step_avg:148.15ms step:5375/5625 val_loss:3.28558 train_time:796.249s step_avg:148.14ms step:5500/5625 val_loss:3.28145 train_time:814.734s step_avg:148.13ms step:5625/5625 val_loss:3.27903 train_time:833.211s step_avg:148.13ms