手撕Transformer:用PyTorch从token到训练循环的完整实现

手撕Transformer:用PyTorch从token到训练循环的完整实现
我第一次完整跑通 Transformer 的训练循环是在一次需要把定长序列做分类的小项目里。虽然调包也能跑但当我试图把序列长度从 128 改成 512模型突然不收敛再后来想加一个因果 mask发现维度怎么都对不上。那时候我才意识到光会调用 Transformer 是不够的。真正把 Transformer“手撕”一遍用 PyTorch 从 token 到训练完整走通不是重复造轮子而是把那些看起来抽象的概念——自注意力、QKV、位置编码、层归一化——变成一个一个能调试、能验证的代码片段。这篇文章我按实际操作顺序来写从数据准备开始一直写到最小训练循环尽量把每一个维度变化和设计取舍讲清楚。1. 先搞清楚 Transformer 真正解决的是什么问题1.1 从 RNN 到自注意力并行与长距离依赖的取舍在 Transformer 出现之前序列建模的主流工具是 RNN 这一类结构。RNN 的问题是它天然是按时间步顺序处理的当前时刻的隐藏状态必须等前一时刻计算完才能得到。这种顺序依赖给并行计算带来了很大麻烦。你可能可以用一些技巧在 batch 内部做并行但序列维度本身很难被真正切断。另一个问题是长距离依赖。虽然 LSTM、GRU 通过门控机制改善了梯度传播但在非常长的序列里信息经过多个时间步之后仍然容易被稀释。CNN 可以并行感受野也能通过堆叠层数扩大但每扩大一层能看到的范围只是线性增长。而且要让两个距离很远的位置直接交互通常需要很多层或者很大的卷积核。这既增加了参数也让建模变得间接。Transformer 的核心转变是让序列里任意两个位置之间可以直接建立联系。自注意力机制做的事就是在一个序列内部每个位置都去计算它和其他所有位置的相关性然后根据相关性把信息聚合过来。这个过程没有顺序依赖可以在一次矩阵乘法里完成所有位置之间的交互。1.2 为什么最后是 Transformer而不是堆更深的 RNN你当然可以把 RNN 堆得很深也可以把 CNN 堆得很宽。但在大规模训练和超长序列场景里Transformer 的几个优势太明显了第一自注意力可以把整个序列的交互一次性并行计算第二任意两个位置之间的距离都被压缩成一步理论上没有信息衰减第三矩阵乘法的计算模式非常契合 GPU 的并行架构。需要说明的是这不代表 RNN 没有价值。PyTorch 里仍然保留着 RNN、GRU、LSTM 的实现在很多小规模、强时序场景下它们仍然有用。但如果你关注的是预训练语言模型、图像 Transformer 这类主流方向几乎都是基于 Transformer 或它的变体。所以“为什么最后是 Transformer”答案不是一句“它更好”就能说完而是它在并行性、长距离建模、训练稳定性和扩展性上取得了更好的综合平衡。手撕 Transformer 之前先理解这个背景很重要。因为后面写代码时你会发现很多设计都是为了保住这几点优势能并行、能传播梯度、能处理长距离。理解了为什么代码就不会是死记硬背。2. 数据准备从原始文本到 token、embedding 与注意力掩码2.1 最简 tokenization 与 embedding 层的实现零基础手撕 Transformer不建议一开始就上 BPE 或 SentencePiece。建议先用一个最简单的字符级 tokenizer 把流程跑通。字符级的意思是把一段文本拆成一个个字符然后给每个字符分配一个 id。这样词表很小处理逻辑也很直观。import torch import torch.nn as nn import math # 一个极简字符级词表 text the quick brown fox jumps over the lazy dog chars sorted(set(text)) vocab_size len(chars) char2idx {c: i for i, c in enumerate(chars)} idx2char {i: c for c, i in char2idx.items()} # 文本转 id token_ids [char2idx[c] for c in text]这里有一个关键点Transformer 的输入不是一串文本而是一个形状为(batch_size, seq_len)的整数张量。每个整数是一个 token id。然后通过nn.Embedding把 token id 映射成稠密向量d_model 64 embedding nn.Embedding(vocab_size, d_model) x torch.tensor(token_ids).unsqueeze(0) # (1, seq_len) embedded embedding(x) # (1, seq_len, d_model)nn.Embedding本质上就是一张查找表。输入 id 是什么就取出对应的一行向量。真正进入 Transformer 的不是原始文本也不是整数 id而是这些稠密向量。d_model是模型内部每个 token 的特征维度后面所有子层的输入输出都会围绕这个维度展开。2.2 padding mask让模型知道哪些位置不该被注意真实场景里一个 batch 里的序列长度往往不一样。最常见的手段是把短序列用pad补齐到和长序列一样长。但 pad 不是真实内容如果模型在注意力计算时把 pad 当作正常 token 来关注就会学到一堆没有意义的模式。所以我们需要一个 mask用来标记哪些位置是真实 token哪些位置是 pad。假设input_ids是 batch 数据pad_id 是 0def build_padding_mask(input_ids, pad_id0): # 返回 True 表示有效位置False 表示 pad 位置 return input_ids ! pad_id这个 mask 后续要传给注意力函数。在注意力计算中我们会把 pad 位置对应的 attention score 替换成一个非常大的负数让 softmax 之后权重趋近于 0。这里的核心原则是pad 位置可以接收信息但不应被其他位置关注。有一种常见误区是以为 padding mask 要在整个模型里全局使用。实际上它主要影响注意力分数。在 Loss 计算时如果你用到了 padding还要记住把目标位置上 pad token 的损失忽略掉不然模型会一直在“预测下一个 token 是 pad”这件事上浪费容量。后面训练部分我会专门再说。3. 手写多头自注意力QKV 不是三个神秘矩阵3.1 从单头缩放点积注意力开始自注意力的第一步是为每个 token 生成三个向量Query、Key、Value。用一句话理解它们的分工Query 是你想查的信息Key 是候选信息的索引Value 是候选信息的内容。你可以把整个序列当成一个图书馆。每个 token 都会发出一个查询想找到和它相关的信息每个 token 也会提供一个钥匙Key和一份资料Value。注意力机制做的事就是让查询和所有钥匙做匹配然后把匹配到的资料按权重汇总。在 PyTorch 里最基础的缩放点积注意力可以这样写def scaled_dot_product_attention(Q, K, V, maskNone): d_k Q.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, -1e9) weights torch.softmax(scores, dim-1) return torch.matmul(weights, V), weights这里有几个细节值得展开。第一Q和K的形状通常是(batch_size, ..., seq_len, d_k)。为了算两两之间的相关性要把K的最后两个维度转置变成(batch_size, ..., d_k, seq_len)这样矩阵相乘结果的形状就是(batch_size, ..., seq_len, seq_len)。这个矩阵里的第 i 行第 j 列就是第 i 个 Query 和第 j 个 Key 的相似度。第二除以sqrt(d_k)非常重要。如果d_k很大点积结果的数值范围会变大softmax 的梯度会变得非常小模型训练会很不稳定。缩放之后点积的方差被控制住梯度能更平稳地回传。第三masked_fill是掩码操作。mask 里为 0 的位置score 被替换成一个很大的负数。经过 softmax 后这个位置的概率会趋近于 0。这就是 padding mask 或因果 mask 真正起作用的地方。3.2 多头是如何并行且不丢信息的单头注意力的问题在于模型只能有一种“查询方式”。但真实语言中一个词可能同时和句法角色、语义内容、指代关系相关。多头注意力允许模型同时使用多组 QKV 投影每组关注不同的关系最后把所有头的结果拼起来。代码上多头注意力并不是真的并行开多个循环而是通过 reshape 把多个头放在同一个 batch 维度里做矩阵乘法。这样既高效也更符合 PyTorch 的习惯。class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() self.d_model d_model self.num_heads num_heads self.d_k 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() Q self.W_q(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) K self.W_k(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) V self.W_v(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) attn, _ scaled_dot_product_attention(Q, K, V, mask) # attn shape: (batch_size, num_heads, seq_len, d_k) attn attn.transpose(1, 2).contiguous().view(batch_size, seq_len, self.d_model) return self.W_o(attn)这段代码第一次看很容易懵但拆开看其实不复杂。原始输入x的形状是(batch_size, seq_len, d_model)。线性层把每个 token 的向量从d_model映射到d_model。然后我们用.view(batch_size, seq_len, num_heads, d_k)把最后一维拆成num_heads和d_k两部分。注意d_model num_heads * d_k常见的配置是 8 个头、每个头 64 维这样d_model就是 512。紧接着.transpose(1, 2)把num_heads这个维度提到 seq_len 前面让每个头独立做注意力计算。这一步做完Q、K、V 的形状都是(batch_size, num_heads, seq_len, d_k)。在scaled_dot_product_attention里我们只看到Q.size(-1)是d_k所以 mask 必须能广播到最后的 scores 形状。scores 的形状是(batch_size, num_heads, seq_len, seq_len)因此标准做法是给 mask 增加两个维度比如mask.unsqueeze(1).unsqueeze(2)让它变成(batch_size, 1, 1, seq_len)然后 PyTorch 会自动广播。if mask is not None: attn, _ scaled_dot_product_attention(Q, K, V, mask.unsqueeze(1).unsqueeze(2))最后把多头的结果拼回去attn.transpose(1, 2).contiguous().view(batch_size, seq_len, d_model)。这里.contiguous()是为了确保内存布局连续不然.view会报错。view是按顺序展开的而transpose会改变内存跳跃方式所以必须先连续化再 reshape。3.3 因果自注意力让模型只能看过去训练语言模型时如果目标是预测下一个 token模型就不能在预测当前位置时看到后面的 token。这需要一种叫因果 mask 的掩码也叫自回归 mask。因果 mask 是一个下三角矩阵当前位置只能看到当前位置和之前的位置不能看到未来。def build_causal_mask(seq_len): return torch.tril(torch.ones(seq_len, seq_len)).bool()如果 batch 里同时有 padding 和因果 mask要把两个 mask 结合起来。一个简单做法是让 padding mask 和因果 mask 按位与combined_mask padding_mask.unsqueeze(1).unsqueeze(2) causal_mask.unsqueeze(0)这里padding_mask形状是(batch_size, seq_len)扩展后为(batch_size, 1, 1, seq_len)causal_mask形状是(seq_len, seq_len)扩展后为(1, seq_len, seq_len)。广播后得到(batch_size, seq_len, seq_len)再通过广播适配多头最后变成(batch_size, num_heads, seq_len, seq_len)。有一个很容易踩的坑因果 mask 的维度到底应该放在哪一维。如果你的输入序列是(batch, seq)那么 scores 是(batch, heads, seq_q, seq_k)。因果 mask 的大小应该是(seq_q, seq_k)并且要确保mask[i][j]为 True 时表示位置 i 可以看到位置 j。如果你不小心把 mask 转置了模型会在训练时不断看到未来信息loss 会假装很低但生成时立刻崩掉。排查方式很简单打印scores和mask的形状再检查mask[0][0]是不是一个下三角。4. 位置编码给没有顺序感的模型加上坐标4.1 为什么 Transformer 本身完全没有顺序概念自注意力对输入顺序是“对称”的。如果把序列里两个 token 的位置互换attention 计算出来的结果只会相应交换行和列但模型本身不知道“猫追狗”和“狗追猫”到底哪个语义更合理。比如“我打你”和“你打我”token 完全相同但语义完全不同。如果没有位置信息Transformer 会把这两句话当成同一个输入。这就是为什么必须有位置编码在 token embedding 上叠加一个表示位置的信号让模型能够区分不同位置的 token。有一个常见的理解误区是位置编码只是给模型加一个“序号”。实际上它不是简单的 0、1、2、3而是一种能和 token embedding 相加、参与后续注意力计算的高维向量。位置编码的设计会影响模型区分相邻位置和远距离位置的能力。4.2 正余弦位置编码的实现与直觉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) pe pe.unsqueeze(0) # (1, max_len, d_model) self.register_buffer(pe, pe) def forward(self, x): return x self.pe[:, :x.size(1)]代码里div_term的作用是让不同维度的正余弦频率不同。位置 0 的编码、位置 1 的编码、位置 2 的编码各不相同同时位置相差 k 的两个编码之间存在线性关系。这个性质理论上可以让模型更容易学到相对位置。简单解释一下对于固定的偏移 kPE(pos k)可以看作PE(pos)的一个线性变换。因此模型不需要单独为每一个“距离为 k”的关系学习一套参数只要学到了这种线性变换就能把相对位置信息用起来。当然正余弦编码不是唯一选择。工程上也有人用可学习位置编码直接把位置当作一个可训练的 Embedding。它在小规模任务上往往更省心但在序列长度超过训练长度时外推能力通常不如正余弦编码。两者各有取舍。手撕阶段我建议先用正余弦编码因为它不增加可训练参数量也能让你理解位置信号为什么要以向量的形式存在。4.3 可学习位置编码与相对位置编码的取舍可学习位置编码的实现也很简单self.pos_embedding nn.Parameter(torch.zeros(1, max_len, d_model))但要注意如果 max_len 是训练时固定的 512推理时来了一个长度为 600 的序列pos_embedding就取不到第 512 之后的位置。你需要做截断、插值或者换用相对位置编码。这也是很多现代模型使用旋转位置编码、ALiBi 这类方案的原因。手撕 Transformer 时不需要在位置编码上做太复杂的选择。先用正余弦把流程跑通理解它为什么存在之后再去看更高级的位置编码方案就会轻松很多。5. 一个完整的 Transformer 编码器块5.1 残差连接与层归一化为什么能保住深层训练Transformer 不是把多头注意力结果直接输出而是把它和输入相加再做 LayerNorm。这个过程叫残差连接。残差连接让梯度可以绕过注意力层直接流回更浅层避免深层模型中的梯度消失。Transformer 一般会堆很多层如果没有残差连接网络会很深很难训练。层归一化是对每个 token 的特征维度做归一化让均值接近 0、方差接近 1。它与 BatchNorm 的区别是不依赖 batch 内其他样本因此在变长序列和在线推理场景下更稳定。在代码里最常见的两种顺序是 Post-Norm 和 Pre-Norm。Post-Norm 是原始 Transformer 论文里的写法先注意力再加残差再 LayerNorm。Pre-Norm 则是先 LayerNorm再注意力再加残差。工程经验里Pre-Norm 通常更容易稳定训练尤其是层数较多的时候。所以你会在很多开源模型里看到 Pre-Norm 的实现。我们在下面的代码里用 Pre-Norm。5.2 FFN自注意力之后的“信息加工间”自注意力输出的是每个 token 融合了其他 token 信息后的向量。接下来通常会过一个前馈神经网络 FFN两个线性层加一个激活函数。class FeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) self.linear2 nn.Linear(d_ff, d_model) self.dropout nn.Dropout(dropout) def forward(self, x): return self.linear2(self.dropout(torch.relu(self.linear1(x))))d_ff通常比d_model大很多常见的是 4 倍。比如d_model512d_ff2048。这个结构的直觉是自注意力已经在序列维度上做了信息交换FFN 再对每个位置逐点做一次非线性变换。它不像注意力那样在 token 之间通信而是对每个 token 自己进行“加工”。5.3 Pre-Norm 还是 Post-Norm工程里的一个常见选择把多头注意力、FFN、残差连接和 LayerNorm 组合起来就是一个完整的 Transformer 编码器块。class TransformerBlock(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout0.1): super().__init__() self.attn MultiHeadAttention(d_model, num_heads) self.ffn FeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, x, maskNone): # Pre-Norm x x self.dropout1(self.attn(self.norm1(x), mask)) x x self.dropout2(self.ffn(self.norm2(x))) return x这里有一个值得注意的细节self.norm1(x)是在注意力之前做归一化而不是之后。如果你用的是 Post-Norm顺序会变成x self.norm1(x self.dropout1(self.attn(x, mask))) x self.norm2(x self.dropout2(self.ffn(x)))两种写法在不太深的模型里都能用。但如果你发现训练深层 Transformer 时 loss 不稳定可以先检查是不是 Pre-Norm。在实现阶段我建议直接固定用 Pre-Norm减少一个变量。6. 用最小 Transformer 完成一次训练6.1 准备一个最简单的任务预测下一个 token要验证手撕的 Transformer 真的能工作不需要准备一个大语料。我们可以用字符级输入做一个“预测下一个字符”的任务。这里我用模型结构上的 decoder-only接收输入 token 序列输出每个位置下一个 token 的 logits。训练时通过因果 mask 让模型只能看当前位置之前的内容。先定义一个小型模型class MiniTransformer(nn.Module): def __init__(self, vocab_size, d_model64, num_heads4, num_layers2, d_ff256, max_len128, dropout0.1): super().__init__() self.embed nn.Embedding(vocab_size, d_model) self.pos PositionalEncoding(d_model, max_len) self.blocks nn.ModuleList([ TransformerBlock(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.ln_f nn.LayerNorm(d_model) self.head nn.Linear(d_model, vocab_size) def forward(self, x, maskNone): x self.embed(x) x self.pos(x) for block in self.blocks: x block(x, mask) x self.ln_f(x) logits self.head(x) return logits注意我们在这里没有做 padding所以 mask 只需要因果 mask。如果你的训练数据是变长 batch就需要把 padding mask 和因果 mask 组合起来。数据部分我们把字符序列切成固定长度的“窗口”。比如每次取 128 个字符作为输入下一个字符作为预测目标然后窗口每次向后滑动一个字符。def create_sequences(token_ids, seq_len): inputs, targets [], [] for i in range(0, len(token_ids) - seq_len - 1, seq_len): inputs.append(token_ids[i:iseq_len]) targets.append(token_ids[i1:iseq_len1]) return torch.tensor(inputs), torch.tensor(targets)把 token_ids 传进来后inputs[i]的内容是tokens[0:128]targets[i]是tokens[1:129]。也就是说模型看到位置 0 到 127要在每个位置预测位置 1 到 128 的 token。这和 decoder-only 语言模型的标准训练方式一致。6.2 训练循环、损失函数与优化器训练循环本身不复杂关键是要把输出和目标对齐。model MiniTransformer(vocab_sizevocab_size, d_model64, num_heads4, num_layers2, d_ff256) criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr1e-3) seq_len 128 inputs, targets create_sequences(token_ids, seq_len) causal_mask torch.tril(torch.ones(seq_len, seq_len)).bool() model.train() for epoch in range(500): optimizer.zero_grad() logits model(inputs, causal_mask) logits logits.view(-1, vocab_size) targets_flat targets.view(-1) loss criterion(logits, targets_flat) loss.backward() optimizer.step() if epoch % 50 0: print(fepoch {epoch}, loss {loss.item():.4f})这里没有用到 padding所以 loss 不需要忽略 pad。如果你构造了 padding目标里的 pad 位置应该设为pad_id并且在CrossEntropyLoss里设置ignore_indexpad_id这样模型不会在 pad 位置计算损失。还有一点logits.view(-1, vocab_size)会把(batch_size, seq_len, vocab_size)压平成(batch_size * seq_len, vocab_size)目标也压平成(batch_size * seq_len,)。这一步是 PyTorch 常见的“把序列当成批次”的写法可以省掉循环。6.3 怎么判断模型真的学到了而不是在过拟合如果你的数据只有一句话500 步之后 loss 降到非常低不代表模型学会了语言它只是记住了训练序列。真正判断模型学到什么的简单方法是做一次生成给定一个启动字符用模型贪心生成后面的字符。一个最简单的贪心生成前向推理def generate(model, start_str, char2idx, idx2char, max_new_tokens30): model.eval() tokens [char2idx[c] for c in start_str] input_seq torch.tensor([tokens]) with torch.no_grad(): for _ in range(max_new_tokens): current_len input_seq.size(1) mask torch.tril(torch.ones(current_len, current_len)).bool() logits model(input_seq, mask) next_logits logits[0, -1, :] next_token torch.argmax(next_logits).item() tokens.append(next_token) input_seq torch.cat([input_seq, torch.tensor([[next_token]])], dim1) return .join(idx2char[t] for t in tokens)如果生成的字符和训练文本风格接近说明模型至少学会了“看前文猜下一字符”。如果生成结果乱串先不要急着改模型结构先检查 loss 是否在持续下降、因果 mask 是否有效、学习率是否合适。有一个容易被忽略的点在递归生成时每次都要根据当前长度重新构造因果 mask。因为input_seq在变长mask 不能固定不变。如果你在训练时用的是固定seq_len的 mask推理时长度变了mask 的维度也要跟着变。7. 常见报错与排查链路先看维度再看 mask最后看数据7.1 最常出现的三类问题第一类是维度不匹配。常见报错是mat1 and mat2 shapes cannot be multiplied。出现这种问题优先打印每一层的输入输出形状。比如 Q、K、V 在view和transpose前后分别是什么形状scores 是什么形状最后拼回多头结果时有没有连续化。第二类是 mask 广播失败。scores是(batch, heads, seq_q, seq_k)mask 是(batch, seq)。如果你直接传给masked_fill大概率会报错。正确做法是mask.unsqueeze(1).unsqueeze(2)让 mask 变成(batch, 1, 1, seq)。如果 mask 是纯因果 mask维度是(seq, seq)就要先广播到 batch 维度。第三类是训练不收敛或 loss 为 NaN。这里要先检查数据里有没有nan再检查学习率。Transformer 对学习率比较敏感常见做法是先用 warmup再逐步衰减。手撕阶段可以先用一个很小的学习率比如1e-3然后把 batch size 调小一点看看 loss 会不会稳定下降。如果 loss 剧烈震荡先降低学习率。7.2 一份从现象到根因的排查顺序我自己排障时会按这个顺序来先看前向输出能不能算出来。写一个固定小输入跑一次model(x)看 logits 的形状是不是(batch, seq, vocab)。再看 mask 是否正确。打印 mask 的形状和具体值确认 pad 位置为 0 或 False有效位置为 1 或 True。因果 mask 要确认是下三角。再单步检查 loss。如果 loss 第一次迭代就极高或极低要看损失函数是否用了ignore_index目标是否 shift 正确。再检查梯度。在loss.backward()后打印几个参数的梯度范数。如果梯度很大说明需要梯度裁剪或降低学习率。最后看数据量。如果训练文本太短loss 可能无法反映真实效果换一个稍微大一点的数据集或者用更小的d_model和层数先跑通。这组顺序背后的逻辑是从模型结构到数据从静态到动态。模型如果根本跑不出输出再调数据没有意义输出正常再检查 maskmask 没问题再看训练过程这样能避免在错误层面浪费时间。7.3 手撕之后用官方实现还是保留手写版手撕完模型之后你可能会想我到底要不要在生产代码里用自己写的这个版本我的建议很明确学习阶段可以手写生产阶段优先用官方实现或成熟框架。PyTorch 自带了nn.TransformerEncoderLayer和nn.TransformerDecoderLayerHugging Face 的 transformers 也提供了大量预训练模型和训练工具。它们已经处理好了很多工程细节比如注意力 mask、数值稳定性、内存优化、推理加速等。但这不意味着手撕白费了。恰恰相反正因为你手撕过你才能看懂官方实现里那些参数是什么意思才知道为什么nn.TransformerEncoderLayer里有一个src_key_padding_mask、还有一个src_mask它们分别对应 padding mask 和因果 mask。在遇到 bug 时你也不会只是把报错贴给搜索引擎而是能大致判断问题出在哪一层。8. 从最小实现到真实工程还差什么8.1 手写版的边界在哪里手写版本适合教学和小规模验证但它离生产环境还有一些距离。首先是性能。我们手写的多头注意力是一次性计算所有位置的注意力分数时间和内存复杂度都是序列长度的平方。真实场景中长序列会直接打爆显存所以工程上需要 Flash Attention、稀疏注意力、分块计算等优化手段。其次是推理优化。自回归生成时每一步都要重新计算所有位置的 Key 和 Value这在工程上很不划算。真实模型会使用 KV Cache把已经计算过的 Key 和 Value 缓存下来避免重复计算。第三是训练稳定性。真实模型通常需要学习率调度、梯度裁剪、混合精度、分布式训练等机制。手写版本为了简洁很多细节都被省略了。如果你要训练一个真正的模型这些都不能少。8.2 用理解去读更高级的实现当你已经能把一个最小 Transformer 从 token 到训练完整跑通接下来可以做的事是去读一份更成熟的实现对照着看有哪些设计是你在手写时没考虑到的。比如官方nn.MultiheadAttention需要你传入query、key、value并且支持attn_mask和key_padding_mask两个 mask。attn_mask负责控制注意力可以看到哪些位置key_padding_mask负责忽略 padding。如果你手拆过 mask这个接口就很容易理解。再比如很多开源模型里的RotaryPositionEmbedding是在 Q/K 向量上做旋转操作来实现位置信息而不是简单的加法。当你理解了“位置编码是给模型提供顺序坐标”之后你会更容易接受旋转编码的做法它不是把位置信息加在 embedding 上而是直接调整 Q/K 的方向让两个向量的点积自然包含相对位置信息。8.3 回到手撕的价值所以零基础手撕 Transformer最终目标不是“我不用官方库了”而是“我知道官方库在做什么”。当你看到一个维度不对的报错当你遇到 attention mask 广播失败当你发现模型生成结果重复你不会再觉得 Transformer 是一个黑盒。你能顺着输入、embedding、位置编码、QKV、注意力分数、mask、FFN、loss 的路径一级一级往下排查。这篇文章里所有代码都是为了让你亲手把这条路走通。更具体地说你可以从最基础的字符级任务开始先把单头注意力跑通再加上多头再加上位置编码再堆两个 block最后训练一个小模型。每一步都打印 shape都检查 loss都试着 generate 几个 token。等你真的看完一次 loss 下降你对 Transformer 的理解会比看十篇论文更扎实。回到开头那个问题为什么需要手撕因为“调包”只能让你知道 API不能让你知道为什么 sequence length 变了就不收敛为什么 attention mask 会挡住结果为什么位置编码不是多余的。真正跑过一遍这些坑才会变成你自己的经验。

最新新闻

日新闻

周新闻

月新闻