1. 项目概述
作为一名从传统软件开发转型AI的工程师,我深刻理解学习Transformer架构时的困惑。这个看似复杂的模型,其实核心思想非常优雅。今天我将用最接地气的方式,带大家手撕Transformer代码,同时保证每个模块都能独立运行测试。
注意:本文假设读者已经掌握Python和PyTorch基础,但对Transformer原理尚不熟悉。我们会从最基础的矩阵运算开始构建,而非直接调用现成的nn.Transformer模块。
2. 核心概念解析
2.1 注意力机制的本质
想象你在阅读一篇技术文档时,眼睛会不自觉地聚焦在关键词上——这就是注意力的生物学基础。在NLP中,注意力机制让模型能够动态决定应该"关注"输入序列的哪些部分。
数学上,注意力计算分为三步:
- 计算查询(Query)与键(Key)的相似度
- 用softmax归一化得到注意力权重
- 对值(Value)进行加权求和
# 最基础的注意力计算示例 def attention(query, key, value): scores = torch.matmul(query, key.transpose(-2, -1)) weights = torch.softmax(scores, dim=-1) return torch.matmul(weights, value)2.2 Transformer的架构创新
传统RNN的序列处理是串行的,而Transformer的突破在于:
- 完全基于自注意力机制
- 并行处理整个序列
- 引入位置编码(Positional Encoding)保留序列信息
下图展示了Transformer的标准架构(编码器-解码器结构):
[输入嵌入] → [位置编码] → [N×编码器层] → [N×解码器层] → [输出概率]3. 手写实现详解
3.1 基础组件实现
3.1.1 位置编码
由于Transformer没有递归结构,需要显式注入位置信息:
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)]技巧:位置编码的维度(d_model)必须与词嵌入维度一致,这样才能直接相加。
3.1.2 多头注意力
将注意力机制并行化,提升模型容量:
class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads == 0 self.d_k = d_model // num_heads self.num_heads = num_heads self.linears = nn.ModuleList([nn.Linear(d_model, d_model) for _ in range(4)]) def forward(self, query, key, value, mask=None): batch_size = query.size(0) # 线性变换后切分为多头 query, key, value = [ lin(x).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) for lin, x in zip(self.linears, (query, key, value)) ] # 计算缩放点积注意力 scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores = scores.masked_fill(mask == 0, -1e9) attn = torch.softmax(scores, dim=-1) x = torch.matmul(attn, value) # 合并多头结果 x = x.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) return self.linears[-1](x)3.2 编码器层实现
每个编码器层包含:
- 多头自注意力
- 前馈网络
- 残差连接和层归一化
class EncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, num_heads) self.feed_forward = nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def forward(self, x, mask): attn_output = self.self_attn(x, x, x, mask) x = self.norm1(x + self.dropout(attn_output)) ff_output = self.feed_forward(x) return self.norm2(x + self.dropout(ff_output))3.3 解码器层实现
解码器比编码器多一个交叉注意力层:
class DecoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, num_heads) self.cross_attn = MultiHeadAttention(d_model, num_heads) self.feed_forward = nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.norm3 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def forward(self, x, memory, src_mask, tgt_mask): # 自注意力处理目标序列 attn_output = self.self_attn(x, x, x, tgt_mask) x = self.norm1(x + self.dropout(attn_output)) # 交叉注意力连接编码器输出 attn_output = self.cross_attn(x, memory, memory, src_mask) x = self.norm2(x + self.dropout(attn_output)) ff_output = self.feed_forward(x) return self.norm3(x + self.dropout(ff_output))4. 完整模型组装
4.1 编码器堆叠
class Encoder(nn.Module): def __init__(self, num_layers, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.layers = nn.ModuleList([ EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) def forward(self, x, mask): for layer in self.layers: x = layer(x, mask) return x4.2 解码器堆叠
class Decoder(nn.Module): def __init__(self, num_layers, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.layers = nn.ModuleList([ DecoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) def forward(self, x, memory, src_mask, tgt_mask): for layer in self.layers: x = layer(x, memory, src_mask, tgt_mask) return x4.3 完整Transformer
class Transformer(nn.Module): def __init__(self, src_vocab, tgt_vocab, num_layers=6, d_model=512, num_heads=8, d_ff=2048, dropout=0.1): super().__init__() self.encoder = Encoder(num_layers, d_model, num_heads, d_ff, dropout) self.decoder = Decoder(num_layers, d_model, num_heads, d_ff, dropout) self.src_embed = nn.Sequential( nn.Embedding(src_vocab, d_model), PositionalEncoding(d_model) ) self.tgt_embed = nn.Sequential( nn.Embedding(tgt_vocab, d_model), PositionalEncoding(d_model) ) self.final_linear = nn.Linear(d_model, tgt_vocab) def forward(self, src, tgt, src_mask, tgt_mask): src = self.src_embed(src) memory = self.encoder(src, src_mask) tgt = self.tgt_embed(tgt) output = self.decoder(tgt, memory, src_mask, tgt_mask) return self.final_linear(output)5. 训练技巧与实战建议
5.1 学习率调度
Transformer通常使用带热启动的学习率调度:
def get_lr_scheduler(optimizer, warmup_steps=4000, d_model=512): def lr_lambda(step): arg1 = step ** -0.5 arg2 = step * (warmup_steps ** -1.5) return (d_model ** -0.5) * min(arg1, arg2) return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)5.2 掩码生成
处理变长序列时需要正确生成掩码:
def create_mask(src, tgt, pad_idx): # 源序列填充掩码 src_mask = (src != pad_idx).unsqueeze(1).unsqueeze(2) # 目标序列填充掩码 tgt_mask = (tgt != pad_idx).unsqueeze(1).unsqueeze(3) seq_len = tgt.size(1) # 防止解码器看到未来信息 nopeak_mask = torch.triu(torch.ones(1, seq_len, seq_len), diagonal=1).bool() tgt_mask = tgt_mask & ~nopeak_mask return src_mask, tgt_mask5.3 常见问题排查
梯度消失/爆炸:
- 检查残差连接是否正确实现
- 验证层归一化的位置
- 尝试梯度裁剪
过拟合:
- 增加dropout比例
- 使用标签平滑(Label Smoothing)
- 早停(Early Stopping)
训练不稳定:
- 检查学习率是否合适
- 验证输入数据的归一化
- 尝试更小的初始化范围
6. 扩展思考
6.1 计算效率优化
原始Transformer的计算复杂度是O(n²),对于长序列可以考虑:
- 局部窗口注意力
- 稀疏注意力模式
- 线性注意力变体
6.2 变体架构探索
现代Transformer的改进方向:
- 相对位置编码(Relative Position)
- 深度可分离卷积替代前馈网络
- 共享参数的多任务学习
6.3 实际部署考量
生产环境中需要注意:
- 量化感知训练
- 动态批处理
- 缓存机制优化
我在实际项目中发现,理解Transformer的最好方式就是亲手实现它。虽然PyTorch已经提供了现成的nn.Transformer模块,但通过从零构建,你会对每个矩阵运算的意义有更直观的认识。建议读者在完成基础版本后,尝试添加以下功能:
- 混合精度训练
- 模型并行
- 自定义注意力模式