ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

手撕Transformer:PyTorch从零实现可训练的最小模型

手撕Transformer:PyTorch从零实现可训练的最小模型 最近两年各种 Transformer 图解和源码解析铺天盖地但我观察到一个很有意思的现象很多人看别人的代码觉得“懂了”自己动手写却卡在第一步。比如Q、K、V到底怎么从输入向量里变出来attention算完之后接LayerNorm还是先接残差位置编码加上去之后模型的输入到底是什么形状这些问题如果不亲手写一遍靠看图是记不住的。这篇文章的思路很直接不贴别人的代码不用TransformerEncoder一行完事而是从 Token 到训练用 PyTorch 把 Transformer 的每一个核心模块都自己实现出来。我们不做学术级复现只做一件最有价值的事——跑通一个可以训练的最小 Transformer。读完这篇文章你会得到三层收获从零写出位置编码、单头注意力、多头注意力、前馈网络和 Encoder Layer明白 QKV、自注意力、位置编码在实际代码里到底怎么流转有一个可以直接运行的训练脚本能亲眼看到 loss 下降验证自己写的是对的。Transformer 看起来吓人拆开之后核心其实就几块积木。下面我们开始。1. 为什么要“手撕”Transformer如果你只是想在项目里用 Transformer最省事的方式是直接调 PyTorch 的nn.TransformerEncoder或者 HuggingFace 的现成模型。但很多人在调库的过程中会遇到同一个问题模型跑起来了但效果差、loss 不降、维度报错自己完全不知道从哪里排查。这就是“手撕”的价值所在。当你能从零写出多头注意力你会对维度变化、掩码机制、残差连接这些细节形成肌肉记忆。之后再去看任何开源 Transformer 代码都不会觉得是黑盒。还有一点很重要现在大模型的核心架构依然是 Transformer。不管是 BERT、GPT 还是 LLaMA底层用的都是 QKV 注意力、多头计算、位置编码、LayerNorm、FFN 这些模块。你今天花半天时间手撕一个最小实现未来看大模型源码时能省下大量时间。这篇文章不追求讲全论文里的所有细节更希望给你一条“先跑通再理解再深入”的路径。所有代码都基于 PyTorchCPU 也能运行零基础完全可以直接起步。2. 先从架构图看起Transformer 到底由哪几块组成在动手写代码之前我们先建立整体地图。Transformer 最初是用于机器翻译的分为 Encoder编码器和 Decoder解码器。本文只实现 Encoder 部分但它涵盖了 Transformer 最核心的几个组件。整个 Encoder 流程可以概括为一句话Token 序列 - 词向量 位置编码 - 多层 Encoder Layer - 输出而每一个 Encoder Layer 内部又包含四样东西多头自注意力模块Multi-Head Self-Attention残差连接Residual ConnectionLayerNorm 层归一化前馈网络Feed-Forward Network简称 FFN。用一个生活化的类比来理解自注意力像是一场全员会议每个 Token 都会“提问”、会“展示自己”、也会“吸收别人的信息”会议结束后每个 Token 带着新的信息进入前馈网络相当于“会后个人整理笔记”。残差连接和 LayerNorm 则是让这场会议稳定进行的基础设施避免信息传递过程中丢失或者震荡。很多初学者以为 Transformer 就是“一堆注意力”这是误解。注意力的确是最重要的组件但没有残差、归一化和前馈网络深层 Transformer 根本训练不动。从代码实现的角度我们需要依次解决四件事文本怎么变成模型能计算的向量Token Embedding怎么让向量携带位置信息位置编码怎么让向量之间互相交换信息自注意力 多头怎么把这些模块稳定地拼成深层网络残差 LayerNorm FFN。下面就从环境开始一步步来做。3. 环境准备PyTorch 版本与安装建议手撕 Transformer 需要的基本依赖很少只需要 PyTorch 和一个可用的 Python 环境。建议使用 Python 3.8 以上版本PyTorch 2.x 的任意稳定版都可以本文代码不依赖 PyTorch 2.x 的新特性老版本 1.13 也能跑通。CPU 完全可以运行因为我们的最小模型很小GPU 只是让训练更快。推荐用 conda 创建独立环境避免依赖冲突# 创建并激活虚拟环境 conda create -n transformer python3.10 -y conda activate transformerPyTorch 的安装命令在不同系统和 CUDA 版本下不一样最稳妥的方式是到 PyTorch 官网选择对应选项后复制安装命令。如果不需要 GPU可以选择 CPU 版本# CPU 版本示例具体以官网为准 pip install torch torchvision torchaudio如果默认官方源下载很慢可以使用国内镜像源例如# 解决下载慢问题时可以临时指定镜像源 pip install torch -i https://mirrors.aliyun.com/pypi/simple/安装完成后用下面两条命令验证环境python -c import torch; print(torch.__version__) python -c import torch; print(torch.cuda.is_available())第一条会输出 PyTorch 版本号第二条输出True或False。CPU 环境输出False完全正常不需要纠结。如果你在安装时遇到“手机开热点下载还是很慢”这类问题本质是网络问题优先换镜像源。4. 从 Token 到 Embedding让模型“吃”数字机器学习的底层逻辑是数值计算模型不认识汉字和英文单词只认识数字。所以无论用什么模型第一步都是把文本拆成 Token再把 Token 映射成数字 id最后把 id 变成稠密向量。先看一个最小示例感受一下这个流程import torch import torch.nn as nn # 1. 分词这里用最简单的空格切分 sentence 我 爱 学习 tokens sentence.split() print(分词结果:, tokens) # 2. 构建词表给每个词分配一个 id vocab {pad: 0, unk: 1} for token in tokens: if token not in vocab: vocab[token] len(vocab) print(词表:, vocab) # 3. 文本 - id 序列 ids [vocab[token] for token in tokens] print(id 序列:, ids) # 4. 定义 Embedding 层把 id 映射为稠密向量 embedding nn.Embedding(num_embeddingslen(vocab), embedding_dim64) embedded embedding(torch.tensor(ids)) print(词向量形状:, embedded.shape) # [seq_len, embedding_dim]这里nn.Embedding做的事很简单它维护一个形状为[词表大小, embedding_dim]的查找表每个 id 对应一行向量。Embedding层本身是可学习的训练过程中会不断更新。这里要理解一个关键点Embedding 输出的最后一个维度在 Transformer 里叫d_model代表每个 Token 的向量维度。比如d_model64那么每个 Token 都会变成一个 64 维的向量。Transformer 内部的所有计算都是围绕这个维度展开的。实际训练中我们通常是一次处理一个 batch 的句子。如果句子长度不一样还要做 padding把短句子补齐到 batch 内最长长度否则无法组成矩阵计算。后面训练脚本中会看到具体做法。5. 位置编码为什么正余弦公式能告诉模型顺序先思考一个问题自注意力机制在计算时对每个 Token 都是一视同仁的。它不在乎“我”是第一个词还是第五个词这让模型天然拥有“并行计算”的优势但也带来一个严重问题——模型完全感知不到词的顺序。RNN 天然按照时间步逐个处理单词所以顺序信息是隐含的。Transformer 放弃了这个结构就必须用别的方式把“位置”注入进去。这就是位置编码存在的意义。最初论文《Attention Is All You Need》中使用的是正余弦位置编码公式如下PE(pos, 2i) sin(pos / 10000^(2i/d_model)) PE(pos, 2i1) cos(pos / 10000^(2i/d_model))其中pos是 Token 在序列中的位置i是向量的维度下标。这个公式看起来复杂但代码实现却很简单import math import torch import torch.nn as nn class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000): super().__init__() # 初始化一个 [max_len, d_model] 的位置编码矩阵 pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) div_term torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) # 偶数下标用 sin奇数下标用 cos pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) # 增加 batch 维度方便后续直接相加 pe pe.unsqueeze(0) # shape: [1, max_len, d_model] self.register_buffer(pe, pe, persistentFalse) def forward(self, x): # x: [batch, seq_len, d_model] return x self.pe[:, :x.size(1)]这里有两个容易被问到的点。第一为什么要用register_buffer因为pe不是模型参数不需要梯度更新但需要随着模型一起搬到 GPU。register_buffer恰好就是“非参数的张量但会随模型设备移动”的语义。第二为什么选正余弦而不是直接学一个位置向量最初论文作者推断正余弦函数能够帮助模型“更容易”感知相对位置。不过也必须说清楚后来的 GPT 等模型大量使用了“可学习位置向量”效果也很好。所以位置编码并不是必须用正余弦而是“必须要有”。正余弦作为理解 Transformer 的经典实现是绕不开的一块内容。6. 自注意力机制与 QKV最核心的一块拼图自注意力是整个 Transformer 的灵魂。很多人第一次学到这里会被 Q、K、V 三个字母吓到但其实它们背后是一个特别好懂的逻辑。把 QKV 放到一个场景里理解假设你在参加一场技术会议周围有很多人其他 Token。QQuery是“提问向量”它代表“我现在想了解什么”KKey是“标签向量”它代表“我能提供哪方面的信息”VValue是“内容向量”它代表“我实际给出的内容是什么”。每个 Token 都同时扮演三个角色它会提问也会给别人提供标签和内容。注意力计算的过程就是拿自己的 Q 去和所有 Token 的 K 做匹配算出注意力分数然后用分数对所有 V 做加权求和得到当前 Token 的新向量。公式是Attention(Q, K, V) softmax(Q * K^T / sqrt(d_k)) * V其中d_k是 K 向量的维度。为什么要除以sqrt(d_k)因为 Q 和 K 做点积后数值会随着维度变大而变大容易让 softmax 进入梯度很小的区域缩放之后计算的梯度更稳定。下面用 PyTorch 实现单头缩放点积注意力import torch import torch.nn as nn import math def scaled_dot_product_attention(Q, K, V, maskNone): # Q/K/V shape: [batch, seq_len, d_k] d_k K.size(-1) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn_weights torch.softmax(scores, dim-1) output torch.matmul(attn_weights, V) return output, attn_weights注意这里的 Q、K、V 已经是输入向量经过线性变换之后的结果。它们不是原始输入而是三个可学习矩阵W_Q、W_K、W_V映射出来的Q x W_Q.T K x W_K.T V x W_V.T这也是很多人写代码时容易混淆的地方输入的x同时被映射成三份而不是直接把原始x当作 QKV 使用。实际实现时我们会用nn.Linear来承担这种线性映射。最简单的单头注意力模块可以写成class SingleHeadAttention(nn.Module): def __init__(self, d_model): super().__init__() 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) def forward(self, x, maskNone): Q self.W_Q(x) # [batch, seq_len, d_model] K self.W_K(x) V self.W_V(x) output, weights scaled_dot_product_attention(Q, K, V, mask) return output到这一步你已经把自注意力机制的核心写出来了。接下来要解决的是为什么一个注意力不够还要搞多头7. 多头注意力一次关注多组关系单头注意力的问题在于它只能“在一种语义关系上”看输入。比如一个句子中可能有语法关系、词义关系、指代关系单头注意力往往只能捕获其中一部分。多头注意力的思路很简单把d_model拆成num_heads份每一份做一次独立的自注意力再把所有结果拼回去最后经过一次线性投影。每个头关注不同的子空间整体上模型就能同时看到多种关系。具体来说如果d_model64num_heads4那么每个头负责 16 维的 Q/K/V 计算最后把 4 个[batch, seq_len, 16]的头拼接成[batch, seq_len, 64]。实现如下class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads 0, d_model 必须能被 num_heads 整除 self.num_heads num_heads self.head_dim d_model // num_heads 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, x, maskNone): batch_size, seq_len, _ x.size() # 经过线性映射后拆成 num_heads 份 # Q/K/V: [batch, seq_len, num_heads, head_dim] - [batch, num_heads, seq_len, head_dim] Q self.W_Q(x).view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) K self.W_K(x).view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) V self.W_V(x).view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) # 对每个头单独做注意力 attn_output, _ scaled_dot_product_attention(Q, K, V, mask) # 合并所有头先换回 [batch, seq_len, num_heads, head_dim] attn_output attn_output.transpose(1, 2).contiguous() attn_output attn_output.view(batch_size, seq_len, self.num_heads * self.head_dim) # 最后的线性投影 return self.W_O(attn_output)很多初学者第一次写多头注意力最容易栽在view和transpose的顺序上。这里的关键是view把最后一个维度切成[num_heads, head_dim]然后transpose把num_heads挪到第 1 维这样注意力分数计算自动作用于每个头。一个常见的困惑是“多头注意力会不会让参数翻倍”这里不会。因为W_Q是从d_model映射到d_model参数量和单头注意力的W_Q一样只是把线性变换后的结果人为拆成了多份。多头不是一个额外增加参数的模块而是对同一个特征空间做了分组计算。为了更直观这里整理一下单头和多头的区别对比项单头注意力多头注意力关注关系一种关系多种关系每个头特征维度d_modeld_model / num_heads参数量线性映射 d_model - d_model相同表达能力较弱更强且具有并行性8. 残差连接、LayerNorm 和 FFN组装 Encoder Layer有了多头注意力我们只是造出了 Transformer 的“核心零件”还不能直接堆成深层网络。要让多个注意力层叠加起来稳定训练还需要三样东西残差连接、LayerNorm、前馈网络。残差连接的思路特别简单输出 输入 子层输出。这样一来即使子层内部效果不好信息也能直接绕过它传递到下一层避免深层网络出现梯度消失或退化问题。LayerNorm 和 BatchNorm 经常被拿来对比。BatchNorm 在一个 batch 的所有样本之间做归一化适合 CV 任务LayerNorm 则是在每一个样本内部对所有特征做归一化不依赖 batch 大小所以更适合 NLP 中长度不固定、batch 较小的情况。Transformer 选择 LayerNorm 是合理的。前馈网络是一个简单的两全连接层class FeedForwardNetwork(nn.Module): def __init__(self, d_model, d_ff2048): super().__init__() self.fc1 nn.Linear(d_model, d_ff) self.fc2 nn.Linear(d_ff, d_model) self.relu nn.ReLU() def forward(self, x): return self.fc2(self.relu(self.fc1(x)))d_ff通常比d_model大很多这样网络可以在更高维空间做非线性变换增强表达能力。大模型里常见的d_ff4 * d_model就是这个思路。把上面的零件组合起来就得到了一个完整的 Encoder Layerclass EncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff2048, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads) self.ffn FeedForwardNetwork(d_model, d_ff) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): # 第一部分自注意力 残差 LayerNorm attn_out self.self_attn(x, mask) x self.norm1(x self.dropout(attn_out)) # 第二部分前馈网络 残差 LayerNorm ffn_out self.ffn(x) x self.norm2(x self.dropout(ffn_out)) return x这里代码顺序是x dropout(attn_out)的组合再送入 LayerNorm。这是经典的 Post-LN 结构也是原始 Transformer 论文中的写法。后续很多模型也采用 Pre-LN也就是先归一化再进入子层。两者对训练稳定性有不同影响本文不做深入读者先掌握其中一种即可。9. 组装最小 Transformer 并训练现在我们把所有模块串起来做一个完整的小型 Transformer。模型结构很简单Embedding - 位置编码 - N 个 Encoder Layer - 输出分类层。class SimpleTransformer(nn.Module): def __init__(self, vocab_size, d_model, num_heads, num_layers, d_ff, max_len64): super().__init__() self.embedding nn.Embedding(vocab_size, d_model) self.positional_encoding PositionalEncoding(d_model, max_len) self.encoder_layers nn.ModuleList([ EncoderLayer(d_model, num_heads, d_ff) for _ in range(num_layers) ]) self.output_fc nn.Linear(d_model, vocab_size) def forward(self, x, maskNone): # x: [batch, seq_len] x self.embedding(x) x self.positional_encoding(x) for layer in self.encoder_layers: x layer(x, mask) return self.output_fc(x) # [batch, seq_len, vocab_size]注意self.output_fc会把每个 Token 的d_model维度的向量映射回词表大小得到每个位置属于每个词的得分。这就是一个 Encoder-only 的 Transformer可以用来做句子表示、分类或者按位置的预测任务。为了训练这个模型我设计一个非常简单的任务copy task。也就是输入一个 Token 序列期望模型原样输出同样的序列。在这个任务里我们只要看 loss 能不能降下来就能验证自注意力和整个架构是否实现正确。下面是一个完整的训练脚本。请把前面写的PositionalEncoding、MultiHeadAttention、FeedForwardNetwork、EncoderLayer、SimpleTransformer按顺序放入同一个 Python 文件中再运行训练代码# 文件路径train_transformer.py import torch import torch.nn as nn import torch.optim as optim # 将前文定义好的类复制到此处 # PositionalEncoding # scaled_dot_product_attention # MultiHeadAttention # FeedForwardNetwork # EncoderLayer # SimpleTransformer # 1. 构造一个小语料目标是把输入原样输出 corpus [i love code, i love ai, code is fun, ai is future] # 2. 构建词表 vocab {pad: 0, unk: 1} for sent in corpus: for token in sent.split(): if token not in vocab: vocab[token] len(vocab) print(词表:, vocab) max_len 4 # 3. 文本转 id并 padding 到定长 def encode(sent): ids [vocab[token] for token in sent.split()] ids [vocab[pad]] * (max_len - len(ids)) return ids batch torch.tensor([encode(sent) for sent in corpus]) # 4. 将数据放到 GPU如果可用 device torch.device(cuda if torch.cuda.is_available() else cpu) batch batch.to(device) # 5. 初始化模型 model SimpleTransformer( vocab_sizelen(vocab), d_model32, num_heads4, num_layers2, d_ff64, max_lenmax_len, ).to(device) criterion nn.CrossEntropyLoss(ignore_indexvocab[pad]) optimizer optim.Adam(model.parameters(), lr3e-4) # 6. 训练 for epoch in range(100): optimizer.zero_grad() logits model(batch) # [batch, seq_len, vocab_size] loss criterion(logits.view(-1, len(vocab)), batch.view(-1)) loss.backward() optimizer.step() if epoch % 10 0: print(fepoch {epoch:3d}, loss {loss.item():.4f}) # 7. 预测验证 model.eval() with torch.no_grad(): pred model(batch).argmax(dim-1) print(\n预测结果:) print(pred) print(期待输出:) print(batch)这段代码的核心逻辑是每一步训练时将当前 batch 输入模型得到每个位置的词表得分然后和原始输入做交叉熵损失。因为任务太简单模型不需要百分百精准只要 loss 持续下降就能说明前向传播、反向传播和整个 Transformer 链路是通的。如果训练时 loss 几乎没有下降优先检查位置编码是否用register_buffer保存多头注意力view/transpose之后维度是否正确模型输出预测层前是否用了正确的维度。10. 运行结果与效果验证运行python train_transformer.py正常情况下你会观察到 loss 呈持续下降趋势。例如当训练到几十个 epoch 后loss 会降到一个比较低的范围。由于每次训练存在随机性不同环境下的 loss 数值不会完全一致关键判断标准是趋势loss 是否稳定下降最终预测结果和输入是否基本一致模型是否能在不到一分钟内完成训练。如果这几个条件都满足说明你手写的 Transformer 是可训练、可收敛的。这一步的意义很大很多人的代码能跑但 loss 完全不降那通常意味着某个模块在逻辑上有 bug比如注意力 mask 写错、位置编码维度对不上、或者损失函数没有忽略 padding 位置。在验证时还可以顺手打印注意力权重直观感受模型的学习效果_, attn_weights scaled_dot_product_attention(Q, K, V) print(attn_weights.shape) # [batch, num_heads, seq_len, seq_len]attn_weights的每一行代表当前位置对所有位置的注意力概率。训练前它可能接近均匀分布训练后会出现明显的侧重这就是大家常说的“注意力可视化”。11. 常见问题与排查思路手写 Transformer 的过程中有几个问题出现频率非常高整理成表格方便对照排查问题现象可能原因排查方式解决方案loss 不下降学习率过大或过小打印每个 epoch 的 loss调整学习率建议从 3e-4 开始尝试loss 为 NaN数据包含异常值或数值溢出检查输入是否含 NaN检查 padding 位置是否被错误计算进损失维度不一致报错view/transpose写错打印每一步 tensor shape重点核对多头注意力的 reshape 顺序效果差但 loss 正常位置编码未加或加错检查 forward 里是否执行了x pe确认位置编码加
返回列表