数据中心/云端

共同设计 AI 模型注意力,实现快速的交互式长上下文推理

随着代理式和长上下文工作负载变得常见,上下文长度增加,注意力消耗更大比例的推理时间 (图 1) 。由于注意力现在主宰了成本,因此模型的设计方式 (而不仅仅是实现方式) 越来越多地决定了模型的推理性能。围绕 GPU 的执行方式塑造模型架构是 AI 模型协同设计的前提。有关模型设计选择如何在不牺牲准确性的情况下影响吞吐量和交互性的讨论,请参阅上一篇文章《 AI 模型协同设计:硬件友好型 LLM 设计》。

本文将探讨组大小 (每个 KV 头的查询头) 、头维度和序列长度如何影响密集注意力的性能,其中每个查询都会关注序列长度上的所有键和值。蒸馏的分析以及注意力在 GPU 上的并行性,形成了四个实用指南:共同设计检查清单,帮助模型开发者提高 NVIDIA GPU 上的推理吞吐量和交互性。敬请关注有关稀疏注意力的博文。

每次分析都基于两个来源:来自 GEMM 形状算法的分析公式,以及使用 FP8 预填充和解码内核的测量数据,用于注意力计算和 KV 缓存。

缩略语 定义
PB 预填充批量大小
DB 解码批量大小
QH 查询头数量
KH KV 头数量 ( MHA 代表 KH = QH,GQA 代表 KH = QH/ G,MQA 代表 KH= 1)
\(G\) 组大小 = QH/ KH (查询头共享一个 KV 头)
高频率 头部尺寸 (通常为 64、128 或 256)
ISL 输入序列长度 (在预填充中查询 token)
KVSL 解码迭代中的平均 KV 缓存序列长度
表 1. 本文中介绍的方程中使用的符号 

如何预填充和解码两个不同的问题?

预填充并行处理完整的提示词,生成受计算限制的大型 GEMM-M (= ISL x \(G\)) 矩阵。无需预测性解码,解码一次即可生成一个 token,生成小型 GEMM-M (= \(G\)) matmuls,并通过从高带宽显存 (HBM) 读取 KV 缓存成为受内存限制的对象。

预测性解码可提升 GEMM-M 的性能,并可将解码转变为受计算限制的解码方式。由于预填充和解码的查询长度、KV 访问和瓶颈各不相同 (表 2) ,因此会针对每个阶段单独分析每个参数。

预填充 解码
查询长度 完整输入 ( ISL token) 1 个 token
KV 环境 提示 ( ISL token) 完整的 KV 缓存 ( KVSL token)
Attention GEMM-M ISL × \(G\) (大) \(G\) (小型)
主要瓶颈 计算 ( matmul + softmax) HBM 带宽 (显存)
表 2. 预填充和解码在查询长度、KV 上下文、注意力 GEMM-M 和主要瓶颈方面有所不同

注意:借助代理式应用和多回合应用中常见的前缀缓存,在处理大型前缀缓存时,新回合可能会出现较短的 ISL。使用短 ISL 但长的前缀缓存时,预填充的行为类似于解码。

算术强度如何控制计算受限行为与内存受限行为

如前文所述,计算和带宽上限限制了 GPU 性能。算术强度决定了哪些绑定 (方程 1) :

算术强度 = 总 FLOPS/ 访问的总字节数

岭点标志着从内存受限到计算受限的过渡。预填充位于其上方且受计算限制,而解码位于其下方且受内存限制 (图 2) 。预测解码会提高解码的算术强度,并使其向脊移动。

FlashAttention 内核如何在 GPU 上计算注意力?

FlashAttention 在不实现完整注意力矩阵的情况下计算注意力。它将 \(Q\)、\(K\) 和 \(V\) 的图块从 HBM 流式传输到片上 SRAM,并将三个步骤融合为一个通道:

  • 首先,batched matmul (BMM1) 对关键帧的查询进行评分
  • 其次,在线 Softmax 使用运行中的最大值和总和对分数进行归一化
  • 第三,第二批 matmul (BMM2) 对值进行加权

BMM 在 Tensor Core 上运行,而 softmax 指数在特殊功能单元上运行。BMM 形状驱动随后的算术强度分析。

GEMM 形状

注意力表现取决于两个 Matmul 的形状。表 3 列出了 BMM1 和 BMM2 的每相位 (批量、M、N、K) 维度。

BMM 阶段 批量 M N K 含义
BMM1 预填充 PB × KH ISL × \(G\) ISL 高频率 Q K:对关键帧进行评分查询
解码 DB × KH 1 × \(G\) KVSL 高频率
BMM2 预填充 PB × KH ISL × \(G\) 高频率 ISL 权重 = V:聚合值
解码 DB × KH 1 × \(G\) 高频率 KVSL
表 3. BMM1 和 BMM2 的 GEMM (批量、M、N、K) 形状

对于解码,GEMM-M = \(G\) (通常为 8-16) 远低于 GPU 图块 M ( 64 或 128) ,这限制了每个图块的并行工作。更大的 \(G\) 可降低每个 token 的 KV 负载,并在更多查询头之间摊销每个负载,从而提高利用率。下一节将量化此效果。

组规模

组大小 (\(G\)) 是共享一个 KV 头的查询头数量。MHA 具有 \(G\)+ 1,GQA 具有 \(G\)+ 4、8、16、…,MQA 具有 \(G\) = QH。

算术强度作为 \(G\) 的函数。在以下公式中,“字节”是指移动的 HBM 字节。为简单起见,假设每个元件 1 字节 (即 FP8 KV 缓存) 。

预填充

随着 \(G\) 的增长,1/ \(G\) 消失,运算强度接近 2* ISL。在 ISL = 32K 时,将 \(G\) 从 8 提高到 16 可将运算强度提升 6% 以下。换言之,预填充以 ISL 为主,而非 \(G\)。图 4 证实了这一点,通过将 \(G\) 从 1 (MHA) 改为 64 (MQA) ,预填充运行时的变化幅度不到 1%。方程 2、3 和 4:

FLOPS = 4 × PB × QH × ISL² × Hsz ( \(G\) 中的常量) 字节 = 2 × PB × KH × Hsz × ISL × (\(G\) + 1) 算术强度 = 2 × \(G\) × ISL / (\(G\) + 1) = 2 × ISL / (1 + 1/\(G\)) → 2 × ISL 作为 \(G\) → ∞

解码 (GEMM-M+ \(G\))

\(G\) 的解码运算强度翻倍。将 \(G\) 从 1 提高到 8 可减少内存流量并提高 GPU 计算利用率,从而将性能提升 8 倍。它独立于 KVSL:算术强度保持在 2% \(G\) 附近,因此解码仍然受内存限制,除非 \(G\) 非常大。NVIDIA Nemotron 3 等模型采用具有两个 KV 头的 GQA,从而提高解码效率。方程 5、6 和 7:

FLOPS = 4 × DB × QH × KVSL × Hsz ( \(G\) 中的常数)
字节 = 2 × DB × KH × Hsz × (\(G\) + KVSL)
算术强度 = 2 × \(G\) × KVSL / (\(G\) + KVSL) ≈ 2 × \(G\) (当 KVSL ≫ \(G\) 时)

图 4 显示,由于将 KV 头数量减半,每个 token 加载的数据减半,因此 \(G\) 的解码运行时间每翻一番大约会下降 2 倍。在 \(G\) = 16 之外,KVSL = 32K 曲线会变平。它的每步内核足够小,主要有两种成本:固定的设置和后处理开销,以及通过在 SM 中拆分 KV 以与几个 KV 头保持并行,从而减少闪存解码。KVSL+ 128K 内核的时间越长,就能更好地分摊这些成本,并继续跟踪 2 倍趋势。

注意:预测解码会将有效的 GEMM-M 提升至 (1 = \(D\)) = \(G\),其中 \(D\) 是草稿 token 的数量。一旦足够大,可以填充计算图块,解码就会受到计算限制。

准则 1:选择 \(G\) 以获得解码效率,并将其调高。在 \(G\) 中,预填充运行时间是固定的,而解码算术强度 ~2 \(G\),因此 \(G\) 越高,解码速度和 GPU 利用率就越高。预测解码是在给定 \(G\) 下提升性能的另一个手段。

头部尺寸

与组大小不同,头部维度 (Hsz) 不会影响算术强度。将 Hsz 翻倍会使 FLOP (方程 2 和 5) 和字节 (方程 3 和 6) 翻倍,两者的比率保持不变。然而,图 5 显示运行时间随着 Hsz 的增加而增加,因为注意力内核执行三种不同 Hsz 扩展的作业。

  • Matmul:随 Hsz 生长,但步骤一致。上篇文章推荐模型维度为 128 的倍数,以与 GPU 图块大小和缓存行宽度保持一致。部分填充的图块的成本与完整图块的成本相当,因此 Hsz+ 64 可支付 128。Hsz = 512 接近张量内存 (TMEM) 容量限制。这使 128 和 256 成为高效选择。上篇文章
  • 内存 ( KV 状态) :也会随 Hsz 增长。由于 GPU 以 128 字节的单位移动数据,因此当 Hsz 是 128 的倍数时,内存访问 (如 matmul) 的效率最高。
  • Softmax 的不同之处在于:其代价独立于 Hsz,因为它在无头维度的注意力评分矩阵 (查询 token + 键) 上运行。方程 8 和 9:

Softmax 运算 (预填充) ≈ PB × QH × ISL²
Softmax 运算 (解码) ≈ DB × QH × KVSL

两者之间的平衡决定了每个阶段的 Hsz 成本 (图 5) 。

预填充受计算限制 ( matmul 和 softmax) 。随着 Hsz 的增大,matmul FLOPS 也在增加,而 softmax 保持不变。如果预填充是纯 matmul,则双倍 Hsz 会使运行时间翻倍;但固定的 softmax 不会缩放,因此运行时间的增加幅度小于 Hsz。图 5 证实了这一点:预填充爬升采用 Hsz,但速度低于 Hsz。更宽的 Hsz 会摊销 softmax,将更多内核时间转移到 matmul,并减少预填充对 softmax 的限制。

解码受内存限制 (串流 KV 缓存) 。Hsz 越大,每个 token 的 KV 字节数就越多 (方程 6) ,因此运行时应随 Hsz 进行扩展。图 5 证实了这一点,不过这只是略微次线性的,因为设置、后处理和闪存解码减少用度不会随 Hsz 扩展,而在较短的 KVSL = 32K 内核上会更重。

准则 2:使用 128 或 256.Hsz 的 Hsz 不会改变算术强度,但必须与硬件保持一致。通常,Hsz = 64 仍可支付 128 宽的图块,而 Hsz = 512 则接近 TMEM 容量限制。因此,128 和 256 是最佳选择。

序列长度

序列长度 (ISL/ KVSL) 对预填充和解码的影响不同,因为它通过不同的变量进入每个阶段:预填充一起处理所有 ISL 输入 token,而每个解码步骤读取长度为 KVSL 的 KV 缓存。因此,两者的比例各不相同 (图 6) 。

预填充扩展

预填充执行 ISL+ 工作 (每个 token 关注每个 token) ,而 KV 流量仅与 ISL (等式) 成正比增长。2、3) 。因此,算术强度随 ISL 呈线性上升,使预填充远高于岭点并受计算限制。ISL 翻倍应使运行时间大约翻两番 (图 6) 。简而言之,由于固定设置和后处理开销占主导地位,因此扩展低于 4 倍;一旦 ISL 足够大,可以摊销,就会出现二次扩展。

线性解码

每个步骤都会读取完整的 KV 缓存以生成一个 token,因此字节数会随着 KVSL 的增长而增长,而每个步骤的工作量仍然很小 (方程 5 和 6) 。算术强度保持在 2 × \(G\) 附近,远低于岭点,因此解码始终受内存限制。KVSL 翻倍应使运行时翻倍,图 6 证实了这一点。简而言之,KVSL 的扩展性低于 2 倍,因为设置、后处理和闪存解码减少开销不会随着 KVSL 的扩展而扩展。随着 KVSL 的增长,他们所占的比例也在下降。

准则 3:尽可能降低有效 KV 状态。用例设置序列长度,但代价是不对称的:预填充随 ISL+ 增长,而解码随 KVSL 线性增长。通过 KV-cache 压缩、稀疏或滑动窗口注意力或混合模型架构 (如 Nemotron 3) 降低有效的 KV 状态,其中只有部分层承载不断增长的全局 KV 状态。

张量并行性将注意力头分割到 GPU 上

张量并行 (TP) 将注意力头分割到 GPU 上,从而为每个 GPU 提供 QH/ TP 查询头和 KH/ TP KV 头。它可以分割头部,而不是 token。在表 3 中,只有包含 KH/ TP 的批量维度缩小;每个 GPU GEMM 形状和运算强度保持不变。

TP 有一个实际限制:KV 头必须在 GPU 之间平均分配。一旦 TP > KH,组的查询头就会跨越多个 rank,每个 rank 都需要一个共享 KV 头的副本。这复制了 KV 状态,增加了内存和带宽开销,但却无济于事 (图 7) 。因此,保持 TP = KH,以便每个 GPU 至少拥有一个完整组:一个 KV 头及其 \(G\) 查询头。

KV 头数较少的模型 (例如,Nemotron 3 有两个 KV 头的模型) 可以快速耗尽 TP,因为缓存无法在每个 GPU 低于一个 KV 头的情况下进行分片,而不会出现重复。注意力必须以不同的方式扩展:注意力数据并行 (ADP) 分片请求,而 KV 并行 (KVP) 分片跨 GPU 的长序列 KV 缓存。FFN 通过专家并行 (EP) 单独扩展。

TensorRT-LLM 结合了 Wide EP (用于 FFN 的注意力加 EP 的 ADP) 和 Helix Parallelism (用于 FFN 的注意力加 EP 的 KVP) 。在这两种情况下,KH 决定了高效扩展。

准则 4:让 KH 确定并行策略。保持 TP = KH,以便每个 GPU 都有一个完整的 KV 头。KV 头数较少的模型 ( MQA 有 1 个,GQA 有 2 个) 可快速排气 TP,并且 ADP 或 KVP 更好地服务于 MoE FFN (在 TensorRT-LLM 中作为 Wide EP 和 Helix Parallelism 实现) 的注意力加 EP。

开始共同设计 AI 模型注意力

使用以下总结的四项准则作为模型设计检查清单,开始共同设计 AI 模型注意力。这些选择可以在同一硬件上提高 GPU 利用率、推理速度、吞吐量和交互性。

  • 准则 1:选择用于解码的组大小 (\(G\)) 并调高。预填充对群组大小不敏感。
  • 准则 2:使用头部维度 (Hsz) = 128 或 256 与 GPU 图块和 128 字节传输保持一致,同时保持在 TMEM 预算范围内。较大的打印头还会在预填充时隐藏 softmax。
  • 准则 3:通过 KV-cache 压缩、稀疏或滑动窗口注意力或混合模型降低有效的 KV 状态。
  • 准则 4:将并行性与 KV 头数 (KH) 相匹配。保持 TP ≤ KH,并使用 Wide EPHelix Parallelism 扩展几个 KH 模型。

致谢

本文由 NVIDIA 跨团队撰写。我们非常感谢 Timmy Liu、Jatin Mitra、Tiyasa Mitra、Bhargava Gopireddy、Brian Pharris、Julien Demouth 和 Eduardo Alvarez 提供的帮助。

标签