- 人工智能
- 深度学习
- 机器学习
- 教程
【免费下载链接】d2l-zh
《动手学深度学习》:面向中文读者、能运行、可讨论。中英文版被70多个国家的500多所大学用于教学。
多头注意力(Multi-Head Attention)是注意力机制家族中最重要的变体之一,也是 Transformer 架构的核心组件。本篇文章以《动手学深度学习》d2l-zh 仓库中 multihead-attention 章节 为主体,完整讲解多头注意力的数学定义、从零实现(涵盖 MXNet / PyTorch / TensorFlow / Paddle 四个框架)、并行计算所需的张量转置技巧,并结合仓库源码揭示其与缩放点积注意力、Transformer 编码器/解码器的实际调用关系。读完本文,你将能够独立实现一个可用、可扩展的多头注意力模块,并理解它在现代序列模型中的定位。
为什么需要多头注意力
在实践中,当给定相同的查询(queries)、键(keys)和值(values)集合时,我们希望模型能够基于同一套注意力机制学习到不同的行为,并把不同的行为作为知识组合起来,从而捕获序列内各种范围的依赖关系——例如短距离依赖和长距离依赖。仅使用一个注意力汇聚(attention pooling)只能得到一种加权平均模式,表达能力有限。
为此,与其只用单独一个注意力汇聚,我们可以用独立学习得到的 $h$ 组不同的线性投影(linear projections)来变换查询、键和值,然后把 $h$ 组变换后的查询、键和值并行地送入注意力汇聚,最后将 $h$ 个注意力汇聚的输出拼接起来,再通过另一个可学习的线性投影变换产生最终输出。这种设计称为多头注意力,其中 $h$ 个注意力汇聚输出中的每一个都被称为一个头(head),该概念出自 Vaswani et al., 2017 的注意力机制论文。如上图所示,实现时用全连接层完成这些可学习的线性变换。
多头注意力的核心价值在于:每个头都可能关注输入的不同部分,多个头联合可以表示比简单加权平均值更复杂的函数。
模型:数学形式化描述
在实现多头注意力之前,先用数学语言将模型形式化。给定查询 $\mathbf{q} \in \mathbb{R}^{d_q}$、键 $\mathbf{k} \in \mathbb{R}^{d_k}$ 和值 $\mathbf{v} \in \mathbb{R}^{d_v}$,每个注意力头 $\mathbf{h}_i$($i = 1, \ldots, h$)的计算方法为:
$$\mathbf{h}_i = f(\mathbf W_i^{(q)}\mathbf q, \mathbf W_i^{(k)}\mathbf k,\mathbf W_i^{(v)}\mathbf v) \in \mathbb R^{p_v},$$
其中可学习的参数包括:
- $\mathbf W_i^{(q)}\in\mathbb R^{p_q\times d_q}$:查询的投影矩阵;
- $\mathbf W_i^{(k)}\in\mathbb R^{p_k\times d_k}$:键的投影矩阵;
- $\mathbf W_i^{(v)}\in\mathbb R^{p_v\times d_v}$:值的投影矩阵;
- $f$:注意力汇聚函数,可以是加性注意力(Additive Attention)或缩放点积注意力(Scaled Dot-Product Attention),详见 attention-scoring-functions 章节。
多头注意力的输出需要经过另一个线性变换,它作用于 $h$ 个头拼接后的结果,其可学习参数为 $\mathbf W_o\in\mathbb R^{p_o\times h p_v}$:
$$\mathbf W_o \begin{bmatrix}\mathbf h_1\\vdots\\mathbf h_h\end{bmatrix} \in \mathbb{R}^{p_o}.$$
基于这种设计,每个头都能关注输入的不同部分,模型得以表达远超简单加权平均的复杂函数。
核心实现:四个框架的多头注意力类
实现中通常选择缩放点积注意力作为每个注意力头内部的注意力汇聚。为避免计算代价和参数代价大幅增长,设定 $p_q = p_k = p_v = p_o / h$,即每个头的维度为总输出维度除以头数。若将查询、键和值线性变换的输出数量统一设为 $p_q h = p_k h = p_v h = p_o$,则 $h$ 个头可以并行计算。在实现中,$p_o$ 通过参数num_hiddens指定。
PyTorch 实现
#@save class MultiHeadAttention(nn.Module): """多头注意力""" def __init__(self, key_size, query_size, value_size, num_hiddens, num_heads, dropout, bias=False, **kwargs): super(MultiHeadAttention, self).__init__(**kwargs) self.num_heads = num_heads self.attention = d2l.DotProductAttention(dropout) self.W_q = nn.Linear(query_size, num_hiddens, bias=bias) self.W_k = nn.Linear(key_size, num_hiddens, bias=bias) self.W_v = nn.Linear(value_size, num_hiddens, bias=bias) self.W_o = nn.Linear(num_hiddens, num_hiddens, bias=bias) def forward(self, queries, keys, values, valid_lens): # queries,keys,values的形状: # (batch_size,查询或者“键-值”对的个数,num_hiddens) # valid_lens 的形状: # (batch_size,)或(batch_size,查询的个数) # 经过变换后,输出的queries,keys,values 的形状: # (batch_size*num_heads,查询或者“键-值”对的个数, # num_hiddens/num_heads) queries = transpose_qkv(self.W_q(queries), self.num_heads) keys = transpose_qkv(self.W_k(keys), self.num_heads) values = transpose_qkv(self.W_v(values), self.num_heads) if valid_lens is not None: # 在轴0,将第一项(标量或者矢量)复制num_heads次, # 然后如此复制第二项,然后诸如此类。 valid_lens = torch.repeat_interleave( valid_lens, repeats=self.num_heads, dim=0) # output的形状:(batch_size*num_heads,查询的个数, # num_hiddens/num_heads) output = self.attention(queries, keys, values, valid_lens) # output_concat的形状:(batch_size,查询的个数,num_hiddens) output_concat = transpose_output(output, self.num_heads) return self.W_o(output_concat)该实现同样保存在仓库的 d2l/torch.py 中,供各章节复用。
MXNet 实现
MXNet(Gluon)版本将四个投影层统一为输出num_hiddens维的nn.Dense层,并通过flatten=False保持三维张量的批维度:
#@save class MultiHeadAttention(nn.Block): """多头注意力""" def __init__(self, num_hiddens, num_heads, dropout, use_bias=False, **kwargs): super(MultiHeadAttention, self).__init__(**kwargs) self.num_heads = num_heads self.attention = d2l.DotProductAttention(dropout) self.W_q = nn.Dense(num_hiddens, use_bias=use_bias, flatten=False) self.W_k = nn.Dense(num_hiddens, use_bias=use_bias, flatten=False) self.W_v = nn.Dense(num_hiddens, use_bias=use_bias, flatten=False) self.W_o = nn.Dense(num_hiddens, use_bias=use_bias, flatten=False) def forward(self, queries, keys, values, valid_lens): queries = transpose_qkv(self.W_q(queries), self.num_heads) keys = transpose_qkv(self.W_k(keys), self.num_heads) values = transpose_qkv(self.W_v(values), self.num_heads) if valid_lens is not None: valid_lens = valid_lens.repeat(self.num_heads, axis=0) output = self.attention(queries, keys, values, valid_lens) output_concat = transpose_output(output, self.num_heads) return self.W_o(output_concat)TensorFlow 与 Paddle 实现
TensorFlow 版本继承tf.keras.layers.Layer,四个投影层使用tf.keras.layers.Dense,前向方法名为call,并且需要在调用内部注意力时透传**kwargs(例如training标志),以保证 dropout 行为在训练/推理模式下正确切换:
#@save class MultiHeadAttention(tf.keras.layers.Layer): """多头注意力""" def __init__(self, key_size, query_size, value_size, num_hiddens, num_heads, dropout, bias=False, **kwargs): super().__init__(**kwargs) self.num_heads = num_heads self.attention = d2l.DotProductAttention(dropout) self.W_q = tf.keras.layers.Dense(num_hiddens, use_bias=bias) self.W_k = tf.keras.layers.Dense(num_hiddens, use_bias=bias) self.W_v = tf.keras.layers.Dense(num_hiddens, use_bias=bias) self.W_o = tf.keras.layers.Dense(num_hiddens, use_bias=bias) def call(self, queries, keys, values, valid_lens, **kwargs): queries = transpose_qkv(self.W_q(queries), self.num_heads) keys = transpose_qkv(self.W_k(keys), self.num_heads) values = transpose_qkv(self.W_v(values), self.num_heads) if valid_lens is not None: valid_lens = tf.repeat(valid_lens, repeats=self.num_heads, axis=0) output = self.attention(queries, keys, values, valid_lens, **kwargs) output_concat = transpose_output(output, self.num_heads) return self.W_o(output_concat)Paddle 版本继承nn.Layer,四个投影层使用nn.Linear,通过bias_attr=bias控制偏置,valid_lens的扩展使用paddle.repeat_interleave:
#@save class MultiHeadAttention(nn.Layer): def __init__(self, key_size, query_size, value_size, num_hiddens, num_heads, dropout, bias=False, **kwargs): super(MultiHeadAttention, self).__init__(**kwargs) self.num_heads = num_heads self.attention = d2l.DotProductAttention(dropout) self.W_q = nn.Linear(query_size, num_hiddens, bias_attr=bias) self.W_k = nn.Linear(key_size, num_hiddens, bias_attr=bias) self.W_v = nn.Linear(value_size, num_hiddens, bias_attr=bias) self.W_o = nn.Linear(num_hiddens, num_hiddens, bias_attr=bias) def forward(self, queries, keys, values, valid_lens): queries = transpose_qkv(self.W_q(queries), self.num_heads) keys = transpose_qkv(self.W_k(keys), self.num_heads) values = transpose_qkv(self.W_v(values), self.num_heads) if valid_lens is not None: valid_lens = paddle.repeat_interleave( valid_lens, repeats=self.num_heads, axis=0) output = self.attention(queries, keys, values, valid_lens) output_concat = transpose_output(output, self.num_heads) return self.W_o(output_concat)并行计算的关键:两个张量转置函数
为了使多个头能够并行计算,MultiHeadAttention类使用两个转置函数:transpose_output逆转transpose_qkv的操作。这两个函数是实现层面最值得细读的部分。
以 PyTorch 版本为例,transpose_qkv将输入从(batch_size, 序列长度, num_hiddens)逐步变形为(batch_size, num_heads, 序列长度, num_hiddens/num_heads),最终展平为(batch_size * num_heads, 序列长度, num_hiddens/num_heads),从而把每个头变成独立的"批"样本,交给批量矩阵乘法处理:
#@save def transpose_qkv(X, num_heads): """为了多注意力头的并行计算而变换形状""" # 输入X的形状:(batch_size,查询或者“键-值”对的个数,num_hiddens) # 输出X的形状:(batch_size,查询或者“键-值”对的个数,num_heads, # num_hiddens/num_heads) X = X.reshape(X.shape[0], X.shape[1], num_heads, -1) # 输出X的形状:(batch_size,num_heads,查询或者“键-值”对的个数, # num_hiddens/num_heads) X = X.permute(0, 2, 1, 3) # 最终输出的形状:(batch_size*num_heads,查询或者“键-值”对的个数, # num_hiddens/num_heads) return X.reshape(-1, X.shape[2], X.shape[3]) #@save def transpose_output(X, num_heads): """逆转transpose_qkv函数的操作""" X = X.reshape(-1, num_heads, X.shape[1], X.shape[2]) X = X.permute(0, 2, 1, 3) return X.reshape(X.shape[0], X.shape[1], -1)transpose_output恰好是上述过程的逆运算:先把(batch_size * num_heads, 查询数, num_hiddens/num_heads)还原为(batch_size, num_heads, 查询数, num_hiddens/num_heads),再转置并拼回(batch_size, 查询数, num_hiddens)。MXNet 版本使用transpose(0, 2, 1, 3),TensorFlow 版本使用tf.transpose(X, perm=(0, 2, 1, 3)),Paddle 版本使用X.transpose((0, 2, 1, 3)),逻辑完全一致。
每个头内部的缩放点积注意力
多头注意力类内部复用的是d2l.DotProductAttention,即缩放点积注意力。查看仓库源码 d2l/torch.py 可以看到其完整实现:
class DotProductAttention(nn.Module): """缩放点积注意力 Defined in :numref:`subsec_additive-attention`""" def __init__(self, dropout, **kwargs): super(DotProductAttention, self).__init__(**kwargs) self.dropout = nn.Dropout(dropout) def forward(self, queries, keys, values, valid_lens=None): d = queries.shape[-1] scores = torch.bmm(queries, keys.transpose(1,2)) / math.sqrt(d) self.attention_weights = masked_softmax(scores, valid_lens) return torch.bmm(self.dropout(self.attention_weights), values)它通过批量矩阵乘法计算查询与键的点积,除以 $\sqrt{d}$ 进行缩放,再用 masked_softmax 依据valid_lens将无效位置的注意力分数掩蔽为 $0$,最后对值加权求和。这也解释了为什么多头注意力要求num_hiddens能被num_heads整除:每个头的维度为num_hiddens/num_heads,缩放因子取该头的维度。
测试与运行示例
仓库使用一个键和值相同的小例子来测试MultiHeadAttention类:设置num_hiddens=100、num_heads=5,dropout 取0.5,构造批量大小为 2、查询数为 4、键值对数为 6 的全 1 张量,并用valid_lens=[3, 2]分别限制两个样本的有效键值对数量。多头注意力输出的形状应为(batch_size, num_queries, num_hiddens),即(2, 4, 100)。
PyTorch 下的实例化与调用:
num_hiddens, num_heads = 100, 5 attention = MultiHeadAttention(num_hiddens, num_hiddens, num_hiddens, num_hiddens, num_heads, 0.5) attention.eval() batch_size, num_queries = 2, 4 num_kvpairs, valid_lens = 6, d2l.tensor([3, 2]) X = d2l.ones((batch_size, num_queries, num_hiddens)) Y = d2l.ones((batch_size, num_kvpairs, num_hiddens)) attention(X, Y, Y, valid_lens).shapeMXNet 版本构造参数更简洁(只需num_hiddens, num_heads, dropout三个参数,初始化后直接调用):
num_hiddens, num_heads = 100, 5 attention = MultiHeadAttention(num_hiddens, num_heads, 0.5) attention.initialize()TensorFlow 版本调用时需要显式传入training=False以关闭 dropout 的随机性,Paddle 版本则先调用attention.eval()再前向,细节差异体现了各框架对训练/推理模式的处理习惯。
在 Transformer 中的实际应用
多头注意力在仓库中不是孤立的模块,而是 Transformer 编码器和解码器的核心组件。在 transformer 章节 中,EncoderBlock将MultiHeadAttention与残差连接、层规范化组合:
class EncoderBlock(nn.Module): """Transformer编码器块""" def __init__(self, key_size, query_size, value_size, num_hiddens, norm_shape, ffn_num_input, ffn_num_hiddens, num_heads, dropout, use_bias=False, **kwargs): super(EncoderBlock, self).__init__(**kwargs) self.attention = d2l.MultiHeadAttention( key_size, query_size, value_size, num_hiddens, num_heads, dropout, use_bias) self.addnorm1 = AddNorm(norm_shape, dropout) self.ffn = PositionWiseFFN(ffn_num_input, ffn_num_hiddens, num_hiddens) self.addnorm2 = AddNorm(norm_shape, dropout) def forward(self, X, valid_lens): Y = self.addnorm1(X, self.attention(X, X, X, valid_lens)) return self.addnorm2(Y, self.ffn(Y))这里self.attention(X, X, X, valid_lens)就是典型的多头自注意力用法:查询、键、值全部来自同一个输入序列X。从源码结构看,编码器各层复用同一套MultiHeadAttention,而解码器则同时包含两个多头注意力实例——一个做自注意力(attention1),另一个做编码器-解码器注意力(attention2),后者以编码器输出作为键和值、以解码器自身状态作为查询。可见,本仓库 multihead-attention 章节 实现的多头注意力类,是后续 Transformer、机器翻译等模型(见 seq2seq 章节 与 transformer 章节)直接复用的公共组件。
小结
- 多头注意力融合了来自多个注意力汇聚的不同知识,这些知识的不同来源于同一套查询、键和值的不同子空间表示。
- 基于适当的张量操作(
transpose_qkv与transpose_output),可以实现多头注意力的并行计算,显著提升效率。
练习
- 分别可视化本实验中的多个头的注意力权重(可通过
attention.attention_weights或DotProductAttention中保存的attention_weights属性观察每个头对输入的关注分布)。 - 假设有一个完成训练的基于多头注意力的模型,现在希望修剪最不重要的注意力头以提高预测速度。思考如何设计实验来衡量每个注意力头的重要性——例如逐头屏蔽后评估模型性能的下降幅度,或分析注意力权重的熵、梯度统计等指标来量化各头的贡献。
- 人工智能
- 深度学习
- 机器学习
- 教程
【免费下载链接】d2l-zh
《动手学深度学习》:面向中文读者、能运行、可讨论。中英文版被70多个国家的500多所大学用于教学。
相关推荐
《动手学深度学习》d2l-zh 详解:多头注意力(Multi-Head Attention)原理与四框架实现
《动手学深度学习》d2l zh 详解:多头注意力(Multi Head Attention)原理与四框架实现 多头注意力(multi head attentio
人工智能深度学习机器学习教程深入解析 D2L 中的多头注意力(Multi-Head Attention):原理、实现与并行化技巧
深入解析 D2L 中的多头注意力(Multi Head Attention):原理、实现与并行化技巧 导读 多头注意力(Multi Head Attention
文档教程人工智能深度学习NLP计算机视觉强化学习注意力机制实战指南:《动手学深度学习》d2l-zh 中从注意力提示到 Transformer 的完整脉络
注意力机制实战指南:《动手学深度学习》d2l zh 中从注意力提示到 Transformer 的完整脉络 导读 本文基于《动手学深度学习》中文版仓库(d2l z
人工智能深度学习机器学习教程
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考