The goal is to reproduce the 124 million parameter version of GPT-2 from scratch, leveraging insights from both the original GPT-2 and GPT-3 papers. This involves understanding the decoder-only Transformer architecture and its components. The process aims to achieve performance comparable to or better than the original model, using modern tools and techniques. The initial step involves loading the pre-trained GPT-2 weights to understand the target architecture and parameter structure. This foundational step ensures alignment before embarking on training from scratch. The final sentence of this claim is that this meticulous reproduction serves as a crucial educational step for understanding large language models.
Karpathy: Implementing the GPT-2 Architecture
The implementation of the GPT-2 architecture begins by defining the core `Transformer` module, which includes token and positional embeddings (`wte`, `pe`), a list of 12 Transformer blocks (`h`), a final layer normalization (`lnf`), and the language model head (`lm_head`). Each block consists of pre-layer normalization, multi-head self-attention, and a feed-forward network (MLP) with a `GELU` nonlinearity. The MLP uses the approximate `GELU` activation function as used in the original GPT-2. This structured approach ensures that the custom implementation closely mirrors the architecture described in the papers and used by libraries like Hugging Face Transformers, facilitating weight loading and understanding. The final sentence of this claim is that this detailed structural mapping is essential for a faithful reproduction.
Efficient GPT-2 Implementation
The implementation of the GPT-2 architecture in PyTorch is optimized for efficiency by treating the number of attention heads as a batch dimension, allowing parallel operations across heads and batches. This approach, while algorithmically equivalent to previous methods, leverages PyTorch's capabilities for faster execution. Variable naming conventions are aligned with Hugging Face's Transformers library to facilitate weight porting.
Loading Hugging Face GPT-2 Parameters
To initialize our custom GPT-2 model with pre-trained weights, we load parameters from a Hugging Face checkpoint. This involves creating state dictionaries for both our model and the Hugging Face model, then copying tensors over. Some buffers are ignored, and specific weights that are transposed from PyTorch's expected format (originating from TensorFlow) are manually transposed back.
Data Preparation for Training
To prepare data for training, sequences of tokens are fetched with an extra token to serve as the target for the last token in the input sequence. These sequences are then reshaped into batches of (B, T) for input and (B, T) for targets, ensuring that each input token has a corresponding target token for loss calculation. This process is fundamental for supervised learning in language models.
Implementing Cross-Entropy Loss
The cross-entropy loss is calculated by flattening the logits (B, T, vocab_size) and targets (B, T) into two-dimensional tensors, which are then passed to PyTorch's functional cross-entropy function. This loss quantifies the difference between the model's predicted probability distribution for the next token and the actual next token, guiding the optimization process.
Karpathy: GPT-2 Initialization Nuances
The initialization of weights in the GPT-2 model is critical, with the paper suggesting a standard deviation of 0.02. This value is roughly consistent with theoretical calculations based on the model's internal dimensions. However, a more advanced initialization scales weights in residual layers by 1/sqrt(N) to control activation variance growth, a detail implemented by scaling the standard deviation.
Karpathy: The Quest for Speed - GPU Utilization
To maximize training speed, one must understand the hardware capabilities. Karpathy showcases his setup with eight A100 80GB GPUs, emphasizing the importance of checking GPU utilization (e.g., via `nvidia-smi`). He notes that deep learning training is often memory-bound, meaning tensor cores can be idle waiting for data, making memory bandwidth a critical bottleneck.
Mixed Precision Training
Utilizing BFloat16 with PyTorch's AutoCast context manager allows for mixed-precision training, where certain operations run in lower precision (BFloat16) while others remain in Float32. This is enabled by Ampere GPUs and significantly speeds up computation by leveraging Tensor Cores, though it may slightly impact accuracy. The key is to selectively apply this to operations like matrix multiplications while keeping sensitive operations like normalization in higher precision. This optimization reduced iteration time from 333ms to 300ms.
The Power of torch.compile
Introducing torch.compile, a compiler for neural networks, dramatically reduces Python overhead and GPU read/write operations. By analyzing the entire network structure, it fuses operations and eliminates the interpreter's step-by-step execution. This single-line addition to the code resulted in a significant speedup, reducing iteration time from 300ms to 129ms, a 2.3x improvement, by optimizing memory access patterns and enabling kernel fusion.
Karpathy: Implementing Cosine Decay LR Schedule
A cosine decay learning rate schedule with warmup is implemented, mirroring GPT-3's approach. The learning rate linearly increases during warmup, then decays following a cosine curve to 10% of its maximum value over the training horizon. This sophisticated schedule aims to balance rapid initial learning with fine-tuning later in training.
Karpathy: Weight Decay and Fused AdamW
Weight decay of 0.1 is applied, primarily to embeddings and matrix multiplication weights, excluding biases and layer norm parameters. This regularization technique encourages the model to distribute learning across more parameters. Additionally, a fused implementation of AdamW is utilized for performance gains on CUDA, consolidating multiple update kernels into one.
Fused AdamW Optimizer
Implementing a fused AdamW optimizer, which combines multiple operations into a single kernel, leads to performance improvements. This optimization reduced the per-step running time from 93 milliseconds to 90 milliseconds, demonstrating the benefits of hardware-level optimizations for training speed.
Gradient Accumulation Explained
To simulate a large batch size (e.g., 0.5 million tokens) on limited GPU memory, gradient accumulation is employed. This technique involves performing multiple forward-backward passes with smaller 'micro-batches' and accumulating their gradients before performing a single optimizer update, effectively simulating a larger batch size serially.
Distributed Data Parallel (DDP) Implementation
Implementing Distributed Data Parallel (DDP) in PyTorch requires wrapping the model and carefully managing gradient synchronization. While the forward pass remains unchanged, DDP synchronizes gradients during the backward pass via an all-reduce operation. For gradient accumulation, synchronization is intentionally skipped until the final micro-step to avoid performance overhead, achieved by directly toggling PyTorch's internal gradient synchronization flag.
Synchronizing Loss Accumulation with Gradients
After averaging gradients with DDP, the accumulated loss (loss_AUM) also needs to be synchronized across all processes. This is achieved using `torch.distributed.all_reduce` on the loss_AUM tensor. This ensures that the reported loss accurately reflects the average loss across all GPUs, maintaining consistency with the averaged gradients.
HellaSwag Evaluation Explained
HellaSwag is a sentence completion benchmark designed to test world knowledge. It presents a context and four multiple-choice options, where only one is a natural continuation. Models are evaluated by their ability to predict the most likely completion, with humans achieving 95% accuracy historically, though modern models surpass this. The evaluation method involves constructing batches of four options and assessing the average probability of tokens within each option to determine the most likely completion.
Training Script Modifications and HellaSwag Integration
The main training script is updated to optionally disable `torch.compile` due to issues with the evaluation and sampling code, impacting speed. A log directory is created for `log.txt` to record training loss, validation loss, and HellaSwag accuracies. Periodic evaluation of validation loss and HellaSwag accuracy (every 250 iterations, if `torch.compile` is off) and sampling are incorporated into the training loop.
Hyperparameter Tuning and GPT-3 Parity
The hyperparameters inherited from the GPT-3 paper are quite conservative; for instance, the maximum learning rate can be almost tripled, leading to faster training. To achieve exact parity with GPT-3, the sequence length should be increased to 2048, and the batch size decreased to 32 to maintain the same total number of tokens. This adjustment ensures the model's sequence length matches GPT-3's, making the models virtually identical in architecture.