深度系列教程手把手讲解 Transformer 核心机制,帮助程序员透彻理解 LLM 原理和实现细节。
这是我跟随 Sebastian Raschka 著作《Build a Large Language Model (from Scratch)》的第八篇博文。我在博客中分享一些引起我兴趣的片段,以及让我费脑筋的内容——既是为了梳理自己的思路,也希望能帮助其他正在学习这本书的人。距离我上次更新已经过了将近一个月——如果你怀疑我在写关于写博文的文章,并花时间在这个网站上让 LaTeX 工作,实际上是在拖延因为下一部分注定会很难,那你 100% 说对了!好消息是——这种情况往往就是这样——当我真正深入时,它的难度并没有那么大。我重拾了势头。
如果你通过那些关于博写的文章找到这个博客,欢迎!那些文章不是特别代表我的常规风格,我希望你会喜欢我回到正常形式的这一篇。
这次我讲解的是第 3.4 节,"用可训练的权重实现自注意力"。我们如何创建一个系统,它能学会如何解释在查看句子中的其他词时应该对这些词给予多少注意力——例如,学会在"the fat cat sat on the mat"中,当你查看"cat"时,单词"fat"很重要,但当你查看"mat"时,"fat"就不那么重要了?
在深入探讨之前,特别是考虑到距上一篇的时间差,让我们先从 1000 英尺的高度看一下 GPT 类仅解码器 Transformer 基础的 LLM(以下简称"LLM"以避免手腕疲劳)的工作原理。对于每一步,我都链接到我详细讲解的那些文章。
你从一个字符串开始,通常是单词。(第 2 部分)
你将其拆分为令牌("the"这样的单词,或"semi"这样的块)。(第 2 部分)
LLM 的任务是预测下一个令牌,基于到目前为止字符串中的所有令牌。(第 1 部分)
步骤 1:将令牌映射到称为令牌嵌入的向量序列。一个特定的令牌,比如"the",会有一个特定的嵌入——这些一开始是随机的,但 LLM 在训练时会学到有用的嵌入。(第 3 部分)
步骤 2:生成另一个位置嵌入序列——与令牌嵌入大小相同的向量,也一开始是随机的但可训练,代表"这是第一个令牌"、"这是第二个令牌"等等。(第 3.1 部分)
步骤 3:将两个序列相加生成新的输入嵌入序列。第一个输入嵌入是第一个令牌嵌入加上第一个位置嵌入(按元素相加),第二个是第二个令牌嵌入加上第二个位置嵌入,以此类推。(第 3 部分)
步骤 4:自注意力。取输入嵌入,对于每一个,生成一个注意力分数列表。这些数字代表在考虑相关令牌时,应该对每个其他令牌给予多少注意力。所以(假设每个词一个令牌)在"the fat cat sat on the mat"中,令牌"cat"需要一个 7 个注意力分数的列表——对第一个"the"给予多少注意力,对"fat"给予多少,对它自己"cat"给予多少,对"sat"给予多少,等等。它究竟如何做到这一点正是这一章节的内容——到现在为止我们一直在使用一个"玩具"示例计算。(第 4 部分,第 5 部分,第 6 部分,第 7 部分)。
步骤 5:将注意力分数归一化为注意力权重。我们希望每个令牌的注意力权重列表加起来等于 1——我们通过运行每个列表通过 softmax 函数来实现这一点。(第 4 部分,第 5 部分,第 6 部分,第 7 部分)。
步骤 6:生成新的上下文向量序列。在我们迄今为止构建的系统中,这为每个令牌包含所有输入嵌入乘以其各自的注意力权重,然后将结果相加的结果。所以在上面的例子中,"cat"的上下文向量将是第一个"the"的输入嵌入乘以"cat"对那个"the"的注意力分数,加上"fat"的输入嵌入乘以"cat"对"fat"的注意力分数,以此类推序列中的每个其他令牌。(第 4 部分,第 5 部分,第 6 部分,第 7 部分)。
完成所有这些后,我们有一个上下文向量序列,每一个都应该以某种方式代表其各自令牌的含义,包括它从所有其他令牌获得的那些含义片段。所以"cat"的上下文向量会包含它的某种"fatness"的暗示。
接下来对这些上下文向量会发生什么,使得 LLM 能使用它们来预测下一个令牌可能是什么?这部分还有待解释,所以我们必须等着看。但首先要学习的是我们如何创建一个可训练的注意力机制,它能取输入向量并生成注意力分数,这样我们就可以计算出上下文向量。
Raschka 在这一部分给出的答案叫做缩放点积注意力。他对代码的讲解清晰明了,但我花了整个周末才真正理解它的工作原理。所以,与其逐段讲解这一部分,我会呈现我自己对它如何工作的解释——这样可以避免我将来试图回忆它时头痛,也许也能拯救其他人的额头免受同样的折磨。
提前的总结
我长期以来喜欢 Pimsleur 风格的语言课程,他们会在每个教程开始时播放大约一分钟的你正在学习的语言对话,然后说"在 30 分钟后,你会再次听到这个,你会理解它"。你完成课程,他们再次播放对话,你确实理解了。
所以这是一个关于自注意力如何工作的压缩总结,用我自己的话,基于 Raschka 的解释。现在看起来可能像一堵术语的墙,但(希望)当你读完这篇博文时,你会重新阅读它,一切都会变得有意义。
我们有一个长度为 n 的输入令牌序列。我们已将其转换为一个输入嵌入序列,每个都是长度为 d 的向量——其中每一个都可以被视为 d 维空间中的一个点。让我们用这样的值表示该嵌入序列:x1, x2, x3, ...xn。我们的目标是生成一个长度为 n 的序列,由上下文向量组成,每个都代表各自输入令牌在整个输入语境中的含义。这些上下文向量中的每一个的长度都是 c(实际中通常等于 d,但理论上可以是任何长度)。
我们定义三个矩阵,查询权重矩阵 Wq、键权重矩阵 Wk 和值权重矩阵 Wv。它们由可训练的权重组成;每一个的大小都是 d×c。由于这些维度,我们可以将它们视为将长度为 d 的向量——d 维空间中的一个点——投影到长度为 c 的向量——c 维空间中的一个点——的操作。我们将这些投影空间称为键空间、查询空间和值空间。例如,要将输入向量 xm 转换为查询空间,我们只需将其乘以 Wq,像这样 qm=xmWq。
当我们考虑输入 xm 时,我们想为序列中的每个输入(包括它自己)计算其注意力权重。第一步是计算注意力分数,当考虑另一个输入 xp 时,通过取 xm 投影到查询空间的点积与 xp 投影到键空间的点积来计算。对所有输入执行此操作为 xm 提供了每个其他令牌的注意力分数。我们随后将这些分数除以我们投影到的空间维度的平方根,即 c,并运行结果列表通过 softmax 函数使它们全部加起来等于 1。这个列表就是 xm 的注意力权重。这个过程被称为缩放点积注意力。
下一步是为 xm 生成一个上下文向量。这就是所有输入投影到值空间的总和,每一个乘以其关联的注意力权重。
通过对每个输入向量执行这些操作,我们可以生成一个长度为 n 的列表,由长度为 c 的上下文向量组成,每一个都代表一个输入令牌在整个输入语境中的含义。
重要的是,巧妙地使用矩阵乘法,所有这些都可以为序列中的所有输入完成,为每一个生成一个上下文向量,只需五次矩阵乘法和一次转置。
首先,如果有人能不提前了解注意力机制就理解了上面的全部内容,那我向你致敬!内容确实很密集,希望它不像我朋友 Jonathan 写的那些令人费解的 git 使用指南。对我而言,要达到理解的程度,我花了八遍去读 Raschka(非常清晰易懂)的解释。我认为这也值得注意的是,它是一个非常"机制性"的解释——它说明我们如何进行这些计算,但没有说明为什么。我认为"为什么"其实不在这本书的范围内,但这让我很着迷,我很快会在博客上讨论它。[更新:这是"为什么"的文章。]不过,要理解"为什么",我认为我们需要对"如何"有坚实的基础,所以让我们在这篇文章中深入探讨。
到本书的这一部分为止,我们一直通过对输入嵌入相互进行点积来计算注意力分数——也就是说,当你看 xm 时,xp 的注意力分数就是 xm·xp。我之前怀疑 Raschka 在他的"玩具"自注意力中使用该特定操作的原因是实际实现是相似的,这被证明是对的,因为我们在这里做的是缩放点积。但我们要做的是先调整它们——被考虑的 xm 先乘以查询权重矩阵 Wq,而另一个 xp 乘以键权重矩阵 Wk。Raschka 把这称为投影,对我来说这是看待它的一个非常好的方式。但他的参考只是顺带一提,对我来说需要更多的挖掘。
如果你的矩阵数学有点生疏——就像我的一样——而且你还没有读我上周发布的入门指南,那么你现在可能想看一下。
从你的学生时代,你可能还记得矩阵可以用来应用几何变换。例如,如果你取一个代表点的向量,你可以将它乘以一个矩阵来旋转该点围绕原点。你可以使用这样的矩阵来逆时针旋转θ度:
由于这是矩阵乘法,你可以加上更多的点——也就是说,如果第一个矩阵有更多行,每一行都是你想旋转的点,同样的乘法会将它们全部旋转θ。所以你可以把矩阵看作是一个函数,它将点的集合映射到它们旋转后的等价物。这也适用于更高维度——在 3d 图形中,人们用这样的 2×2 矩阵来表示 2 维空间中的变换,例如 3×3 矩阵来对构成 3d 对象的点进行类似的变换。²
看待这个 2×2 矩阵的另一种方式是,它是一个将点从一个 2 维空间投影到另一个空间的函数,目标空间是第一个空间逆时针旋转θ度。对于这样一个简单的 2d 例子,或者甚至是 3d 例子,这不一定是更好的看待方式。这是一个哲学差异而不是实践差异。
但想象一下,如果矩阵不是正方形——也就是说,它的行数与列数不同。如果你有一个 3×2 矩阵,它可以用来乘以 3d 空间中的向量矩阵,并产生 2d 空间中的矩阵。记住矩阵乘法的规则:一个 n×3 矩阵乘以一个 3×2 矩阵会给你一个 n×2 的。
这实际上非常有用;如果你做过任何 3d 图形工作,你可能记得视锥体矩阵,它用于将你处理的 3d 点转换为屏幕上的 2d 点。不详细说明,它允许你用一次矩阵乘法将那些 3d 点投影到 2d 空间。
所以:一个 d×c 矩阵可以被看作是一种方式,它将代表 d 维空间中的点的向量投影到代表不同 c 维空间中的点的向量。
我们在自注意力中做的是拿我们的 d 维向量,这些向量构成输入嵌入序列,然后将它们投影到三个不同的 c 维空间,并使用投影后的版本工作。我们为什么这么做?这是我想在未来的"为什么"文章中探讨的问题,但现在,我认为有一点相当清楚的是,因为这些投影是作为训练的一部分学习的(记住,我们用于投影的三个矩阵由可训练的权重组成),它把某种间接性混入了进来,我们之前使用的简单点积注意力没有这种间接性。
坚持这个机制性的观点——"如何"而不是"为什么"——现在,让我们看看计算以及矩阵乘法如何使它们高效。我大致会遵循 Raschka 的解释,但使用数学符号而不是代码,因为(对我这个职业技术人来说不寻常)我发现这样更容易理解正在发生的事情。
我们坚持考虑标记 xm 并尝试计算它对 xp 的注意力分数的情况。我们做的第一件事是将 xm 投影到查询空间,我们通过将它乘以查询权重矩阵 Wq 来做:
现在,让我们通过将 xp 乘以键权重矩阵 Wk 将 xp 投影到键空间:
我们的注意力分数定义为这两个向量的点积:
所以我们可以写一个简单的循环,迭代所有输入 x1...xn 一次,为每个生成到查询空间的投影,然后在该循环内迭代 x1...xn 第二次,将它们投影到键空间,进行点积,并将这些存储为注意力分数。
但这样做会很浪费!我们在做矩阵乘法,所以我们可以批量处理。让我们首先考虑输入投影到键空间的情况;在我们假设的循环中,每次都会是相同的。所以我们可以一次性完成。让我们将输入序列视为这样的矩阵 X:
我们的输入序列中每个输入嵌入有一行 x1、x2 等等,行由该嵌入中的元素组成。所以它有 n 行,输入序列中每个元素一行,d 列,每个输入嵌入的维度一列,所以是 n×d。(我在这里使用 d=3 作为例子,就像 Raschka 在书中做的那样。)
这就像上面旋转矩阵例子中点的矩阵一样,所以我们可以一次性将其投影到键空间,只需将它乘以 Wk。让我们称结果为 K:
它看起来像这样(再次,像 Raschka 一样,我使用一个 2 维的键空间——也就是说,c=2——所以很容易看出矩阵是在原始 3d 输入嵌入空间还是 2d 投影的空间中):
...其中每一行都是输入 xn 投影到键空间。这只是所有投影堆叠在一起。
现在,让我们考虑那个点积——这是早前的那一位:
我们现在有一个包含所有 kn 值的矩阵 K。当你进行矩阵乘法时,输出矩阵中第 i 行第 j 列的元素 Mi,j 是第一个矩阵第 i 行(作为向量考虑)与第二个矩阵第 j 列(同样作为向量考虑)的点积。
听起来我们可以利用它批量进行所有的点积。让我们将 qm(我们的第 m 个输入标记投影到查询空间的结果)视为一个单行矩阵。我们能像这样乘以键矩阵吗
很遗憾不能。qm 是一个单行矩阵(大小为 1×c)而 K 是我们的 n×c 键矩阵。在矩阵乘法中,第一个矩阵的列数——在这个情况下是 c——需要与第二个矩阵的行数相匹配,即 n。但是,如果我们转置 K,基本上是交换行和列:
...那么我们有一个 1×c 矩阵乘以一个 c×n 矩阵,这是有意义的——而且更好的是,它是所有 p 值对所有 (qm, kp) 对的每个点积——也就是说,通过两个矩阵乘法——一个来计算 K,这个来计算,加上转置,我们已经计算出我们的输入序列中元素 xm 的所有注意力分数。
首先,让我们做与我们投影输入序列到键空间相同的事情,将其全部投影到查询空间。我们计算了 K=XWk 来计算键矩阵,所以我们可以用相同的方式计算查询矩阵,Q=XWq。就像 K 是所有投影到键空间的输入向量"堆叠"在一起一样,Q 是所有投影到查询空间的输入向量。
现在,如果我们将其乘以转置的键矩阵会发生什么呢?
好的,我们的Q矩阵每个输入占一行,投影空间中的每个维度占一列,所以是n×c的。而我们知道,转置后的K矩阵是c×n的。所以我们的结果是n×n的——由于矩阵乘法是按点积定义的,它包含的是Q中每一行(输入转换到查询空间)与K^T中每一列(输入转换到键空间)的点积。
我们的计划正是通过计算这些点积来生成注意力分数!
所以用三次矩阵乘法,我们就做到了:
...其中我用大写Ω表示一个矩阵,其中每一行代表序列中的一个输入,该行中的每一列代表该输入的一个注意力分数。元素Ωm,p表示在计算xm的上下文向量时,对输入xp要付出多少注意力。它通过计算xm投影到查询空间与xp投影到键空间的点积来实现这一点。
"缩放点积注意力"的"点积"部分就完成了 :-)
所以我们已经计算出了注意力分数。接下来我们需要对它们进行归一化;过去我们使用了softmax函数。这个函数接受一个列表,调整其中的值使它们都加起来等于1,但会提升较大的数字,降低较小的数字。我想它被命名为"soft" "max"是因为它就像找最大值,但在某种意义上更"软",因为它保留了其他较小的数字并降低它们。
Raschka解释说,当我们处理大量维度时——在现实的LLMs中,d和c很容易达到数千——使用纯softmax可能导致梯度变小——他说它可能开始表现得"像阶跃函数",我理解这意味着你最后会让列表中除了最大数字外的所有数字都缩放到极小,而最大的数字占主导。所以,作为一个解决方案,我们将数字除以投影空间维度数c的平方根,然后再通过softmax。3
记住Ω是一个注意力分数矩阵,每个输入令牌占一行,所以我们需要分别对每一行应用softmax函数。我们最后得到的是:
(axis=1不是真正的数学记号,它只是我从PyTorch借来的用法,表示我们在矩阵上按行应用softmax。)
一旦我们做到这一点,我们就有了归一化的注意力分数——也就是注意力权重。下一个也是最后一个步骤,是使用这些来计算上下文向量。
让我们重申一下如何计算上下文向量。在之前的玩具示例中,对于每个令牌,我们取输入嵌入,将其中每一个乘以其注意力权重,逐元素求和,这就是结果。现在我们做同样的事情,但首先将输入嵌入投影到另一个空间——值空间。所以让我们首先进行投影,就像我们对其他空间所做的那样,作为一个简单的矩阵乘法:
现在,从上面我们有我们的注意力权重矩阵A,它在第m行包含对于输入xm的输入序列中每个令牌的注意力权重——也就是说,在Am,p处,我们有当计算输入m的上下文向量时输入p的注意力权重。这意味着对于长度为n的输入序列,它是一个n×n的矩阵。
在我们的值矩阵V中,我们也有每个输入占一行。第m行中的值,视为一个向量,是输入xm投影到值空间中的结果。所以它是一个n×c的矩阵。
如果我们进行矩阵乘法会怎样?根据矩阵乘法的规则,我们会得到某种n×c的矩阵,但它意味着什么呢?
重申一下,矩阵乘法的规则是Mi,j的值——也就是输出矩阵中第i行第j列的元素——是第一个矩阵中第i行(视为向量)与第二个矩阵中第j列(也视为向量)的点积。
所以,在位置(1,1)——第一行第一列,我们有A中第一行的点积——当我们考虑第一个令牌时输入序列中每个令牌的注意力权重——和V中的第一列,这是每个输入嵌入的第一个元素,投影到值空间中。所以,这是每个输入嵌入的第一个元素乘以第一个令牌的注意力权重。换句话说,它是第一个令牌的上下文向量的第一个元素!
在位置(1,2)——第一行第二列——我们会做同样的计算,但针对的是每个输入嵌入的第二个元素。这是第一个令牌的上下文向量的第二个元素。
...对其余列也是如此。到第一行的末尾,我们会得到某个东西,它(视为向量)是所有输入嵌入的和,乘以第一个输入的权重。这是该输入的上下文向量!
当然,对每一行都是同样的情况。这一次矩阵乘法的结果是一个矩阵,其中第m行是输入xm的上下文向量。
让我们把这些步骤综合在一起。我们从输入矩阵X开始,这是我们之前为长度为n的令牌序列生成的输入嵌入。每一行是一个嵌入,有d列,其中d是我们嵌入的维度。
我们还有权重矩阵,用于将输入嵌入映射到不同的空间:查询权重矩阵Wq、键权重矩阵Wk和值权重矩阵Wv。
所以,我们用三次矩阵乘法将输入矩阵投影到这些空间中:
...得到我们的查询矩阵、键矩阵和值矩阵。
然后我们用一个额外的矩阵乘法和一个转置来计算我们的注意力分数以得出点积:
我们通过将它们缩放c的平方根,然后应用softmax来将这些归一化为注意力权重:
...然后我们使用最后一个矩阵乘法来使用这些来计算上下文向量:
这就是我们的自注意力机制 :-)
现在,如果你回到开始的解释,希望它会有意义。
书中第3.4节用PyTorch代码演示了上面的内容,并给出了一个很好的简单nn.Module子类,它正好做了这些矩阵操作。这然后被改进了——第一个版本对三个权重矩阵使用了通用的nn.Parameter对象