FP8 推理精度与性能分析
FP8(8-bit 浮点)是 NVIDIA H100 引入的新数据格式,在保持浮点动态范围的同时将精度从 FP16 的 16-bit 降低到 8-bit,实现显存减半和计算加速。
FP8 格式定义
FP8 有两种编码格式:
| 格式 | 指数位 | 尾数位 | 动态范围 | 精度 |
|---|---|---|---|---|
| E4M3 | 4 | 3 | ±448 | 较高(用于前向) |
| E5M2 | 5 | 2 | ±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) * scaleFP8 推理流程
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 结果 | 差异 |
|---|---|---|---|
| MMLU | 45.3 | 45.1 | -0.4% |
| HumanEval | 12.8 | 12.5 | -2.3% |
| GSM8K | 52.0 | 51.5 | -1.0% |
| 翻译 (BLEU) | 42.1 | 41.8 | -0.7% |
FP8 精度结论
FP8 对大多数 NLU 任务的影响在 1% 以内,对生成任务的影响略大但仍可接受。对于精度要求极高的场景(如数学推理),建议使用 FP8 + 残差补偿。
性能提升
在 NVIDIA H100 上的性能对比:
| 模型 | FP16 吞吐 | FP8 吞吐 | 加速比 |
|---|---|---|---|
| Llama-2-7B | 3100 tok/s | 5800 tok/s | 1.87x |
| Llama-2-13B | 1800 tok/s | 3400 tok/s | 1.89x |
| Llama-2-70B (4TP) | 850 tok/s | 1550 tok/s | 1.82x |
硬件要求
FP8 的硬件加速仅在 NVIDIA H100/H200 和 AMD MI300X 上可用。在 A100/V100 等 GPU 上,FP8 会退化为软件模拟,无法获得加速收益。