Prepare the dataset
First, download and tokenize the OpenWebText dataset:This downloads the OpenWebText dataset, an open reproduction of OpenAI’s private WebText dataset, and tokenizes it using GPT-2 BPE encoding.
Dataset statistics
After preparation, you’ll have:train.bin: ~17GB, ~9B tokens (9,035,582,198)val.bin: ~8.5MB, ~4M tokens (4,434,897)- Split: 8,009,762 training documents, 4,007 validation documents
uint16 bytes containing GPT-2 BPE token IDs.
Training GPT-2 (124M)
To reproduce GPT-2 with 124M parameters, you need at least an 8x A100 40GB node:Training configuration
Theconfig/train_gpt2.py file contains the hyperparameters:
Model architecture
Fromtrain.py defaults:
Expected results
- Training time: ~4 days on 8x A100 40GB
- Final validation loss: ~2.85
- Tokens per iteration: 491,520
- Total tokens trained: ~300 billion
GPT-2 (124M) evaluated directly on OpenWebText gets a validation loss of ~3.11, but finetuning brings it down to ~2.85. This indicates a domain gap between OpenWebText and the original (closed) WebText dataset.
Distributed training
Single node, multiple GPUs
Thetorchrun command automatically sets up PyTorch Distributed Data Parallel (DDP):
1
DDP initialization
From
train.py:82-95, DDP is automatically detected and initialized:2
Gradient accumulation scaling
Gradient accumulation steps are divided by world size:
3
Model wrapping
The model is wrapped with DDP at
train.py:210-212:Multi-node training
For training across multiple nodes with Infiniband interconnect:Benchmark your interconnect
Before running multi-node training, test your network speed:Baseline comparisons
OpenAI GPT-2 checkpoints provide baselines for OpenWebText:
Evaluate these baselines yourself:
The domain gap between WebText (closed) and OpenWebText means a direct GPT-2 (124M) evaluation gives 3.11 validation loss. After finetuning on OpenWebText, it reaches ~2.85, matching our reproduction target.
Performance optimizations
PyTorch 2.0 compile
By default, nanoGPT usestorch.compile() for significant speedups:
Mixed precision training
Automatic mixed precision is enabled by default:Efficient data loading
The “poor man’s data loader” fromtrain.py:114-131 uses memory-mapped files:
Monitor training progress
Enable Weights & Biases logging:- Training and validation loss
- Learning rate
- Model FLOPs Utilization (MFU)
- Iteration time
Sample from the model
After training, generate samples:Next steps
Distributed training
Deep dive into DDP setup and multi-node training
Finetuning
Finetune your trained model on custom datasets