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:

Comparison of floating-point format bit layouts Sign bit Exponent bits Mantissa bits FP32 32 bits S 1 Exponent 8 bits Mantissa 23 bits Range: ±3.4×10³⁸ Precision: ~7 significant digits FP16 16 bits S Exp 5 bits Mantissa 10 bits Range: ±65504 Precision: ~3.3 significant digits ⚠ Fewer exponent bits, prone to overflow BF16 16 bits S Exponent 8 bits (same as FP32) Mantissa 7 bits Range: ±3.4×10³⁸ Precision: ~2.4 significant digits ✓ Same range as FP32 BF16 retains FP32's exponent bits → same numerical range → more stable training, usually no GradScaler needed

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:

AMP single-step training flow Input data FP32 autocast Forward propagation Automatically select FP16/FP32 Compute loss FP32 GradScaler Scaling + backward propagation loss × scale_factor Unscale gradients + gradient clipping grad / scale_factor Optimizer update FP32 weights Update Scaler Adjust scale autocast control — automatically selects precision GradScaler control — prevents gradient underflow FP32 weight update — maintains precision Key: weights are always stored and updated in FP32, and are only temporarily converted to FP16 during computation

2.2 Basic Usage Example

Example

import torch
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 marktorch.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 approach (PyTorch 1.10+) ────────────────────
# 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:

Dynamic Loss Scaling Scale × Loss Backward pass Gradients contain Inf / NaN? Yes Reduce scale × backoff (default 0.5) Skip this update no Unscale + update parameters Consecutive successes +1 Consecutive successes ≥ growth_interval? (default 2000 steps) Increase scale × growth growth_factor defaults to 2.0 · backoff_factor defaults to 0.5 · Dynamically balancing precision and stability

GradScaler automatically manages the above feedback loop; you only need to fine-tune its behavior via parameters:

Example

from torch.cuda.amp import GradScaler

# 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

# ── Gradient accumulation + mixed precision ────────────────
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

# ── Using autocast for inference ──────────────────────────
@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

# ── Method 1: Wrap a loss function that forces FP32 ────────
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

# ── Correct gradient clipping order ──────────────────────────
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

import time

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

def compare_memory(model_class, train_loader, device):
    """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 intotorch.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

# ── Combined use with PyTorch 2.0+ ───────────────────────
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

# ── Distributed training + mixed precision ────────────────────────
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
Other Extensions