共计 2360 个字符,预计需要花费 6 分钟才能阅读完成。
为什么自注意力是 Transformer 的核心?
自注意力机制通过动态计算 token 间关联权重,彻底解决了 RNN 的长程依赖问题。它允许模型直接捕获任意位置的关系,为并行计算提供基础架构。正是这种特性让 Transformer 在捕捉复杂语义模式时展现出惊人效果。

数学原理拆解
缩放点积注意力公式
核心计算公式如下:
$$Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V$$
- $Q/K/V$ 分别表示查询 (Query)、键(Key)、值(Value) 矩阵
- $d_k$ 是 key 向量的维度,缩放因子 $\sqrt{d_k}$ 用于防止点积结果过大导致 softmax 梯度消失
- 计算过程可分为三步:
- 计算 Q 与 K 的点积得到相似度分数
- 缩放分数并做 softmax 归一化
- 用注意力权重加权求和 V 矩阵
多头机制实现原理
多头注意力的关键在于:
$$MultiHead = Concat(head_1,…,head_h)W^O$$
其中每个头的计算为:
$$head_i = Attention(QW_i^Q,KW_i^K,VW_i^V)$$
- 通过将 $Q/K/V$ 投影到 $h$ 个不同子空间(通常 $h=8$ 或 $12$)
- 每个头学习不同的注意力模式(如局部 / 全局、语法 / 语义特征)
- 最后拼接所有头输出并通过线性层 $W^O$ 融合
PyTorch 实战实现
import torch
import torch.nn as nn
import math
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, num_heads=8, dropout=0.1):
super().__init__()
assert d_model % num_heads == 0
self.d_k = d_model // num_heads
self.num_heads = num_heads
# 线性变换层
self.wq = nn.Linear(d_model, d_model)
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)
def forward(self, q, k, v, mask=None):
# q/k/v shape: [batch, seq_len, d_model]
batch_size = q.size(0)
# 线性投影 + 分头 [batch, seq_len, num_heads, d_k]
q = self.wq(q).view(batch_size, -1, self.num_heads, self.d_k)
k = self.wk(k).view(batch_size, -1, self.num_heads, self.d_k)
v = self.wv(v).view(batch_size, -1, self.num_heads, self.d_k)
# 转置为 [batch, num_heads, seq_len, d_k]
q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)
# 计算注意力分数 [batch, num_heads, q_len, k_len]
scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
# 掩码处理(padding/sequence mask)if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# softmax 归一化
attn = torch.softmax(scores, dim=-1)
attn = self.dropout(attn)
# 加权求和 [batch, num_heads, seq_len, d_k]
output = torch.matmul(attn, v)
# 拼接所有头 [batch, seq_len, d_model]
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, -1, self.num_heads * self.d_k)
return self.wo(output)
关键维度说明:
– 输入输出保持 [batch, seq_len, d_model] 统一维度
– 分头后每个头的维度为d_k = d_model // num_heads
– 注意力分数矩阵形状为[batch, num_heads, q_len, k_len]
性能优化策略
计算复杂度分析
- 时间复杂度:$O(n^2 \cdot d)$(n 为序列长度)
- 空间复杂度:$O(n^2)$(需存储注意力矩阵)
优化方案:
1. 滑动窗口注意力:限制每个 token 只关注局部邻域
2. 内存优化:
– 梯度检查点(gradient checkpointing)
– KV 缓存(解码时重复利用已计算的 K /V)
常见问题解决方案
梯度爆炸预防
- 在残差连接前使用 LayerNorm(Post-LN 结构)
- 初始化时缩小线性层权重范围
混合精度训练
with torch.cuda.amp.autocast():
# 前向计算时自动转为 FP16
output = attention_layer(q, k, v)
# 损失计算需保持 FP32
loss = loss_fn(output.float(), target)
思考题
- 注意力头差异化验证:
- 可视化各头的注意力分布热力图
-
计算不同头注意力矩阵的相似度
-
长序列处理方法:
- 位置编码外推(如 ALiBi)
- 动态稀疏注意力(如 Longformer 的局部 + 全局注意力)
正文完
