混合精度训练中,AMP与BF16应该怎么选?

来源:JS脚本作者:IT柏拉图头衔:草根站长
导读:本期聚焦于IT柏拉图创作的《混合精度训练中,AMP与BF16应该怎么选?》,敬请观看详情。深度学习训练中,显存和算力经常成为瓶颈,混合精度训练能同时缓解这两个问题。它把一部分计算从FP32切换到FP16或BF16,借助GPU的张量核心获得更高吞吐,同时减少参数、梯度与中间激活的显存占用。不过AMP与BF16并不是同一条技术路线。AMP通常指FP16自动混合精度,需要GradScaler做损失缩放来防止梯度下溢;BF16保留了FP32的指数位,数值范围大,一般不用损失缩放,但对硬件有要求,且尾数精度低于FP16。本文从底层表示、训练稳定性、框架配置和硬件兼容性几个方面展开,帮助读者在卷积网络、Transformer和大模型训练中做出合适选择,避开常见的数值溢出与精度下降问题。

混合精度训练并不是简单把所有张量都改成半精度,而是根据不同算子的数值敏感度,保留部分关键权重与梯度在FP32中,把矩阵乘法等算术密集部分放到低精度格式上。当前主要的两条路线分别为FP16 AMP和BF16。前者出现更早,在NVIDIA Volta架构之后广泛使用,需要损失缩放来稳定训练;后者则依靠更宽的指数范围,在训练大模型时简化了数值管理。

混合精度训练中,AMP与BF16应该怎么选?

一、混合精度训练为什么有效

在默认的FP32训练中,每个浮点数占用4个字节,前向传播、反向传播和优化器状态都会累积大量显存。以一个大模型为例,参数本身是一部分,梯度和Adam优化器的一阶、二阶动量往往比参数更占空间。把这些张量中的一部分改成FP16或BF16后,显存占用可以下降三分之一到一半,同样的显存容量下可以训练更大的模型或者使用更大的批次。

更重要的是算力收益。现代NVIDIA GPU的Tensor Core对FP16和BF16矩阵乘法有专门加速,吞吐远超FP32。比如V100的FP16 Tensor Core算力是FP32的8倍左右,A100、H100上BF16也有类似优势。因此把卷积、全连接和注意力中的矩阵乘法切到低精度格式,训练吞吐会明显提升。但低精度格式的数值范围有限,如果直接替换所有计算,梯度很容易变成零或者出现NaN。混合精度训练的做法是:矩阵乘法等吃算力的算子用FP16或BF16,归一化、Softmax、权重更新等对精度敏感的算子保留FP32,同时维护一份FP32的主权重。

import torch
from torch.cuda.amp import GradScaler, autocast

model = MyModel().cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
scaler = GradScaler()

for data, target in dataloader:
    optimizer.zero_grad()
    with autocast():
        output = model(data)
        loss = criterion(output, target)
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

上面的代码展示了PyTorch里AMP的典型用法。autocast会按照算子白名单自动选择计算精度,GradScaler负责在反向传播前对loss做缩放,避免FP16梯度下溢。整体改动非常小,但训练速度和显存占用会得到立竿见影的改善。

二、AMP与BF16的底层差异

FP16和BF16虽然都占16位存储,但内部结构差异很大。FP16采用1个符号位、5个指数位和10个尾数位,能表示的数值范围大约从1e-8到65504。BF16采用1个符号位、8个指数位和7个尾数位,指数位与FP32完全一致,数值范围大约从1e-38到3.4e38。换句话说,BF16牺牲了尾数精度换来了和FP32一样的动态范围,而FP16保留了更多尾数精度但动态范围非常窄。

这个差异直接导致了AMP和BF16在工程实现上的不同。FP16 AMP因为动态范围窄,反向传播时许多小梯度会直接下溢成零,所以必须用GradScaler在求导前放大loss,反向后再把梯度缩放回来。这个过程需要额外监控scale值,如果出现inf或NaN,还要自动降低scale并跳过当前批次更新。BF16由于动态范围足够大,大部分任务可以不做损失缩放,训练代码更简洁。但BF16只有7位尾数,做累加时舍入误差比FP16更大,某些数值敏感的任务可能看到损失曲线抖动或最终指标下降。

特性FP16 AMPBF16
指数位5位8位
尾数位10位7位
数值范围约1e-8到65504约1e-38到3.4e38
损失缩放需要一般不需要
硬件要求支持Tensor Core的GPUAmpere及以上GPU、TPU等

还有一点容易被忽略:BF16对部分CPU以及特定AI芯片的亲和度更高。很多训练加速器直接以BF16作为推荐格式,因为它的指数位宽更容易设计运算单元。但在NVIDIA旧卡上,BF16无法获得硬件加速,如果用这些卡强行跑BF16,速度会比FP16 AMP慢很多。

import torch

model = MyModel().cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)

for data, target in dataloader:
    optimizer.zero_grad()
    with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
        output = model(data)
        loss = criterion(output, target)
    loss.backward()
    optimizer.step()

使用BF16的代码甚至比AMP更简单,它不需要GradScaler,只要在autocast中声明dtype即可。不过代码简单不代表不需要验证精度,尤其是学习率较大时,BF16的尾数误差可能让权重更新方向产生偏差。

三、硬件与模型如何决定选择

硬件条件是第一道门槛。NVIDIA的V100、T4等基于Volta和Turing架构的GPU只支持FP16 Tensor Core,无法加速BF16,所以这些环境下AMP几乎是唯一选择。从Ampere架构开始,包括A100、A10、A30、RTX 3090、RTX 4090等,同时支持FP16和BF16加速。到了Hopper架构的H100,BF16的硬件利用率非常高,训练大语言模型时通常直接采用BF16。Google TPU则从很早开始就把BF16作为主要训练格式,因此在TPU生态中BF16是默认方案。

模型结构同样影响选择。对于ResNet、YOLO、U-Net这类卷积网络,FP16 AMP已经经历了大量验证,社区默认配置成熟,损失缩放也基本不会引发额外问题。对于Transformer类模型,尤其是参数量较大的预训练语言模型,BF16的优势更明显。因为这类模型在训练初期经常出现attention分数过大或梯度范围较宽的情况,FP16的动态范围容易触发溢出,而BF16的宽指数范围天然更稳定。反过来,一些对权重微小更新很敏感的模型,例如图像超分、语音合成或者强化学习中的策略网络,BF16的7位尾数可能造成精度不足,这时FP16 AMP凭借10位尾数会更合适。

此外还要考虑下游部署。如果最终要把模型导出到TensorRT或ONNX Runtime,需要确认推理端对两种格式的支持。FP16在多数推理框架中支持广泛,BF16虽然也在逐步普及,但在部分边缘设备上仍然需要转换回FP32,这会增加工程成本。

import tensorflow as tf

tf.keras.mixed_precision.set_global_policy('mixed_bfloat16')

model = create_model()
optimizer = tf.keras.optimizers.AdamW()
model.compile(optimizer=optimizer, loss='categorical_crossentropy')
model.fit(train_dataset, epochs=10)

TensorFlow中启用的方式也很直接,通过set_global_policy把全局策略设置为mixed_bfloat16即可。框架会自动把适合低精度的算子放到BF16中,同时保留BatchNorm、LayerNorm和Softmax等操作在FP32。需要注意的是,输出层通常仍然保持FP32,否则交叉熵损失的计算容易因为BF16尾数不足而出现较大误差。

四、实战中常见问题与调优建议

FP16 AMP最常见的两个问题是scale值异常和梯度溢出。GradScaler默认会动态调整scale,在连续多步没有出现inf时会逐步增大scale,发现inf或NaN时会跳过当前批次并降低scale。如果训练过程中频繁报出overflow,可以先检查初始学习率是否过大,或者查看模型输出层是否存在指数增长。调整方法是降低学习率、在损失函数前加梯度裁剪,以及把容易溢出的激活层手动排除在autocast之外。

BF16训练则要重点监控最终指标。因为损失曲线可能看起来正常,但模型精度比FP32低零点几个点。这种情况通常不是因为数值溢出,而是累加误差在长时间训练中被逐步放大。可以尝试把优化器状态保持在FP32,或者只对部分模块启用BF16,诸如Embedding、LayerNorm和分类头仍然使用FP32。对于大模型,也可以采用BF16加FP32主权重的方案,这样前向计算快,权重更新又相对稳定。

# 检查GradScaler的scale变化
if scaler.get_scale() < 128:
    print("scale过小,注意检查梯度是否频繁溢出")

混合精度训练不是单纯的开关切换,需要配合验证流程。建议先跑几百步,对比FP32、FP16 AMP和BF16在验证集上的损失、梯度范数分布和最终指标。短时间看不出问题,但长训练中的稳定性差异会逐渐显现。对于分布式训练,低精度还能减少梯度通信量,如果带宽压力大,可以考虑把通信部分也切到FP16或BF16,但要确认集合通信库支持对应的dtype。

最终的选择可以用一个简单原则概括:看硬件是否支持BF16加速,再看任务是否对尾数精度敏感。老GPU上老老实实用AMP,Ampere以上且训练大模型优先试BF16,数值敏感的中小模型继续用FP16 AMP。两条路线没有绝对的优劣,关键是理解数值格式对训练动态的影响,并在自己的任务上做完整对比。

混合精度训练AMPBF16修改时间:2026-10-02 15:38:10

免责声明:已尽一切努力确保本网站所含信息的准确性。网站作品多为原创整理与精心创作,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们进行处理Email:chomcom@qq.com。
引用或转载本作品时,请注明当前出处:https://www.ipipp.com/html/1002/64725.html,基于非商业用途的前提下,欢迎转载或二创本作品。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。