10.混合精度训练
第 10 章 — 混合精度训练(AMP)
第 9 章用数据增强解决了"模型记不住"的问题,但训练还有另一道更硬的坎:显存不够、训练太慢。模型一大、batch 一大,FP32 的权重、激活、梯度瞬间撑爆显存,
CUDA out of memory随手就来;而现代 GPU 对 FP16 有专门的"张量核心",算同一份数据的速度接近翻倍。本章用 PyTorch 自带的自动混合精度(AMP),不改一行模型结构,就同时拿到"省显存 + 提速"。
一、本章要解决的问题
- 之前(第 9 章):代码能跑、能防过拟合,但只用了 FP32——权重、激活、梯度每个数占 4 字节,显存压力大;
batch_size想调大一点就报CUDA out of memory。 - 现在:引入 AMP(Automatic Mixed Precision,自动混合精度),让大部分计算用占一半字节的 FP16 跑,关键环节保持 FP32,从而省显存、提速度。
- 判断标准:GPU 上
use_amp: true与false各跑一次,前者显存占用明显更低、单 epoch 耗时更短,而 val_acc / test_acc 几乎不变;并且 CPU 上、关闭开关时,行为与第 9 章完全一致。
二、核心概念速览
下面 7 个概念是本章的全部"生词"。先花 5 分钟读完,再看代码会轻松很多。
1. FP32 / FP16 精度与显存的关系:两种"记账本"
FP32(32 位浮点数)占 4 字节、小数点后约 7 位有效数字;FP16(16 位浮点数)只占 2 字节、约 3 位有效数字。把它想成记账本:FP32 是一本能精确记到"分"的厚账本,每个数字占 4 格;FP16 是只记到"元"的简装账本,每个数字占 2 格。同样的本子,用 FP16 能记两倍的账——这就是省显存的来源:显存就是"本子的页数",数据从 4 字节变 2 字节,同样显存能装下更大的模型、更大的 batch,代价是数字记得更粗糙。
2. 为什么全 FP16 会出问题:梯度下溢 + 动态范围太小
FP16 能表示的数字范围很小:最大约 65504,最小正数约 6e-8;而 FP32 最小能到 1e-38。反向传播时梯度会被层层相乘越变越小,经常小于 6e-8——这时 FP16 直接把它"记成 0",参数再也不更新,训练原地踏步,这就是梯度下溢。类比:FP32 像一台能称出 0.0000001 克的精密天平,FP16 像只精确到"克"的厨房秤,一粒灰尘放上去读数直接是 0。所以不能"全 FP16",得让损失、梯度、归一化这些关键环节保持 FP32,其余用 FP16——"混合"两个字就是这么来的。
3. autocast:自动决定谁用 FP16、谁用 FP32 的"自动挡"
torch.amp.autocast 是一个上下文管理器,用 with 把前向传播包起来。它像一个自动挡变速箱:根据"当前是哪种操作"自动换挡——卷积、矩阵乘这类计算量大、对精度不敏感的操作切到 FP16;BatchNorm、softmax、损失函数这类对精度敏感的操作保持 FP32。你不需要手动把某个张量转成 FP16,模型代码一行不改,它还会自动把输入转成合适的类型。
4. GradScaler:给梯度配一副"放大镜"
GradScaler 专门治第 2 点的"梯度下溢":反向传播前先把 loss 放大一个倍数(默认 2^16 = 65536),梯度也跟着同比例放大,就不会掉进 FP16 记不出的"看不见区";真正更新参数前再缩回去,等价于没放大。它还有第二重保险:如果某一步放大后梯度出现了 inf 或 nan(说明这一步数值爆了),就跳过这次参数更新,并把放大倍数降一半,防止一错再错。类比:照片太小看不清细节,先用放大镜看,看完记得还原成真实大小。
5. scaler 三连:顺序不能乱的三道工序
每步训练固定的三连:scaler.scale(loss).backward() → scaler.step(optimizer) → scaler.update()。像工厂流水线:第一道把 loss 放大后反向传播(得到放大的梯度);第二道 step(optimizer) 内部先检查有没有 inf/nan,没有才把梯度缩回去并更新参数,有就跳过;第三道 update() 放在最后,根据"这一步有没有溢出"动态调整放大倍数(溢出就减半,稳定就慢慢回升)。顺序写反(比如先 optimizer.step() 再 scaler.step())等于绕过了放大和检查,前面的功夫全白费。
6. checkpoint 里为什么要存 scaler_state:训练日志不能丢
scaler 不是固定放大 65536 就完事,它的缩放因子会随训练动态变化(溢出减半、稳定回升)。如果 checkpoint 只存模型和优化器状态、恢复训练时丢掉 scaler_state,放大倍数会从默认值重新起步——但之前可能已经降到了 1024 或涨回了 131072,数值突变会让恢复后的训练曲线出现"跳变"。类比健身教练的训练日志:停练两周回来,不能凭印象里的重量直接上,得翻日志从上次的实际重量继续。所以本章 checkpoint 里多存一个 scaler_state,--resume 时用 load_state_dict 恢复,训练无缝衔接。
7. enabled 开关:一按就变回第 9 章
代码里的 use_amp = cfg.train.use_amp and device.type == "cuda",然后所有 AMP 组件都带 enabled=use_amp:只要配置里关了 AMP,或者机器是 CPU,autocast 不切精度、GradScaler 不放大,行为与第 9 章完全一致。这就像自动挡上的"手动模式"开关——关掉就是普通驾驶。好处是:一套代码在 GPU / CPU、开 / 关 AMP 之间随便切换,不会因为换台机器就报错或悄悄变了行为。
三、解决思路
- 不加新依赖、不改模型:PyTorch 内置
torch.amp,只需要在训练循环里加"一个上下文 + 一个 scaler",模型结构一行不动。 - 改动点只有三个:训练循环的前向/反向/更新换成 autocast + scaler 三连;验证循环也包一层 autocast(只加速、不放大);checkpoint 里新增并恢复
scaler_state。 - 全开关控制:
use_amp由"配置是否开启"和"是否 CUDA"共同决定,CPU 或关闭时零成本退化到第 9 章行为,不删代码、不写两套分支。 - 用数据说话:
nvidia-smi观察显存、记录每 epoch 耗时,对比开/关两组——省多少、快多少,让数字回答。
trade-off:AMP 省的是显存和算力,代价是引入了"溢出"这种新风险(需要 scaler 兜底),而且收益只在有张量核心的 GPU 上兑现——CPU 上开 AMP 没收益,还可能更慢。对 CIFAR-10 + SimpleCNN 这种小模型,显存收益不夸张,但原理是通用的:换大模型、大 batch 时收益会非常直观。
四、代码变更
相对第 9 章的改动点:
# config/config.yaml:train 段新增 use_amp 字段(第 9 章没有)
train:
epochs: 30
lr: 0.001
lr_scheduler: step
lr_step_size: 15
lr_gamma: 0.1
patience: 7
+ use_amp: true # 混合精度开关(仅 CUDA 生效)
# train.py:调度器定义之后新增 AMP 初始化
if cfg.train.lr_scheduler == "step":
scheduler = optim.lr_scheduler.StepLR(
optimizer, step_size=cfg.train.lr_step_size, gamma=cfg.train.lr_gamma)
elif cfg.train.lr_scheduler == "cosine":
scheduler = optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=cfg.train.epochs)
else:
scheduler = None
+ # 混合精度:仅 CUDA 上启用;enabled=False 时所有调用退化为普通 FP32 流程
+ use_amp = cfg.train.use_amp and device.type == "cuda"
+ scaler = torch.amp.GradScaler("cuda", enabled=use_amp)
+ if use_amp:
+ logger.info("AMP 已启用(FP16 混合精度)")
# train.py:训练循环里,前向传播包进 autocast,反向/更新换成 scaler 三连
def train_one_epoch(model, loader, criterion, optimizer, device):
model.train()
total_loss, correct, total = 0.0, 0, 0
for images, labels in tqdm(loader, desc="Train", leave=False):
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
- outputs = model(images)
- loss = criterion(outputs, labels)
- loss.backward()
- optimizer.step()
+ # autocast:上下文内自动把卷积/矩阵乘等算子切到 FP16
+ with torch.amp.autocast(device_type=device.type, enabled=use_amp):
+ outputs = model(images)
+ loss = criterion(outputs, labels)
+ # scaler 三连:放大 loss -> 反传 -> 缩回更新 -> 调整缩放因子
+ scaler.scale(loss).backward()
+ scaler.step(optimizer)
+ scaler.update()
# train.py:验证循环同样包一层 autocast(只加速,不经过 scaler)
def validate(model, loader, criterion, device):
model.eval()
total_loss, correct, total = 0.0, 0, 0
with torch.no_grad():
- for images, labels in tqdm(loader, desc="Val", leave=False):
+ # 验证阶段也可以 autocast 加速,但不经过 scaler
+ with torch.amp.autocast(device_type=device.type, enabled=use_amp):
+ for images, labels in tqdm(loader, desc="Val", leave=False):
...
# train.py:恢复训练时加载 scaler_state,checkpoint 存档时也保存它
if args.resume:
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model_state"])
optimizer.load_state_dict(ckpt["optimizer_state"])
if scheduler is not None and ckpt.get("scheduler_state") is not None:
scheduler.load_state_dict(ckpt["scheduler_state"])
+ # AMP 的 scaler 状态也要恢复,否则放大倍数从默认重新起步
+ if ckpt.get("scaler_state") is not None:
+ scaler.load_state_dict(ckpt["scaler_state"])
...
state = {
"model_state": model.state_dict(),
"optimizer_state": optimizer.state_dict(),
"scheduler_state": scheduler.state_dict() if scheduler else None,
+ "scaler_state": scaler.state_dict(), # AMP 状态一并存档
"epoch": epoch,
"best_acc": best_acc,
"metrics": dict(metrics),
}其余部分(Config 类、EarlyStopping、数据增强、调度器、日志)与第 9 章完全一致,改动只集中在"训练循环怎么跑"这一条线上。
五、完整代码
创建 config/config.yaml:
# CIFAR-10 图像分类实验配置
device: auto
data:
root: ./data
batch_size: 64
num_workers: 2
val_ratio: 0.1
augmentation:
random_crop_padding: 4
horizontal_flip: true
color_jitter: [0.2, 0.2, 0.2, 0.1]
model:
num_classes: 10
train:
epochs: 30
lr: 0.001
lr_scheduler: step
lr_step_size: 15
lr_gamma: 0.1
patience: 7
use_amp: true
checkpoint:
dir: ./checkpoints
log:
dir: ./runs/cifar10创建 train.py:
import os
import logging
import argparse
import yaml
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, random_split
from torch.utils.tensorboard import SummaryWriter
from torchvision import datasets, transforms
from collections import defaultdict
from tqdm import tqdm
from rich.logging import RichHandler
def parse_args():
parser = argparse.ArgumentParser(description="CIFAR-10 图像分类训练")
parser.add_argument("--config", type=str, default="config/config.yaml",
help="YAML 配置文件路径")
parser.add_argument("--lr", type=float, default=None, help="覆盖配置文件中的 train.lr")
parser.add_argument("--epochs", type=int, default=None, help="覆盖配置文件中的 train.epochs")
parser.add_argument("--batch-size", type=int, default=None,
help="覆盖配置文件中的 data.batch_size")
parser.add_argument("--resume", type=str, default=None,
help="从指定 checkpoint 恢复训练")
return parser.parse_args()
args = parse_args()
class Config(dict):
def __getattr__(self, key):
try:
return self[key]
except KeyError as e:
raise AttributeError(key) from e
def _to_config(obj):
if isinstance(obj, dict):
return Config({k: _to_config(v) for k, v in obj.items()})
if isinstance(obj, list):
return [_to_config(v) for v in obj]
return obj
def load_config(path):
with open(path, "r", encoding="utf-8") as f:
return _to_config(yaml.safe_load(f))
cfg = load_config(args.config)
if args.lr is not None:
cfg.train.lr = args.lr
if args.epochs is not None:
cfg.train.epochs = args.epochs
if args.batch_size is not None:
cfg.data.batch_size = args.batch_size
def setup_logger(log_dir):
os.makedirs(log_dir, exist_ok=True)
logger = logging.getLogger("train")
logger.handlers.clear()
logger.setLevel(logging.INFO)
console = RichHandler(rich_tracebacks=True, markup=True)
console.setLevel(logging.INFO)
logger.addHandler(console)
file_handler = logging.FileHandler(os.path.join(log_dir, "train.log"),
encoding="utf-8")
file_handler.setFormatter(logging.Formatter(
"%(asctime)s | %(levelname)s | %(message)s", datefmt="%Y-%m-%d %H:%M:%S"))
logger.addHandler(file_handler)
return logger
logger = setup_logger(cfg.log.dir)
writer = SummaryWriter(log_dir=cfg.log.dir)
if cfg.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(cfg.device)
logger.info(f"Using device: {device}")
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(64, 128, kernel_size=3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
)
self.classifier = nn.Sequential(
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(128, num_classes),
)
def forward(self, x):
return self.classifier(self.features(x))
def build_transforms(cfg, train=False):
aug = cfg.data.augmentation
transforms_list = []
if train:
if aug.random_crop_padding > 0:
transforms_list.append(
transforms.RandomCrop(32, padding=aug.random_crop_padding))
if aug.horizontal_flip:
transforms_list.append(transforms.RandomHorizontalFlip())
if aug.color_jitter is not None:
transforms_list.append(transforms.ColorJitter(*aug.color_jitter))
transforms_list += [
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
]
return transforms.Compose(transforms_list)
train_dataset = datasets.CIFAR10(root=cfg.data.root, train=True, download=True,
transform=build_transforms(cfg, train=True))
test_dataset = datasets.CIFAR10(root=cfg.data.root, train=False, download=True,
transform=build_transforms(cfg, train=False))
val_size = int(len(train_dataset) * cfg.data.val_ratio)
train_dataset, val_dataset = random_split(
train_dataset, [len(train_dataset) - val_size, val_size]
)
train_loader = DataLoader(train_dataset, batch_size=cfg.data.batch_size,
shuffle=True, num_workers=cfg.data.num_workers)
val_loader = DataLoader(val_dataset, batch_size=cfg.data.batch_size,
shuffle=False, num_workers=cfg.data.num_workers)
test_loader = DataLoader(test_dataset, batch_size=cfg.data.batch_size,
shuffle=False, num_workers=cfg.data.num_workers)
model = SimpleCNN(num_classes=cfg.model.num_classes).to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=cfg.train.lr)
if cfg.train.lr_scheduler == "step":
scheduler = optim.lr_scheduler.StepLR(
optimizer, step_size=cfg.train.lr_step_size, gamma=cfg.train.lr_gamma)
elif cfg.train.lr_scheduler == "cosine":
scheduler = optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=cfg.train.epochs)
else:
scheduler = None
# 混合精度:仅 CUDA 上启用;enabled=False 时所有调用退化为普通 FP32 流程
use_amp = cfg.train.use_amp and device.type == "cuda"
scaler = torch.amp.GradScaler("cuda", enabled=use_amp)
if use_amp:
logger.info("AMP 已启用(FP16 混合精度)")
def train_one_epoch(model, loader, criterion, optimizer, device):
model.train()
total_loss, correct, total = 0.0, 0, 0
for images, labels in tqdm(loader, desc="Train", leave=False):
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
# autocast:上下文内自动把卷积/矩阵乘等算子切到 FP16
with torch.amp.autocast(device_type=device.type, enabled=use_amp):
outputs = model(images)
loss = criterion(outputs, labels)
# scaler 三连:放大 loss -> 反传 -> 缩回更新 -> 调整缩放因子
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def validate(model, loader, criterion, device):
model.eval()
total_loss, correct, total = 0.0, 0, 0
with torch.no_grad():
# 验证阶段也可以 autocast 加速,但不经过 scaler
with torch.amp.autocast(device_type=device.type, enabled=use_amp):
for images, labels in tqdm(loader, desc="Val", leave=False):
images, labels = images.to(device), labels.to(device)
outputs = model(images)
loss = criterion(outputs, labels)
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def save_checkpoint(state, path):
os.makedirs(os.path.dirname(path), exist_ok=True)
torch.save(state, path)
logger.info(f"[Checkpoint] 已保存到 {path}")
class EarlyStopping:
def __init__(self, patience=7, delta=0.0, verbose=False):
self.patience = patience
self.delta = delta
self.verbose = verbose
self.best_score = None
self.counter = 0
self.early_stop = False
def __call__(self, val_acc):
score = val_acc
if self.best_score is None:
self.best_score = score
elif score < self.best_score + self.delta:
self.counter += 1
if self.verbose:
logger.info(f"[EarlyStopping] 连续 {self.counter}/{self.patience} 轮未提升")
if self.counter >= self.patience:
self.early_stop = True
else:
self.best_score = score
self.counter = 0
metrics = defaultdict(list)
start_epoch, best_acc = 1, 0.0
if args.resume:
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model_state"])
optimizer.load_state_dict(ckpt["optimizer_state"])
if scheduler is not None and ckpt.get("scheduler_state") is not None:
scheduler.load_state_dict(ckpt["scheduler_state"])
# AMP 的 scaler 状态也要恢复,否则放大倍数从默认重新起步
if ckpt.get("scaler_state") is not None:
scaler.load_state_dict(ckpt["scaler_state"])
start_epoch = ckpt["epoch"] + 1
best_acc = ckpt.get("best_acc", 0.0)
metrics = defaultdict(list, ckpt.get("metrics", {}))
logger.info(f"[Resume] 从 epoch {ckpt['epoch']} 恢复,历史最优 val_acc={best_acc:.4f}")
early_stopping = EarlyStopping(patience=cfg.train.patience, verbose=True)
logger.info("=" * 60)
logger.info("最终生效配置:")
logger.info(yaml.safe_dump(cfg, allow_unicode=True, sort_keys=False))
logger.info("=" * 60)
for epoch in range(start_epoch, cfg.train.epochs + 1):
train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, device)
val_loss, val_acc = validate(model, val_loader, criterion, device)
if scheduler is not None:
scheduler.step()
current_lr = optimizer.param_groups[0]["lr"]
metrics["train_loss"].append(train_loss)
metrics["train_acc"].append(train_acc)
metrics["val_loss"].append(val_loss)
metrics["val_acc"].append(val_acc)
logger.info(f"Epoch {epoch:02d}/{cfg.train.epochs} | "
f"Train loss {train_loss:.4f} acc {train_acc:.4f} | "
f"Val loss {val_loss:.4f} acc {val_acc:.4f} | "
f"lr {current_lr:.2e}")
writer.add_scalar("train/loss", train_loss, epoch)
writer.add_scalar("train/acc", train_acc, epoch)
writer.add_scalar("val/loss", val_loss, epoch)
writer.add_scalar("val/acc", val_acc, epoch)
writer.add_scalar("lr", current_lr, epoch)
state = {
"model_state": model.state_dict(),
"optimizer_state": optimizer.state_dict(),
"scheduler_state": scheduler.state_dict() if scheduler else None,
"scaler_state": scaler.state_dict(),
"epoch": epoch,
"best_acc": best_acc,
"metrics": dict(metrics),
}
save_checkpoint(state, os.path.join(cfg.checkpoint.dir, "last.pth"))
if val_acc > best_acc:
best_acc = val_acc
state["best_acc"] = best_acc
save_checkpoint(state, os.path.join(cfg.checkpoint.dir, "best.pth"))
early_stopping(val_acc)
if early_stopping.early_stop:
logger.info(f"[EarlyStopping] 连续 {cfg.train.patience} 轮无提升,在 epoch {epoch} 停止")
break
test_loss, test_acc = validate(model, test_loader, criterion, device)
logger.info(f"[Test] loss {test_loss:.4f} acc {test_acc:.4f}")
writer.close()AMP 新增代码逐行解读(本章全部改动就是下面这几处):
use_amp = cfg.train.use_amp and device.type == "cuda":配置开关和硬件开关用and焊在一起。配置里写了true,但机器是 CPU,结果还是False——CPU 上没有张量核心,开了也没收益。scaler = torch.amp.GradScaler("cuda", enabled=use_amp):创建缩放器。enabled=False时它是个"透明人":scale(loss)原样返回 loss、step(optimizer)直接调用optimizer.step()、update()什么都不做。所以 CPU 上放心写三连,不会报错,行为就是普通的 FP32 训练。- 训练循环里
autocast只包住"前向 + 算 loss"两行:outputs = model(images)在上下文里自动决定每个算子用 FP16 还是 FP32;反向传播、指标统计都在上下文外面,不受影响。 scaler.scale(loss).backward():等价于(loss * 65536).backward(),让梯度进入 FP16 能表示的范围;scaler.step(optimizer)内部先检查有没有inf/nan,安全才更新参数;scaler.update()最后动态调整放大倍数。三连顺序就是"放大 → 检查更新 → 调倍数",缺一不可。- 验证循环只包
autocast、不碰 scaler:推理阶段没有"梯度下溢"问题,只需享受 FP16 的加速。 - 恢复训练:
if ckpt.get("scaler_state") is not None:用get而不是直接索引,是为了兼容第 9 章存下来的旧 checkpoint(里面没有这个字段,直接索引会KeyError)。 - checkpoint 里
"scaler_state": scaler.state_dict():把当前放大倍数连同模型、优化器、调度器一起存档,恢复时无缝衔接。
运行与观察:
# 安装新增依赖(第 9 章已装过可跳过)
pip install tqdm rich
# GPU 上:开启 AMP 训练
python train.py
# 另开一个终端,训练期间反复执行,观察显存
nvidia-sminvidia-smi 输出的 Memory-Usage 一列就是显存占用(MiB 为单位)。等第一个 epoch 跑起来、显存稳定后再记录数值,取稳定期的峰值做对比。
对比实验:use_amp true vs false
# 实验 A:开启 AMP(默认配置,config/config.yaml 里 use_amp: true)
python train.py
# 实验 B:关闭 AMP
# 把 config/config.yaml 里的 use_amp 改成 false,再跑一次
python train.py两组实验(其余配置完全相同)的对比要点:
| 观察项 | use_amp: false | use_amp: true |
|---|---|---|
| 显存峰值(nvidia-smi) | 基准(约 100%) | 约 50%~70%,大模型更明显 |
| 单 epoch 耗时 | 基准 | 有张量核心的 GPU 上明显下降 |
| val_acc / test_acc | 基准 | 几乎不变 |
读这张表的三个要点:
- 看显存:FP16 数据占 2 字节、FP32 占 4 字节,权重、激活、梯度的存储直接减半,同样显存能塞更大的模型或更大的 batch。
- 看耗时:现代 GPU(V100/A100 及 RTX 20/30/40 系)对 FP16 有专门的张量核心,卷积、矩阵乘这类主力算子在 FP16 下吞吐接近翻倍。CPU 上没有张量核心,这个收益看不到。
- 看精度:autocast 只把"安全"的算子切到 FP16,scaler 又兜住了梯度下溢,所以准确率几乎不掉——这正是"混合"设计的精妙之处。
如果启动日志里看到
AMP 已启用(FP16 混合精度),说明use_amp生效;CPU 上不会出现这行日志,属于正常现象。
六、本章小结
- 学到了什么
- AMP =
autocast(自动给每个算子选精度)+GradScaler(防梯度下溢 + 溢出时跳过更新),模型结构一行不用改。 - scaler 三连是固定顺序:
scale(loss).backward()→step(optimizer)→update(),顺序不能乱。 - FP16 省显存的原理:每个数从 4 字节变 2 字节,存储减半;但全 FP16 会梯度下溢,所以关键环节(loss、梯度、归一化)必须保持 FP32。
- checkpoint 要额外存
scaler_state,恢复训练才能无缝衔接;用get兼容旧 checkpoint。 enabled=use_amp让 CPU / 关闭开关时行为完全退化为第 9 章,一套代码通吃所有环境。
- AMP =
- 常见坑
- 三连顺序写错或漏调
update():放大倍数永不调整,scaler 白配;update()必须在每次 step 后调用。 - 在 CPU 上写死
torch.amp.autocast("cuda"):会直接报错。统一用enabled=use_amp+device_type=device.type,别在 CPU 分支里删代码。 - 手动把数据
.half()全转 FP16:那是"全 FP16",不是 AMP,梯度会下溢成 0;AMP 靠 autocast 自动转类型,手动转是画蛇添足。 - checkpoint 忘存
scaler_state:resume 后放大倍数从默认值重新起步,训练曲线跳变。 - 与梯度裁剪(
ClipGradNorm)同时用:必须先把scaler.unscale_(optimizer)再裁剪,否则裁剪作用在"放大后的梯度"上,结果完全不对——本章没用到裁剪,但这是 AMP 最常见的连环坑,先记住。
- 三连顺序写错或漏调
- 下一章预告:
train.py已经 400 多行,模型、数据、日志、存档、早停、AMP、训练循环全挤在一个文件里,改模型要去翻训练循环——第 11 章按职责拆成独立模块,让train.py只做编排。
七、动手练习
- 显存与耗时对比:分别以
use_amp: true / false各训练 5 个 epoch,用nvidia-smi记录显存峰值、用日志里的 epoch 耗时记录时长,画一张对比表量化收益(CPU 上无 GPU 的同学可跳过显存部分,只对比精度是否一致)。 - 断点续训:AMP 开启时训练 5 轮后
--resume继续,确认scaler_state恢复后训练曲线没有跳变;再试着把 checkpoint 里的scaler_state删掉后 resume,观察曲线是否异常,验证"为什么要存它"。 - 制造溢出:把
lr调到 0.1 制造 loss 爆炸,训练中途打印scaler.get_scale(),观察它下降、以及损失/准确率是否因跳过更新而"原地不动"——理解防溢出机制。 - 手动全 FP16 实验:在
train_one_epoch里把images手动.half()、去掉autocast和 scaler,跑几个 epoch 观察 loss 是否变成nan——亲手验证"为什么不能全 FP16"。 - 脑内推导:不运行代码,回答——
scaler.scale(loss).backward()里 loss 被放大了 65536 倍,参数更新的数值为什么和没放大时一样?如果漏掉scaler.update(),放大倍数会发生什么?
