1. 项目概述:手写实现最简Transformer
第一次看到"草履虫级Transformer"这个说法时,我忍不住笑出了声。这个比喻实在太贴切了——就像生物学实验中用草履虫研究细胞基础机制一样,我们要用最精简的代码揭示Transformer的核心运作原理。不同于直接调用PyTorch的nn.Transformer,这次我们从零开始,用不到200行Python代码实现一个能真实运行的微型Transformer。
这个项目的独特价值在于:当你亲手实现过每个矩阵乘法,调试过每个attention分数,才能真正理解为什么Transformer能在NLP领域所向披靡。我见过太多人虽然能背诵"self-attention"的定义,却说不清楚QKV矩阵究竟如何相互作用。通过这个极简实现,你将获得三个关键收获:
- 掌握Transformer每个组件的数学实现细节
- 理解位置编码等设计背后的物理意义
- 获得可自由修改的实验平台
2. 核心架构拆解
2.1 输入处理层
我们的微型Transformer从这两个组件开始:
class InputEmbedding(nn.Module): def __init__(self, d_model: int, vocab_size: int): super().__init__() self.embedding = nn.Embedding(vocab_size, d_model) self.d_model = d_model def forward(self, x): return self.embedding(x) * math.sqrt(self.d_model) # 缩放因子很重要!位置编码的实现尤为精妙:
def positional_encoding(seq_len, d_model): pe = torch.zeros(seq_len, d_model) position = torch.arange(0, seq_len).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) return pe # (seq_len, d_model)关键细节:频率递减的正余弦函数组合,让模型既能识别绝对位置又能学习相对位置关系。div_term中的10000这个魔法数决定了波长范围。
2.2 自注意力机制
这才是真正的精华所在:
class SelfAttention(nn.Module): def __init__(self, d_model: int, h: int): super().__init__() self.d_k = d_model // h self.h = h self.qkv = nn.Linear(d_model, d_model * 3) # 天才的QKV同源设计 self.out = nn.Linear(d_model, d_model) def forward(self, x, mask=None): batch_size = x.size(0) qkv = self.qkv(x).chunk(3, dim=-1) # 并行计算QKV q, k, v = [t.view(batch_size, -1, self.h, self.d_k).transpose(1, 2) for t in qkv] 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) attn = torch.softmax(scores, dim=-1) return self.out(torch.matmul(attn, v).transpose(1, 2).contiguous() .view(batch_size, -1, self.h * self.d_k))避坑指南:注意力分数一定要除以√d_k!这是保证梯度稳定的关键。我曾在早期版本漏掉这步,导致模型完全无法收敛。
2.3 前馈网络与残差连接
FFN的实现看似简单却暗藏玄机:
class FeedForward(nn.Module): def __init__(self, d_model: int, d_ff: int): super().__init__() self.linear1 = nn.Linear(d_model, d_ff) self.linear2 = nn.Linear(d_ff, d_model) def forward(self, x): return self.linear2(F.relu(self.linear1(x))) # 原论文使用GELU残差连接需要特别小心:
class ResidualConnection(nn.Module): def __init__(self, dropout: float): super().__init__() self.dropout = nn.Dropout(dropout) self.norm = nn.LayerNorm() def forward(self, x, sublayer): return x + self.dropout(sublayer(self.norm(x))) # 注意norm在sublayer前3. 完整组装与训练技巧
3.1 模型组装流水线
class TransformerBlock(nn.Module): def __init__(self, d_model: int, h: int, d_ff: int, dropout: float): super().__init__() self.attention = SelfAttention(d_model, h) self.ffn = FeedForward(d_model, d_ff) self.res1 = ResidualConnection(dropout) self.res2 = ResidualConnection(dropout) def forward(self, x, mask): x = self.res1(x, lambda x: self.attention(x, mask)) return self.res2(x, self.ffn)3.2 训练配置秘籍
这些超参组合经实测有效:
config = { 'batch_size': 32, 'd_model': 512, # 嵌入维度 'h': 8, # 注意力头数 'd_ff': 2048, # FFN隐藏层维度 'dropout': 0.1, # 最佳实践值 'lr': 1e-4, # 比CNN更小的学习率 'epochs': 20 # 早停很关键 }3.3 数据预处理要点
制作自己的微型数据集时注意:
def create_mock_data(vocab_size=5000, seq_len=20, samples=10000): src = torch.randint(1, vocab_size, (samples, seq_len)) trg = torch.roll(src, shifts=-1, dims=1) # 简单的移位任务 trg[:, -1] = 1 # 用1表示序列结束 return src, trg实战建议:先用这种确定性任务验证模型能学习到基础模式,再尝试真实语料
4. 调试与优化实录
4.1 常见报错解决方案
| 错误现象 | 可能原因 | 修复方案 |
|---|---|---|
| NaN损失 | 梯度爆炸 | 检查attention分数缩放,添加梯度裁剪 |
| 输出全零 | 残差连接错误 | 确认是x+sublayer而非sublayer(x) |
| 性能震荡 | 学习率过高 | 尝试3e-5到1e-4范围 |
4.2 性能优化技巧
- 内存优化:使用
torch.utils.checkpoint分段计算attention - 速度优化:用
torch.jit.script编译关键模块 - 精度技巧:混合精度训练需单独设置LN层为fp32
4.3 可视化诊断工具
def plot_attention(attention_map, sentence): fig = plt.figure(figsize=(12,8)) sns.heatmap(attention_map[0,0].detach().numpy(), xticklabels=sentence, yticklabels=sentence) plt.show()这个简单的热力图能直观显示模型到底在关注什么——我曾发现某个头专门捕捉句末标点,这就是Transformer自发形成的分工机制。
5. 扩展实验建议
尝试修改这些部分会有意外收获:
- 将位置编码改为可学习的参数
- 在FFN中尝试Swish激活函数
- 给attention添加相对位置偏置
- 实现多头注意力的不同头共享参数
每次修改后运行相同的测试用例,观察BLEU分数变化。我最惊喜的发现是:当把d_model缩减到128时,模型在简单任务上仍有85%准确率——这说明Transformer的鲁棒性远超预期。