Skip to main content

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

Loads a batch of training or validation data using memory-mapped files.
str
required
Data split to load from: 'train' or 'val'
Returns: Tuple of (x, y)
  • 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

Computes accurate loss estimates over multiple batches for both training and validation splits. Returns: Dictionary with keys 'train' and 'val', each containing the mean loss
int
default:"200"
Number of iterations to average over (configured globally)

get_lr

Computes the learning rate for a given iteration using cosine decay with linear warmup.
int
required
Current iteration number
Returns: float - Learning rate for this iteration

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:
To resume training from a checkpoint, set 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
When enabled, the following metrics are logged:
  • Training loss
  • Validation loss
  • Learning rate
  • Model FLOPs Utilization (MFU)

Performance tips

Use gradient accumulation to simulate larger batch sizes without running out of memory. Effective batch size = batch_size * gradient_accumulation_steps * num_gpus.
Enable compilation with compile=True to use PyTorch 2.0’s optimizations for faster training (requires PyTorch >= 2.0).
Use bfloat16 if your GPU supports it (requires Ampere or newer). It provides better numerical stability than float16 without requiring gradient scaling.
Allow TF32 (enabled by default) for ~20% speedup on Ampere GPUs without loss of accuracy.