PyTorch Mixed Precision Training (AMP)
Mixed precision training is one of the most important performance optimization techniques in deep learning. By simultaneously using FP32 (single-precision) and FP16 (half-precision) floating-point numbers for computation, it can significantly improve training speed and reduce memory usage with almost no loss in model accuracy. This section details the Automatic Mixed Precision (AMP) technique in PyTorch.
Applicable versions:The code in this article is written using the PyTorch 1.6+
torch.cuda.ampAPI. PyTorch 2.4+ recommends usingtorch.amp.autocastandtorch.amp.GradScaler, and the usage is basically the same; differences will be noted in the article.
1. Fundamentals of Mixed Precision Training
1.1 Why Do We Need Mixed Precision
Deep learning model training involves a large number of matrix operations. Traditional FP32 (32-bit floating point) has high computational precision, but takes up more memory and is slower. FP16 (16-bit floating point) is faster and uses less memory, but its numerical range is smaller and it is prone to gradient underflow problems.
The core idea of mixed precision training is:Use FP32 for operations that require high precision, and FP16 for operations that don't require high precision.. This way, you can enjoy the speed advantage of FP16 while avoiding precision issues.
The following figure shows the bit layout differences of the three floating-point formats—Exponent bitsdetermine the numerical range.Mantissa bitsdetermine the precision:
1.2 Advantages of Mixed Precision
| Metric | Improvement | Description |
|---|---|---|
| Training speed | 2-3x improvement | Depends on GPU Tensor Core support |
| Memory usage | Reduced by about 50% | Activations and intermediate results stored in FP16 |
| Memory bandwidth | Reduced by about 50% | Smaller data size means less data transfer |
| Communication overhead | Reduced by about 50% | Gradient transfer volume halved in distributed training |
1.3 Tensor Core Acceleration Principle
NVIDIA's Tensor Core is a hardware unit dedicated to matrix operations. It can complete a 4×4 matrix multiply-accumulate operation (D = A × B + C) in a single clock cycle, which is the main source of FP16 training acceleration. Compared with ordinary CUDA cores that require multiple instructions to complete the same operation, Tensor Core compresses it into a single instruction.
GPUs that support Tensor Core include:
- Volta architecture(V100) — First-generation Tensor Core, only supports FP16
- Turing architecture(RTX 20 series) — Supports FP16 / INT8 / INT4
- Ampere architecture(RTX 30 series, A100) — Added BF16 / TF32 support
- Ada Lovelace architecture(RTX 40 series) — Added FP8 support
- Hopper architecture(H100) — Added FP8 Transformer Engine
Consumer-grade RTX GPUs also support Tensor Core; for example, RTX 3060 and above models can enjoy AMP acceleration.
2. PyTorch AMP Basic Usage
2.1 autocast and GradScaler
PyTorch's AMP API mainly includes two core components:
autocast: A context manager that automatically switches operations within its scope to FP16 (precision-sensitive operations automatically fall back to FP32)GradScaler: Dynamically adjusts the gradient scaling factor to amplify FP16 gradients and prevent underflow (only needed for FP16; usually not needed for BF16)
The following figure shows the data flow of one complete training step with AMP:
2.2 Basic Usage Example
Example
import torch.nn as nn
import torch.optim as optim
from torch.cuda.amp import autocast, GradScaler
# Recommended approach for PyTorch 2.4+:
# from torch.amp import autocast, GradScaler
# Check if CUDA is available
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
if torch.cuda.is_available():
print(f"GPU: {torch.cuda.get_device_name(0)}")
print(f"BF16 supported: {torch.cuda.is_bf16_supported()}")
# ── Model definition ──────────────────────────────────────
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Linear(128, 256),
nn.ReLU(),
nn.Linear(256, 256),
nn.ReLU(),
nn.Linear(256, 10)
)
def forward(self, x):
return self.net(x)
model = SimpleModel().to(device)
# Loss function and optimizer
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=1e-3)
# ── Key components for mixed precision training ──────────────────────────
# GradScaler: scales the loss to avoid FP16 gradient underflow
scaler = GradScaler()
# Training loop
def train_epoch_amp(model, loader, criterion, optimizer, scaler, device):
model.train()
total_loss = 0
correct = 0
total = 0
for inputs, labels in loader:
inputs = inputs.to(device, non_blocking=True)
labels = labels.to(device, non_blocking=True)
optimizer.zero_grad()
# ── Core: use the autocast context manager ──────
# Operations inside the autocast scope automatically use FP16
# Precision-sensitive operations (e.g., softmax, loss) automatically fall back to FP32
with autocast(device_type='cuda'):
outputs = model(inputs)
loss = criterion(outputs, labels)
# ── Use scaler for backward propagation ─────────────
# 1. Scale the loss (multiply by scale_factor)
# 2. Backward propagation (computed on the scaled gradients)
# 3. scaler.step internally unscales gradients and checks for Inf/NaN
scaler.scale(loss).backward()
# Update parameters
scaler.step(optimizer)
# Update scaler's scale factor
scaler.update()
# Statistics
total_loss += loss.item() * inputs.size(0)
_, predicted = outputs.max(1)
correct += predicted.eq(labels).sum().item()
total += labels.size(0)
return total_loss / total, correct / total
# Simulated data
train_loader = [
(torch.randn(32, 128), torch.randint(0, 10, (32,))) for _ in range(10)
]
# Start training
for epoch in range(3):
loss, acc = train_epoch_amp(
model, train_loader, criterion, optimizer, scaler, device
)
print(f"Epoch {epoch+1}: Loss={loss:.4f}, Acc={acc:.4f}")
print("Mixed precision training complete!")
API Migration Note:PyTorch 2.4 will mark
torch.cuda.amp.autocastas deprecated; it is recommended to usetorch.amp.autocast('cuda'). The usage is exactly the same; only the import path differs.
2.3 Choosing Between BF16 and FP16
Modern GPUs support two half-precision formats. Their core difference lies in the allocation strategy of exponent bits and mantissa bits:
| Format | Exponent bits | Mantissa bits | Numeric range | Precision | Use case |
|---|---|---|---|---|---|
| FP16 | 5 bits | 10 bits | ±65504 | Higher (~3.3 digits) | Requires compatibility with older GPUs (V100/RTX 20 series) |
| BF16 | 8 bits | 7 bits | ±3.4×10³⁸ | Lower (~2.4 digits) | For stability (A100/RTX 30 series+) |
Recommendation: If the hardware supports BF16,prefer BF16. Its numeric range is exactly the same as FP32, so overflow issues rarely occur during training, and GradScaler is not needed.
Example
# BF16 does not require GradScaler because its numeric range is the same as FP32
from torch.cuda.amp import autocast
# Method 1: Specify dtype in autocast
with autocast(device_type='cuda', dtype=torch.bfloat16):
outputs = model(inputs)
loss = criterion(outputs, labels)
# Method 2: Globally enable BF16 by default (if hardware supports it)
# torch.backends.cuda.matmul.allow_bf16_reduced_precision = True
# torch.backends.cudnn.allow_bf16_reduced_precision = True
# ── Check hardware support ─────────────────────────────
print(f"Hardware supports BF16: {torch.cuda.is_bf16_supported()}")
print(f"Current matmul allows BF16: {torch.backends.cuda.matmul.allow_bf16}")
print(f"Current cuDNN allows BF16: {torch.backends.cudnn.allow_bf16}")
# ── Complete BF16 training example (no GradScaler needed) ──────────
scaler_bf16 = None # BF16 does not need a scaler
for inputs, labels in train_loader:
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
with autocast(device_type='cuda', dtype=torch.bfloat16):
outputs = model(inputs)
loss = criterion(outputs, labels)
# Backpropagate directly, no scaler needed
loss.backward()
optimizer.step()
3. Advanced Tips and Optimization
3.1 Dynamic Loss Scaling
The core mechanism of GradScaler isdynamic loss scaling—it automatically adjusts the scale factor based on the training state, preventing gradient underflow while gradually reducing scaling when training stabilizes to minimize precision loss:
GradScaler automatically manages the above feedback loop; you only need to fine-tune its behavior via parameters:
Example
# Custom GradScaler parameters
scaler = GradScaler(
init_scale=2**16, # Initial scale factor, default 65536
growth_factor=2.0, # Scale factor growth multiplier, default 2.0
backoff_factor=0.5, # Scale factor fallback multiplier, default 0.5
growth_interval=2000, # Number of consecutive successful steps before growth
enabled=True # Whether enabled (can be toggled dynamically)
)
# Workflow:
# 1. Initial scale = 65536
# 2. If a step has Inf/NaN → scale × 0.5 (reduce), skip this step
# 3. If 2000 consecutive steps have no Inf/NaN → scale × 2.0 (increase)
# 4. Always find the maximum usable scale within a safe range
# View the current scale factor
print(f"Current scale factor: {scaler.get_scale()}")
# Determine whether scaler considers the last step successful
print(f"Recently successful: {scaler._found_inf.item() == 0}")
3.2 AMP with Gradient Accumulation
Gradient accumulation simulates a larger batch size under limited memory. When using AMP, note that the loss of each sub-batch must be divided by the number of accumulation steps; otherwise, the final gradient will be amplified:
Example
accumulation_steps = 8
scaler = GradScaler()
model.train()
for batch_idx, (inputs, labels) in enumerate(train_loader):
inputs, labels = inputs.to(device), labels.to(device)
with autocast(device_type='cuda'):
outputs = model(inputs)
# Key: divide each sub-batch's loss by the accumulation steps
# This way, the accumulated gradient is equivalent to the average gradient of a large batch
loss = criterion(outputs, labels) / accumulation_steps
# Accumulate the scaled gradients (note: do not clear here)
scaler.scale(loss).backward()
# Update parameters every accumulation_steps batches
if (batch_idx + 1) % accumulation_steps == 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
# ── Handling gradients at the end that are less than a full accumulation cycle ──────────
remainder = len(train_loader) % accumulation_steps
if remainder != 0:
# The accumulated gradients need to be re-scaled according to the actual number of steps
# Simple approach: still execute step; gradients will be slightly larger but the impact is usually minimal
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
3.3 AMP in Validation and Inference
AMP is also recommended for validation and inference. Since backpropagation and gradient scaling are not needed, the code is simpler:
Example
@torch.no_grad()
def inference_amp(model, inputs, device):
model.eval()
# Inference uses FP16, no GradScaler needed
with autocast(device_type='cuda'):
outputs = model(inputs.to(device))
return outputs
# ── inference_mode is faster than no_grad ──────────────
# inference_mode disables more tracing overhead, suitable for pure inference
@torch.inference_mode()
def inference_fast(model, inputs, device):
model.eval()
with autocast(device_type='cuda'):
outputs = model(inputs.to(device))
return outputs
# ── Batch inference example ───────────────────────────────
def batch_inference(model, dataloader, device):
model.eval()
all_outputs = []
with torch.inference_mode():
for inputs in dataloader:
with autocast(device_type='cuda'):
outputs = model(inputs.to(device))
all_outputs.append(outputs.cpu())
return torch.cat(all_outputs, dim=0)
4. Common Issues and Solutions
4.1 Unstable Training
| Problem | Cause | Solution |
|---|---|---|
| Loss becomes NaN / Inf | Gradient overflow (scale too large or learning rate too high) | Reduceinit_scale, add gradient clipping, lower learning rate |
| Loss does not decrease / oscillates | Gradient underflow (scale too small, gradients truncated to 0) | Increaseinit_scale, check gradient statistics, try BF16 |
| Significant precision drop | Certain layers (e.g., LayerNorm, Softmax) are sensitive to FP16 | Manually keep these layers in FP32 (see 4.2) |
| Instability in late training | Scale mismatch due to changes in model parameter value ranges | Increase appropriatelygrowth_interval, making scale adjustments more conservative |
4.2 Manual Precision Control
Some operations suffer significant precision loss under FP16 and need to be manually specified to use FP32. PyTorch's autocast has built-in protection for these operations, but custom operations may need manual handling:
Example
class FP32Loss(nn.Module):
"""Loss function wrapper that forces FP32 computation"""
def __init__(self, base_criterion):
super().__init__()
self.base_criterion = base_criterion
def forward(self, input, target):
# Explicitly convert to FP32, outside autocast's influence
return self.base_criterion(input.float(), target.float())
# ── Method 2: Locally disable within autocast ──────────────
with autocast(device_type='cuda'):
outputs = model(inputs)
# For precision-sensitive operations, temporarily disable autocast
with autocast(enabled=False):
loss = criterion(outputs.float(), labels)
# ── Method 3: Keep specific layers in FP32 in the model ────────
class CustomModel(nn.Module):
def __init__(self):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 64, 3, padding=1),
nn.BatchNorm2d(64), # BN is precision-sensitive
nn.ReLU(),
nn.Conv2d(64, 128, 3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(),
)
self.classifier = nn.Linear(128, 10)
def forward(self, x):
x = self.features[0](x) # Conv: FP16 (controlled by autocast)
x = self.features[1](x.float()) # BN: force FP32
x = self.features[2](x) # ReLU: FP16
x = self.features[3](x)
x = self.features[4](x.float()) # BN: force FP32
x = self.features[5](x)
x = x.mean(dim=[2, 3]) # Global Average Pooling
x = self.classifier(x)
return x
Tip:PyTorch's autocast automatically keeps the following operations in FP32: softmax, log_softmax, cross_entropy, layer_norm, batch_norm, etc. Usually only custom operators need manual handling.
4.3 Gradient Clipping and AMP
Gradient clipping must bescaler.step()before, andscaler.unscale_()after execution. This is because clipping needs to operate ontrue gradient valuesrather than the scaled gradient:
Example
for inputs, labels in train_loader:
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
with autocast(device_type='cuda'):
outputs = model(inputs)
loss = criterion(outputs, labels)
# Step 1: Scale loss and backpropagate
scaler.scale(loss).backward()
# Step 2: Unscale gradients (restore gradients from scale to true values)
# Must be called before clipping and step
scaler.unscale_(optimizer)
# Step 3: Clip the true gradients
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
# Step 4: Update parameters (internally checks whether gradients contain Inf/NaN)
scaler.step(optimizer)
# Step 5: Update the scale factor
scaler.update()
# ── Common mistake: clipping without unscale_ ──────────────
# Wrong! Clipping is applied to the amplified gradients, and the threshold is affected by scale
# scaler.scale(loss).backward()
# torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) # ← Wrong
# scaler.step(optimizer)
5. Performance Comparison and Best Practices
5.1 Performance Benchmark
Example
def benchmark_training(model, train_loader, device, use_amp=True,
dtype=torch.float16, num_iterations=100):
"""Compare the training speed of FP32 and AMP"""
model = model.to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=1e-3)
scaler = GradScaler(enabled=use_amp and dtype == torch.float16)
# ── Warm-up (first 10 steps not timed) ────────────────────
for i, (inputs, labels) in enumerate(train_loader):
if i >= 10:
break
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
if use_amp:
with autocast(device_type='cuda', dtype=dtype):
outputs = model(inputs)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
else:
outputs = model(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
# ── Formal test ─────────────────────────────────
torch.cuda.synchronize()
start = time.time()
for i, (inputs, labels) in enumerate(train_loader):
if i >= num_iterations:
break
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
if use_amp:
with autocast(device_type='cuda', dtype=dtype):
outputs = model(inputs)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
else:
outputs = model(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
torch.cuda.synchronize()
elapsed = time.time() - start
return elapsed / num_iterations
# ── Run comparison ─────────────────────────────────────
model = SimpleModel()
fp32_time = benchmark_training(model, train_loader, device, use_amp=False)
fp16_time = benchmark_training(model, train_loader, device, use_amp=True, dtype=torch.float16)
print(f"FP32 average time: {fp32_time*1000:.2f} ms/batch")
print(f"FP16 average time: {fp16_time*1000:.2f} ms/batch")
print(f"Speedup: {fp32_time / fp16_time:.2f}x")
# If BF16 is supported, test it as well
if torch.cuda.is_bf16_supported():
bf16_time = benchmark_training(
model, train_loader, device, use_amp=True, dtype=torch.bfloat16
)
print(f"BF16 average time: {bf16_time*1000:.2f} ms/batch")
print(f"BF16 speedup: {fp32_time / bf16_time:.2f}x")
5.2 Memory Usage Comparison
Example
"""Compare GPU memory usage of FP32 and AMP"""
def reset():
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
# ── FP32 memory test ────────────────────────────
reset()
model_fp32 = model_class().to(device)
optimizer_fp32 = optim.Adam(model_fp32.parameters())
for inputs, labels in list(train_loader)[:5]:
inputs, labels = inputs.to(device), labels.to(device)
optimizer_fp32.zero_grad()
outputs = model_fp32(inputs)
loss = nn.CrossEntropyLoss()(outputs, labels)
loss.backward()
optimizer_fp32.step()
fp32_peak = torch.cuda.max_memory_allocated() / 1024**2
# ── AMP memory test ─────────────────────────────
reset()
model_amp = model_class().to(device)
optimizer_amp = optim.Adam(model_amp.parameters())
scaler = GradScaler()
for inputs, labels in list(train_loader)[:5]:
inputs, labels = inputs.to(device), labels.to(device)
optimizer_amp.zero_grad()
with autocast(device_type='cuda'):
outputs = model_amp(inputs)
loss = nn.CrossEntropyLoss()(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer_amp)
scaler.update()
amp_peak = torch.cuda.max_memory_allocated() / 1024**2
print(f"FP32 peak memory: {fp32_peak:.1f} MB")
print(f"AMP peak memory: {amp_peak:.1f} MB")
print(f"Memory savings: {(fp32_peak - amp_peak) / fp32_peak * 100:.1f}%")
compare_memory(SimpleModel, train_loader, device)
5.3 Best Practices Summary
| Scenario | Recommended Configuration | Reason |
|---|---|---|
| A100 / H100 / RTX 40 series | BF16, no GradScaler needed | Same numeric range as FP32, most stable |
| V100 / RTX 20 series | FP16 + GradScaler | Hardware does not support BF16; a scaler is needed to prevent overflow |
| RTX 30 series | BF16 preferred, FP16 as alternative | Ampere architecture supports BF16 |
| Large model training (limited GPU memory) | AMP + gradient accumulation + gradient checkpointing | Combining all three maximizes GPU memory utilization |
| Inference deployment | autocast + inference_mode | No scaler needed, fastest inference |
PyTorch 2.0+ has integrated AMP into
torch.compileit, which automatically applies mixed precision optimization. When usingtorch.compile(model)it, the compiler automatically determines which operations are suitable for FP16.
6. Combining with Other Optimization Techniques
6.1 AMP + torch.compile
Example
model = model.to(device)
# Method 1: Compile first, then train with AMP
# torch.compile automatically fuses operators and optimizes memory access
model_compiled = torch.compile(model, mode="reduce-overhead")
scaler = GradScaler()
for inputs, labels in train_loader:
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
with autocast(device_type='cuda'):
outputs = model_compiled(inputs)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
# Method 2: Enable TF32 (extra acceleration on Ampere+ architectures)
# TF32 uses 19-bit precision, with speed close to FP16 and accuracy close to FP32
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.backends.cudnn.benchmark = True # Accelerates convolution when input size is fixed
6.2 AMP + Distributed Training
Example
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
def setup(rank, world_size):
dist.init_process_group("nccl", rank=rank, world_size=world_size)
torch.cuda.set_device(rank)
def train_ddp_amp(rank, world_size):
setup(rank, world_size)
device = torch.device(f"cuda:{rank}")
model = SimpleModel().to(device)
model = DDP(model, device_ids=[rank])
optimizer = optim.Adam(model.parameters(), lr=1e-3)
scaler = GradScaler()
for inputs, labels in train_loader:
inputs = inputs.to(device, non_blocking=True)
labels = labels.to(device, non_blocking=True)
optimizer.zero_grad()
with autocast(device_type='cuda'):
outputs = model(inputs)
loss = criterion(outputs, labels)
# DDP automatically synchronizes gradients during backward
scaler.scale(loss).backward()
# Gradient clipping (must unscale first)
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
dist.destroy_process_group()
Summary
Mixed precision training is astandard practice。
Key points:
- Choose precision format: prefer BF16 when hardware supports it, otherwise use FP16 + GradScaler
- Understand autocast: it automatically manages precision switching, and in most cases no manual intervention is needed
- Understand GradScaler: required only for FP16, prevents gradient underflow through dynamic scaling
- Pay attention to clipping order: unscale → clip → step → update, the order cannot be reversed
- Leverage combined optimization: AMP can be seamlessly combined with torch.compile, gradient accumulation, and distributed training