Transformer 中的多头注意力机制:一次多视角的“聚焦力”
自从 2017 年 “Attention Is All You Need”(Vaswani 等)提出 Transformer 架构以来,多头注意力(Multi‑Head Attention, MHA)便成为 NLP 领域最具革命性的创新之一。下面,我们将分步剖析这一机制的原理、实现与实际意义,并通过代码和案例让技术细节更易上手。
一、注意力机制回顾
在深入多头注意力之前,让我们快速回顾一下缩放点积自注意力(Scaled Dot‑Product Self‑Attention)的基本原理:
给定输入序列,将其映射为 Query(Q)、Key(K)、Value(V) 三组向量:
QKᵀ 计算相似性;
除以 控制数值范围;
softmax 得到注意力权重;
权重加权 V 得到输出。
这种机制使每个 token 可以“看”到序列内的其他 token,从而捕获上下文语义。Transformer 将其称为 Self-Attention。
二、为什么要用“多头”?
多头注意力的核心动机有以下几点:
多视角观察:不同 head 可以捕捉不同类型的语义关系(例如主谓一致、定语修饰、时间顺序等)。
子空间优势:每个头在 Q/K/V 的不同子空间操作,可提高表达能力
并行计算:多个头同时计算,自然适配现代硬件并行处理,大幅提升效率
三、多头注意力的数学推导
设输入维度为 d_model,头数为 h,每个头的维度为 d_k = d_model / h。整个流程如下:
-
输入
X ∈ R^{L×d_model}; -
线性映射得到 Q, K, V 并 reshape 为
(h, L, d_k); -
每个 head 独立计算注意力;
-
将 h 个 Head 的输出拼接为
(L, d_model); -
最后一层线性映射得到最终输出。
简要公式:
head_i = Attention(XW_i^Q, XW_i^K, XW_i^V)
MHA(X) = Concat(head_i) · W^O
图示详解见下:

【插图:Q/K/V 分 head → 并行 Attention → 拼接 → 线性变换】。
四、代码 实现
1.示例
下面带你手写一个简易版 MHA,以加深理解:
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, num_heads=8):
super().__init__()
assert d_model % num_heads == 0, "d_model must be divisible by num_heads"
self.d_k = d_model // num_heads
self.h = num_heads
self.w_q = nn.Linear(d_model, d_model)
self.w_k = nn.Linear(d_model, d_model)
self.w_v = nn.Linear(d_model, d_model)
self.w_o = nn.Linear(d_model, d_model)
def forward(self, x, mask=None):
B, L, D = x.size()
# (B, L, D) → (B, h, L, d_k)
q = self.w_q(x).view(B, L, self.h, self.d_k).transpose(1,2)
k = self.w_k(x).view(B, L, self.h, self.d_k).transpose(1,2)
v = self.w_v(x).view(B, L, self.h, self.d_k).transpose(1,2)
# Scaled dot-product attention
scores = torch.matmul(q, k.transpose(-2,-1)) / (self.d_k ** 0.5)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
weights = F.softmax(scores, dim=-1)
context = torch.matmul(weights, v) # (B, h, L, d_k)
# 拼接 + 线性层
context = context.transpose(1,2).contiguous().view(B, L, D)
out = self.w_o(context)
return out, weights
测试运行:
x = torch.rand(2, 5, 512) # batch=2, seq_len=5
mha = MultiHeadAttention(512,8)
o, w = mha(x)
print(o.shape, w.shape) # (2,5,512), (2,8,5,5)
-
o: 模型输出 -
w: 每个 head 的注意力权重,可用于可视化理解。
2.案例代码
# 多头注意力机制
import copy
def clones(module, N):
return nn.ModuleList([copy.deepcopy(module) for _ in range(N)])
class MultiHeadedAttention(nn.Module):
def __init__(self, head, embedding_dim, dropout=0.1):
super(MultiHeadedAttention, self).__init__()
assert embedding_dim%head== 0
self.d_k = embedding_dim // head
self.head = head
self.linears = clones(nn.Linear(embedding_dim, embedding_dim), 4)
self.attn = None
self.dropout = nn.Dropout(p=dropout)
def forward(self, query, key, value, mask=None):
if mask is not None:
mask = mask.unsqueeze(0)
# print('multmaskshape===', mask.shape) #multmaskshape=== torch.Size([1, 8, 4, 4])
batch_size = query.size(0)
# view中的四个参数的意义
# batch_size: 批次的样本数量
# -1这个位置应该是: 每个句子的长度
# self.head*self.d_k应该是embedding的维度, 这里把词嵌入的维度分到了每个头中, 即每个头中分到了词的部分维度的特征
# query, key, value形状torch.Size([2, 8, 4, 64])
query, key, value = [model(x).view(batch_size, -1, self.head, self.d_k).transpose(1, 2) for model, x in zip(self.linears, (query, key, value))]
# query, key, value = [model(x) for model, x in zip(self.linears, (query, key, value))]
# print('-=-=', query.shape)
# print('-=-=', key.shape)
# print('-=-=', value.shape)
'''
-=-= torch.Size([2, 4, 512])
-=-= torch.Size([2, 4, 512])
-=-= torch.Size([2, 4, 512])
'''
# 所以mask的形状 torch.Size([1, 8, 4, 4]) 这里的所有参数都是4维度的 进过dropout的也是4维度的
x, self.attn = attention(query, key, value, mask=mask, dropout=self.dropout)
# contiguous解释:https://zhuanlan.zhihu.com/p/64551412
# 这里相当于图中concat过程
x = x.transpose(1, 2).contiguous().view(batch_size, -1, self.head*self.d_k)
return self.linears[-1](x)
五、可视化 & 案例解析
案例:句子 “The cat sat on the mat.”
不同 head 会关注不同类型的关联:
-
head1:冠词与名词关系 (“The”→“cat”,“the”→“mat”);
-
head2:动词修饰 (“sat”→“on”, “on”→“mat”);
-
head3:跨距远距离关注 (“cat”→“mat”);
六、多头注意力在 Transformer 中的位置
Transformer Encoder 中,每一层包括:
x → MHA(x,x,x) → 残差 + LayerNorm → FFN → 残差 + LayerNorm
MHA 的输出为 FFN 输入,而残差结构确保原始位置信息得以保留。
七、多头数 h 的选取与实践
-
通常 h=8 或 16;
-
增加 h 提升表达能力,但也会分割维度,导致每 head 可用特征减少;
-
实际效果 h 与 d_model、数据量、任务相关,是个超参。
八、端到端示例:将 MHA 应用于分类任务
class TransformerClassifier(nn.Module):
def __init__(self, d_model=512, heads=8, num_layers=2, num_classes=10):
super().__init__()
self.embed = nn.Embedding(10000, d_model)
self.pe = PositionalEncoding(d_model)
self.layers = nn.ModuleList([
EncoderBlock(d_model, heads) for _ in range(num_layers)
])
self.fc = nn.Linear(d_model, num_classes)
def forward(self, x):
x = self.embed(x)
x = self.pe(x)
for layer in self.layers:
x = layer(x)
out = x.mean(dim=1)
return self.fc(out)
说明:EncoderBlock 包含 MHA + FFN + LayerNorm + 残差。
九、为什么 MHA 能更懂“语言”
-
捕获复杂语义模式:多头允许关注时序、语法、实体等多种联系;
-
并行训练:加速模型训练,摆脱 RNN 串行限制;
-
强泛化能力:适配翻译、分类、摘要、视觉等多种任务。
十、总结与展望
🔸 核心思想一览
| 特性 | 说明 |
|---|---|
| Q/K/V 子空间 | 将原始特征投射到多个子空间,学习多样关注模式 |
| 并行头 | 多头并行计算,效率高、表达丰富 |
| 残差整合 | 最终拼接 + 映射,将多路信息融合输出 |
🔸 应用拓展
-
NLP:翻译、摘要、问答等;
-
CV:Vision Transformer 的核心模块;
-
多模态:支持文本-图像、音频-文本等交互任务。
🔸 未来趋势
-
头剪枝:去除无用头以减小模型;
-
可解释性增强:深入研究各 head 的关注模式;
-
高效 Transformer:稀疏注意力、更少资源消耗设计。
结语
多头注意力机制是 Transformer 成为主流 AI 架构的关键,它赋予模型多角度理解数据的能力,同时兼顾并行效率。本文从原理、代码、案例、应用层面深入剖析,帮助你掌握这颗 AI 大脑中的“聚焦引擎”。
更多推荐
所有评论(0)