共计 2863 个字符,预计需要花费 8 分钟才能阅读完成。
引言:Transformer 架构中的注意力机制
Transformer 模型在自然语言处理领域取得了巨大成功,而其中的核心组件就是注意力机制。简单来说,注意力机制可以让模型在处理序列数据时,动态地关注不同位置的信息。想象一下你在阅读一篇文章时,大脑会自动聚焦于当前最相关的词语或句子——注意力机制的作用与此类似。

单头注意力回顾
在进入多头注意力之前,我们先简单回顾下单头注意力的基本原理。单头注意力通过三个关键向量进行计算:
- Query(Q):表示当前要查询的内容
- Key(K):表示待匹配的内容
- Value(V):表示最终要提取的信息
计算过程可以概括为:
1. 计算 Q 和 K 的相似度
2. 通过 softmax 归一化得到注意力权重
3. 用权重对 V 进行加权求和
数学表达式为:
Attention(Q, K, V) = softmax(QK^T/√d_k)V
其中 d_k 是 Key 的维度,√d_k 用于缩放防止 softmax 值过大。
多头注意力的动机和优势
单头注意力有一个明显局限:它只能学习一种注意力模式。就像我们人类理解复杂问题时,会从多个角度思考一样,多头注意力通过并行计算多个注意力头,可以捕获不同子空间的信息。
多头注意力的主要优势包括:
- 能够同时关注不同位置的不同信息模式
- 提高了模型的表示能力
- 允许信息在不同子空间中进行处理
数学原理详解
多头注意力的计算过程可以分为以下几个步骤:
- 线性投影:将 Q、K、V 分别投影 h 次(h 是头数)
- 缩放点积注意力:对每个头独立计算注意力
- 拼接输出:将所有头的输出拼接起来
- 最终投影:通过线性层得到最终输出
数学表达式为:
MultiHead(Q, K, V) = Concat(head_1, ..., head_h)W^O
where head_i = Attention(QW_i^Q, KW_i^K, VW_i^V)
PyTorch 实现
下面是一个完整的 PyTorch 实现,包含类型注解和详细注释:
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple
class MultiHeadAttention(nn.Module):
def __init__(self, embed_dim: int, num_heads: int):
super().__init__()
assert embed_dim % num_heads == 0, "Embed dim must be divisible by num_heads"
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
# 定义 Q、K、V 和输出的投影矩阵
self.q_proj = nn.Linear(embed_dim, embed_dim)
self.k_proj = nn.Linear(embed_dim, embed_dim)
self.v_proj = nn.Linear(embed_dim, embed_dim)
self.out_proj = nn.Linear(embed_dim, embed_dim)
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
mask: Optional[torch.Tensor] = None
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
前向传播
Args:
query: [batch_size, seq_len, embed_dim]
key: [batch_size, seq_len, embed_dim]
value: [batch_size, seq_len, embed_dim]
mask: [batch_size, seq_len, seq_len] (可选)
Returns:
output: [batch_size, seq_len, embed_dim]
attention_weights: [batch_size, num_heads, seq_len, seq_len]
"""
batch_size = query.size(0)
# 线性投影
Q = self.q_proj(query) # [batch_size, seq_len, embed_dim]
K = self.k_proj(key) # [batch_size, seq_len, embed_dim]
V = self.v_proj(value) # [batch_size, seq_len, embed_dim]
# 重排形状以分离注意力头
Q = Q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
K = K.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
V = V.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)
# 计算缩放点积注意力
attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.head_dim ** 0.5)
# 应用 mask(如果有)
if mask is not None:
attn_scores = attn_scores.masked_fill(mask == 0, float('-inf'))
attn_weights = F.softmax(attn_scores, dim=-1)
output = torch.matmul(attn_weights, V)
# 拼接所有头并做最终投影
output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.embed_dim)
output = self.out_proj(output)
return output, attn_weights
性能分析和优化建议
多头注意力的计算复杂度主要来自于矩阵乘法。对于长度为 n 的序列:
- 时间复杂度:O(n^2 * d) 其中 d 是 embedding 维度
- 空间复杂度:O(n^2 * h) 其中 h 是头数
优化建议:
- 对于长序列,考虑使用稀疏注意力或局部注意力
- 在 GPU 上,确保 batch_size 足够大以充分利用并行计算
- 适当选择头的数量 (通常 4 -16)
常见问题解答
Q: 为什么需要除以√d_k?
A: 防止点积结果过大导致 softmax 梯度消失。
Q: 如何选择头的数量?
A: 通常 embed_dim 能被头的数量整除。常见选择是 4 -16 个。
Q: 多头注意力层可以堆叠吗?
A: 可以,Transformer 就是堆叠多个多头注意力层。
课后练习
尝试实现一个改进版的多头注意力层,加入以下特性:
- 相对位置编码
- 残差连接和层归一化
- 注意力权重的 dropout
提示:可以参考 HuggingFace 的 Transformer 实现,但尽量自己动手实现核心部分。
