【转载】理解并实现大语言模型中的自注意力、多头注意力、因果注意力与交叉注意力

Understanding and Coding Self-Attention, Multi-Head Attention, Causal-Attention, and Cross-Attention in LLMs,by Sebastian Raschka, 2024-01-14

理解并实现大语言模型中的自注意力、多头注意力、因果注意力与交叉注意力

本文将为你讲解Transformer架构与GPT-4、Llama等大语言模型(LLM)中所使用的自注意力机制。自注意力及相关机制是大语言模型的核心组成部分,在从事相关模型开发工作时,理解这一主题非常有价值。

不过,本文不会只停留在理论层面讲解自注意力机制,而是会带你用Python和PyTorch从零开始编码实现。在我看来,从零手写实现算法、模型与技术,是最高效的学习方式!

补充说明:本文是《从零理解并编码实现大语言模型的自注意力机制》的更新扩展版,那篇文章发布在我旧博客上,距今刚好差不多一年。我本人非常喜欢写(也喜欢读)这类“从零实现”的文章,因此特意将这篇内容更新优化,发布在Ahead of AI平台。

此外,这篇文章也促使我动笔撰写《从零构建大语言模型》一书,目前该书仍在创作中。下方是梳理全书脉络的思维模型,展示了自注意力机制在整个大语言模型体系中的位置。

figure01

《从零构建大语言模型》一书涵盖主题概览

为控制文章篇幅,本文默认你已经对大语言模型有基础了解,也掌握了注意力机制的基本概念。本文的核心目标,是通过Python与PyTorch的代码逐行讲解,帮你彻底理解注意力机制的运行原理。

自注意力简介

自注意力机制自《Attention Is All You Need》这篇Transformer开山论文提出以来,已成为众多顶尖深度学习模型的基石,在自然语言处理(NLP)领域尤其如此。如今自注意力的应用无处不在,理解其工作原理至关重要。

figure02

原始Transformer架构图,来源:https://arxiv\.org/abs/1706\.03762

深度学习中“注意力”的概念,最早源于对循环神经网络(RNN)的改进——解决RNN难以处理长序列、长句子的问题。举个例子,把一句话从一种语言翻译成另一种语言时,逐词直译通常行不通,因为它忽略了每种语言独有的复杂语法结构与习语表达,最终会得到不准确、甚至不通顺的译文。

figure03

错误的逐词直译(上)与正确翻译(下)对比

为解决这个问题,研究者提出了注意力机制:让模型在每个时间步都能访问序列中的所有元素,核心是做到“有选择地关注”,判断特定上下文中哪些单词最重要。2017年问世的Transformer架构,提出了独立的自注意力机制,彻底摆脱了对循环神经网络的依赖。

(为精简篇幅、聚焦自注意力的技术细节,本文只对背景动机做简要介绍,把重点放在代码实现上。)

figure04

摘自《Attention Is All You Need》论文的可视化图:通过注意力权重展示单词“making”对输入中其他单词的依赖与关注程度(颜色深浅与注意力权重的数值大小正相关)。

我们可以把自注意力理解为一种信息增强机制:它通过融入输入的上下文信息,丰富输入嵌入的信息量。换句话说,自注意力机制让模型能够权衡输入序列中不同元素的重要性,并动态调整它们对输出的影响。这在语言处理任务中尤为关键——同一个单词的含义,会随着它在句子、文档中的上下文变化而改变。

需要注意的是,自注意力有很多变体,其中一个主流研究方向是提升自注意力的计算效率。但绝大多数论文仍沿用《Attention Is All You Need》中提出的**缩放点积注意力(scaled-dot product attention)**原始实现;对大多数训练大规模Transformer的公司而言,自注意力本身通常并不是计算瓶颈。

因此,本文将聚焦最经典、实际应用最广泛的原始缩放点积注意力机制(下文统称自注意力)。如果你对其他类型的注意力机制感兴趣,可以参考2020年的《Efficient Transformers: A Survey》、2023年的《A Survey on Efficient Training of Transformers》综述,以及近年提出的FlashAttention与FlashAttention-v2相关论文。

输入句子的嵌入处理

正式开始前,我们先以句子Life is short, eat dessert first为例,看看如何将它输入自注意力机制。和其他文本建模方法(比如循环神经网络、卷积神经网络)一样,我们首先要生成句子的嵌入表示。

为简化演示,这里我们的词典dc只包含输入句子中出现的单词。在实际应用中,词典会覆盖训练数据集中的所有单词,典型的词表规模在3万到5万之间。

输入代码:

sentence = 'Life is short, eat dessert first'

dc = {s:i for i,s 
      in enumerate(sorted(sentence.replace(',', '').split()))}

print(dc)

输出:

{'Life': 0, 'dessert': 1, 'eat': 2, 'first': 3, 'is': 4, 'short': 5}

接下来,我们用这个词典为每个单词分配一个整数索引:

输入代码:

import torch

sentence_int = torch.tensor(
    [dc[s] for s in sentence.replace(',', '').split()]
)
print(sentence_int)

输出:

tensor([0, 4, 5, 2, 1, 3])

得到句子的整数向量表示后,我们就可以通过嵌入层,将输入编码为实数向量形式的嵌入。这里我们使用极小的3维嵌入,也就是每个输入单词对应一个3维向量。

注意:实际场景中的嵌入维度通常在数百到数千维,比如Llama 2的嵌入维度就达到了4096。这里使用3维纯粹是为了演示方便,方便我们查看每个向量的具体数值,不会让页面被数字占满。

句子包含6个单词,因此最终会得到一个6×3的嵌入矩阵:

输入代码:

vocab_size = 50_000

torch.manual_seed(123)
embed = torch.nn.Embedding(vocab_size, 3)
embedded_sentence = embed(sentence_int).detach()

print(embedded_sentence)
print(embedded_sentence.shape)

输出:

tensor([[ 0.3374, -0.1778, -0.3035],
        [ 0.1794,  1.8951,  0.4954],
        [ 0.2692, -0.0770, -1.0205],
        [-0.2196, -0.3792,  0.7671],
        [-0.5880,  0.3486,  0.6603],
        [-1.1925,  0.6984, -1.4097]])
torch.Size([6, 3])

定义权重矩阵

接下来我们讲解应用最广泛的自注意力机制——缩放点积注意力,它是Transformer架构的核心组成部分。

自注意力会用到三个权重矩阵,分别记作$W_q$、$W_k$、$W_v$,它们是模型的可训练参数,会在训练过程中不断更新。这三个矩阵的作用,分别是将输入投影为序列的查询(query)、键(key)和值(value)分量。

对应的查询、键、值序列,通过权重矩阵$W$与嵌入输入$x$做矩阵乘法得到:

  • 查询序列:$q(i) = x(i)W_q$,其中$i$为序列中第1到第$T$个位置

  • 键序列:$k(i) = x(i)W_k$,其中$i$为序列中第1到第$T$个位置

  • 值序列:$v(i) = x(i)W_v$,其中$i$为序列中第1到第$T$个位置

下标$i$代表输入序列中的token位置索引,序列总长度为$T$。

figure05

通过输入$x$与权重$W$计算查询、键、值向量

其中,$q(i)$和$k(i)$都是维度为$d_k$的向量。投影矩阵$W_q$和$W_k$的形状为 $d \times d_k$,而$W_v$的形状为 $d \times d_v$。

(注意:$d$代表每个单词向量$x$的维度。)

由于我们要计算查询向量和键向量的点积,因此这两个向量的元素数量必须相等($d_q = d_k$)。在很多大语言模型中,值向量也会采用相同的维度,也就是 $d_q = d_k = d_v$。不过值向量$v(i)$的元素数(决定了最终上下文向量的维度)本身是可以自由设置的。

在接下来的代码演示中,我们设置 $d_q = d_k = 2$,$d_v = 4$,并按如下方式初始化投影矩阵:

输入代码:

torch.manual_seed(123)

d = embedded_sentence.shape[1]

d_q, d_k, d_v = 2, 2, 4

W_query = torch.nn.Parameter(torch.rand(d, d_q))
W_key = torch.nn.Parameter(torch.rand(d, d_k))
W_value = torch.nn.Parameter(torch.rand(d, d_v))

(和前面的词嵌入向量一样,实际场景中$d_q$、$d_k$、$d_v$的维度会大得多,这里用小数值只是为了方便演示。)

figure06

计算未归一化注意力权重

现在,假设我们要计算第二个输入元素的注意力向量——此时第二个输入元素就作为查询。

figure07

接下来的内容,我们都以第二个输入$x(2)$为例

代码实现如下:

输入代码:

x_2 = embedded_sentence[1]
query_2 = x_2 @ W_query
key_2 = x_2 @ W_key
value_2 = x_2 @ W_value

print(query_2.shape)
print(key_2.shape)
print(value_2.shape)

输出:

torch.Size([2])
torch.Size([2])
torch.Size([4])

我们可以把这个逻辑推广到所有输入,计算出所有输入对应的键和值,因为后续计算未归一化注意力权重时会用到它们:

输入代码:

keys = embedded_sentence @ W_key
values = embedded_sentence @ W_value

print("keys.shape:", keys.shape)
print("values.shape:", values.shape)

输出:

keys.shape: torch.Size([6, 2])
values.shape: torch.Size([6, 4])

有了全部的键和值,我们就可以进入下一步:计算未归一化的注意力权重$\omega$(omega),如下图所示:

figure08

计算未归一化注意力权重$\omega$(omega)

如图所示,$\omega_{i,j}$ 等于查询向量与键向量的点积,即 $\omega_{i,j} = q(i) \cdot k(j)$。

举个例子,我们可以计算当前查询与第5个输入元素(对应索引4)之间的未归一化注意力权重:

输入代码:

omega_24 = query_2.dot(keys[4])
print(omega_24)

(注:$\omega$是希腊字母“omega”,上面代码中的变量名也由此而来。)

输出:

tensor(1.2903)

后续计算正式注意力权重时需要用到这些未归一化的$\omega$值,我们把所有输入token对应的$\omega$都计算出来,就像上图演示的那样:

输入代码:

omega_2 = query_2 @ keys.T
print(omega_2)

输出:

tensor([-0.6004,  3.4707, -1.5023,  0.4991,  1.2903, -1.3374])

计算注意力权重

自注意力的下一步,是对未归一化的注意力权重$\omega$做归一化,通过softmax函数得到归一化后的注意力权重$\alpha$(alpha)。此外,在输入softmax之前,我们会先将$\omega$乘以 $1/\sqrt{d_k}$ 做缩放,如下所示:

figure09

计算归一化注意力权重$\alpha$

除以$d_k$的平方根进行缩放,是为了保证权重向量的欧氏长度保持在相近的量级,避免注意力权重过大或过小,防止数值不稳定,同时保障模型的训练收敛效果。

为什么偏偏是$\sqrt{d_k}$?因为$q$和$k$的点积是$d_k$个独立项的和,每一项的方差约为1,这意味着原始得分的方差会随$d_k$线性增长。除以$\sqrt{d_k}$后,就能抵消这种增长,让方差回到约1的水平。

代码中的注意力权重计算实现如下:

输入代码:

import torch.nn.functional as F

attention_weights_2 = F.softmax(omega_2 / d_k**0.5, dim=0)
print(attention_weights_2)

输出:

tensor([0.0386, 0.6870, 0.0204, 0.0840, 0.1470, 0.0229])

最后一步,是计算上下文向量$z(2)$。它是原始查询输入$x(2)$的注意力加权版本,通过注意力权重融合了所有其他输入元素的上下文信息:

figure10

注意力权重是针对特定输入元素计算的。这里我们选择的是输入元素$x(2)$。

代码实现如下:

输入代码:

context_vector_2 = attention_weights_2 @ values

print(context_vector_2.shape)
print(context_vector_2)

输出:

torch.Size([4])
tensor([0.5313, 1.3607, 0.7891, 1.3110])

注意,这个输出向量的维度($d_v = 4$)比原始输入向量的维度($d = 3$)更高,因为我们前面设置了$d_v > d$;实际上,嵌入维度$d_v$的取值是任意的。

自注意力完整实现

结合前面几节的内容,我们可以把自注意力机制的代码整合起来,封装成一个简洁的SelfAttention类:

输入代码:

import torch.nn as nn

class SelfAttention(nn.Module):

    def __init__(self, d_in, d_out_kq, d_out_v):
        super().__init__()
        self.d_out_kq = d_out_kq
        self.W_query = nn.Parameter(torch.rand(d_in, d_out_kq))
        self.W_key   = nn.Parameter(torch.rand(d_in, d_out_kq))
        self.W_value = nn.Parameter(torch.rand(d_in, d_out_v))

    def forward(self, x):
        keys = x @ self.W_key
        queries = x @ self.W_query
        values = x @ self.W_value
        
        attn_scores = queries @ keys.T  # 未归一化注意力权重    
        attn_weights = torch.softmax(
            attn_scores / self.d_out_kq**0.5, dim=-1
        )
        
        context_vec = attn_weights @ values
        return context_vec

遵循PyTorch的编码惯例,上面的SelfAttention类在__init__方法中初始化自注意力的参数,在forward方法中完成所有输入的注意力权重与上下文向量计算。我们可以这样使用这个类:

输入代码:

torch.manual_seed(123)

d_in, d_out_kq, d_out_v = 3, 2, 4

sa = SelfAttention(d_in, d_out_kq, d_out_v)
print(sa(embedded_sentence))

输出:

tensor([[-0.1564,  0.1028, -0.0763, -0.0764],
        [ 0.5313,  1.3607,  0.7891,  1.3110],
        [-0.3542, -0.1234, -0.2627, -0.3706],
        [ 0.0071,  0.3345,  0.0969,  0.1998],
        [ 0.1008,  0.4780,  0.2021,  0.3674],
        [-0.5296, -0.2799, -0.4107, -0.6006]], grad_fn=<MmBackward0>)

可以看到,输出的第二行和我们上一节算出的context_vector_2完全一致:tensor([0.5313, 1.3607, 0.7891, 1.3110])

多头注意力

在文章开头的第一张图里(这里再放一次方便查看),我们能看到Transformer中使用了一个叫“多头注意力”的模块。

figure11

原始Transformer架构中的多头注意力模块,来源:https://arxiv\.org/abs/1706\.03762

这个“多头”注意力模块,和我们前面讲解的自注意力机制(缩放点积注意力)是什么关系呢?

在缩放点积注意力中,输入序列通过查询、键、值三个矩阵做变换。在多头注意力的语境里,这一套矩阵就对应一个注意力头。下图总结了我们前面实现的单个注意力头:

figure12

前文实现的自注意力机制总结

顾名思义,多头注意力就是包含多个这样的“头”,每个头都有自己的查询、键、值矩阵。这个概念和卷积神经网络中使用多个卷积核的思路类似,最终会输出多通道的特征图。

figure13

多头注意力:包含多个头的自注意力

用代码演示的话,我们可以基于前面的SelfAttention类,写一个MultiHeadAttentionWrapper包装类:

class MultiHeadAttentionWrapper(nn.Module):

    def __init__(self, d_in, d_out_kq, d_out_v, num_heads):
        super().__init__()
        self.heads = nn.ModuleList(
            [SelfAttention(d_in, d_out_kq, d_out_v) 
             for _ in range(num_heads)]
        )

    def forward(self, x):
        return torch.cat([head(x) for head in self.heads], dim=-1)

其中d_*参数和SelfAttention类中的定义完全一致,唯一新增的参数是注意力头的数量:

  • d_in:输入特征向量的维度

  • d_out_kq:查询与键的输出维度

  • d_out_v:值的输出维度

  • num_heads:注意力头的数量

我们用这些参数初始化num_headsSelfAttention实例,并用PyTorch的nn.ModuleList存储这些实例。

在前向传播时,每个存储在self.heads中的SelfAttention头都会独立处理输入x,然后将所有头的结果在最后一维(dim=-1)拼接起来。我们来实际运行一下:

首先,为了简化演示,我们用单个自注意力头,输出维度设为1:

输入代码:

torch.manual_seed(123)

d_in, d_out_kq, d_out_v = 3, 2, 1

sa = SelfAttention(d_in, d_out_kq, d_out_v)
print(sa(embedded_sentence))

输出:

tensor([[-0.0185],
        [ 0.4003],
        [-0.1103],
        [ 0.0668],
        [ 0.1180],
        [-0.1827]], grad_fn=<MmBackward0>)

现在我们把它扩展为4个注意力头:

输入代码:

torch.manual_seed(123)

mha = MultiHeadAttentionWrapper(
    d_in, d_out_kq, d_out_v, num_heads=4
)

context_vecs = mha(embedded_sentence)

print(context_vecs)
print("context_vecs.shape:", context_vecs.shape)

输出:

tensor([[-0.0185,  0.0170,  0.1999, -0.0860],
        [ 0.4003,  1.7137,  1.3981,  1.0497],
        [-0.1103, -0.1609,  0.0079, -0.2416],
        [ 0.0668,  0.3534,  0.2322,  0.1008],
        [ 0.1180,  0.6949,  0.3157,  0.2807],
        [-0.1827, -0.2060, -0.2393, -0.3167]], grad_fn=<CatBackward0>)
context_vecs.shape: torch.Size([6, 4])

从输出可以看到,前面单个自注意力头的结果,正好对应输出张量的第一列。

可以看到,多头注意力的结果是一个6×4的张量:我们有6个输入token,4个自注意力头,每个自注意力头输出1维结果。而前面单头自注意力部分,我们也得到过6×4的张量——那是因为我们把单头的输出维度设成了4,而不是1。

那么问题来了:既然单头自注意力本身就能调整输出嵌入的大小,为什么还需要多头注意力呢?

“提升单个自注意力头的输出维度”和“使用多个注意力头”,两者的本质区别在于模型处理、学习数据的方式不同。虽然两种方式都能提升模型表征不同特征、不同数据维度的能力,但底层逻辑完全不同。

比如,多头注意力中的每个头,都有可能学会关注输入序列的不同部分,捕捉数据中不同层面的关系。这种表征的多样性,正是多头注意力成功的关键。

此外,多头注意力的计算效率更高,尤其适合并行计算。每个头都可以独立处理,非常适配GPU、TPU这类擅长并行运算的现代硬件加速设备。

简而言之,使用多头注意力不只是为了提升模型容量,更是为了增强模型学习数据中多样化特征与关联的能力。举个例子,70亿参数的Llama 2模型就使用了32个注意力头。

因果自注意力

本节我们将前面讲解的自注意力机制,改造为因果自注意力机制——它主要用于GPT这类(解码器风格的)文本生成大语言模型。因果自注意力也常被称为“掩码自注意力”,在原始Transformer架构中,它对应“掩码多头注意力”模块。为了简化讲解,本节我们只看单个注意力头,其原理可以直接推广到多头场景。

figure14

原始Transformer架构中的因果自注意力模块,来源:《Attention Is All You Need》, https://arxiv\.org/abs/1706\.03762

因果自注意力的作用是:序列中每个位置的输出,只能依赖当前位置之前的已知输出,不能看到未来位置的信息。简单来说,就是保证下一个单词的预测,只由它前面的单词决定。

在GPT类大语言模型中,为了实现这一点,处理每个token时,我们都会把输入文本中位于当前token之后的“未来token”掩码掉。

下图演示了如何对注意力权重施加因果掩码,隐藏输入中的未来token:

为了演示并实现因果自注意力,我们沿用前面的未归一化注意力得分与注意力权重。首先快速回顾一下自注意力部分的注意力得分计算:

输入代码:

torch.manual_seed(123)

d_in, d_out_kq, d_out_v = 3, 2, 4

W_query = nn.Parameter(torch.rand(d_in, d_out_kq))
W_key   = nn.Parameter(torch.rand(d_in, d_out_kq))
W_value = nn.Parameter(torch.rand(d_in, d_out_v))

x = embedded_sentence

keys = x @ W_key
queries = x @ W_query
values = x @ W_value

# attn_scores 就是前面的 "omegas",即未归一化注意力权重
attn_scores = queries @ keys.T 

print(attn_scores)
print(attn_scores.shape)

输出:

tensor([[ 0.0613, -0.3491,  0.1443, -0.0437, -0.1303,  0.1076],
        [-0.6004,  3.4707, -1.5023,  0.4991,  1.2903, -1.3374],
        [ 0.2432, -1.3934,  0.5869, -0.1851, -0.5191,  0.4730],
        [-0.0794,  0.4487, -0.1807,  0.0518,  0.1677, -0.1197],
        [-0.1510,  0.8626, -0.3597,  0.1112,  0.3216, -0.2787],
        [ 0.4344, -2.5037,  1.0740, -0.3509, -0.9315,  0.9265]],
       grad_fn=<MmBackward0>)
torch.Size([6, 6])

和之前自注意力部分一样,上面的输出是6×6的张量,保存了6个输入token两两之间的未归一化注意力权重(也叫注意力得分)。

之前我们是这样通过softmax计算缩放点积注意力的:

输入代码:

attn_weights = torch.softmax(attn_scores / d_out_kq**0.5, dim=1)
print(attn_weights)

输出:

tensor([[0.1772, 0.1326, 0.1879, 0.1645, 0.1547, 0.1831],
        [0.0386, 0.6870, 0.0204, 0.0840, 0.1470, 0.0229],
        [0.1965, 0.0618, 0.2506, 0.1452, 0.1146, 0.2312],
        [0.1505, 0.2187, 0.1401, 0.1651, 0.1793, 0.1463],
        [0.1347, 0.2758, 0.1162, 0.1621, 0.1881, 0.1231],
        [0.1973, 0.0247, 0.3102, 0.1132, 0.0751, 0.2794]],
       grad_fn=<SoftmaxBackward0>)

这个6×6的输出就是注意力权重,和我们之前自注意力部分计算的结果一致。

而在GPT类大语言模型中,模型的训练方式是从左到右,逐个读取、生成token(单词)。如果我们有一条训练文本Life is short eat dessert first,就会形成如下的训练模式:箭头右侧单词的上下文向量,只能融合自身和箭头左侧的单词信息。

  • "Life""is"

  • "Life is""short"

  • "Life is short""eat"

  • "Life is short eat""dessert"

  • "Life is short eat dessert""first"

实现这个模式最简单的方法,就是对注意力权重矩阵的上三角(主对角线以上的部分)做掩码,如下图所示。这样一来,在计算上下文向量(即输入的注意力加权和)时,“未来”的单词就不会被纳入计算。

figure15

主对角线以上的注意力权重需要被掩码掉

代码层面,我们可以用PyTorch的tril函数实现:先生成一个由0和1组成的掩码矩阵。

输入代码:

block_size = attn_scores.shape[0]
mask_simple = torch.tril(torch.ones(block_size, block_size))
print(mask_simple)

输出:

tensor([[1., 0., 0., 0., 0., 0.],
        [1., 1., 0., 0., 0., 0.],
        [1., 1., 1., 0., 0., 0.],
        [1., 1., 1., 1., 0., 0.],
        [1., 1., 1., 1., 1., 0.],
        [1., 1., 1., 1., 1., 1.]])

接下来,我们把注意力权重和这个掩码相乘,将主对角线以上的所有注意力权重置零:

输入代码:

masked_simple = attn_weights*mask_simple
print(masked_simple)

输出:

tensor([[0.1772, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.0386, 0.6870, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.1965, 0.0618, 0.2506, 0.0000, 0.0000, 0.0000],
        [0.1505, 0.2187, 0.1401, 0.1651, 0.0000, 0.0000],
        [0.1347, 0.2758, 0.1162, 0.1621, 0.1881, 0.0000],
        [0.1973, 0.0247, 0.3102, 0.1132, 0.0751, 0.2794]],
       grad_fn=<MulBackward0>)

这是掩码未来单词的一种实现方式,但你会发现:此时每一行的注意力权重不再求和等于1。为了解决这个问题,我们可以对每一行重新做归一化,让它们的和重新变为1——这是注意力权重的标准约定。

输入代码:

row_sums = masked_simple.sum(dim=1, keepdim=True)
masked_simple_norm = masked_simple / row_sums
print(masked_simple_norm)

输出:

tensor([[1.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.0532, 0.9468, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.3862, 0.1214, 0.4924, 0.0000, 0.0000, 0.0000],
        [0.2232, 0.3242, 0.2078, 0.2449, 0.0000, 0.0000],
        [0.1536, 0.3145, 0.1325, 0.1849, 0.2145, 0.0000],
        [0.1973, 0.0247, 0.3102, 0.1132, 0.0751, 0.2794]],
       grad_fn=<DivBackward0>)

可以看到,现在每一行的注意力权重之和都为1了。

在Transformer这类模型中,使用归一化的注意力权重主要有两点优势:第一,求和为1的归一化注意力权重类似概率分布,更直观地体现了模型对输入不同部分的关注比例;第二,约束权重和为1,可以控制权重与梯度的量级,优化训练过程的稳定性。

更高效的掩码方式:无需二次归一化

上面我们实现的因果自注意力流程是:先计算注意力得分,再计算注意力权重,接着掩码掉上三角的权重,最后重新归一化。整个流程总结如下图:

figure16

前文实现的因果自注意力流程

其实还有一种更高效的方式,可以得到完全相同的结果。这种方法是在输入softmax计算注意力权重之前,就把注意力得分中主对角线以上的值替换为负无穷。流程总结如下:

figure17

另一种更高效的因果自注意力实现方案

我们可以用PyTorch写出这个流程,首先对注意力得分的上三角做掩码:

输入代码:

mask = torch.triu(torch.ones(block_size, block_size), diagonal=1)
masked = attn_scores.masked_fill(mask.bool(), -torch.inf)
print(masked)

上面的代码先生成一个掩码矩阵:主对角线及以下为0,以上为1。其中torch.triu(上三角)会保留矩阵主对角线及以上的元素,将下方元素置零,保留上三角部分;对应的torch.tril(下三角)则保留主对角线及以下的元素。

然后masked_fill方法会把掩码中为1的位置(主对角线以上)的所有元素替换为-torch.inf(负无穷),结果如下:

输出:

tensor([[ 0.0613,    -inf,    -inf,    -inf,    -inf,    -inf],
        [-0.6004,  3.4707,    -inf,    -inf,    -inf,    -inf],
        [ 0.2432, -1.3934,  0.5869,    -inf,    -inf,    -inf],
        [-0.0794,  0.4487, -0.1807,  0.0518,    -inf,    -inf],
        [-0.1510,  0.8626, -0.3597,  0.1112,  0.3216,    -inf],
        [ 0.4344, -2.5037,  1.0740, -0.3509, -0.9315,  0.9265]],
       grad_fn=<MaskedFillBackward0>)

接下来,我们照常应用softmax函数,就能得到经过掩码且归一化的注意力权重:

输入代码:

attn_weights = torch.softmax(masked / d_out_kq**0.5, dim=1)
print(attn_weights)

输出:

tensor([[1.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.0532, 0.9468, 0.0000, 0.0000, 0.0000, 0.0000],
        [0.3862, 0.1214, 0.4924, 0.0000, 0.0000, 0.0000],
        [0.2232, 0.3242, 0.2078, 0.2449, 0.0000, 0.0000],
        [0.1536, 0.3145, 0.1325, 0.1849, 0.2145, 0.0000],
        [0.1973, 0.0247, 0.3102, 0.1132, 0.0751, 0.2794]],
       grad_fn=<SoftmaxBackward0>)  

为什么这种方法有效?因为最后一步的softmax函数会把输入值转化为概率分布。当输入中存在负无穷时,$e^{-\infty}$趋近于0,因此这些位置对输出概率没有贡献,相当于概率为0。

结语

本文通过逐行编码的方式,拆解了自注意力的内部工作原理。在此基础上,我们进一步学习了多头注意力——这是Transformer大语言模型的核心组件。

随后我们还编码实现了交叉注意力(自注意力的一种变体,在处理两个独立序列时效果突出),以及因果自注意力——这是GPT、Llama等解码器风格大语言模型生成连贯、符合上下文的序列的关键概念。

通过从零编码实现这些复杂机制,希望你已经对Transformer与大语言模型中使用的自注意力机制有了透彻的理解。

(注:本文中的代码仅用于演示原理。如果你要在大语言模型训练中实现自注意力,推荐使用Flash Attention这类优化实现,它们能显著降低内存占用与计算开销。)

附加主题:交叉注意力

在前面自注意力与因果注意力的代码讲解中,我们设置了 $d_q = d_k = 2$,$d_v = 4$。也就是说,查询序列和键序列使用了相同的维度。虽然值矩阵$W_v$通常也会设置成和查询、键矩阵相同的维度(比如PyTorch自带的MultiHeadAttention类就是如此),但值的维度其实可以自由设置……

【以下为付费内容,请跳转至原文付费订阅,6$,两杯咖啡钱】
Understanding and Coding Self-Attention, Multi-Head Attention, Causal-Attention, and Cross-Attention in LLMs,by Sebastian Raschka, 2024-01-14

=== 以下为付费内容 ===

Leave a Reply

Your email address will not be published. Required fields are marked *

*