From 9f08882051426ac5d0fe30db283b195c20dad34b Mon Sep 17 00:00:00 2001 From: Andrej Karpathy Date: Sat, 25 May 2024 00:14:10 +0000 Subject: [PATCH] add weight decay, but only for 2D tensors, as done in GPT series and in general too. this forces us to break up our adamw kernel again into one call per tensor, so there is a small throughput hit, of about 0.5% for me. but we have to break up this kernel in near future anyway --- train_gpt2.cu | 30 +++++++++++++++++++++++++----- train_gpt2.py | 35 ++++++++++++++++++++++++++++++++--- 2 files changed, 57 insertions(+), 8 deletions(-) diff --git a/train_gpt2.cu b/train_gpt2.cu index 26d2992..3e89646 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -2711,14 +2711,34 @@ float gpt2_update(GPT2 *model, float learning_rate, float beta1, float beta2, fl // AdamW update int block_size = 512; - int num_blocks = CEIL_DIV(num_parameters, block_size); float beta1_correction = 1.0f - powf(beta1, t); float beta2_correction = 1.0f - powf(beta2, t); unsigned int seed = random_u32(&model->rng_state); - adamw_kernel3<<>>(params_memory, model->master_weights, grads_memory, - model->m_memory, model->v_memory, num_parameters, - learning_rate, beta1, beta2, beta1_correction, beta2_correction, eps, weight_decay, - grad_scale, seed); + + // individually call the adamw_kernel3 on all parameter tensors separately + floatX* params_memory_iter = params_memory; + float* master_weights_iter = model->master_weights; + floatX* grads_memory_iter = grads_memory; + float* m_memory_iter = (float*)model->m_memory; + float* v_memory_iter = (float*)model->v_memory; + for (int i = 0; i < NUM_PARAMETER_TENSORS; i++) { + size_t num_parameters = model->param_elements[i]; + int num_blocks = CEIL_DIV(num_parameters, block_size); + // we only want to weight decay the 2D tensors and leave all 1D tensors alone + // in particular this also decays the embedding weights, but this is ok: + // - the token embeddings are weight shared and participate in the final projection to logits + // - the position embeddings actively participate at every forward/backward pass + float wd = (i == 0 || i == 1 || i == 4 || i == 6 || i == 10 || i == 12) ? weight_decay : 0.0f; + adamw_kernel3<<>>(params_memory_iter, master_weights_iter, grads_memory_iter, + m_memory_iter, v_memory_iter, num_parameters, + learning_rate, beta1, beta2, beta1_correction, beta2_correction, eps, wd, + grad_scale, seed); + params_memory_iter += num_parameters; + if (master_weights_iter != NULL) { master_weights_iter += num_parameters; } + grads_memory_iter += num_parameters; + m_memory_iter += num_parameters; + v_memory_iter += num_parameters; + } cudaCheck(cudaGetLastError()); return grad_norm_cpu; } diff --git a/train_gpt2.py b/train_gpt2.py index d32ab50..f2b0ae3 100644 --- a/train_gpt2.py +++ b/train_gpt2.py @@ -19,6 +19,7 @@ torchrun --standalone --nproc_per_node=4 train_gpt2.py --write_tensors=0 --num_i import os import math import struct +import inspect from contextlib import nullcontext from dataclasses import dataclass @@ -228,6 +229,31 @@ class GPT(nn.Module): return model + def configure_optimizers(self, weight_decay, learning_rate, betas, device_type): + # start with all of the candidate parameters + param_dict = {pn: p for pn, p in self.named_parameters()} + # filter out those that do not require grad + param_dict = {pn: p for pn, p in param_dict.items() if p.requires_grad} + # create optim groups. Any parameters that is 2D will be weight decayed, otherwise no. + # i.e. all weight tensors in matmuls + embeddings decay, all biases and layernorms don't. + 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") + # Create AdamW optimizer and use the fused version if it is available + fused_available = 'fused' in inspect.signature(torch.optim.AdamW).parameters + use_fused = fused_available and device_type == 'cuda' + extra_args = dict(fused=True) if use_fused else dict() + optimizer = torch.optim.AdamW(optim_groups, lr=learning_rate, betas=betas, **extra_args) + print(f"using fused AdamW: {use_fused}") + return optimizer + @torch.no_grad() def generate(self, idx, max_new_tokens, temperature=1.0, top_k=None): """ @@ -429,6 +455,7 @@ if __name__ == "__main__": parser.add_argument("--sequence_length", type=int, default=64, help="sequence length") parser.add_argument("--total_batch_size", type=int, default=256, help="total desired batch size, in units of #tokens") parser.add_argument("--grad_clip", type=float, default=1.0, help="maximum gradient magnitude") + parser.add_argument("--weight_decay", type=float, default=0.0, help="weight decay") parser.add_argument("--overfit_single_batch", type=int, default=1, help="overfit just one batch of data") args = parser.parse_args() B, T = args.batch_size, args.sequence_length @@ -467,6 +494,7 @@ if __name__ == "__main__": elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): device = "mps" print(f"using device: {device}") + device_type = 'cuda' if 'cuda' in device else 'cpu' # calculate gradient accumulation from the desired total batch size and the current run configuration tokens_per_fwdbwd = B * T * ddp_world_size @@ -478,7 +506,7 @@ if __name__ == "__main__": # set up a context manager following the desired dtype and device ptdtype = {'float32': torch.float32, 'bfloat16': torch.bfloat16, 'float16': torch.float16}[args.dtype] - ctx = torch.amp.autocast(device_type="cuda", dtype=ptdtype) if device == "cuda" else nullcontext() + ctx = torch.amp.autocast(device_type=device_type, dtype=ptdtype) if device_type == "cuda" else nullcontext() # seed the random number generators (in DDP we want different processes to use different offsets) # in the code below we don't actually use random numbers because there is no active dataloader @@ -604,8 +632,9 @@ if __name__ == "__main__": raw_model = model.module if ddp else model # always contains the "raw" unwrapped model # init the optimizer - adam_use_fused = device == "cuda" # only works on CUDA (?) - optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, betas=(0.9, 0.95), weight_decay=0.0, fused=adam_use_fused) + optimizer = raw_model.configure_optimizers(weight_decay=args.weight_decay, + learning_rate=1e-4, betas=(0.9, 0.95), + device_type=device) if device == "cuda": torch.cuda.reset_peak_memory_stats()