共计 2856 个字符,预计需要花费 8 分钟才能阅读完成。
为什么需要 Transformer?
在自然语言处理(NLP)领域,传统的 RNN 和 LSTM 模型曾长期占据主导地位。但这些模型存在两个致命缺陷:

- 顺序计算限制:必须逐个处理序列中的词元,无法充分利用 GPU 并行计算能力
- 长程依赖丢失:随着序列长度增加,早期时间步的信息会逐渐衰减(即便 LSTM 有所改善,但问题依然存在)
2017 年 Google 提出的 Transformer 架构彻底改变了这一局面。BERT(Bidirectional Encoder Representations from Transformers)作为其典型代表,通过完全基于注意力机制的设计,实现了:
- 真正的双向上下文建模
- 可并行计算的序列处理
- 对长距离依赖关系的直接捕捉
自注意力机制详解
核心思想:动态权重分配
自注意力(Self-Attention)的本质是让序列中的每个词元都能 ” 看到 ” 其他所有词元,并动态决定关注哪些部分。这个过程通过三个关键向量实现:
- 查询向量(Query):当前词元想要查询什么信息
- 键向量(Key):其他词元能提供什么信息
- 值向量(Value):实际传递的信息内容
计算过程分步拆解
-
输入表示:假设输入序列有 n 个词元,每个词元的嵌入维度是 d。则输入矩阵 X ∈ ℝ^(n×d)
-
线性变换:通过可学习的权重矩阵生成 Q /K/V
# PyTorch 实现(假设 d_model=512)W_Q = nn.Linear(512, 64) # 通常使 dk=64 W_K = nn.Linear(512, 64) W_V = nn.Linear(512, 64) Q = W_Q(X) # (n, 64) K = W_K(X) # (n, 64) V = W_V(X) # (n, 64) -
注意力分数计算(Scaled Dot-Product Attention):
Attention(Q,K,V) = softmax(QK^T/√dk)V这里除以√dk(键向量维度)是为了防止点积结果过大导致 softmax 梯度消失
-
多头注意力:将这个过程并行执行 h 次(BERT-base 中 h =12),拼接后通过线性层融合
维度变化可视化
以 ” 我爱自然语言处理 ” 这句话为例(n=7,d=512,h=12):
- 输入 X: (7, 512)
- 单头 Q /K/V: (7, 64)
- 注意力矩阵: (7, 7) ← 每个单元格表示词元间关联度
- 多头拼接: (7, 768) ← 12 头×64 维
- 输出投影: (7, 512) ← 保持维度一致
关键实现细节
位置编码(Positional Encoding)
由于自注意力本身没有位置概念,必须显式添加位置信息。常用正弦 / 余弦函数生成:
def positional_encoding(max_len, d_model):
position = torch.arange(max_len).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * -(math.log(10000.0) / d_model))
pe = torch.zeros(max_len, d_model)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
return pe
常见错误:直接使用可学习的位置嵌入(learned positional embeddings)时,如果训练数据长度不足,在推理长文本时会出现未初始化位置的问题
注意力掩码(Attention Mask)
两种主要应用场景:
- 填充掩码(Padding Mask):遮盖无效的 padding 位置(通常在 batch 处理时使用)
- 因果掩码(Causal Mask):防止解码器看到未来信息(BERT 不需要,但 GPT 类模型需要)
# 生成上三角矩阵(用于解码器)mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).bool()
attention_scores.masked_fill_(mask, -1e9)
性能优化技巧
- 内存优化:当序列长度 >512 时,可以考虑:
- 使用稀疏注意力(如 Longformer 的滑动窗口模式)
-
分块计算(Reformer 的 LSH 注意力)
-
计算加速:
- 使用
torch.baddbmm替代矩阵乘法链 - 混合精度训练(AMP)
动手实验
推荐在 Colab 上运行这个简化版的多头注意力实现:
import torch
import torch.nn as nn
import math
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, h=8):
super().__init__()
assert d_model % h == 0, "d_model 必须能被 h 整除"
self.d_k = d_model // h
self.h = h
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):
batch_size, seq_len, _ = x.shape
# 线性变换并分头 [batch, seq_len, h, d_k]
Q = self.W_Q(x).view(batch_size, seq_len, self.h, self.d_k)
K = self.W_K(x).view(batch_size, seq_len, self.h, self.d_k)
V = self.W_V(x).view(batch_size, seq_len, self.h, self.d_k)
# 转置为[batch, h, seq_len, d_k]
Q, K, V = Q.transpose(1, 2), K.transpose(1, 2), V.transpose(1, 2)
# 计算缩放点积注意力
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attention = torch.softmax(scores, dim=-1)
# 输出拼接和投影
output = torch.matmul(attention, V)
output = output.transpose(1, 2).contiguous().view(batch_size, seq_len, -1)
return self.W_O(output)
实践建议
- 调试时可以可视化注意力矩阵,观察模型关注的重点是否合理
- 开始时使用小头数(如 4 头),逐渐增加观察效果变化
- 在自定义任务中,尝试调整
d_k维度(通常 64-256 之间)
通过理解这些基本原理,你就能更好地调整 BERT 模型适应自己的任务。下次我们将深入探讨 BERT 的预训练技巧和微调策略。
