Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

7.3 使用单一矩阵 𝐗 并行化计算

前面对多头注意力和 Transformer 块其余部分的描述,采用的是在单条残差流中、单个时间步 ii 上计算单个输出的视角。不过,正如前文指出的,为每个词元计算 ai\mathbf{a}_i 时所执行的注意力计算,与其他每个词元的计算相互独立;对于 Transformer 块中从输入 xi\mathbf{x}_i 计算 hi\mathbf { h } _ { i } 的全部过程也是如此。这意味着,我们可以利用高效的矩阵乘法程序,轻松地将整个计算并行化。

具体做法是把输入序列中 NN 个词元的输入嵌入装入一个大小为 [N×d][ N \times d ] 的矩阵 X\mathbf{X}X\mathbf{X} 的每一行都是一个输入词元的嵌入。大语言模型所用 Transformer 的输入长度 NN 可以达到数十万词元;借助本文不予讨论的长上下文机制,还可以达到数百万词元。因此,对于普通 Transformer,可以把 X\mathbf{X} 看作一个有数十万行的矩阵,其中每一行的维度都是嵌入维度 dd(即模型维度)。

并行化注意力。 我们先考察单个注意力头的情形,然后转向多个注意力头,最后加入 Transformer 块中的其余组件。对于一个注意力头,将 X\mathbf{X} 分别乘以形状为 [d×dk][ d \times d _ { k } ] 的查询矩阵 WQ\mathbf { W } ^ { Q }、形状为 [d×dk][ d \times d _ { k } ] 的键矩阵 WK\mathbf { W } ^ { K },以及形状为 [d×dv][ d \times d _ { v } ] 的值矩阵 WV\mathbf { W } ^ { V },即可得到形状为 [N×dk][ N \times d _ { k } ] 的矩阵 QQ、形状为 [N×dk][ N \times d _ { k } ] 的矩阵 KK 和形状为 [N×dv][ N \times d _ { v } ] 的矩阵 VV;它们分别包含全部查询向量、键向量和值向量:

Q=XWQ;K=XWK;V=XWV(7.33)Q = \mathbf {X} \mathbf {W} ^ {Q}; K = \mathbf {X} \mathbf {W} ^ {K}; V = \mathbf {X} \mathbf {W} ^ {V}\tag{7.33}

有了这些矩阵,只需执行一次矩阵乘法,将 QQKK ^ { \top } 相乘,就能同时完成所有必要的查询—键比较。乘积的形状为 N×NN \times N,如图 7.9 所示。

Nq1·k1q1·k2q1·k3q1·k4
q2·k1q2·k2q2·k3q2·k4
q3·k1q3·k2q3·k3q3·k4
q4·k1q4·k2q4·k3q4·k4

图 7.9 形状为 N×NN \times NQKQK^\top矩阵,展示了它如何通过一次矩阵乘法完成所有 qikjq _ { i } \cdot k _ { j } 比较。

获得 QKQK^\top 矩阵后,我们可以非常高效地缩放这些分数、应用 softmax,再将结果乘以 VV,最终得到一个形状为 N×dN \times d 的矩阵,即输入中每个词元各有一个向量嵌入表示。这样一来,针对 NN 个词元的完整序列和单个注意力头,整个自注意力步骤便简化为如下计算:

head= softmax (mask(QKdk))V(7.34)\mathbf {h e a d} = \text { softmax } \left(\operatorname{mask} \left(\frac {Q K ^ {\top}}{\sqrt {d _ {k}}}\right)\right) V\tag{7.34}
A= head WO(7.35)\mathbf {A} = \text { head } \mathbf {W} ^ {O}\tag{7.35}

遮蔽未来信息。 你可能已经注意到,我们在上面的式 7.34 中引入了掩码函数。这是因为,按照目前的描述,自注意力计算存在一个问题:QKQ K ^ { \top } 会为每个查询与所有键之间的组合计算分数,其中也包括位于该查询之后的键。在语言建模场景中,这显然不合适——如果已经知道下一个词是什么,猜出它就太容易了!为解决这个问题,需要把矩阵上三角部分的元素设为 -\infty,softmax 会将其变为零,从而排除对序列后续词语的任何了解。实际做法是加入一个掩码矩阵 MM:当 j>ij > i 时(即上三角部分),Mij=M _ { i j } = - \infty;否则 Mij=0M _ { i j } = 0。图 7.10 展示了遮蔽后的 QKQ K ^ { \top } 矩阵。正如 7.7 节将要说明的,正是掩码的使用,使窗口中的每个位置都能充当一个训练样本,让 Transformer 的训练非常高效。第 9 章还将介绍,如何为需要利用未来词语的任务调整掩码。

图 7.11 以矩阵形式示意了单个注意力头中全部计算的并行化过程。

图 7.9 和图 7.10 还清楚地表明,注意力计算量相对于输入长度呈二次增长,因为在每一层中,都需要计算输入中每对词元之间的点积。这使得在非常长的文档(例如整部小说)上计算注意力代价高昂。尽管如此,现代大语言模型仍设法使用了包含数千乃至数万词元的相当长的上下文。

图 7.10 形状为 N×NN \times NQKQK^\top 矩阵,其中显示了 qikjq _ { i } \cdot k _ { j } 的值;比较矩阵的上三角部分被置零(先设为 -\infty,随后 softmax 会将其变为零)。

图 7.11 单个注意力头中并行注意力计算的示意图。第一行展示 QQKKVV 矩阵的计算;第二行展示 QKQ K ^ { \top } 的计算与遮蔽(图中省略了 softmax 计算和按维度归一化),以及随后对值向量进行加权求和以得到最终注意力向量的过程。

并行化多头注意力。 与自注意力相同,在多头注意力中,输入和输出的维度都是模型维度 dd,键嵌入和查询嵌入的维度为 dkd _ { k },值嵌入的维度为 dvd _ { v }(同样,在最初的 Transformer 论文中,dk=dv=64d _ { k } = d _ { v } = 64A=8A = 8,且 d=512d = 512)。因此,对于每个头 cc,都有形状为 [d×dk][ d \times d _ { k } ] 的权重层 WcQ\mathbf { W } ^ { Q } _ { c }、形状为 [d×dk][ d \times d _ { k } ]WcK\mathbf { W } ^ { K } _ { c },以及形状为 [d×dv][ d \times d _ { v } ]WcV\mathbf { W } ^ { V } _ { c }。它们分别与装入 X\mathbf{X} 的输入相乘,生成形状为 [N×dk][ N \times d _ { k } ]QQ、形状为 [N×dk][ N \times d _ { k } ]KK,以及形状为 [N×dv][ N \times d _ { v } ]VVAA 个注意力头中每个头的输出形状都是 [N×dv][ N \times d _ { v } ],所以拥有 AA 个头的多头层会产生 AA 个形状为 [N×dv][ N \times d _ { v } ] 的矩阵。为了在后续处理中使用这些矩阵,需要把它们拼接起来,生成维度为 [N×Adv][ N \times A d _ { v } ] 的单个输出。最后,再使用形状为 [Adv×d][ A d _ { v } \times d ] 的最终线性投影 WO\mathbf { W } ^ { O },将每个词元的输出恢复到原始维度。

把拼接后形状为 [N×Adv][ N \times A d _ { v } ] 的矩阵输出乘以形状为 [Adv×d][ A d _ { v } \times d ]WO\mathbf { W } ^ { O },即可得到形状为 [N×d][ N \times d ] 的自注意力输出 A\mathbf{A}

Qi=XWQi;Ki=XWKi;Vi=XWVi(7.36)Q ^ {i} = \mathbf {X} \mathbf {W} ^ {Qi}; \quad K ^ {i} = \mathbf {X} \mathbf {W} ^ {Ki}; \quad V ^ {i} = \mathbf {X} \mathbf {W} ^ {Vi}\tag{7.36}
headi= SelfAttention (Qi,Ki,Vi)= softmax ( mask (QiKiTdk))Vi(7.37)\mathbf {h e a d} _ {i} = \text { SelfAttention } (Q ^ {i}, K ^ {i}, V ^ {i}) = \text { softmax } \left(\text { mask } \left(\frac {Q ^ {i} K ^ {\mathrm{iT}}}{\sqrt {d _ {k}}}\right)\right) V ^ {i}\tag{7.37}
 MultiHeadAttention (X)=(head1head2headA)WO(7.38)\text { MultiHeadAttention } (\mathbf {X}) = \left(\mathbf {h e a d} _ {1} \oplus \mathbf {h e a d} _ {2} \dots \oplus \mathbf {h e a d} _ {A}\right) \mathbf {W} ^ {O} \tag {7.38}

结合并行输入矩阵。 使用 X\mathbf{X} 时,由完整的一层 NN 个 Transformer 块(每个块对应 NN 个输入词元之一)并行计算的函数,可以表示为:

O=X+ MultiHeadAttention ( LayerNorm (X))(7.39)\mathbf {O} = \mathbf {X} + \text { MultiHeadAttention } (\text { LayerNorm } (\mathbf {X}))\tag{7.39}
H=O+FFN( LayerNorm (O))(7.40)\mathbf {H} = \mathbf {O} + \operatorname{FFN} (\text { LayerNorm } (\mathbf {O}))\tag{7.40}

请注意,在式 7.39 中,我们用 X\mathbf{X} 表示该层的输入,无论这个输入来自哪里。对于第一层,正如下一节将说明的,这个输入就是我们一直用 X\mathbf{X} 表示的初始词嵌入与位置嵌入向量之和。但对于后续的第 kk 层,输入则是前一层的输出 Hk1\mathbf { \mathbf { H } } ^ {k - 1 }。我们也可以进一步分解一个 Transformer 层中的计算,用一个方程表示每项组件计算。下面用 T\mathbf{T}(形状为 [N×d][ N \times d ])代表 Transformer,并用上标区分块内的每步计算;同时仍用 X\mathbf{X} 表示来自前一层的输入或初始嵌入:

T1= LayerNorm (X)(7.41)\mathbf {T} ^ {1} = \text { LayerNorm } (\mathbf {X})\tag{7.41}
T2= MultiHeadAttention (T1)(7.42)\mathbf {T} ^ {2} = \text { MultiHeadAttention } (\mathbf {T} ^ {1})\tag{7.42}
T3=T2+X(7.43)\mathbf {T} ^ {3} = \mathbf {T} ^ {2} + \mathbf {X}\tag{7.43}
T4=LayerNorm(T3)(7.44)\mathbf {T} ^ {4} = \operatorname{LayerNorm} (\mathbf {T} ^ {3})\tag{7.44}
T5=FFN(T4)(7.45)\mathbf {T} ^ {5} = \operatorname{FFN} (\mathbf {T} ^ {4})\tag{7.45}
H=T5+T3(7.46)\mathbf {H} = \mathbf {T} ^ {5} + \mathbf {T} ^ {3}\tag{7.46}

这里,当我们使用 FFN(T3)\operatorname { FFN } ( \mathbf { T } ^ { 3 } ) 这样的记号时,指的是把同一个 FFN 并行应用于窗口中的 NN 个嵌入向量。同理,NN 个词元也各自在 LayerNorm 中并行完成归一化。至关重要的是,Transformer 块的输入和输出维度保持一致,因此这些块可以堆叠。由于块输入端的每个词元 xi\mathbf { x } _ { i } 都由维度为 [1×d][ 1 \times d ] 的嵌入表示,因此输入 X\mathbf{X} 和输出 H\mathbf{H} 的形状都是 [N×d][ N \times d ]