共计 3114 个字符,预计需要花费 8 分钟才能阅读完成。
Transformer 的革命性意义
Transformer 彻底改变了 NLP 领域处理序列数据的方式,完全摒弃了传统的循环神经网络结构。它通过自注意力机制实现了对长距离依赖的高效建模,使得并行计算成为可能。最重要的是,这种架构为后来的 BERT、GPT 等预训练模型奠定了基础,成为现代 NLP 的核心组件。

模型构建详解
1. 词嵌入与位置编码
Transformer 首先需要将单词转换为固定维度的向量表示。这里我们使用 PyTorch 的 nn.Embedding 来实现词嵌入层:
import torch.nn as nn
vocab_size = 10000 # 词汇表大小
d_model = 512 # 嵌入维度
embedding = nn.Embedding(vocab_size, d_model)
由于 Transformer 没有循环结构,我们需要显式地加入位置信息。下面是正弦位置编码的实现:
import torch
import math
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super().__init__()
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)
self.register_buffer('pe', pe)
def forward(self, x):
return x + self.pe[:x.size(1)]
2. 自注意力机制实现
自注意力的核心是 QKV(Query-Key-Value)三元组计算。我们先定义线性变换层:
self.query = nn.Linear(d_model, d_model)
self.key = nn.Linear(d_model, d_model)
self.value = nn.Linear(d_model, d_model)
然后实现缩放点积注意力:
def scaled_dot_product_attention(q, k, v, mask=None):
matmul_qk = torch.matmul(q, k.transpose(-2, -1))
d_k = q.size(-1)
scaled_attention_logits = matmul_qk / math.sqrt(d_k)
if mask is not None:
scaled_attention_logits += (mask * -1e9)
attention_weights = torch.softmax(scaled_attention_logits, dim=-1)
output = torch.matmul(attention_weights, v)
return output
3. 多头注意力并行化
将注意力头拆分成多个子空间可以捕捉不同的语义信息。关键实现技巧是使用 view 和transpose操作:
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, num_heads):
super().__init__()
self.num_heads = num_heads
self.d_model = d_model
assert d_model % num_heads == 0
self.depth = d_model // 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.dense = nn.Linear(d_model, d_model)
def split_heads(self, x, batch_size):
x = x.view(batch_size, -1, self.num_heads, self.depth)
return x.transpose(1, 2)
def forward(self, q, k, v, mask=None):
batch_size = q.size(0)
q = self.split_heads(self.wq(q), batch_size)
k = self.split_heads(self.wk(k), batch_size)
v = self.split_heads(self.wv(v), batch_size)
scaled_attention = scaled_dot_product_attention(q, k, v, mask)
concat_attention = scaled_attention.transpose(1, 2).contiguous()
concat_attention = concat_attention.view(batch_size, -1, self.d_model)
output = self.dense(concat_attention)
return output
完整文本分类实现
数据准备
我们使用 HuggingFace 的 IMDB 数据集:
from datasets import load_dataset
imdb = load_dataset('imdb')
tokenizer = ... # 实现分词器
模型训练技巧
为了防止训练初期的不稳定,我们实现学习率预热和梯度裁剪:
optimizer = torch.optim.Adam(model.parameters(), lr=0, betas=(0.9, 0.98), eps=1e-9)
# 学习率预热
lr_scheduler = torch.optim.lr_scheduler.LambdaLR(
optimizer,
lambda step: min((step+1)**-0.5, (step+1)*warmup_steps**-1.5)
)
# 梯度裁剪
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
避坑指南
数值稳定性
计算注意力分数时,较大的输入会导致 softmax 进入梯度饱和区。解决方案:
scaled_attention_logits = matmul_qk / math.sqrt(d_k)
显存管理
Transformer 的显存占用主要来自注意力矩阵,其空间复杂度为 O(n²)。估算公式:
显存(B) ≈ batch_size × seq_len² × num_heads × 4 (float32)
对于长文本,可以尝试:
– 减小 batch_size
– 使用梯度累积
– 采用线性注意力变体
进阶思考
- 如何修改架构实现字符级 Transformer?
- 将词嵌入层替换为字符卷积
-
调整位置编码的最大长度
-
对比 CNN 在长文本处理的优劣:
- CNN 的局部感受野限制长距离建模
- Transformer 的自注意力具有全局视野
-
CNN 在短文本上可能有计算效率优势
-
解释 LayerNorm 对梯度传播的影响:
- 保持各层输入的分布稳定
- 缓解梯度消失 / 爆炸问题
- 与 BatchNorm 相比更适合变长输入
结语
通过本文的实践,你应该已经掌握了 Transformer 的核心实现要点。建议在实际项目中先从小规模数据开始实验,逐步调整模型规模和训练策略。Transformer 虽然强大,但也需要根据具体任务进行适当的调整和优化。
