使用 ONNX Runtime 加速 Phi-2、CodeLlama、Gemma 及其他生成式 AI 模型

作者:

Parinita Rahi, Sunghoon Choi, Yufeng Li, Kshama Pawar, Ashwini Khade, Ye Wang

2024年2月26日

在速度和效率至关重要的快速发展领域中,ONNX Runtime (ORT) 允许用户轻松地将生成式 AI 模型的功能集成到其应用程序和服务中,通过改进的优化带来更快的推理速度并有效降低成本。这些优化包括最先进的融合和内核优化,有助于提高模型性能。最近发布的 ONNX Runtime 1.17 提升了多个生成式 AI 模型(包括 Phi-2、Mistral、CodeLlama、Orca-2 等)的推理性能。ONNX Runtime 是小型语言模型 (SLM) 从训练到推理的完整解决方案,与其他框架相比显示出显著的加速效果。通过支持 float32、float16 和 int4,ONNX Runtime 的推理增强功能提供了最大的灵活性和性能。

在本博客中,我们将介绍针对 Phi-2、Mistral、CodeLlama、SD-Turbo、SDXL-Turbo、Llama2 和 Orca-2 等最新 GenAI 模型在训练和推理方面的显著优化提速。对于这些模型架构,与 PyTorch 和 Llama.cpp 等其他框架相比,ONNX Runtime 在各种批次大小和提示词长度范围内显著提升了性能。现在,使用 Olive 也可以实现这些基于 ONNX Runtime 的优化。

快速链接

Phi-2

Phi-2 是由微软开发的一个拥有 27 亿参数的 Transformer 模型。它是一个展现出优秀推理和语言理解能力的小型语言模型 (SLM)。由于其体积小巧,Phi-2 是研究人员的一个极佳平台,他们可以探索机制可解释性、安全性改进以及针对不同任务的微调实验等各个方面。

ONNX Runtime 1.17 引入了支持 Phi-2 模型的内核更改,包括针对 Phi-2 的 Attention、Multi-Head Attention、Grouped-Query Attention 和 RotaryEmbedding 的优化。具体而言,已添加对以下内容的支持

  • Multi-Head Attention CPU 内核中的因果掩码 (causal mask)
  • Attention 和 Rotary Embedding 内核中的 rotary_embedding_dim
  • Grouped-Query Attention 内核中的 bfloat16

支持基于 TorchDynamo 的 Phi-2 ONNX 导出,并且优化脚本在此基础上构建。

对于 Phi-2 推理,在所有提示词长度下,使用 float16 和 int4 量化的 ORT 的表现均优于使用 float32 的 ORT、PyTorch 以及 Llama.cpp。

推理

使用 float16 带来的 ORT 收益

针对提示词吞吐量(即模型根据输入提示词处理和生成响应的速率)优化的 CUDA 性能比 PyTorch Compile 快高达 7.39 倍。我们还观察到,与 Llama.cpp 相比,对于更大的批次大小和提示词长度,ONNX Runtime 的速度显著更快。例如,当批次大小 = 16、提示词长度 = 2048 时,它快高达 13.08 倍

Token 生成吞吐量是前 256 个生成的 token 的平均吞吐量。带有 float16 的 ONNX Runtime 平均比 torch.compile 快 6.6 倍,最高可达 18.55 倍。它也比 Llama.cpp 快高达 1.64 倍

Phi2 float16 prompt throughput comparison Phi2 float16 token generation throughput comparison

使用 int4 带来的 ORT 收益

ORT 提供了对 int4 量化的支持。与 PyTorch 相比,带有 int4 量化的 ORT 可以提供高达 20.48 倍的性能提升。它平均比 Llama.cpp 好 3.9 倍,对于长序列长度,速度快高达 13.42 倍。由于具有针对 GemV 的特殊内核,带有 int4 量化的 ONNX Runtime 通常在批次大小为 1 时表现最佳。

Phi2 int4 prompt throughput comparison Phi2 int4 token generation throughput comparison
注意:torch.compile 在 4 位量化下的表现不佳。此外,Llama.cpp 不使用 FlashAttention,其 attention 实现对于长序列长度来说速度较慢。

  • Phi-2 基准测试是在 1 个 A100 GPU(SKU: Standard_ND96amsr_A100_v4)上进行的。软件包:torch: 2.3.0. dev20231221+cu121; pytorch-triton: 2.2.0+e28a256d71;ort-nightly-gpu: 1.17.0.dev20240118001;deepspeed: 0.12
  • 批次(Batch)是一组长度不同的输入句子;提示词长度指输入文本的大小或长度。

以下是 Olive 的 Phi-2 优化示例,该示例利用了本博客中强调的 ONNX Runtime 优化,并使用了易于使用的硬件感知模型优化工具 Olive

训练

除了推理之外,ONNX Runtime 还为 Phi-2 和其他 LLM 提供训练加速。ORT 训练是 PyTorch 生态系统的一部分,可通过 torch-ort python 包获取,该包是 Azure Container for PyTorch (ACPT) 的一部分。它提供灵活且可扩展的硬件支持,相同的模型和 API 既可用于 NVIDIA GPU,也可用于 AMD GPU。ORT 通过优化的内核和内存优化来加速训练,在缩短大型模型训练的端到端时间方面表现出显著的增益。这只需在模型中更改几行代码,用 ORTModule API 将其包装即可。它还可以与 DeepSpeed 和 Megatron 等流行的加速库组合使用,以实现更快、更高效的训练。

OpenAI 的 Triton 是一种领域特定语言和编译器,用于编写高效的自定义深度学习原语。ORT 支持 OpenAI Triton 集成(ORT+Triton),其中所有按元素运算符都被转换为 Triton 运算,并且 ORT 在 Triton 中创建自定义融合内核。

ORT 还执行稀疏性优化,以评估输入数据的稀疏性并利用此稀疏性执行图优化。这减少了计算 FLOP 需求并提高了性能。

基于低秩适应(LoRA)的微调通过仅训练少量附加参数(适配器)同时冻结原始模型的权重,使训练更加高效。这些适配器使模型适应特定任务。量化感知 LoRA (QLoRA) 将量化与 LoRA 相结合,使用较少的位数来表示权重,同时保持模型的性能和质量。ONNX Runtime 训练可与 LoRA 和 QLoRA 相结合,从而提高 LLM 的内存效率并加速训练时间。LoRA 和 QLoRA 技术使像 LLM 这样非常大的模型能够放入 GPU 内存中以高效完成训练。

使用 ORT 训练的 Phi-2 模型在性能上优于 PyTorch Eager 模式和 torch.compile。Phi-2 使用合成数据集和网络数据集的混合进行训练。我们针对 ORT 和 ORT+Triton 模式测量了收益,并且随着批次大小的增加,收益也在增加。该模型使用 DeepSpeed Stage-2 在 wikitext 数据集上进行了 5 个周期的训练,批次大小逐渐增加。V100 和 A100 的收益总结在下表中。

训练基准测试在 8 个 V100 上运行,并以每秒迭代次数衡量吞吐量(越高越好)

Phi2 training throughput comparison

下面的训练基准测试在 2 个 A100 上运行,并以每秒迭代次数衡量吞吐量(越高越好)

Phi2 training benchmarks on 2 A100 注意:使用了 PyTorch Stable 2.2.0 和 ONNXRuntime Training: Stable 1.17.0 版本。

Mistral

推理

Mistral7B 是一个拥有 70 亿参数的预训练生成式文本 LLM。对于 float16 和 int4 模型,ONNX Runtime 都显著提升了 Mistral 的推理性能。使用 float16 时,与 Llama.cpp 相比,ONNX Runtime 高达 9.46 倍。对于批次大小 1,使用 int4 量化时,token 生成吞吐量显著提高,比 PyTorch Eager 快高达 18.25 倍

Mistral float16 prompt throughput comparison Mistral float16 token generation throughput comparison Mistral int4 prompt throughput comparison Mistral int4 token generation throughput comparison

您现在可以在 Huggingface 上访问优化的 Mistral 模型:点击此处。

训练

与 Phi-2 类似,Mistral 也受益于使用 ORT 的训练加速。我们使用以下配置训练了 Mistral-7B,以观察 ORT 带来的收益(包括与 LoRA 和 QLoRA 组合使用时)。该模型使用 DeepSpeed Stage-2 在 wikitext 数据集上以批次大小 1 训练了 5 个周期。

Mistral training benchmarks

CodeLlama

Codellama-70B 是在 Llama-2 平台上开发的专注于编程的模型。该模型可以编写代码并使用自然语言生成围绕代码的讨论。由于 CodeLlama-70B 是微调后的 Llama 模型,因此可以直接应用现有的优化。我们将 4 位量化的 ONNX 模型与 PyTorch Eager 和 Llama.cpp 进行了比较。对于提示词吞吐量,对于所有批次大小,ONNX Runtime 都比 PyTorch Eager 至少快 1.4 倍。对于任何批次大小,ONNX Runtime 生成 token 的平均速度比 PyTorch Eager 高 3.4 倍;对于批次大小 1,比 Llama.cpp 高 1.5 倍

CodeLLama int4 prompt throughput comparison CodeLLama int4 token generation throughput comparison

SD-Turbo 和 SDXL-Turbo

当与 SD TurboSDXL Turbo 一起使用时,ONNX Runtime 提供了推理性能优势,并且还使这些模型能够在 Python 之外的其他语言(如 C# 和 Java)中访问。对于评估的所有(批次大小,步数)组合,ONNX Runtime 的吞吐量均高于 PyTorch,其中 SDXL Turbo 模型的吞吐量提升了高达 229%,SD Turbo 模型提升了 120%。ONNX Runtime CUDA 在处理动态形状方面特别出色,但在静态形状方面它也显示出比 PyTorch 显著的优势。

Stable Diffusion XL Turbo Speedup

要了解有关使用 ONNX Runtime 加速 SD-Turbo 和 SDXL-Turbo 推理的更多信息,请查看我们最近与 Hugging Face 合作发布的博客

Llama-2

我们针对 ORT 推理中 Llama-2 的改进单独发布了一篇博客,详见此处。此外,Llama-2-7B 和 Llama-2-13B 在使用 ORT 进行训练时也表现出良好的收益,特别是当与 LoRA 和 QLoRA 结合使用时。这些脚本可用作使用 Optimum 通过 ORT 微调 Llama-2 的示例。以下数据是使用 DeepSpeed Stage-2 在 wikitext 数据集上以批次大小 1 训练 5 个周期的使用 ORT 训练的 Llama-2 模型。

Llama2 training benchmarks

Orca-2

推理

Orca-2 是一个仅用于研究的系统,它在诸如使用用户提供的数据进行推理、理解文本、解决数学问题和总结文本等任务中提供一次性答案。Orca-2 有两个版本(70 亿和 130 亿参数);它们都是通过在定制的高质量人工数据上微调各自的 Llama-2 基础模型制成的。ONNX Runtime 有助于优化 Orca-2 推理,使用了类似于 Llama-2 的图融合和内核优化。

使用 int4 带来的 ORT 收益

Orca-2-7B int4 量化性能对比表明,与 PyTorch 相比,提示词吞吐量性能提升了高达 26 倍,token 生成吞吐量提升了高达 16.5 倍。与 Llama.cpp 相比,它还显示提示词吞吐量提升了 4.75 倍以上,token 生成吞吐量提升了 3.64 倍

Orca2 7b int4 prompt throughput comparison Orca2 7b int4 token generation throughput comparison Orca2 13b int4 prompt throughput comparison Orca2 13b int4 token generation throughput comparison

Orca-2 7b 与 ONNX runtime float16 性能对比也显示出提示词和 token 生成吞吐量的显著提升。

Orca2 7b float16 prompt throughput comparison Orca2 7b float16 token generation throughput comparison Orca2 13b float16 prompt throughput comparison Orca2 13b float16 token generation throughput comparison

Orca-2 基准测试在 1 个 A100 GPU 上进行,SKU: Standard_ND96amsr_A100_v4,软件包:torch 2.2.0, triton 2.2.0, onnxruntime-gpu 1.17.0, deepspeed 0.13.2, llama.cpp - commit 594fca3fefe27b8e95cfb1656eb0e160ad15a793, transformers 4.37.2

训练

Orca-2-7B 也受益于使用 ORT 的训练加速。我们使用 LoRA 并在启用稀疏性优化的条件下训练了序列长度为 512 的 Orca-2-7B 模型,并看到了性能的大幅提升。以下数据是使用 DeepSpeed Stage-2 在 wikitext 数据集上以批次大小 1 训练 5 个周期的使用 ORT 训练的 Orca-2-7B 模型。

Orca2 training benchmarks 使用 ACPT 镜像:nightly-ubuntu2004-cu118-py38-torch230dev:20240131

Gemma

Gemma 是一系列轻量级开源模型,基于 Google 用于创建 Gemini 模型的研究所和技术构建。它提供两种尺寸:2B 和 7B。每种尺寸都发布了预训练和指令微调变体。ONNX Runtime 可用于优化和高效运行任何开源模型。我们针对 Gemma-2B 模型进行了基准测试,带有 float16 的 ONNX Runtime 比 PyTorch Compile 快高达 7.47 倍,比 Llama.cpp 快高达 3.47 倍。带有 int4 量化的 ORT 比 PyTorch Eager 快高达 19.81 倍,比 Llama.cpp 快 2.62 倍

Gemma2b int4 token generation throughput comparison Gemma2b token generation throughput comparison

结论

总之,ONNX Runtime (ORT) 为多个模型(包括 Phi-2、Mistral、CodeLlama、SDXL-Turbo、Llama-2、Orca-2 和 Gemma)提供了显著的性能提升。ORT 提供最先进的融合和内核优化,包括对 float16 和 int4 量化的支持,从而实现更快的推理速度和更低的成本。在提示词和 token 生成吞吐量方面,ORT 的表现优于 PyTorch 和 Llama.cpp 等其他框架。ORT 在训练 LLM 方面也显示出巨大的优势,随着批次大小的增加,收益也随之增加,并且它可以与最先进的技术很好地结合,以实现高效的大模型训练。