PTPC-FP8:提升 AMD ROCm 上的 vLLM 性能

9 分钟阅读
AMD 与嵌入式 LLM

TL;DR:AMD ROCm 上的 vLLM 现在拥有了更佳的 FP8 性能!

  • 有什么新功能? AMD ROCm 上的 vLLM (v0.7.3+) 现在支持 PTPC-FP8 量化
  • 它有什么好处? 你可以获得与其他 FP8 方法相当的速度,但精度却非常接近原始模型(BF16)的质量。这是目前 ROCm 平台上最好的 FP8 选项。
  • 如何使用
    1. 安装 ROCm。
    2. 获取最新版 vLLM (v0.7.3 或更高版本)。
    3. 运行 Hugging Face 模型时添加 --quantization ptpc_fp8 标志。无需预先量化!
What is PTPC-FP8
什么是 PTPC-FP8
什么是 PTPC-FP8? 这是一种同时针对 FP8 权重激活值的量化方法。它对激活使用按 token 缩放(Per-token scaling),对权重使用按通道缩放(Per-channel scaling),从而比传统的按张量(Per-tensor)FP8 提供更好的精度。

简介

大型语言模型(LLM)正在彻底改变我们与技术互动的方式,但其巨大的计算需求可能成为一道门槛。如果能在不牺牲精度的情况下,在 AMD GPU 上更快、更高效地运行这些强大的模型呢?现在你可以做到了!本文介绍了一项突破:vLLM 中的 PTPC-FP8 量化,专为 AMD 的 ROCm 平台优化。准备好直接使用 Hugging Face 模型,以 FP8 的速度获得近乎 BF16 的精度——无需预量化!我们将向你展示它是如何工作的,对其性能进行基准测试,并引导你开始使用。

LLM 量化挑战与 PTPC-FP8 解决方案

运行大模型计算成本高昂。FP8(8 位浮点数)通过减少内存占用和加速矩阵乘法提供了一种引人注目的解决方案,但传统的量化方法在处理 LLM 时面临严峻挑战。

离群值(Outlier)问题

随着模型规模增大,LLM 会出现激活值离群点。这些异常大的数值给量化带来了巨大挑战。

  • 当使用按张量量化时,大多数数值获得的有效精度位很少。
  • 离群值在不同 token 的特定通道中持续出现。
  • 权重相对均匀且易于量化,但激活值则不然。

PTPC:一种精度针对性方法

PTPC-FP8(按 Token 激活、按通道权重 FP8)通过基于以下三个关键观察结果的定制缩放因子解决了这一挑战:

  1. 离群值持续出现在相同的通道中。
  2. 单个 token 内的通道量级差异巨大。
  3. 同一通道在不同 token 间的量级保持相对稳定。

这种见解引出了双粒度方法:

  • 按 Token 激活量化:每个输入 token 拥有自己的缩放因子。
  • 按通道权重量化:每个权重列拥有独特的缩放因子。
Per-Token Activation + Per-Channel Weight Quantization
按 Token 激活 + 按通道权重量化

理解示意图

该图示展示了两种量化方法:

张量维度(两种方法均适用)

  • ...:输入激活张量()
  • ...:权重张量()
  • ...:Token 序列长度
  • ...:输入/输出通道
  • ...:矩阵乘法

缩放因子

  • 顶部(按张量):为整个张量使用单个标量。针对整个张量
  • 底部(PTPC):向量每个 token 一个比例,并且每个输入通道一个比例

这种细粒度的缩放方法使得 PTPC-FP8 能够在保持 8 位计算的速度和内存优势的同时,达到接近 BF16 的精度。

深入探讨:PTPC-FP8 在 vLLM 中如何工作(以及融合算子)

如果没有适当的优化,PTPC-FP8 的细粒度缩放可能会降低速度。保持速度的关键在于 AMD ROCm 对融合 FP8 按行缩放 GEMM(矩阵乘法)算子的实现。

挑战:两步走方法 vs. 融合方法

如果不优化,带有按 token 和按通道缩放的矩阵乘法将需要两个昂贵的步骤:

# Naive 2-step approach:
output = torch._scaled_mm(input, weight)       # Step 1: FP8 GEMM
output = output * token_scales * channel_scales  # Step 2: Apply scaling factors

这造成了性能瓶颈:

  • 将巨大的中间结果写入内存;
  • 将其读回以进行缩放操作;
  • 浪费内存带宽和计算周期。

解决方案:算子融合

融合方法将矩阵乘法和缩放整合为一个单一的硬件操作:

# Optimized fused operation:
output = torch._scaled_mm(input, weight, 
                         scale_a=token_scales, 
                         scale_b=channel_scales)
Fused GEMM Operation
融合 GEMM 操作

为什么这很重要

这种融合利用了 AMD GPU 的专用硬件(特别是在具有原生 FP8 支持的 MI300X 上):

  • 内存效率:在将结果写入内存之前,缩放会在片上内存中进行。
  • 计算效率:消除了冗余操作。
  • 性能提升:我们的测试显示,与朴素实现相比,性能提升高达 2.5 倍。

融合操作使 PTPC-FP8 能够用于实际部署,消除了使用更细粒度缩放因子带来的性能损失,同时保持了精度优势。

PTPC-FP8 基准测试:MI300X 上的速度与精度

我们使用 vLLM 在 AMD MI300X GPU(提交记录 4ea48fb35cf67d61a1c3f18e3981c362e1d8e26f)上对 PTPC-FP8 进行了广泛的基准测试。以下是我们的发现:

1. 吞吐量比较(PTPC-FP8 vs. 按张量/Per-Tensor FP8)

  • 模型: Llama-3.1-70B-Instruct
  • 数据集: SharedGPT
  • GPU: 1x MI300X
  • 结果: PTPC-FP8 实现了与按张量 FP8 几乎相同的吞吐量(甚至略好——提升了 1.01 倍)。这证明了融合算子完全克服了 PTPC-FP8 更复杂缩放所带来的潜在开销。
Throughput in Reqs/s across various input-output sequence length of Llama-3.1-70B-Instruct
Llama-3.1-70B-Instruct 在不同输入输出序列长度下的吞吐量(请求/秒)
Request/s Throughput gain over FP8 per-tensor quantization
across different input token length - output token length
不同输入 token 长度 - 输出 token 长度下,相对于 FP8 按张量量化的吞吐量增益

2.1. 精度:困惑度(Perplexity,越低越好)

  • 模型: Llama-3.1-8B-Instruct
  • 数据集: Wikitext
  • 配置: 2× MI300X GPU,使用张量并行

理解困惑度:预测能力测试

将困惑度视为模型在预测文本时有多“困惑”的衡量标准。就像学生参加测验一样:

  • 较低的困惑度 = 更好的预测(模型自信地为正确的下一个单词分配高概率)
  • 较高的困惑度 = 更多的不确定性(模型经常对接下来出现的内容感到惊讶)

困惑度的小幅增加(即使是 0.1)也可能意味着模型质量的显著下降,对于已经过深度优化的大型语言模型尤其如此。

结果:PTPC-FP8 保持了接近 BF16 的质量

bits and byte perplexity
比特和字节困惑度
Word Perplexity Comparison
单词困惑度比较
精度单词困惑度% 退化
BF16(基准)9.4281-
PTPC-FP89.50930.86%
标准 FP89.51240.89%

如表格和图表所示:

  1. PTPC-FP8 优于标准 FP8 量化(9.5093 vs 9.5124)
  2. 与 BF16 的差距极小 - 相较于全精度基准,退化仅为 0.86%
  3. 字节级指标(bits_per_byte 和 byte_perplexity)显示出相同的模式

为什么这很重要: 虽然标准 FP8 已经提供了不错的结果,但 PTPC-FP8 更低的困惑度表明它更好地保持了模型进行准确预测的能力。这对于复杂的推理和生成任务尤为重要,因为微小的质量下降在这些任务中可能会累积,从而导致输出质量上的明显差异。

2.2. GSM8K 上的精度:测试数学推理能力**

什么是 GSM8K 及其重要性

GSM8K 测试模型解决小学数学应用题的能力,这是 LLM 最具挑战性的任务之一。与简单的文本预测不同,这些问题需要:

  • 多步推理
  • 数值准确性
  • 逻辑一致性

该基准测试是衡量量化是否保留了模型推理能力的重要指标。

理解结果

我们使用两种方法衡量精度:

  • 灵活提取(Flexible-extract):如果正确数字出现在响应的任何位置,即视为答案正确。
  • 严格匹配(Strict-match):要求以预期格式给出精确答案。
Accuracy Comparison on Llama-3.1-8B
Llama-3.1-8B 上的精度比较
8B 模型结果一览
方法严格匹配精度% 相当于 BF16 性能
BF16(基准)73.2%100%
PTPC-FP870.8%96.7%
标准 FP869.2%94.5%

70B 模型结果

Accuracy Comparison on Llama-3.1-70B
Llama-3.1-70B 上的精度比较
对于更大的 70B 模型:
  • PTPC-FP8 达到了 87.3% 的严格匹配精度
  • 这实际上 略好于 BF16 的 86.3%
  • 两者在严格匹配条件下均优于标准 FP8

为什么这些结果很重要

  1. 保留推理能力:数学推理往往是量化后首先下降的能力

  2. PTPC-FP8 在两种模型规模上均持续优于标准 FP8

  3. 接近 BF16 的质量,且内存大幅减少,性能得到提升

  4. 扩展优势:随着模型规模的增加,量化方法之间的性能差距会缩小,这表明 PTPC-FP8 对于大型模型尤为有价值

这些结果证明,PTPC-FP8 量化在保留模型执行复杂推理任务能力的同时,提供了 8 位精度的速度和效率优势。

入门指南

  1. 安装 ROCm: 确保你拥有较新版本。
  2. 立即克隆最新的 vLLM 提交!设置并开始探索这项新功能!
$ git clone https://github.com/vllm-project/vllm.git
$ cd vllm  
$ DOCKER_BUILDKIT=1 docker build -f Dockerfile.rocm -t vllm-rocm .  
$ docker run -it \
   --network=host \
   --group-add=video \
   --ipc=host \
   --cap-add=SYS_PTRACE \
   --security-opt seccomp=unconfined \
   --device /dev/kfd \
   --device /dev/dri \
   -v <path/to/model>:/app/model \
   vllm-rocm \
   bash
  1. 运行 vLLM 并使用 --quantization ptpc_fp8 标志。
VLLM_USE_TRITON_FLASH_ATTN=0 vllm serve <your-model> --max-seq-len-to-capture 16384 --enable-chunked-prefill=False --num-scheduler-steps 15 --max-num-seqs 1024 --quantization ptpc_fp8

(将 <your-model> 替换为任何 Hugging Face 模型;它将自动在运行时实时量化权重。)

结论:精度与速度的完美平衡点

AMD ROCm 上 vLLM 中的 PTPC-FP8 量化代表了在大众化获取强大 LLM 方面迈出的重要一步。通过以 FP8 的速度实现接近 BF16 的精度,我们正在打破限制更广泛采用的计算障碍。这一进步为更广泛的社区——从个人研究人员到资源受限的组织——赋能,使他们能够在可访问的 AMD 硬件上利用大模型的力量。我们邀请你探索 PTPC-FP8,分享你的经验,为 vLLM 项目做出贡献,并帮助我们构建一个让每个人都能使用高效且准确的 AI 的未来。

附录

lm-evaluation-harness 命令

# Unquantized (Bfloat16)
MODEL=meta-llama/Llama-3.1-8B-Instruct
HIP_VISIBLE_DEVICES=0,1 lm_eval \
  --model vllm \
  --model_args pretrained=$MODEL,add_bos_token=True,tensor_parallel_size=2,kv_cache_dtype=auto,max_model_len=2048,gpu_memory_utilization=0.6 \
  --tasks wikitext --batch_size 16
 
# Per-Tensor FP8 Quantization
MODEL=meta-llama/Llama-3.1-8B-Instruct
HIP_VISIBLE_DEVICES=0,1 lm_eval \
  --model vllm \
  --model_args pretrained=$MODEL,add_bos_token=True,tensor_parallel_size=2,quantization=fp8,kv_cache_dtype=fp8_e4m3,max_model_len=2048,gpu_memory_utilization=0.6 \
  --tasks wikitext --batch_size 16
 
# Per-Token-Activation Per-Channel-Weight FP8 Quantization
MODEL=meta-llama/Llama-3.1-8B-Instruct
HIP_VISIBLE_DEVICES=0,1 lm_eval \
  --model vllm \
  --model_args pretrained=$MODEL,add_bos_token=True,tensor_parallel_size=2,quantization=ptpc_fp8,kv_cache_dtype=fp8_e4m3,max_model_len=2048,gpu_memory_utilization=0.6 \
  --tasks wikitext --batch_size 16

lm-evaluation-harness 命令(8B 模型 - 70B 模型请相应调整)

# FP8 (Per-Tensor)
MODEL=/app/model/Llama-3.1-8B-Instruct/  # Or Llama-3.1-70B-Instruct
lm_eval \
  --model vllm \
  --model_args pretrained=$MODEL,add_bos_token=True,quantization=fp8,kv_cache_dtype=fp8_e4m3 \
  --tasks gsm8k  --num_fewshot 5 --batch_size auto --limit 250
 
# PTPC FP8
MODEL=/app/model/Llama-3.1-8B-Instruct/  # Or Llama-3.1-70B-Instruct
lm_eval \
  --model vllm \
  --model_args pretrained=$MODEL,add_bos_token=True,quantization=ptpc_fp8,kv_cache_dtype=fp8_e4m3 \
  --tasks gsm8k  --num_fewshot 5 --batch_size auto --limit 250
 
# BF16
MODEL=/app/model/Llama-3.1-8B-Instruct/  # Or Llama-3.1-70B-Instruct
lm_eval \
  --model vllm \
  --model_args pretrained=$MODEL,add_bos_token=True,kv_cache_dtype=auto \
  --tasks gsm8k  --num_fewshot 5 --batch_size auto --limit 250