自注意力机制与多头注意力机制

1.引言

自注意力(self-attention )机制是Tramsformer 架构的一个关键底层model ,本篇文章-作为自用备忘录-将会着重讲解自注意力机制和多头注意力(multi head-attention)机制的思想和一般形式。

2.自注意力机制

2.1 原始数据输入

  对于文本翻译和文字生成类任务而言,输入通常是一个文本序列。将文本序列拆分为更小的词元(token )后,通过词嵌入(Word Embedding )的方式将词元投射到更高特征维度空间形成固定维度的稠密向量,这一步主要是为了提取出词语的高维特征。

2.2 算法核心思路

  自注意力层通过计算输入序列中每个位置与其他所有位置之间的相关性(由 Query 和 Key 决定),并将这些相关性归一化为权重,然后用这些权重对所有位置的 Value 向量进行加权求和,从而为每个位置生成一个融合了全局上下文信息的新表示。每个位置的输入向量 ai a_i ai都有对应一组 qi q_i qi, ki k_i ki和 vi v_i vi。下文将先以单个向量为例推导自注意力机制,再给出便于并行计算的矩阵形式

  对于输入序列中位置 ii i 的向量 ai∈Rdmodel a_i \in \mathbb{R}^{d_{\text{model}}} ai∈Rdmodel,其对应的 Query、Key、Value 向量分别通过三个可学习的参数矩阵 WQ W_Q WQ、 WK W_K WK、 WV W_V WV 进行线性变换得到:
qi=aiWQ(WQ∈R dmodel×dk ) q_i = a_i W_Q \quad (W_Q \in \mathbb{R}^{d_{\text{model}} \times d_k}) qi=aiWQ(WQ∈Rdmodel×dk)
ki=aiWK(WK∈R dmodel×dk ) k_i = a_i W_K \quad (W_K \in \mathbb{R}^{d_{\text{model}} \times d_k}) ki=aiWK(WK∈Rdmodel×dk)
vi=aiWV(WV∈R dmodel×dv ) v_i = a_i W_V \quad (W_V \in \mathbb{R}^{d_{\text{model}} \times d_v}) vi=aiWV(WV∈Rdmodel×dv)

  其中, dk d_k dk 和 dv d_v dv 分别是 Query/Key 和 Value 的维度,通常 dk=dv d_k = d_v dk=dv。(在后续的多头注意力机制中 dk=dv=dmodel/h d_k = d_v = d_{\text{model}} / h dk=dv=dmodel/h)

   qi q_i qi, ki k_i ki, vi v_i vi均来自同一个向量 ai a_i ai,但是由于 WQ W_Q WQ、 WK W_K WK、 WV W_V WV的不同被赋予了截然不同的角色。

  Query(查询) 表示当前词正在寻找的信息。每个位置都会用自己的 Query 去和其他位置的 Key 进行匹配,从而决定应该关注哪些位置。

  Key(键) 表示当前词能够提供什么匹配特征。它相当于一个标签,Query 与 Key 的点积结果反映了二者之间的相关程度,也就是注意力分数。

  Value(值) 表示当前词实际承载的语义内容。当 Query 与某个 Key 匹配后,模型就会从该位置取出对应的 Value,并将其融入当前位置的输出中。

  自注意力机制最重要的作用是能够让输出序列中的每个向量(包含自身)都融合所有位置的输入信息,而实现这一点的前提是先计算出当前位置输入 ai a_i ai与所有位置(包括自己在内)输入的相关性 αi \alpha_i αi,再利用相关性权重 αi \alpha_i αi计算得到的最终注意力权重计算输出。在《Attention is all you need》中相关性权重 α\alpha α的计算使用如下点积(Dot-product)的形式

  上图中展示的是点积计算相关性 α\alpha α的方法, qi,kj q_i, k_j qi,kj均由位置输入 ai,aj a_i, a_j ai,aj乘矩阵 WQ,Wk W_Q,W_k WQ,Wk得到,再将 qi,kj q_i, k_j qi,kj点乘即可求得 ai a_i ai位置相对 aj a_j aj位置的相关性权重 αi,j \alpha_{i,j} αi,j

  按照这样的方式,我们可以计算出 ai a_i ai相对所有输入位置的相关性(包括 ai a_i ai自己)

  在Transformer论文中,计算出的相关性权重还需要进行一层softmax层以获得最终的注意力权重 α′\alpha' α′

  在得到了 a1 a_1 a1相对所有位置的相关性权重后,就可以计算 a1 a_1 a1位置的输出 b1 b_1 b1。如下图所示,将 v1,v2,v3,v4 v_1,v_2,v_3, v_4 v1,v2,v3,v4分别按各自的相关性权重加权求和后便得到了 b1 b_1 b1。按照同样的方法可以求得各位置的输出 b2,b3,b4 b_2,b_3,b_4 b2,b3,b4

2.3 算法的矩阵形式

  假设输入序列长度为 nn n,模型维度为 dmodel d_{\text{model}} dmodel,输入矩阵为:
X∈R n×dmodel X \in \mathbb{R}^{n \times d_{\text{model}}} X∈Rn×dmodel

  其中每一行对应一个位置的输入向量。

  通过三个可学习的参数矩阵 WQ W_Q WQ、 WK W_K WK、 WV W_V WV,分别得到 Query、Key、Value 矩阵:
Q=XWQ(WQ∈R dmodel×dk ) Q = X W_Q \quad (W_Q \in \mathbb{R}^{d_{\text{model}} \times d_k}) Q=XWQ(WQ∈Rdmodel×dk)
K=XWK(WK∈R dmodel×dk ) K = X W_K \quad (W_K \in \mathbb{R}^{d_{\text{model}} \times d_k}) K=XWK(WK∈Rdmodel×dk)
V=XWV(WV∈R dmodel×dv ) V = X W_V \quad (W_V \in \mathbb{R}^{d_{\text{model}} \times d_v}) V=XWV(WV∈Rdmodel×dv)

得到的矩阵形状分别为:
Q∈R n×dk Q \in \mathbb{R}^{n \times d_k} Q∈Rn×dk
K∈R n×dk K \in \mathbb{R}^{n \times d_k} K∈Rn×dk
V∈R n×dv V \in \mathbb{R}^{n \times d_v} V∈Rn×dv

其中:

  • nn n 是序列长度;

  • dk d_k dk 是 Query 和 Key 的维度,二者必须相同,因为需要计算 QKTQK^T QKT;

  • dv d_v dv 是 Value 的维度,可以不同于 dk d_k dk,它决定了注意力输出的维度。

    自注意力层的输出可以表示为如下形式:


Attention(Q,K,V)=softmax ( QKT dk ) V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{Q K^T}{\sqrt{d_k}}\right) V Attention(Q,K,V)=softmax(dk QKT)V

上式与先前推导的唯一不同就是在softmax前除以 dk \sqrt{d_k} dk 以排除因维度 dk d_k dk过大而导致的softmax饱和,保证梯度能够有效传播(也可以理解为让注意力分数的方差与维度 dk d_k dk无关)

3 多头注意力机制(Multi-Head Attention)

多头注意力机制的公式为:
MultiHead(Q,K,V)=Concat(head1,head2,...,headh)WO\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \text{head}_2, \dots, \text{head}_h) W_O MultiHead(Q,K,V)=Concat(head1,head2,...,headh)WO

其中每个头的计算为:
headi=Attention(QWQi, KWKi, VWVi) \text{head}_i = \text{Attention}(Q W_Q^i,\ K W_K^i,\ V W_V^i) headi=Attention(QWQi, KWKi, VWVi)

Attention 函数为缩放点积注意力:
Attention(Q,K,V)=softmax ( QKT dk ) V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{Q K^T}{\sqrt{d_k}}\right) V Attention(Q,K,V)=softmax(dk QKT)V

参数矩阵与维度:
WQi∈R dmodel×dk ,WKi∈R dmodel×dk ,WVi∈R dmodel×dv W_Q^i \in \mathbb{R}^{d_{\text{model}} \times d_k},\quad W_K^i \in \mathbb{R}^{d_{\text{model}} \times d_k},\quad W_V^i \in \mathbb{R}^{d_{\text{model}} \times d_v} WQi∈Rdmodel×dk,WKi∈Rdmodel×dk,WVi∈Rdmodel×dv
WO∈R hdv×dmodel W_O \in \mathbb{R}^{h d_v \times d_{\text{model}}} WO∈Rhdv×dmodel

通常设 dk=dv=dmodel/h d_k = d_v = d_{\text{model}} / h dk=dv=dmodel/h,其中 hh h 是注意力头的数量。

多头注意力机制可以理解为:将输入分别投影到n个低维子空间,在每一个空间分别独立进行一次维度为 dmodel/h d_{\text{model}} / h dmodel/h的注意力计算,将所有头的输出在特征维度上拼接,并最终乘以一个线性变换矩阵得到最终输出。

4 参考文献

【Transformer论文逐段精读【论文精读】】 www.bilibili.com/video/BV1pu...

【(2025版)李宏毅机器学习深度学习系列课程全集,公认体验感最好的入门课程!--人工智能/机器学习/深度学习】 www.bilibili.com/video/BV1TA...

Vaswani, Ashish et al. "Attention is All you Need." Neural Information Processing Systems (2017).

相关推荐
linx2952 小时前
单元七 · 零基础路线图-第 1–3 周·变量、类型与字符串
c语言·开发语言·数据结构·c++·算法
Benny_Tang2 小时前
P7514 [省选联考 2021 A/B 卷] 卡牌游戏 题解
c++·算法
residual_fan3 小时前
经验模态重构:直向工业时间序列数据的数据扩增方法
人工智能·算法·重构·数据挖掘·数据分析
aramae3 小时前
模拟实现strcpy(字符串拷贝)(C语言)
java·c语言·开发语言·算法
Benny_Tang3 小时前
「雅礼集训 2018 Day7」A 题解
数据结构·c++·算法
Doubbbbbbble云3 小时前
跳表结构在高并发系统中的应用与优势分析4
算法
TAN-90°-3 小时前
Deep Learning for Computer Vision——Large Scale Distributed Training
人工智能·深度学习·神经网络·算法·机器学习·计算机视觉·语言模型
Lyyaoo.3 小时前
【回溯】【中等】全排列
java·数据结构·算法
写后端的胖头鱼3 小时前
一文讲懂JVM与调优
jvm·后端·算法·架构·jvm调优