首页 > 教程攻略 > ai资讯 >FP8 低精度训练:Transformer Engine 简析

FP8 低精度训练:Transformer Engine 简析

来源:互联网 时间:2026-08-16 14:05:13

训练大模型时,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)中,定义了两种格式:

  • E4M3

    :表示范围[-448, 448],更适合权重(weight)和激活值(activation)。
  • E5M2

    :表示范围[-57334, 57334],更适合梯度(gradient)。

简单说,E4M3精度高但范围窄,适合数据分布集中的张量;E5M2范围宽但精度低,适合梯度这种变化幅度大的数据。这一分工,是FP8训练能保持精度的关键。

为什么是FP8,不是int8?

int8在数值空间是均匀分布的,而大模型中的参数分布往往极不均匀。FP8的浮点表示能提供更宽的动态范围,更好地捕获这些分布的细节。另外,虽然FP16/BF16已经是主流,但大模型对精度损失的容忍度相对较高,FP8能在更小的精度误差下(通过Per-tensor Scaling控制),获得比16位更快的速度和更低的资源占用。

FP8训练效果如何?

从NVIDIA公布的测试结果来看,FP8在绝大多数训练任务上都能保持与FP16相当的精度,仅在少数困难任务(如数学运算)上有微小差距。

CV模型在FP8下的分类精度:

FP8 低精度训练:Transformer Engine 简析

NLP预训练任务:

LLM Benchmark:

SFT微调效果:

数据和图表清楚表明,FP8训练在主流场景下完全可行。

FP8的应用案例

  • Inflection AI

    :Inflection-2模型采用FP8混合精度,在5000块Hopper GPU上训练,累计算力达10^25 FLOPs。在MMLU、TriviaQA等基准上,其表现甚至超过了Google的PaLM 2,证明了FP8训练策略能保证模型正常收敛并取得优秀性能。
  • 零一万物

    :与NVIDIA合作,基于Megatron-LM开发了Yi训练框架,并在其中集成了Transformer Engine(TE)。FP8训练相比BF16获得了1.3倍的吞吐提升。
  • Google + NVIDIA

    :将TensorRT-LLM应用于Gemma模型,并用FP8推理加速。在Hopper GPU上,FP8比FP16的吞吐量提升3倍以上,可以在相同时间内使用更大的batch size,GPU利用率更高。

目前,NVIDIA的开源库

Transformer Engine (TE)

已经集成了FP8支持,并且被集成到PyTorch、JAX、PaddlePaddle等基础框架中,Megatron、NeMo、DeepSpeed、HuggingFace、Colossal-AI等LLM专用框架也都提供了FP8示例。

二、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

:每当梯度溢出时降低scale,间歇性地尝试增加scale,从而在不引起溢出的前提下使用最高scale,更好地恢复精度。

三、FP8技术分析

1. 宏观实现框架

FP16的Loss Scaling本质上是一种全局的离线量化。而FP8的数据范围更窄,单一的全局scale显然不够用。因此,FP8将量化的基本单位精细到每个tensor(甚至可以进一步到block-wise,用于更低精度)。

具体来说,每一次前向的GEMM计算需要对3个tensor记录scale值:inputweightoutput;反向计算需要记录2个tensor:grad_outputgrad_input。在TE的Hybrid模式下,前向用E4M3,反向用E5M2。两种格式均采用对称线性量化,只需计算scale值。

那么,如何在训练过程中高效地找到这个scale值?NVIDIA在TE文档中给出了两种方案:

  • Just-in-time scaling

    :先计算高精度的output tensor,再在其上计算amax,然后做量化。这种方式需要将完整的output tensor搬到HBM,导致多次kernel调用,增加了数据传输量,严重削弱FP8的性能优势。
  • Delayed scaling

    :提前知道scale值后,计算过程可以在一个kernel内完成,amax的计算和scale更新与计算独立,不中断计算进程,能完全发挥FP8性能。但需要额外空间记录scale历史,并引入微小误差。

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)参数:

  • margin

    :调整scaling factor。
  • interval

    :多少步更新一次scaling factor。
  • fp8_format

    :前向反向的计算精度(默认前向E4M3,反向E5M2)。
  • amax_history_len

    :amax历史窗口长度。

2. TE及各类框架集成方法

任何集成FP8能力的框架只需做两件事:

  1. 使用TE模块搭建模型

    ,因为计算要用到TE提供的FP8算子。
  2. 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框架

:支持DDP和FSDP的FP8训练。它会检查模型是否为TE结构,若不是则自动转换,并在前向计算时启用fp8_autocast

Megatron Core

:支持Tensor、Sequence、Pipeline并行与FP8训练结合。它定义了TE版本的Transformer层,并在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_checkpointnum_gemmsrecipefp8_groupfp8_max_fwdfp8_max_bwdscaling_fwdscaling_bwd等。

整体计算流程:

  1. 数据、模型先经过BF16 AMP处理,转为BF16精度。
  2. 遇到TE FP8 Module时,在前向/反向方法内,将input和weight转为FP8(如果有cache则跳过),调用fp8_gemm进行FP8计算,读取fp8_meta中的scale值,并将计算出的amax写入amax_history。
  3. 其他非FP8 Module按普通AMP逻辑计算。
  4. 优化器更新权重不属于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/