外观
AMP混合精度训练原理与实战指南
内容整理自学习笔记,仅供面试备考参考;不构成录用、培训或考试承诺。
1. 什么是自动混合精度训练(AMP)?
PyTorch 1.6最大更新即为稳定版AMP(Automatic Mixed Precision)。PyTorch中有10种tensor类型:
| 类型 | 说明 |
|---|---|
| torch.FloatTensor | 32-bit floating point(默认) |
| torch.DoubleTensor | 64-bit floating point |
| torch.HalfTensor | 16-bit floating point(FP16) |
| torch.BFloat16Tensor | 16-bit floating point(BF16) |
| torch.ByteTensor | 8-bit integer (unsigned) |
| torch.CharTensor | 8-bit integer (signed) |
| torch.ShortTensor | 16-bit integer (signed) |
| torch.IntTensor | 32-bit integer (signed) |
| torch.LongTensor | 64-bit integer (signed) |
| torch.BoolTensor | Boolean |
AMP的两个关键词:
- 混合精度: 同时使用
torch.FloatTensor(FP32)和torch.HalfTensor(FP16) - 自动: 框架按需自动调整tensor的dtype
python
from torch.cuda.amp import autocast as autocastAMP只能在CUDA上使用,且需要支持Tensor Core的CUDA硬件(如2080ti)。Tensor Core每个时钟执行64个FP16矩阵乘+FP32累加的混合精度操作。
2. 为什么需要混合精度?
| 精度类型 | 优势 | 劣势 |
|---|---|---|
| torch.HalfTensor(FP16) | 存储小、计算快、利用Tensor Core | 数值范围小(易溢出)、舍入误差(微小梯度信息丢失) |
| torch.FloatTensor(FP32) | 数值范围大、精度高 | 存储大、计算慢 |
两种解决方案消除FP16劣势:
- 梯度scale(GradScaler): 放大loss防止梯度underflow,更新权重时再unscale回去
- 回落到FP32: 混合一词的由来,部分操作自动使用FP32
AMP中自动转为FP16的操作
以下操作在autocast上下文中tensor会被自动转为FP16:
__matmul__, addbmm, addmm, addmv, addr, baddbmm, bmm,
chain_matmul, conv1d, conv2d, conv3d, conv_transpose1d,
conv_transpose2d, conv_transpose3d, linear, matmul, mm, mv, prelu3. 混合精度训练的优点
| 优点 | 说明 |
|---|---|
| 减少显存占用 | FP16占用的显存为FP32的一半 |
| 加快训练速度 | 通信量减半,计算性能翻倍 |
4. 混合精度训练的缺点
| 缺点 | 说明 |
|---|---|
| 数据溢出 | FP16数值范围小,易Overflow/Underflow |
| 舍入误差 | 微小梯度信息可能达不到16bit最低分辨率 |
5. 混合精度训练的关键技术
| 技术 | 说明 |
|---|---|
| FP32主权重备份 | 保持FP32副本用于精确更新 |
| 动态损失缩放 | 放大loss防止梯度underflow |
6. 动态损失缩放(GradScaler)详解
静态损失标度: 将损失乘以某个大数字(如1024),梯度也放大1024倍,大幅降低梯度下溢几率。计算梯度后除以1024恢复准确值。
动态损失标度:
- 发生溢出:跳过优化器更新,损失标度减半
- 连续N个steps没有溢出:损失标度翻倍
python
scaler = GradScaler()
# scaler大小在每次迭代中动态估计
# 原则:在不出现inf/NaN的情况下尽可能增大scaler7. PyTorch中使用AMP
7.1 autocast
python
from torch.cuda.amp import autocast as autocast
model = Net().cuda()
optimizer = optim.SGD(model.parameters(), ...)
for input, target in data:
optimizer.zero_grad()
# 前向过程(model + loss)开启 autocast
with autocast():
output = model(input)
loss = loss_fn(output, target)
# 反向传播在autocast上下文之外
loss.backward()
optimizer.step()注意:
- autocast上下文应只包含前向过程(含loss计算),不包含反向传播
- 不要在model或input上手工调用
.half(),框架会自动处理- 遇到
RuntimeError: expected scalar type float but found c10::Half时,可在tensor上手工调用.float()匹配type
7.2 GradScaler
python
from torch.cuda.amp import autocast as autocast
from torch.cuda.amp import GradScaler
model = Net().cuda()
optimizer = optim.SGD(model.parameters(), ...)
scaler = GradScaler()
for input, target in data:
optimizer.zero_grad()
with autocast():
output = model(input)
loss = loss_fn(output, target)
# Scales loss,为梯度放大
scaler.scale(loss).backward()
# scaler.step() 先unscale梯度,若无inf/NaN则调用optimizer.step()
# 否则忽略step调用,保证权重不被破坏
scaler.step(optimizer)
# 准备是否要增大scaler
scaler.update()GradScaler动态调整逻辑:
| 情况 | 操作 |
|---|---|
| 出现inf或NaN | 忽略此次权重更新,scaler缩小(乘上backoff_factor) |
| 连续多次无inf/NaN | scaler.update()增大scaler(乘上growth_factor) |
8. AMP完整实战代码
8.1 Trainer训练类(AMP部分)
python
class Trainer:
def train(self, train_loader, dev_loader=None, train_sampler=None):
if self.args.use_amp:
scaler = torch.cuda.amp.GradScaler()
for epoch in range(1, self.args.epochs + 1):
train_sampler.set_epoch(epoch)
for step, batch_data in enumerate(train_loader):
self.model.train()
if self.args.use_amp:
with torch.cuda.amp.autocast():
logits, label = self.on_step(batch_data)
loss = self.criterion(logits, label)
torch.distributed.barrier()
scaler.scale(loss).backward()
scaler.step(self.optimizer)
scaler.update()
else:
logits, label = self.on_step(batch_data)
loss = self.criterion(logits, label)
torch.distributed.barrier()
loss.backward()
self.optimizer.step()8.2 启动DDP+AMP训练
python
def main_worker(local_rank, local_world_size):
dist.init_process_group(backend="nccl",
init_method="tcp://localhost:12345",
world_size=local_world_size, rank=local_rank)
torch.cuda.set_device(local_rank)
# 封装模型
model.cuda()
model = torch.nn.parallel.DistributedDataParallel(
model, device_ids=[local_rank])
...
if __name__ == '__main__':
mp.spawn(main_worker, nprocs=p_args.local_world_size,
args=(p_args.local_world_size,))AMP要点总结:
autocast()仅包裹前向+loss计算GradScaler在训练开始前实例化- 使用
scaler.scale(loss).backward()替代loss.backward()- 使用
scaler.step(optimizer)替代optimizer.step()- 每次迭代调用
scaler.update()动态调整scaler大小- AMP可与DDP配合使用