共计 2145 个字符,预计需要花费 6 分钟才能阅读完成。
为什么需要 Transformer?
传统 RNN 处理文本序列时存在明显缺陷:

- 梯度消失问题:长距离依赖难以捕捉(如 ”The cat that ate the fish was hungry” 中 was 与 cat 的关系)
- 顺序计算限制:无法并行处理序列,训练效率低下
2017 年 Google 提出的 Transformer 架构通过自注意力机制完美解决了这些问题。下面我们通过三个关键步骤来理解它的工作原理。
自注意力机制图解
想象你在阅读这句话时,大脑会动态关注不同位置的词:
句子:"The animal didn't cross the street because it was too tired"
^^ ^^^^^^ ^^ ^^^^^^^^^^^
| | | |
| +----------------+ |
+---------------------------------+
自注意力通过三个关键向量实现这种动态聚焦:
- Query(查询):当前词想知道 ” 我应该关注谁 ”
- Key(键):每个词提供的 ” 被关注价值 ”
- Value(值):实际传递的信息内容
计算过程用向量点积实现:
Q · K^T softmax × V
[查询×所有键] → [注意力权重] → [加权求和]
NumPy 手写实现
让我们用代码实现最基本的注意力计算(假设输入维度 =64):
import numpy as np
def attention(Q, K, V, mask=None):
"""
Q: [batch_size, seq_len, d_k]
K: [batch_size, seq_len, d_k]
V: [batch_size, seq_len, d_v]
"""
d_k = Q.shape[-1]
scores = np.matmul(Q, K.transpose(0,2,1)) / np.sqrt(d_k) # [b, seq_len, seq_len]
if mask is not None:
scores = scores + mask * -1e9 # 屏蔽未来信息(解码器用)attn_weights = softmax(scores, axis=-1)
return np.matmul(attn_weights, V) # [b, seq_len, d_v]
关键细节说明:
sqrt(d_k)缩放:防止点积结果过大导致 softmax 梯度消失- 注意力掩码:在解码时阻止当前位置关注后续词
PyTorch 完整实现
现在实现带位置编码的 Transformer 层:
import torch
import torch.nn as nn
class TransformerLayer(nn.Module):
def __init__(self, d_model=512, n_heads=8):
super().__init__()
self.d_head = d_model // n_heads
# QKV 投影矩阵
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)
# 位置编码(正弦函数版本)pe = torch.zeros(5000, d_model)
position = torch.arange(0, 5000).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * -(math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer('pe', pe)
def forward(self, x):
# 添加位置信息 [batch, seq_len, d_model]
x = x + self.pe[:x.size(1)]
# 多头注意力计算
Q = self.W_q(x).view(x.size(0), -1, n_heads, self.d_head) # [b, seq, heads, d_head]
K = self.W_k(x).view(x.size(0), -1, n_heads, self.d_head)
V = self.W_v(x).view(x.size(0), -1, n_heads, self.d_head)
# 缩放点积注意力(实际实现需拆分为多步)attn_output = ...
return attn_output
对比实验与调参技巧
在 IMDb 影评分类任务上测试不同配置:
| 头数 | 验证准确率 | 训练时间 |
|---|---|---|
| 2 | 85.2% | 32min |
| 4 | 86.7% | 41min |
| 8 | 87.1% | 63min |
实际应用中的避坑指南:
- 梯度爆炸预防:在每个子层(注意力 /FFN)后立即做 LayerNorm
- 显存优化 :使用
torch.utils.checkpoint分段计算注意力 - 学习率调整:配合 Warmup 策略逐步增大学习率
为什么 Transformer 适合长文本?
与传统 CNN 相比,Transformer 的优势在于:
- 任意位置直接交互(无需通过多层卷积传递)
- 自注意力复杂度 O(n²)但并行度高
- 位置编码明确保留序列顺序信息
思考题答案提示:CNN 需要 O(n/k)层才能建立 n 距离的关系(k 为卷积核大小),而 Transformer 只需一层即可建立任意距离的直接关联。
正文完
发表至: 人工智能
近一天内
