Overview
The training module provides a complete implementation for training GPT models with support for:- Single GPU and distributed data parallel (DDP) training
- Mixed precision training (float16/bfloat16)
- Gradient accumulation
- Learning rate scheduling with warmup and cosine decay
- Checkpointing and resumption
- WandB integration for experiment tracking
Training modes
You can run the training script in multiple configurations:If your cluster does not have Infiniband, prepend
NCCL_IB_DISABLE=1 to the commands.Key functions
get_batch
str
required
Data split to load from:
'train' or 'val'x: Input token sequences of shape (batch_size, block_size)y: Target token sequences of shape (batch_size, block_size), shifted by one position
estimate_loss
'train' and 'val', each containing the mean loss
int
default:"200"
Number of iterations to average over (configured globally)
get_lr
int
required
Current iteration number
Training loop structure
The main training loop performs the following steps:1. Learning rate scheduling
2. Periodic evaluation
3. Forward and backward pass
With gradient accumulation to simulate larger batch sizes:4. Gradient clipping and optimizer step
5. Logging
Configuration parameters
I/O settings
str
default:"'out'"
Directory for saving checkpoints
int
default:"2000"
How often to evaluate on val set and save checkpoints
int
default:"1"
How often to log training metrics
int
default:"200"
Number of iterations for loss estimation
bool
default:"False"
If True, exit after first evaluation (useful for testing)
bool
default:"True"
If True, save checkpoint after each eval even if val loss didn’t improve
str
default:"'scratch'"
Initialization mode:
'scratch', 'resume', or a GPT-2 variant ('gpt2', 'gpt2-medium', 'gpt2-large', 'gpt2-xl')Data settings
str
default:"'openwebtext'"
Name of dataset (must have corresponding data// directory)
int
default:"40"
Accumulate gradients over this many steps to simulate larger batches
int
default:"12"
Micro-batch size (per GPU if using DDP)
int
default:"1024"
Context length for training sequences
Model architecture
int
default:"12"
Number of transformer layers
int
default:"12"
Number of attention heads
int
default:"768"
Embedding dimension
float
default:"0.0"
Dropout rate (0.0 for pretraining, 0.1+ for finetuning)
bool
default:"False"
Use bias in Linear and LayerNorm layers
Optimizer settings
float
default:"6e-4"
Maximum learning rate
int
default:"600000"
Total number of training iterations
float
default:"1e-1"
Weight decay coefficient
float
default:"0.9"
AdamW beta1 parameter
float
default:"0.95"
AdamW beta2 parameter
float
default:"1.0"
Gradient clipping threshold (0.0 to disable)
Learning rate decay
bool
default:"True"
Enable learning rate decay
int
default:"2000"
Number of warmup iterations
int
default:"600000"
Iterations for learning rate decay (should be ~= max_iters)
float
default:"6e-5"
Minimum learning rate (should be ~= learning_rate/10)
System settings
str
default:"'cuda'"
Device to train on:
'cpu', 'cuda', 'cuda:0', 'cuda:1', 'mps', etc.str
default:"'bfloat16' or 'float16'"
Data type for training:
'float32', 'bfloat16', or 'float16'. Automatically selects bfloat16 if supported.bool
default:"True"
Use PyTorch 2.0 compilation for faster training
DDP settings
str
default:"'nccl'"
DDP backend:
'nccl' (recommended for CUDA) or 'gloo'Checkpointing
Checkpoints are saved to{out_dir}/ckpt.pt and contain:
init_from='resume' and ensure the checkpoint exists in out_dir.
WandB integration
bool
default:"False"
Enable Weights & Biases logging
str
default:"'owt'"
WandB project name
str
default:"'gpt2'"
WandB run name
- Training loss
- Validation loss
- Learning rate
- Model FLOPs Utilization (MFU)