共计 3707 个字符,预计需要花费 10 分钟才能阅读完成。
为什么需要注意力机制
在传统的 RNN 处理序列数据时,存在两个主要问题:一是长距离依赖难以捕捉,二是无法并行计算。Transformer 架构通过自注意力机制完美解决了这两个痛点。自注意力层能够直接计算序列中任意两个位置的关系权重,不受距离限制,且所有位置的注意力权重可以并行计算。

BERT 作为 Transformer 的典型代表,其核心正是多头注意力机制。这种机制可以理解为让模型同时从多个不同的角度(即多个头)来关注输入序列的不同部分,从而捕获更丰富的语义信息。
数学原理拆解
自注意力机制的核心计算涉及三个关键矩阵:查询矩阵 Q(Query)、键矩阵 K(Key)和值矩阵 V(Value)。它们的计算过程如下:
-
线性变换 :
$$Q = XW_Q, \quad K = XW_K, \quad V = XW_V$$
其中 $X \in \mathbb{R}^{n\times d_{model}}$ 是输入序列,$W_Q,W_K,W_V \in \mathbb{R}^{d_{model}\times d_k}$ 是可学习参数矩阵 -
注意力得分 :
$$Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V$$
这里除以 $\sqrt{d_k}$ 是为了防止点积结果过大导致 softmax 梯度消失 -
多头扩展 :
将 Q、K、V 分别拆分为 h 个头(BERT-base 是 12 个头),每个头独立计算注意力:
$$MultiHead = Concat(head_1,…,head_h)W^O$$
$$where \ head_i = Attention(QW_i^Q, KW_i^K, VW_i^V)$$
单头 vs 多头效率分析
-
单头注意力 :
时间复杂度 $O(n^2 \cdot d)$,其中 n 是序列长度,d 是特征维度 -
多头注意力 :
虽然看起来计算量增加了 h 倍,但因为每个头的维度降为 $d/h$,实际总复杂度仍为 $O(n^2 \cdot d)$
关键优势在于:
1. 参数矩阵被分解到多个头,更容易训练
2. 不同头可以关注不同位置的模式(如语法、指代等)
3. 在 GPU 上可以完美并行计算
PyTorch 实现详解
1. 基础模块定义
import torch
import torch.nn as nn
import math
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=768, n_heads=12, dropout=0.1):
super().__init__()
assert d_model % n_heads == 0
self.d_k = d_model // n_heads
self.n_heads = n_heads
# 线性变换矩阵
self.wq = nn.Linear(d_model, d_model) # (768,768)
self.wk = nn.Linear(d_model, d_model)
self.wv = nn.Linear(d_model, d_model)
self.wo = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
self.scale = 1 / math.sqrt(self.d_k)
2. 多头拆分与注意力计算
def split_heads(self, x):
# x 形状: (batch, seq_len, d_model)
batch_size = x.size(0)
return x.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
# 输出形状: (batch, n_heads, seq_len, d_k)
def forward(self, q, k, v, mask=None):
# 1. 线性变换并分头
q = self.split_heads(self.wq(q)) # (batch, heads, q_len, d_k)
k = self.split_heads(self.wk(k))
v = self.split_heads(self.wv(v))
# 2. 计算缩放点积注意力
attn = torch.matmul(q, k.transpose(-2, -1)) * self.scale
if mask is not None:
attn = attn.masked_fill(mask == 0, -1e10)
attn = self.dropout(torch.softmax(attn, dim=-1))
# 3. 加权求和并合并多头
output = torch.matmul(attn, v) # (batch, heads, q_len, d_k)
output = output.transpose(1, 2).contiguous() \
.view(output.size(0), -1, self.d_model)
return self.wo(output)
3. 残差连接与层归一化
class TransformerBlock(nn.Module):
def __init__(self, d_model, n_heads, dropout=0.1):
super().__init__()
self.attention = MultiHeadAttention(d_model, n_heads)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.ffn = nn.Sequential(nn.Linear(d_model, 4*d_model),
nn.GELU(),
nn.Linear(4*d_model, d_model)
)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
# 注意力子层
attn_out = self.attention(x, x, x, mask)
x = self.norm1(x + self.dropout(attn_out))
# 前馈子层
ffn_out = self.ffn(x)
return self.norm2(x + self.dropout(ffn_out))
实战避坑指南
梯度消失问题
- 初始化技巧 :
- 使用 Xavier 初始化注意力层的权重
-
值矩阵 V 的初始方差应设为 $1/\sqrt{d_k}$
-
学习率预热 :
optimizer = AdamW(model.parameters(), lr=5e-5) scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=1000, num_training_steps=total_steps )
注意力掩码应用
-
Padding 掩码 :处理变长序列时屏蔽无效位置
# seq_len=512, valid_len= 实际长度 mask = torch.arange(seq_len)[None, :] < valid_len[:, None] -
因果掩码 :防止解码器看到未来信息
mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).bool()
显存优化
-
梯度检查点 :
from torch.utils.checkpoint import checkpoint def custom_forward(*inputs): x, mask = inputs return transformer_block(x, mask) output = checkpoint(custom_forward, x, mask) -
混合精度训练 :
scaler = GradScaler() with autocast(): outputs = model(inputs) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()
可视化与分析
我准备了一个 Colab notebook,可以直观观察不同注意力头的聚焦模式: 查看 Notebook
示例可视化代码:
import seaborn as sns
import matplotlib.pyplot as plt
def plot_attention(attention_weights, layer_idx, head_idx):
plt.figure(figsize=(10,8))
sns.heatmap(attention_weights[layer_idx][head_idx].cpu().detach())
plt.title(f"Layer {layer_idx} Head {head_idx}")
plt.show()
思考与延伸
最后留给读者一个思考题:如何设计实验验证不同注意力头确实捕获了不同的语义信息?这里给出两个思路方向:
- 模式分析 :对同一层的不同头计算注意力权重的互信息,如果差异大则说明分工明确
- 消融实验 :冻结某些头的参数,观察模型在不同 NLP 任务上的表现变化
多头注意力机制就像团队合作,每个成员(注意力头)各司其职,有的关注局部语法,有的把握全局语义。希望通过本文的讲解,你能真正理解并驾驭这一强大工具。
