ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

从PPT到代码:手撕Transformer最小实现与训练调优实战

从PPT到代码:手撕Transformer最小实现与训练调优实战 简介这份PPT课件面向NLP入门与进阶学习者系统讲解Transformer模型及其论文《Attention is all you need》的核心思想。内容从传统Seq2Seq与RNN的并行计算瓶颈切入引出Attention机制再对Transformer的编码器-解码器架构做宏观与微观双重解读逐步拆解自注意力中Q、K、V的生成、点积缩放、softmax加权求和等完整流程并延伸至多头注意力、位置编码、Padding Mask、残差连接与LayerNorm等关键细节最后补充训练阶段要点与推理阶段的自回归解码策略。资源为单个pptx文件压缩包约17.08MB结构清晰、图文并茂适合对照论文逐节研读。目前已有3548人学习可作为课堂讲义或自学笔记帮助读者建立从Seq2Seq到Transformer的完整知识脉络并为后续理解BERT、GPT等预训练模型打下基础。1. 从一份 PPT 到能跑通的 Transformer先搞清楚它到底在算什么很多人第一次接触 Transformer是在一份名为「Transformer详解.pptx」的分享材料里。幻灯片上画着 Encoder-Decoder 堆叠、多头注意力方框、位置编码的正余弦公式看完觉得懂了合上电脑却写不出一行能跑的代码。问题不在数学而在于 PPT 把「数据怎么流、张量怎么变、梯度怎么回传」这三件事压缩成了静态图。Transformer 的核心只有一句话用自注意力替代循环让序列中任意两个位置直接建立联系。它解决的是 RNN 长距离依赖衰减、CNN 感受野受限的问题适合机器翻译、文本分类、图像分类ViT、时间序列回归等场景。适合谁读写过 PyTorch 基础层、能看懂矩阵乘法、想从「看得懂图」跨到「手撕得出来」的工程师。接下来按「原理立住 → 最小实现 → 训练实战 → 排错进阶」推一遍每步都给可复现的代码和参数。2. Transformer 架构拆解编码器层数、注意力头数与位置编码怎么定2.1 编码器到底有多少层这个数字从哪来热搜里常出现「transformer 编码部分有多少编码器呢」。答案不是固定的原始论文《Attention Is All You Need》里base 模型编码器和解码器各堆 6 层big 模型各 6 层但维度更大。层数N是超参不是架构常量。层数越多感受野和表达能力越强但显存和训练时间线性上涨。常见做法是小数据集分类任务用 24 层机器翻译用 6 层大规模预训练用 1224 层。判断依据是任务复杂度和数据量不是「越多越好」。层数堆过头小数据上直接过拟合验证集 loss 先降后升。2.2 多头注意力头数、维度与缩放因子多头注意力把d_model拆成h个头每个头维度d_k d_model / h。原始论文 base 用d_model512, h8, d_k64。缩放因子1/sqrt(d_k)是为了防止点积过大导致 softmax 梯度消失。import torch import torch.nn as nn import math class MultiHeadAttention(nn.Module): def __init__(self, d_model512, num_heads8): super().__init__() assert d_model % num_heads 0, d_model 必须能被 num_heads 整除 self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads # Q/K/V 和输出各一个线性层一次性算完所有头 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, q, k, v, maskNone): batch q.size(0) # 投影后拆头: [B, L, d_model] - [B, h, L, d_k] Q self.w_q(q).view(batch, -1, self.num_heads, self.d_k).transpose(1, 2) K self.w_k(k).view(batch, -1, self.num_heads, self.d_k).transpose(1, 2) V self.w_v(v).view(batch, -1, self.num_heads, self.d_k).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) attn torch.softmax(scores, dim-1) out torch.matmul(attn, V) # [B, h, L, d_k] out out.transpose(1, 2).contiguous().view(batch, -1, self.d_model) return self.w_o(out)逻辑说明view transpose把d_model切成h份每个头独立算注意力最后拼回。masked_fill用极小值屏蔽 padding 或未来位置。参数说明d_model是模型主维度num_heads决定并行子空间数量d_k由两者相除得到改num_heads时必须保证整除否则断言直接报错。2.3 位置编码正余弦公式与可学习嵌入的取舍Transformer 本身没有顺序概念位置信息必须显式注入。原始论文用固定正余弦class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000): super().__init__() pe torch.zeros(max_len, d_model) position torch.arange(0, max_len).unsqueeze(1).float() div_term torch.exp(torch.arange(0, d_model, 2).float() * (-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.unsqueeze(0)) # [1, max_len, d_model] def forward(self, x): return x self.pe[:, :x.size(1)]register_buffer让pe随模型搬到 GPU 但不参与梯度。偶数维用 sin、奇数维用 cos不同频率保证每个位置编码唯一。常见做法是翻译类任务用固定编码视觉和长序列任务用可学习位置嵌入或相对位置编码后者对超出训练长度的序列泛化更好。配置项原始论文 base小数据分类常用说明d_model512128256主维度越大越吃显存num_heads848必须整除 d_modelencoder layers624数据少就减层d_ff20484×d_model前馈层中间维度dropout0.10.10.3小数据调大3. 手写一个最小 Transformer 并跑通前向传播3.1 编码器层与前馈网络的拼装一个编码器层 多头注意力 残差 LayerNorm 前馈网络 残差 LayerNorm。前馈就是两层线性夹一个激活class EncoderLayer(nn.Module): def __init__(self, d_model512, num_heads8, d_ff2048, dropout0.1): super().__init__() self.attn MultiHeadAttention(d_model, num_heads) self.ffn 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, maskNone): # 子层1注意力 残差 x self.norm1(x self.dropout(self.attn(x, x, x, mask))) # 子层2前馈 残差 x self.norm2(x self.dropout(self.ffn(x))) return x这里用的是 Post-LN先残差后归一化原始论文写法。现在很多实现改成 Pre-LN先归一化再进子层训练更稳、对学习率不敏感。逻辑上残差保证梯度直通LayerNorm 稳定分布。参数d_ff一般取4×d_model改小省显存但削弱非线性表达。3.2 用正弦函数回归验证模型能学热搜里有「transformer 预测正弦函数」这是验证实现是否正确的最小闭环。构造输入序列预测下一时刻值import numpy as np def make_sin_data(n2000, seq_len20): x np.linspace(0, 100, n) y np.sin(x) xs, ys [], [] for i in range(n - seq_len): xs.append(y[i:iseq_len]) ys.append(y[iseq_len]) return torch.tensor(xs, dtypetorch.float32).unsqueeze(-1), torch.tensor(ys, dtypetorch.float32) class SinTransformer(nn.Module): def __init__(self, d_model64, num_heads4, num_layers2): super().__init__() self.input_proj nn.Linear(1, d_model) self.pos PositionalEncoding(d_model) self.layers nn.ModuleList([EncoderLayer(d_model, num_heads, d_model*4) for _ in range(num_layers)]) self.head nn.Linear(d_model, 1) def forward(self, x): x self.pos(self.input_proj(x)) for layer in self.layers: x layer(x) return self.head(x[:, -1, :]).squeeze(-1) # 取最后位置预测训练循环用 MSE 损失、Adam 优化器学习率 1e-3跑 50 个 epoch 就能拟合。逻辑说明input_proj把标量升到d_model位置编码注入顺序编码器层提取时序关系head取最后时间步输出预测值。参数说明seq_len20是回看窗口num_layers2足够拟合正弦层数多了反而震荡。3.3 训练循环与损失下降的判断model SinTransformer() opt torch.optim.Adam(model.parameters(), lr1e-3) loss_fn nn.MSELoss() X, Y make_sin_data() for epoch in range(50): model.train() pred model(X) loss loss_fn(pred, Y) opt.zero_grad() loss.backward() opt.step() if epoch % 10 0: print(fepoch {epoch}, loss {loss.item():.6f})正常情况 loss 从 0.5 量级降到 1e-3 以下。如果 loss 卡在 0.2 不动先查位置编码有没有加、mask 是否误屏蔽了有效位置、学习率是否过大导致震荡。这套流程跑通说明注意力、残差、归一化的连接顺序没错。4. 训练调优与常见报错从 loss 不降到显存溢出4.1 学习率预热与梯度裁剪Transformer 对学习率敏感原始论文用 warmup前若干步线性升温之后按步数平方根衰减。常见做法是 warmup 4000 步、峰值学习率 1e-4 到 1e-3。梯度裁剪clip_grad_norm_(model.parameters(), 1.0)防止梯度爆炸尤其层数多时必加。def lr_lambda(step, d_model512, warmup4000): if step 0: step 1 return (d_model ** -0.5) * min(step ** -0.5, step * warmup ** -1.5) scheduler torch.optim.lr_scheduler.LambdaLR(opt, lr_lambda) # 每个 step 后调用 scheduler.step() 和梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)参数说明warmup太小前期不稳太大收敛慢max_norm一般 0.51.0太小会压制正常梯度。4.2 显存溢出与 batch 组织显存溢出常见原因序列太长导致注意力矩阵[B, h, L, L]爆炸、d_model或层数过大、没开混合精度。排查顺序是先降 batch size再降序列长度再考虑梯度累积模拟大 batch。# 梯度累积小显存模拟大 batch accum_steps 4 for i, batch in enumerate(loader): loss model(batch) / accum_steps loss.backward() if (i 1) % accum_steps 0: torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) opt.step() opt.zero_grad()注意力复杂度是O(L²)序列翻倍显存翻四倍。长序列场景考虑稀疏注意力或分块计算这是 ViT、Swin Transformer 处理高分辨率图像时的核心优化点。4.3 训练不收敛的排查清单现象可能原因处理loss 不降学习率过大/过小加 warmup试 1e-4loss 变 NaN梯度爆炸梯度裁剪降 lr验证集先降后升过拟合加 dropout减层数输出恒定值mask 全屏蔽检查 mask 形状和填充显存溢出序列过长降 batch梯度累积提示mask 形状错误是最隐蔽的坑。padding mask 应为[B, 1, 1, L]因果 mask 为[1, 1, L, L]广播维度对不上会静默算错。5. 进阶技巧把 Transformer 用到图像分类与长序列5.1 ViT 与 Swin Transformer 的 patch 切分差异Vision Transformer 把图像切成固定 patch如 16×16每个 patch 展平后当 token 送进编码器位置编码改成可学习二维嵌入。Swin Transformer 引入窗口注意力和层级下采样把注意力限制在局部窗口内复杂度从O(L²)降到线性适合高分辨率图像分类和检测。# ViT 的 patch 切分核心 patch_size 16 patches x.unfold(2, patch_size, patch_size).unfold(3, patch_size, patch_size) # [B, C, H/P, W/P, P, P] - [B, N, C*P*P] patches patches.contiguous().view(B, C, -1, patch_size * patch_size).transpose(1, 2)逻辑说明unfold按步长滑窗切块再展平成 token 序列。参数说明patch_size越小 token 越多、精度越高但显存越大常见 16 或 32。5.2 长序列的位置外推与相对位置编码固定正余弦编码在超出训练长度时性能骤降。相对位置编码把位置差注入注意力分数对长度外推更友好。实现上在scores上加一个可学习的相对位置偏置表按i-j索引。# 相对位置偏置scores [B, h, L, L] 加 bias[i-j] rel_bias nn.Parameter(torch.zeros(2 * max_len - 1, num_heads)) idx torch.arange(L).unsqueeze(0) - torch.arange(L).unsqueeze(1) max_len - 1 scores scores rel_bias[idx].permute(2, 0, 1).unsqueeze(0)参数说明max_len是支持的最大相对距离rel_bias可学习。这套做法在长文本和语音任务里比绝对位置编码稳。5.3 用注意力权重做可解释性验证训练完可以取出注意力矩阵看模型关注哪里验证是否学到合理模式。分类任务里正确样本的注意力往往集中在关键 token 上如果注意力均匀铺满说明模型没学到东西回去查数据和 mask。with torch.no_grad(): attn model.layers[0].attn Q attn.w_q(x).view(B, -1, attn.num_heads, attn.d_k).transpose(1, 2) K attn.w_k(x).view(B, -1, attn.num_heads, attn.d_k).transpose(1, 2) weights torch.softmax(Q K.transpose(-2, -1) / math.sqrt(attn.d_k), dim-1) print(weights[0, 0]) # 第一个样本第一个头的注意力分布把这份权重画成热力图横纵轴是 token 位置颜色深浅代表关注强度。这是排查「模型到底看没看对地方」最直接的手段比只看 loss 曲线信息量大得多。本文还有配套的精品资源点击获取
返回列表