Skip to content

FP8 推理精度与性能分析

FP8(8-bit 浮点)是 NVIDIA H100 引入的新数据格式,在保持浮点动态范围的同时将精度从 FP16 的 16-bit 降低到 8-bit,实现显存减半和计算加速。

FP8推理精度分析

FP8 格式定义

FP8 有两种编码格式:

格式指数位尾数位动态范围精度
E4M343±448较高(用于前向)
E5M252±57344较低(用于反向/损失缩放)
python
# FP8 量化示意
import torch

def fp8_e4m3_quantize(tensor):
    """将 FP16 张量量化为 FP8 E4M3"""
    # FP8 E4M3 最大值约 448
    max_val = 448.0
    scale = tensor.abs().max() / max_val
    # 量化
    quantized = torch.clamp(tensor / scale, -max_val, max_val)
    quantized = quantized.to(torch.float8_e4m3fn)
    return quantized, scale

def fp8_dequantize(quantized, scale):
    """FP8 反量化"""
    return quantized.to(torch.float16) * scale

FP8 推理流程

FP8 推理采用混合格式策略:

  • 权重:E4M3 格式存储,量化一次
  • 激活值:E4M3 格式,每层动态量化
  • KV Cache:E5M2 或 E4M3 格式
python
# vLLM 中启用 FP8
from vllm import LLM

llm = LLM(
    model="meta-llama/Llama-2-7b-hf",
    dtype="float8_e4m3fn",
    kv_cache_dtype="fp8_e5m2",
)

精度影响分析

FP8 对不同任务的影响:

任务类型FP16 基准FP8 结果差异
MMLU45.345.1-0.4%
HumanEval12.812.5-2.3%
GSM8K52.051.5-1.0%
翻译 (BLEU)42.141.8-0.7%

FP8 精度结论

FP8 对大多数 NLU 任务的影响在 1% 以内,对生成任务的影响略大但仍可接受。对于精度要求极高的场景(如数学推理),建议使用 FP8 + 残差补偿。

性能提升

在 NVIDIA H100 上的性能对比:

模型FP16 吞吐FP8 吞吐加速比
Llama-2-7B3100 tok/s5800 tok/s1.87x
Llama-2-13B1800 tok/s3400 tok/s1.89x
Llama-2-70B (4TP)850 tok/s1550 tok/s1.82x

硬件要求

FP8 的硬件加速仅在 NVIDIA H100/H200 和 AMD MI300X 上可用。在 A100/V100 等 GPU 上,FP8 会退化为软件模拟,无法获得加速收益。

相关资源

最近更新