@vivekgalatage: 谈及推理时,《How to scale your model》一书中的专门章节不容错过。https:/…
摘要
本文分享了《How to scale your model》一书中关于推理的专门章节,涵盖了Transformer模型的关键优化技术,如KV缓存和延迟分析。
查看缓存全文
缓存时间: 2026/09/27 23:26
当涉及到推理时,万万不可错过《如何扩展你的模型》一书中关于该主题的专门章节。https://jax-ml.github.io/scaling-book/inference/…
Transformer 推理全解析
来源:https://jax-ml.github.io/scaling-book/inference/
目录
- 我们真正想要优化的是什么?(https://jax-ml.github.io/scaling-book/inference/#what-do-we-actually-want-to-optimize)
- 线性运算:是什么限制了我们?(https://jax-ml.github.io/scaling-book/inference/#linear-operations-what-bottlenecks-us)
- 注意力机制呢?(https://jax-ml.github.io/scaling-book/inference/#what-about-attention)
- LLM 延迟与吞吐量的理论估算(https://jax-ml.github.io/scaling-book/inference/#theoretical-estimates-for-llm-latency-and-throughput)
- 内存方面呢?(https://jax-ml.github.io/scaling-book/inference/#what-about-memory)
- 建模 LLaMA 2-13B 的吞吐量与延迟(https://jax-ml.github.io/scaling-book/inference/#modeling-throughput-and-latency-for-llama-2-13b)
- 预填充(https://jax-ml.github.io/scaling-book/inference/#prefill)
- 生成(https://jax-ml.github.io/scaling-book/inference/#generation)
- KV 缓存的分片(https://jax-ml.github.io/scaling-book/inference/#sharding-the-kv-cache)
- 连续批处理(https://jax-ml.github.io/scaling-book/inference/#continuous-batching)
- 前缀缓存(https://jax-ml.github.io/scaling-book/inference/#prefix-caching)
- 让我们来看一个实现:JetStream (https://jax-ml.github.io/scaling-book/inference/#let-s-look-at-an-implementation-jetstream)
Transformer 推理基础
假设你已经训练好了一个 Transformer 模型,并希望用它来生成一些新的序列。*说到底,基准分数上升、损失曲线下降,这些都只是代理指标,真正能衡量成果的,还是当它投入实际应用时是否足够出色!*历史上,你可以在Transformer上进行大量研究而完全不涉及推理——基于评分的选择题基准测试无需正确的KV缓存或生成循环实现就能高效运行。这意味着,尤其是在研究代码库中,推理代码路径中往往有很多低垂的果实(easy to fix issues)。 采样在概念上很简单。我们输入一个序列,我们心仪的 Transformer 会输出 (\log p(\text{下一个token}_i | \text{之前的tokens})),即所有可能下一个 token 的对数概率。我们可以从这个分布中采样得到一个新 token。将这个 token 追加上去,重复这个过程,我们就获得了一段延续了提示词内容的 token 序列。 图示: 从 Transformer 进行朴素采样。蓝色的 logits 给出了下一个 token 的分布,我们可以从中采样。请注意,每一步都重新处理了整个前缀,导致该算法的时间复杂度为 \Theta(n^2)。 我们刚刚描述了 Transformer 采样的朴素实现,虽然它有效,但我们实践中从不这样做,因为我们每生成一个 token 都会重新处理整个序列。这个算法在 FFW(前馈网络)上的复杂度是 (O(n^2)),在注意力机制上的复杂度是 (O(n^3)),用于生成 (n) 个 token! 如何避免这个问题? 与其每次都做完整的前向传播,事实证明我们可以保存每次前向传播中的一些中间激活值,从而避免重新处理之前的 token。具体来说,由于给定 token 在点积注意力中只关注之前的 token,我们可以简单地将每个 token 的 key 和 value 投影写入一个新的数据结构,称为 KV 缓存。一旦我们为过去的 token 保存了这些 key/value 投影,未来的 token 可以简单地计算它们的 (q_i \cdot k_j) 乘积,而无需对较早的 token 执行任何新的浮点运算。太棒了! 考虑到这一点,推理包含两个关键部分:
- 预填充:给定一个长提示,我们同时处理提示中的所有 token,并将生成的激活值(具体来说是 key-value 投影)保存到一个 “KV 缓存” 中。我们还会保存最后一个 token 的 logits。
- 生成:给定一个 KV 缓存和上一步的 logits,我们从中增量地采样一个 token,将该 token 反馈回 Transformer,并为下一步产生一组新的 logits。我们还将该新 token 的 KV 激活值追加到 KV 缓存中。我们重复此过程,直到遇到特殊 ``token 或达到某个最大长度限制。 这是一个使用 KV 缓存进行采样的示意图: 图示: 使用 KV 缓存进行高效 Transformer 采样的示意图。预填充处理我们的提示,并将每个 token 的 key-value 激活值保存在缓存中。生成获取此缓存(和最后一个 token 的 logits),采样一个新 token,并将该新 token 通过模型,关注 KV 缓存,并将新 token 的 key-value 投影保存回缓存。这是一个在 MLP 块中时间复杂度为 O(n) 的算法。 通过使用 KV 缓存进行采样,我们将生成 n 个 token 的时间复杂度降低为 FFW 上的 (O(n)) 和注意力上的 (O(n^2)),因为我们不再重新处理之前的 token。然而,生成一个序列仍然需要多次前向传播——当你向 Gemini 或 ChatGPT 提问并看到结果流式返回时,正是如此。每个 token(通常)都是一个单独(但部分缓存)的 Transformer 调用,指向一个庞大的模型。 我们很快会看到,预填充和生成是非常不同的两种任务——Transformer 推理实际上是伪装成一体的两个任务!与训练相比,KV 缓存也是一个全新的、显著增加复杂性的源头。
我们真正想要优化的是什么?
在我们深入之前,值得强调推理的一个全新方面:延迟。虽然训练时我们只关心吞吐量(每芯片每秒处理的总 token 数),但在推理时,我们必须担心生成 token 的速度(包括首个 token 延迟(TTFT)和每 token 延迟)。例如:
- 离线批处理推理用于评估和数据生成,只关心推理的总成本,而不关心单个样本的延迟。
- 聊天界面/流式任务需要在大规模下低成本运行,同时拥有低 TTFT 和足够快的生成速度以超过人类阅读速度。
- 边缘推理(例如在你的笔记本电脑上运行
llama.cpp)只需以尽可能低的延迟一次服务一个用户,可能还受到严格的硬件限制。最大化硬件利用率仍然至关重要,有助于降低成本和 TTFT,但与训练不同的是,这并不必然在所有情况下都意味着单个用户体验会更好。许多在加速器、系统和模型架构层面的优化都在延迟、吞吐量、上下文长度甚至模型质量之间进行权衡。
更细粒度的 Transformer 视图
到目前为止,我们大多将 Transformer 视为一个前馈块的堆叠。虽然从 FLOPs 和内存角度来看这通常是合理的,但这不足以恰当地建模推理。 你会在本节中注意到的一点是,推理比训练的容错性低得多。我们通常拥有的 FLOPs 少得多,批处理的机会更少,并且对延迟的敏感度要高得多。KV 缓存也极大地增加了推理的复杂性。 正如我们在第 4 部分(https://jax-ml.github.io/scaling-book/transformers)所看到的,Transformer 前向传播的主要组件是:
- 一堆线性运算,包括 MLP(W_{in}, W_{out})和注意力中的 QKV 投影以及输出投影(W_Q, W_K, W_V, 和 W_O)。这些都涉及从 HBM 读取参数和一批激活值,进行一些浮点运算,然后将结果写回 HBM。
- 点积注意力。我们需要从 HBM 读取一批 key-value 投影和一批 query 激活值,进行一些内积运算和 softmax 操作,然后将注意力结果写回 HBM。
- 其他所有操作,包括应用层归一化、激活函数、token 采样、更新 KV 缓存和位置嵌入。这些确实需要一些浮点运算,但会被上述运算所主导,或者与它们融合。 在接下来的几个小节中,我们将从预填充和生成的角度审视这些组件,并探讨什么最可能成为我们的性能瓶颈。在单个加速器内,我们是受计算限制还是受内存带宽限制?我们想强调的是,对于预填充和生成,答案会有多么不同。
线性运算:是什么限制了我们?
无论是在 MLP 块还是注意力中,我们所有的线性运算在概念上都是相同的。它们的算术强度取决于批处理大小。我们在第 1 节(https://jax-ml.github.io/scaling-book/roofline)中做过这个计算,但值得重复一下。 让我们来看一个单独的矩阵乘法:一个 $ \text{bf16}[B, D]$ 维的批次与一个 \text{bf16}[D, F] 维的矩阵相乘。这可以是大的 MLP 块(W_{in} 或 W_{out})或者较小的注意力投影(W_Q, W_K, W_V, W_O)之一。为了执行此矩阵乘法,我们需要将这两个数组从 HBM 加载到 MXU 中,进行乘法运算,然后将结果写回 HBM。和之前一样,我们有: [T_{math} = \frac{\text{计算 FLOPs}}{\text{加速器 FLOPs/s}} = \frac{2BDF}{\text{加速器 FLOPs/s}}] [T_{comms} = \frac{\text{通信字节数}}{\text{带宽 Bytes/s}} = \frac{2BD + 2FD + 2BF}{\text{带宽 Bytes/s}}] TPU 或 GPU 可以在进行计算的同时加载数据,从而重叠这些操作。因此,要成为计算受限的,我们需要 (T_{math} \geq T_{comms}),即: [ \frac{2BDF}{2BD + 2DF + 2BF} \geq \frac{\text{加速器 FLOPs/s}}{\text{带宽 Bytes/s}} \underset{\text{TPU v5e}}{=} \frac{1.97E+14}{8.20E+11} = 240 ] 其中右边是硬件的算术强度。 现在,假设 D 和 F 相对于 B 非常大(通常我们的批处理大小最多为 500,而 D 和 F > 10k),我们可以简化分母,利用 \small{2BD + 2DF + 2BF \approx 2DF},得到: [ \begin{align*} \frac{2BDF}{2BD + 2DF + 2BF} \approx \frac{2BDF}{2DF} \geq \frac{\text{加速器 FLOPs/s}}{\text{带宽 Bytes/s}} \ \underset{\text{TPU v5e}}{=} \frac{1.97E+14}{8.20E+11} \implies B \geq 240 = B_{\text{crit}} \end{align*} ] 如果我们量化权重或使用更低精度的 FLOPs 进行矩阵乘法,这个临界批处理大小会发生变化。例如,如果我们将权重量化为 int8 或 fp8,B_{\text{crit}} 会降低 2 倍。如果我们在 int8 或 fp8 下执行 FLOPs,B_{\text{crit}} 会增加 2 倍。因此,如果我们令 \beta = \text{每参数位数} / \text{每激活值位数},\alpha_{\text{hbm}} = C / W_{\text{hbm}},我们的临界批处理大小实际上是 B_{\text{crit}} = \beta \alpha_{\text{hbm}}。 要点: Transformer 矩阵乘法是计算受限的当且仅当每个副本的token批处理大小大于 B_{\text{crit}} = C / W_{\text{hbm}} \cdot (\text{每参数位数} / \text{每激活值位数}) = \beta \cdot \alpha_{\text{hbm}}。对于在 TPU v5e 上的 bf16 激活值,这个值是 240 个 token。对于 H100,大约是 280 个 token。 在训练期间,由于我们在非常大的批次上重复使用相同的权重,所有矩阵乘法都会具有很高的算术强度。这种高强度算术运算能力会延续到预填充阶段,因为用户的提示通常长达数百甚至数千个 token。 如前所述,TPUv5e 的硬件算术强度是 240,因此,如果将一个超过 240 个 token 的序列输入到在此硬件上以 bf16 运行的稠密模型,我们预计会成为计算受限的,一切顺利。技术上,短于此长度的提示可以批处理在一起以提高利用率,但这通常不是必要的。 要点: 在预填充期间,所有矩阵乘法基本上始终是计算受限的。因此,简单地最大化硬件利用率或 MFU(模型 FLOPs 利用率)就足以最大化每芯片吞吐量(成本)和延迟(TTFT 形式)。除非提示非常短,否则按提示级别进行批处理只会增加延迟,而预填充吞吐量只有小幅改善。 然而,在生成阶段,对于每个请求,我们只能一步一步地进行前向传播,因为步骤之间存在顺序依赖关系!因此,我们只能通过(简单地)将多个请求批处理在一起,在批次维度上并行化来(容易地)获得良好的利用率。我们稍后会详细讨论这个问题,但实际上,将许多并发请求批处理在一起而不影响延迟是很难的。因此,在生成阶段,要让硬件 FLOPs 达到饱和要困难得多。 要点: 在生成期间,总 token 批处理大小必须大于 B_{\text{crit}},才能在线性/前馈运算(在 TPU v5e 上,bf16 参数时为 240)上成为计算受限的。因为生成是串行的、逐 token 进行的,这要求我们将多个请求批处理在一起,这很难!值得注意的是这个数字有多大! 生成批大小为 240 意味着 240 个并发请求同时生成,对于稠密模型来说就是 240 个独立的 KV 缓存。这意味着在实践中很难实现,除非在一些批量推理场景下。相比之下,在预填充期间推送超过 240 个 token 是相当常规的,尽管随着稀疏性增加需要一些注意。 请注意,这个确切数字会因量化类型和硬件而异。 加速器通常在低精度下能提供更多的算术运算能力。例如,如果我们有 int8 参数但以 bf16 进行计算,临界批处理大小会降至 120。如果使用 int8 激活值和 int8 参数,它又会跳回 240,因为 TPUv5e 能提供 400 TOPs/s 的 int8 x int8 运算。
注意力机制呢?
当我们查看点积注意力操作时,事情会变得更复杂,特别是我们必须考虑 KV 缓存。让我们只看一个使用纯多头注意力的单个注意力头。在一次 Flash Attention 融合操作中,我们在这里做了相当多的简化,忽略了应用 softmax、掩码等非矩阵乘法 FLOPs。它们应该与计算或 HBM 读取重叠,但在某些 TPU 代数上做到这点可能并不平凡。虽然这些细节不改变主要信息——即 KV 缓存通常是内存带宽受限的——但它们值得关注:
- 从 HBM 读取形状为 \text{bf16}[B, T, D] 的 Q 激活值。
- 从 HBM 读取 KV 缓存,这是一对 \text{bf16}[B, S, D] 的张量。
- 在 QK 矩阵乘法中执行 2BSTD FLOPs。使用 Flash Attention,我们不需要将 \text{bf16}[B, S, T] 注意力矩阵写回 HBM。
- 在注意力 AV 矩阵乘法中执行 2BSTD FLOPs。
- 将结果张量 \text{bf16}[B, T, D] 写回 HBM。 将所有操作整合在一起,我们得到: [ \text{多头注意力算术强度} = \frac{4BSTD}{4BSD + 4BTD} = \frac{ST}{S+T} ] 对于预填充,S=T,因为我们在做自注意力,所以简化为 T^2 / 2T = T / 2。这很好,因为这意味着预填充期间注意力的算术强度是 \Theta(T)。这意味着在注意力上很容易成为计算受限。只要我们的序列长度足够大,我们就会很顺利! 但是,由于生成阶段的序列维度可以忽略不计(T=1),而 B 和 D 维度相互抵消,我们可以近似为: [ S \gg T = 1 \implies \frac{ST}{S+T} \approx 1 ] 这很糟糕,因为它意味着我们无法做任何事情来改善生成期间注意力的算术强度。我们
相似文章
@antiAIvo: 我已经学习了Transformer架构的基础知识、推理和训练流程。接下来,我将开始新一轮的深入探讨…
作者介绍了一门关于大模型推理优化的学习路径,涵盖KV缓存、连续批处理、分页注意力机制等核心议题,并对vLLM与SGLang展开对比分析,突出了该领域在AI部署中的前沿技术地位。
@TeachTheMachine: 使用Transformer模型:从训练到推理
本教程介绍如何使用Transformer模型,从训练到推理,重点讲解自回归生成、prefill(预填充)与decode(解码)阶段,以及用于高效推理的键值缓存。
@sohailmo: 如果你阅读并理解所有这些概念,你就掌握了推理优化基础知识的80/20
NVIDIA 推出了一系列关于 AI 模型协同设计的文章,从模型维度如何影响 GPU 性能开始,作者称这涵盖了推理优化基础知识的80/20。
@Hi_Mrinal: 这是关于KV缓存的最佳阅读,直观上非常好读 https://medium.com/@saad.ahmed1926q/kv-…
对语言模型中KV缓存的直观解释,涵盖token、嵌入、注意力机制,以及为什么KV缓存能提高推理效率。适合没有机器学习背景的读者。
AI推理工程指南(阅读时间约17分钟)
本指南解释了AI推理工程这一学科,涵盖了预填充和解码阶段的划分、从封闭模型到开放模型的转变,以及针对延迟、吞吐量和成本的优化技术。