Skip to content

AMP混合精度训练原理与实战指南

内容整理自学习笔记,仅供面试备考参考;不构成录用、培训或考试承诺。

1. 什么是自动混合精度训练(AMP)?

PyTorch 1.6最大更新即为稳定版AMP(Automatic Mixed Precision)。PyTorch中有10种tensor类型:

类型说明
torch.FloatTensor32-bit floating point(默认
torch.DoubleTensor64-bit floating point
torch.HalfTensor16-bit floating point(FP16)
torch.BFloat16Tensor16-bit floating point(BF16)
torch.ByteTensor8-bit integer (unsigned)
torch.CharTensor8-bit integer (signed)
torch.ShortTensor16-bit integer (signed)
torch.IntTensor32-bit integer (signed)
torch.LongTensor64-bit integer (signed)
torch.BoolTensorBoolean

AMP的两个关键词:

  • 混合精度: 同时使用 torch.FloatTensor(FP32)和 torch.HalfTensor(FP16)
  • 自动: 框架按需自动调整tensor的dtype
python
from torch.cuda.amp import autocast as autocast

AMP只能在CUDA上使用,且需要支持Tensor Core的CUDA硬件(如2080ti)。Tensor Core每个时钟执行64个FP16矩阵乘+FP32累加的混合精度操作。

2. 为什么需要混合精度?

精度类型优势劣势
torch.HalfTensor(FP16)存储小、计算快、利用Tensor Core数值范围小(易溢出)、舍入误差(微小梯度信息丢失)
torch.FloatTensor(FP32)数值范围大、精度高存储大、计算慢

两种解决方案消除FP16劣势:

  1. 梯度scale(GradScaler): 放大loss防止梯度underflow,更新权重时再unscale回去
  2. 回落到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, prelu

3. 混合精度训练的优点

优点说明
减少显存占用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的情况下尽可能增大scaler

7. 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/NaNscaler.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要点总结:

  1. autocast() 仅包裹前向+loss计算
  2. GradScaler 在训练开始前实例化
  3. 使用 scaler.scale(loss).backward() 替代 loss.backward()
  4. 使用 scaler.step(optimizer) 替代 optimizer.step()
  5. 每次迭代调用 scaler.update() 动态调整scaler大小
  6. AMP可与DDP配合使用