FP8 低精度训练:Transformer Engine 简析
训练大模型时,FP16和BF16的混合精度训练(AMP)早已是标配。它能在不牺牲下游任务精度的前提下,显著提升训练效率、降低显存占用。从PyTorch 1.6开始,原生AMP(torch.amp)就取代了曾经的NVIDIA APEX库,成了主流方案。
不过,这股风潮很快就要被FP8取代了。自NVIDIA Hopper架构起,GPU开始支持FP8精度的Tensor Core计算。相比FP16/BF16,FP8带来了肉眼可见的优势:
- :对比A100上的BF16训练,H100用FP8能提速2~3倍;而FP8自身的计算吞吐也达到了FP16的两倍。
更强的计算性能
- :FP8不仅计算快,还能节省50%~75%的内存占用,连通信量也砍掉了一半多。
更低的训练成本
- :FP8的使用天然要求模型在训练和推理中做量化,这反而有助于模型压缩和部署成本降低。
更好的模型优化
当然,在深入技术细节之前,先回答几个最核心的问题。
什么是FP8?
FP8就是8位浮点数,比16位精度少了一半的指数位和尾数位。在NV、Arm、Intel联合发布的白皮书(arXiv:2209.05433)中,定义了两种格式:
- :表示范围[-448, 448],更适合权重(weight)和激活值(activation)。
E4M3
- :表示范围[-57334, 57334],更适合梯度(gradient)。
E5M2
简单说,E4M3精度高但范围窄,适合数据分布集中的张量;E5M2范围宽但精度低,适合梯度这种变化幅度大的数据。这一分工,是FP8训练能保持精度的关键。
为什么是FP8,不是int8?
int8在数值空间是均匀分布的,而大模型中的参数分布往往极不均匀。FP8的浮点表示能提供更宽的动态范围,更好地捕获这些分布的细节。另外,虽然FP16/BF16已经是主流,但大模型对精度损失的容忍度相对较高,FP8能在更小的精度误差下(通过Per-tensor Scaling控制),获得比16位更快的速度和更低的资源占用。
FP8训练效果如何?
从NVIDIA公布的测试结果来看,FP8在绝大多数训练任务上都能保持与FP16相当的精度,仅在少数困难任务(如数学运算)上有微小差距。
CV模型在FP8下的分类精度:

NLP预训练任务:
LLM Benchmark:
SFT微调效果:
数据和图表清楚表明,FP8训练在主流场景下完全可行。
FP8的应用案例
- :Inflection-2模型采用FP8混合精度,在5000块Hopper GPU上训练,累计算力达10^25 FLOPs。在MMLU、TriviaQA等基准上,其表现甚至超过了Google的PaLM 2,证明了FP8训练策略能保证模型正常收敛并取得优秀性能。
Inflection AI
- :与NVIDIA合作,基于Megatron-LM开发了Yi训练框架,并在其中集成了Transformer Engine(TE)。FP8训练相比BF16获得了1.3倍的吞吐提升。
零一万物
- :将TensorRT-LLM应用于Gemma模型,并用FP8推理加速。在Hopper GPU上,FP8比FP16的吞吐量提升3倍以上,可以在相同时间内使用更大的batch size,GPU利用率更高。
Google + NVIDIA
目前,NVIDIA的开源库
Transformer Engine (TE)
二、FP16/BF16 AMP回顾
在深入FP8之前,有必要回顾一下PyTorch AMP的实现原理,因为FP8的许多设计都源于AMP的成熟方案。
1. 计算流程
典型的FP16 AMP流程是这样的:
- 模型和初始数据都是FP32精度。
- 进入
torch.autocast后,前向计算开始:遇到FP16算子时,权重和数据转为FP16(权重通常有FP16 cache);遇到FP32算子时,保持FP32计算。 - 反向计算时,torch会根据前向精度自动确定反向精度,不需要显式设置autocast。
- 优化器更新权重时,利用Tensor Core直接完成FP16+FP32的加法,以FP32精度更新,无需额外转换。
FP16支持的算子列表可以查阅PyTorch官方文档。
示例代码(注意这里用的是BF16):
with torch.cuda.amp.autocast(dtype=torch.bfloat16):
outputs = model(inputs)
loss = loss_func(outputs, targets)
loss.backward()
optimizer.step()
optimizer.zero_grad()
2. 显存分布
FP16 AMP训练时,显存中主要包括:
- 前向模型权重:FP16
- 梯度:FP16
- 优化器:FP32 Master Model Weight + 2×FP32 Adam States(一阶矩、二阶矩)
- 其他中间计算结果
3. Global Loss Scaling
FP16的数值范围有限,容易发生溢出(overflow/underflow),因此需要loss scaling。PyTorch提供了GradScaler:
scaler = torch.cuda.amp.GradScaler()
with torch.cuda.amp.autocast(dtype=torch.float16):
outputs = model(inputs)
loss = loss_func(outputs, targets)
scaler.scaled(loss).backward()
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
数值范围对比:FP32(1-8-23)、BF16(1-8-7)、FP16(1-5-10)。BF16与FP32指数位相同,无需scaler;而FP16只有5个指数位,容易溢出,必须用scaler动态调整。
实际中,我们维护一个全局的scale值,并采用
Dynamic Loss Scaling
三、FP8技术分析
1. 宏观实现框架
FP16的Loss Scaling本质上是一种全局的离线量化。而FP8的数据范围更窄,单一的全局scale显然不够用。因此,FP8将量化的基本单位精细到每个tensor(甚至可以进一步到block-wise,用于更低精度)。
具体来说,每一次前向的GEMM计算需要对3个tensor记录scale值:input、weight、output;反向计算需要记录2个tensor:grad_output和grad_input。在TE的Hybrid模式下,前向用E4M3,反向用E5M2。两种格式均采用对称线性量化,只需计算scale值。
那么,如何在训练过程中高效地找到这个scale值?NVIDIA在TE文档中给出了两种方案:
- :先计算高精度的output tensor,再在其上计算amax,然后做量化。这种方式需要将完整的output tensor搬到HBM,导致多次kernel调用,增加了数据传输量,严重削弱FP8的性能优势。
Just-in-time scaling
- :提前知道scale值后,计算过程可以在一个kernel内完成,amax的计算和scale更新与计算独立,不中断计算进程,能完全发挥FP8性能。但需要额外空间记录scale历史,并引入微小误差。
Delayed scaling
TE采用Delayed scaling方案。对每个GEMM算子用到的tensor,记录一个amax history数组。需要scale值时,从数组中取最近一段时间窗口内amax的最大值,近似当前tensor的amax,然后计算scale:
FP8_MAX = maximum_representable_value(fp8_format)
new_scaling_factor = (FP8_MAX / amax) / (2 ^ margin)
用户可以自定义策略(Recipe)参数:
- :调整scaling factor。
margin
- :多少步更新一次scaling factor。
interval
- :前向反向的计算精度(默认前向E4M3,反向E5M2)。
fp8_format
- :amax历史窗口长度。
amax_history_len
2. TE及各类框架集成方法
任何集成FP8能力的框架只需做两件事:
- ,因为计算要用到TE提供的FP8算子。
使用TE模块搭建模型
- 。
用
fp8_autocast装饰前向计算过程
实际场景中,FP8训练通常要结合BF16混合精度训练。
TE官方案例
import torch
import transformer_engine.pytorch as te
from transformer_engine.common import recipe
model = te.Linear(in_features, out_features, bias=True)
inp = torch.randn(hidden_size, in_features, device="cuda")
fp8_recipe = recipe.DelayedScaling(margin=0, interval=1, fp8_format=recipe.Format.E4M3)
with te.fp8_autocast(enabled=True, fp8_recipe=fp8_recipe):
out = model(inp)
loss = out.sum()
loss.backward()
Accelerate框架
fp8_autocast。
Megatron Core
forward方法中根据配置设置fp8_autocast上下文。
# 简化的核心逻辑
if self.config.fp8:
fp8_format = ... # 根据配置选择E4M3或HYBRID
fp8_recipe = TEDelayedScaling(...)
fp8_context = transformer_engine.pytorch.fp8_autocast(enabled=True, fp8_recipe=fp8_recipe, fp8_group=fp8_group)
else:
fp8_context = nullcontext()
with fp8_context:
# Forward pass
3. FP8框架TE计算流程
fp8_autocast是一个上下文管理器,其核心逻辑:
- 进入时,保存当前FP8状态,并更新训练设置。
- 退出时,恢复原状态,并在
fp8_autocast_exit中reduce各进程的amax,更新amax history和scale值。
每个TE模块都继承自TransformerEngineBaseModule,拥有一个fp8_meta字典,记录FP8关键信息:
fp8_checkpoint、num_gemms、recipe、fp8_group、fp8_max_fwd、fp8_max_bwd、scaling_fwd、scaling_bwd等。
整体计算流程:
- 数据、模型先经过BF16 AMP处理,转为BF16精度。
- 遇到TE FP8 Module时,在前向/反向方法内,将input和weight转为FP8(如果有cache则跳过),调用
fp8_gemm进行FP8计算,读取fp8_meta中的scale值,并将计算出的amax写入amax_history。 - 其他非FP8 Module按普通AMP逻辑计算。
- 优化器更新权重不属于FP8管辖,按普通AMP进行。
4. Tensor Core如何进行FP8训练
FP8精度计算仅能在Tensor Core上运行。Tensor Core的基本运算单元为D=A×B+C,每个时钟周期完成4×4的mma运算。其wmma::mma_syncAPI最小数据单元是16×16矩阵,因此TE要求输入数据各维度为16的倍数。
两个FP8矩阵输入后,Tensor Core输出高精度结果(FP16/FP32),因此存在FP8→FP16/FP32及FP16/FP32→FP8的精度转化过程。
四、总结与展望
FP8训练的局限性
- 在小参数规模(<1B参数)训练中,FP8的量化处理、精度转化等overhead可能超过计算加速带来的收益。
- 当batch size过小时(比如仅4),FP8训练吞吐可能反而低于BF16(实测下降约17%)。
- 部分下游任务(如数学运算、MMLU中的困难任务)微调效果欠佳。
- FP8训练过程中间出现loss spike或NaN时,调试更具挑战性。
FP8及更低精度训练的前景
尽管有这些限制,FP8在大模型训练中的价值已经明确,工业界也不乏成功案例。它有望成为大模型高效训练的标准配置之一。
硬件层面,NVIDIA最新的Blackwell架构开始支持FP6、FP4等更低精度的Tensor Core运算,并可能采用block-wise量化方案。而DeepSpeed也已推出不依赖硬件的FP6运算(GitHub仓库中有详细说明)。
在低精度运算成为常规方案的今天,在保证训练精度的前提下,性能提升的空间可能还远未触及天花板。
参考资料
FP8相关论文:
- 8-BIT NUMERICAL FORMATS FOR DEEP NEURAL NETWORKS(Graphcore 2021)
- Auto-Precision Scaling for Distributed Deep Learning(2021)
- FP8 Formats For Deep Learning(NV、Intel、Arm 2022)
- Mixed Precision Training With 8-bit Floating Point(Intel 2019)
TE代码与文档:
- https://github.com/NVIDIA/TransformerEngine/tree/main
- https://docs.nvidia.com/deeplearning/transformer-engine/user-guide/index.html
NV技术博客:
- https://developer.nvidia.com/zh-cn/blog/nvidia-gpu-fp8-training-inference/
- https://developer.nvidia.com/zh-cn/blog/fp8-precision-performance/
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名