Implementation:Microsoft DeepSpeedExamples BF16 Master Weight Training
| Knowledge Sources | |
|---|---|
| Domains | Mixed Precision Training, Memory Optimization |
| Last Updated | 2026-02-07 12:00 GMT |
Overview
Training script demonstrating DeepSpeed's BF16 master weights and gradients options for reduced GPU memory usage during transformer model training with configurable precision settings.
Description
This module implements a memory benchmarking training script that demonstrates DeepSpeed's bf16_master_weights_and_grads and bf16_optimizer_states configuration options for reducing memory footprint during training. It constructs a SimpleTransformerModel with configurable vocabulary size, hidden dimensions, number of layers, and attention heads, then trains it using different DeepSpeed BF16 configurations to compare memory usage and training performance.
The SimpleTransformerBlock class implements a standard transformer block with multi-head attention, layer normalization, and a GELU-activated feed-forward network. The SimpleTransformerModel wraps multiple blocks with token and positional embeddings and an output projection layer, supporting activation checkpointing for additional memory savings. The model is designed to be representative enough for meaningful memory benchmarking while remaining simple to understand.
The main() function handles the complete training pipeline: DeepSpeed distributed initialization, optional WikiText dataset loading via HuggingFace, model creation, DeepSpeed engine initialization from a JSON config file, and a training loop with detailed memory tracking at each step (allocated, reserved, and peak GPU memory). It supports torch autocast for mixed-precision training and outputs machine-readable summary statistics including peak memory, allocated memory, and average step time. Loss history can be saved to CSV for analysis.
Usage
Use this script to benchmark and demonstrate the memory savings of DeepSpeed's BF16 low-precision master weight options. Run it with different DeepSpeed configuration files (baseline, bf16_master_wg, bf16_full) to compare peak memory usage and training throughput. It supports both random synthetic data and real WikiText data.
Code Reference
Source Location
- Repository: Microsoft_DeepSpeedExamples
- File: training/bf16_master_weight/train.py
- Lines: 1-358
Signature
class SimpleTransformerBlock(nn.Module):
def __init__(self, hidden_dim, num_heads=8, ff_dim=None):
def forward(self, x):
class SimpleTransformerModel(nn.Module):
def __init__(self, vocab_size=50000, hidden_dim=1024, num_layers=12,
num_heads=16, max_seq_len=512):
def enable_activation_checkpointing(self):
def forward(self, input_ids):
def get_args():
def load_wikitext_data(tokenizer_name, seq_length, batch_size, world_size, rank):
def count_parameters(model):
def get_memory_stats():
def format_memory(bytes_val):
def main():
Import
from train import SimpleTransformerModel, SimpleTransformerBlock, main
I/O Contract
Inputs
| Name | Type | Required | Description |
|---|---|---|---|
| deepspeed_config | str | Yes | Path to the DeepSpeed JSON configuration file |
| hidden_dim | int | No | Hidden dimension size (default: 1024) |
| num_layers | int | No | Number of transformer layers (default: 12) |
| num_heads | int | No | Number of attention heads (default: 16) |
| vocab_size | int | No | Vocabulary size (default: 50000) |
| batch_size | int | No | Per-GPU batch size (default: 4) |
| seq_length | int | No | Input sequence length (default: 512) |
| num_steps | int | No | Number of training steps (default: 20) |
| activation_checkpointing | bool | No | Enable activation checkpointing (default: False) |
| use_real_data | bool | No | Use WikiText dataset instead of random data (default: False) |
| loss_log_file | str | No | File path to save loss history as CSV |
Outputs
| Name | Type | Description |
|---|---|---|
| peak_memory | int | Peak GPU memory usage in bytes during training |
| allocated_memory | int | Final allocated GPU memory in bytes |
| avg_step_time | float | Average training step time in seconds (excluding warmup) |
| loss_history | List[Tuple[int, float]] | List of (step, loss) tuples for the training run |
| summary_line | str | Machine-readable summary with config, params, memory, and timing stats |
Usage Examples
# Run with baseline BF16 configuration
# deepspeed --num_gpus=1 train.py --deepspeed_config configs/baseline.json
# Run with BF16 master weights and gradients for memory savings
# deepspeed --num_gpus=1 train.py --deepspeed_config configs/bf16_master_wg.json
# Run with full BF16 optimization (master weights + optimizer states)
# deepspeed --num_gpus=1 train.py --deepspeed_config configs/bf16_full.json
# Run with real data and activation checkpointing
# deepspeed --num_gpus=1 train.py \
# --deepspeed_config configs/bf16_full.json \
# --use_real_data \
# --activation_checkpointing \
# --num_steps 50 \
# --loss_log_file loss.csv