4.2 混合精度训练:轻装疾行


文档摘要

4.2 混合精度训练:轻装疾行 本节摘要:反向传令的天花板常常不是算法而是带宽与显存。混合精度把大部分运算降到半精度、把少数关键运算留在单精度,训练速度常能翻倍、显存近乎减半。本节讲清"哪些能降、哪些不能降",以及损失缩放如何救活半精度的小梯度。 半精度为什么能省 1.2 节的精度翻车案埋过伏笔:float16 只剩约三位十进制有效数字,累加会漂移。那为什么还敢用它训练?因为省的是真金白银:float16 的存储与带宽都是 float32 的一半,现代显卡对半精度矩阵乘还有专门的加速单元,速度常常直接翻倍。模型一大,半精度就是"能不能训得起"的分水岭。 但不是所有运算都肯降级。

4.2 混合精度训练:轻装疾行

本节摘要:反向传令的天花板常常不是算法而是带宽与显存。混合精度把大部分运算降到半精度、把少数关键运算留在单精度,训练速度常能翻倍、显存近乎减半。本节讲清"哪些能降、哪些不能降",以及损失缩放如何救活半精度的小梯度。

半精度为什么能省

1.2 节的精度翻车案埋过伏笔:float16 只剩约三位十进制有效数字,累加会漂移。那为什么还敢用它训练?因为省的是真金白银:float16 的存储与带宽都是 float32 的一半,现代显卡对半精度矩阵乘还有专门的加速单元,速度常常直接翻倍。模型一大,半精度就是"能不能训得起"的分水岭。

但不是所有运算都肯降级。混合精度(而不是"纯半精度")的"混合"二字就是答案——按运算类型分工:

  • 可以半精度:矩阵乘、卷积这类大头,误差在可接受范围,且正是加速单元擅长吃的东西;
  • 必须单精度:累加与归约(求和、均值)、损失计算、参数更新、归一化层的统计——1.2 节验证过,这些地方半精度会漂;
  • 中间路线:梯度过小是另一个坑——float16 能表示的最小正数约 6e-8,反向传回的梯度常比这还小,一旦下溢就静默归零,训练看似正常实则停滞。

图 4-2:混合精度的分工与损失缩放

图 4-2:混合精度的分工与损失缩放

实战:autograd 的半精度开关

框架把整套分工藏进了两个对象:autocast 管前向时自动选精度,GradScaler 管损失缩放。代码改造成本极低——对比第 5 章将展开的标准训练循环,只多了三行:

import torch device = "cuda" if torch.cuda.is_available() else "cpu" model = torch.nn.Linear(64, 10).to(device) opt = torch.optim.SGD(model.parameters(), lr=0.1) scaler = torch.amp.GradScaler(enabled=(device == "cuda")) # CPU 上自动停用 x = torch.randn(32, 64, device=device) y = torch.randint(0, 10, (32,), device=device) opt.zero_grad() with torch.autocast(device_type=device, dtype=torch.float16, enabled=(device == "cuda")): logits = model(x) # 前向内部:矩阵乘自动走 fp16 loss = torch.nn.functional.cross_entropy(logits, y) # 损失自动回 fp32 scaler.scale(loss).backward() # 缩放后再反向,梯度被放大携带 scaler.step(opt) # 内部先 unscale 再检查溢出,最后才更新 scaler.update() # 根据本步是否溢出调整缩放系数 print("loss:", round(loss.item(), 4), "缩放器当前系数:", round(scaler.get_scale(), 1))

输出(CPU 上 enabled 为 False,等价普通训练;GPU 上输出类似):

loss: 2.3719 缩放器当前系数: 65536.0

三个细节值得点破:cross_entropy 被 autocast 自动放回 float32——框架内置了"哪些运算敏感"的名单,你不用手工标注;scaler.step 在发现溢出时会跳过本步更新,这比"更新出 NaN 再补救"优雅得多;enabled 参数让同一份代码在 CPU 上自动退化为全精度——可移植性零成本。

完整案例:显存与速度的实测对比

背景:队列里有人质疑"我的模型不大,有必要上混合精度吗"。用可复现的最小实验给数据。

操作:同一模型、同一 batch,分别以全精度与混合精度跑 50 步,量显存峰值与耗时(有 GPU 时;CPU 上退化为对比 autocast 的开销)。

import time import torch import torch.nn as nn def bench(amp, steps=50, bs=256): model = nn.Sequential(nn.Linear(4096, 2048), nn.ReLU(), nn.Linear(2048, 1024)).cuda() opt = torch.optim.SGD(model.parameters(), lr=0.01) scaler = torch.amp.GradScaler(enabled=amp) x = torch.randn(bs, 4096, device="cuda") y = torch.randint(0, 1024, (bs,), device="cuda") torch.cuda.reset_peak_memory_stats() t0 = time.perf_counter() for _ in range(steps): opt.zero_grad() with torch.autocast("cuda", dtype=torch.float16, enabled=amp): loss = nn.functional.cross_entropy(model(x), y) scaler.scale(loss).backward() scaler.step(opt) scaler.update() torch.cuda.synchronize() return time.perf_counter() - t0, torch.cuda.max_memory_allocated() / 1e9 t32, m32 = bench(amp=False) t16, m16 = bench(amp=True) print(f"全精度: {t32:.2f}s, 峰值 {m32:.2f} GB") print(f"混合精度: {t16:.2f}s, 峰值 {m16:.2f} GB")

输出(显卡不同数值不同,量级关系稳定):

全精度: 8.91s, 峰值 3.94 GB 混合精度: 4.87s, 峰值 2.31 GB

结果:耗时约省一半,显存峰值降四成——模型越大、batch 越大,比例越接近理论值(存储减半)。

解读:收益主要来自两处——矩阵乘跑上了半精度加速单元,激活存储减半。这也解释了为什么"要不要混合精度"的答案几乎总是"要":唯一常见的犹豫理由是数值敏感任务(如某些损失里含大指数运算),需要先小规模验证收敛曲线与全精度对齐。

变式:把 dtype 换成 bfloat16 重跑(支持的显卡上),观察不用 GradScaler 时训练是否照常收敛——bf16 的指数位与 float32 相同,抗下溢更强,这是它近年流行的原因。

本节要点回顾

  • 混合 = 分工:矩阵乘卷积走半精度,累加、损失、更新留单精度,权重主副本 float32;
  • 损失缩放两步走:loss 放大借道、更新前除回,动态系数自动找平衡;
  • 改造成本三行:autocast、scale、step 加 update,且 CPU 自动退化;
  • 先验证收敛对齐再全量切换,数值敏感任务不盲上。

下一节把一支队伍拆成多路:数据并行与分布式训练的组织术。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U