系列文章目录

大模型推理 & memory bandwidth bound (1) - 性能瓶颈与优化概述
大模型推理 & memory bandwidth bound (2) - Multi-Query Attention
大模型推理 & memory bandwidth bound (3) - MLA
大模型推理 & memory bandwidth bound (4) - Speculative Decoding
大模型推理 & memory bandwidth bound (5) - Medusa



前言

“While training these layers is generally fast and simple, due to parallelizability across the length of the sequence, incremental inference (where such paralleization is impossible) is often slow, due to the memory-bandwidth cost of repeatedly loading the large “keys” and “values” tensors.” ——《Fast Transformer Decoding: One Write-Head is All You Need》

经过上一篇关于memory bandwidth bound的铺垫,本篇来讲一下《Fast Transformer Decoding: One Write-Head is All You Need》这篇论文。其在摘要部分的这句表述(如上所示)就强调了大模型在增量推理,也即Decode阶段由于memory bandwidth bound导致的推理效率低下的问题。作者提出了Multi-Query Attention技术,加速了大模型推理。
Multi-Query AttentionMulti-Head Attention的变体,本篇跟随论文的思路,分析对比Multi-Head AttentionMulti-Query Attention的性能,最后根据一个demo实测一下效果。关于注意力机制的前置知识本文不再赘述,如有需要可参考之前写的GLM-4 (4) - SelfAttention


一、Multi-Head Attention性能分析

1.Prefill Phase

如下是批量计算Multi-head Attention的方法,批量体现在1)多个序列,2)一个序列中的多个位置计算注意力,对应着大模型推理的Prefill阶段。
其中矩阵 X X Xquery部分的输入,通过与投影矩阵 P q P_q Pq 相乘得到 Q Q Q ,矩阵 M M Mkeyvalue部分的输入,通过投影矩阵 P k P_k Pk P v P_v Pv 得到 K K K V V V ;根据点积注意力计算公式计算得到 O O O ;再经过投影矩阵 P o P_o Po 得到输出 Y Y Y b b bbatch_size n n nquery的个数, m m mkey的个数, d d d 是输入维度, h h h 是多头注意力头的个数, k k k v v v 分别是keyvalue的在单个头中的维度。
在这里插入图片描述
上述计算过程性能分析如下,采用如下三点假设以简化分析:
m = n m=n m=n
k = v = d h k=v=\frac{d}{h} k=v=hd
n ≤ d n≤d nd

  1. 计算复杂度为 Θ ( b n d 2 ) Θ(bnd^2) Θ(bnd2)
    1)以 Q Q Q 的计算为例,计算复杂度 O ( b n h d k ) O(bnhdk) O(bnhdk) ,由于 k = v = d h k=v=\frac{d}{h} k=v=hd,计算法复杂度为 O ( b n d 2 ) O(bnd^2) O(bnd2)
    2)同理 K K K V V V 以及 Y Y Y 的计算复杂度为 O ( b n d 2 ) O(bnd^2) O(bnd2)
    3)logits O O O 的计算复杂度为 O ( b n 2 d ) O(bn^2d) O(bn2d)
    4)weights的计复杂度为 O ( b h n 2 ) O(bhn^2) O(bhn2)
    5)由于 n ≤ d n≤d nd ,总的计算复杂度就是 Θ ( b n d 2 ) Θ(bnd^2) Θ(bnd2)

  2. 内存访问复杂度 O ( b n d + b h n 2 + d 2 ) O(bnd+bhn^2+d^2) O(bnd+bhn2+d2)
    1)第一项表示 X X X M M M Q Q Q K K K V V V O O O Y Y Y 占用的内存大小;
    2)第二项表示logitsweights的内存大小;
    3)第三项表示 P q P_q Pq P k P_k Pk P v P_v Pv 以及 P o P_o Po 的内存大小;

  3. 内存访问复杂度 / 计算复杂度 = O ( 1 k + 1 b n ) O(\frac{1}{k} + \frac{1}{bn}) O(k1+bn1)

由于GPUTPU计算能力超过内存带宽memory bandwidth两个数量级,内存访问复杂度与计算复杂度的比值应该保持在一个较小数值上。就Multi-head Attention来说,只要挑选合适的参数显然是能做到这一点的,因此并不会出现memory bandwidth bound问题。

2.Decode Phase

对于自回归模型来说,在Decode阶段无法进行并行计算,此时query的个数 n n n应该是1。我们对这种情况进行相应的性能分析。
计算过程如下,函数名中的Incremental表示是自回归模型处于(增量)解码阶段,逐个token生成。可以看到输入 x x x 相较之前已经少了一个维度, p r e v K prev_K prevK p r e v V prev_V prevV 表示使用KV Cache技术缓存的前面tokenskeyvalue。整体计算和前面是一致的。
在这里插入图片描述

  1. 计算复杂度为 Θ ( b n d 2 ) Θ(bnd^2) Θ(bnd2)
    1)假设经过 n n n 个迭代,总的计算复杂度应该与前面相同,因为计算量没有发生变化;

  2. 内存访问复杂度 Θ ( b n 2 d + n d 2 ) Θ(bn^2d + nd^2) Θ(bn2d+nd2)
    1)前一项表示 K K K V V V 的内存大小,因为一个 p r e v K prev_K prevK 内存占用是 O ( b h m k ) = O ( b n d ) O(bhmk)=O(bnd) O(bhmk)=O(bnd) n n n 个迭代加起来就是 O ( b n 2 d ) O(bn^2d) O(bn2d)
    2)后一项表示 P q P_q Pq P k P_k Pk P v P_v Pv 以及 P o P_o Po 的内存大小;
    3)其他项比较小;

  3. 内存访问复杂度 / 计算复杂度 = Θ ( n d + 1 b ) Θ(\frac{n}{d} + \frac{1}{b}) Θ(dn+b1)

这种情况下要使得这一比值保持在一个较小的数是不容易的,特别是当序列长度 n n n 接近维度 d d d 的时候。因此,在大模型解码阶段,确实会受到memory bandwidth bound的影响。

二、Multi-Query Attention性能分析

为了解决上述问题,作者提出了Multi-Query Attention,改进是KeyValue都只保留一份,而不是原来 h h h 份。下图能很好的反映Multi-Head/Multi-Query/Grouped-Query Attention之间的区别。
在这里插入图片描述

1.Prefill Phase

Prefill阶段的计算过程如下,与Multi-Head Attention相比 K K K V V V 少了一个维度。其实到这边已经也可以预料到Muti-Head Attention在预填充阶段是肯定不会memory bandwidth bound的,但我们这边还是会计算一下。
在这里插入图片描述

  1. 计算复杂度为 Θ ( b n d 2 ) Θ(bnd^2) Θ(bnd2) :对比Multi-Head Attention K K K V V V 的计算操作少了一些;
  2. 内存访问复杂度还是 O ( b n d + b h n 2 + d 2 ) O(bnd+bhn^2+d^2) O(bnd+bhn2+d2)
  3. 内存访问复杂度 / 计算复杂度 = O ( 1 k + 1 b n ) O(\frac{1}{k} + \frac{1}{bn}) O(k1+bn1)Multi-Head Attention是一致的。

2.Decode Phase

下面是Decode阶段的计算过程,与Multi-Head Attention相比, K K K V V V 少了一个维度 h h h
在这里插入图片描述
考虑解码阶段的 n n n 个迭代,

  1. 计算复杂度为 Θ ( b n d 2 ) Θ(bnd^2) Θ(bnd2)
  2. 内存访问复杂度为 Θ ( b n d + b n 2 k + n d 2 ) Θ(bnd + bn^2k + nd^2) Θ(bnd+bn2k+nd2)
  3. 内存访问复杂度 / 计算复杂度 = Θ ( 1 d + n d h + 1 b ) Θ(\frac{1}{d} + \frac{n}{dh} +\frac{1}{b}) Θ(d1+dhn+b1)
    Multi-Head Attention Θ ( n d + 1 b ) Θ(\frac{n}{d} + \frac{1}{b}) Θ(dn+b1) 相比,Multi-Query Attention结果的第二项 n d h \frac{n}{dh} dhn 仅为 n d \frac{n}{d} dn 1 h \frac{1}{h} h1 ,也就是 Θ ( 1 d + n d h + 1 b ) Θ(\frac{1}{d} + \frac{n}{dh} +\frac{1}{b}) Θ(d1+dhn+b1) 的前两项控制的不错;因此只要增大批量大小 b b b 就能够突破memory bandwidth bound瓶颈,提供不错的加速效果。

总结

通过分析计算复杂度和内存访问复杂度的方式,确认了Multi-Head AttentionDecode阶段memory bandwidth bound的问题,以及Multi-Query Attention针对此问题的加速优化。论文中还探讨了Multi-Query Attention对模型生成质量的影响,这不是本系列关注的重点,故未展开。

Logo

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。

更多推荐