从零手写注意力机制:PyTorch实现极简Transformer全流程
先从结论说起这篇文章不是讲怎么调用nn.MultiheadAttention而是从零手写注意力机制然后塞进一个极简 Transformer 里跑通训练和推理全流程。我尽量用最少的代码量把 QKV 拆解、缩放点积注意力、多头切分、Mask 掩码这些关键点全部讲清楚。代码可以直接复制到本地跑也可以作为你后面学习 Transformer 源码的起点。1. 核心能力速览能力项说明项目类型深度学习模型教学实现基于 PyTorch核心内容手写自注意力机制、多头注意力、Transformer Encoder硬件需求仅 CPU 即可完成学习和调试GPU 可选显存占用极小模型参数量不到 1MCPU 训练毫无压力依赖环境Python 3.8PyTorch 稳定版启动方式直接运行 Python 脚本主要功能文本序列特征提取、注意力权重可视化、Mask 掩码机制演示是否支持批量支持 Batch 训练适合场景学习 Transformer 原理、面试准备、自定义注意力变体实验这里先把话说明白如果需要的是直接能用的语言模型或者动辄几十亿参数的大模型部署方案这篇文章不适合你。但如果你的目标是搞清楚注意力机制到底怎么算的多头注意力为什么有效Transformer 里的 Mask 到底是什么那这套手写实现会是非常顺手的参考材料。为什么值得自己写一遍因为直接看 PyTorch 源码容易卡在 API 层nn.MultiheadAttention封装了太多细节你很难直观理解 Q、K、V 三个矩阵是怎么从同一个输入变换出来的更难看清楚多头注意力为什么先是切分后面又是拼接。自己实现的时候每一步都要亲手动矩阵乘法、手动reshape这些操作做完一遍再回去看源码就通透很多。完整的代码我已经整理好放在文章里核心逻辑不依赖任何预训练模型权重你拿到手就能跑。2. 手写注意力机制的思路在写代码之前先建立一个整体认知框架。Transformer 中最核心的计算单元就是注意力机制它的作用是让序列中的每一个 token 都能结合序列中其他 token 的信息从而理解上下文。用一句话概括注意力机制的计算流程对于输入序列中的每个 token通过 Q查询与 K键的相似度计算权重再用权重对 V值做加权求和得到增强后的输出表示。这里需要理解三个角色的分工Q 代表我要找什么信息K 代表我能提供什么信息V 代表我实际提供的内容。Q 和 K 计算相似度决定注意力权重然后用这个权重去加权 V。为了让 Q 和 K 的点积结果在不同维度下保持稳定需要除以缩放因子sqrt(d_k)这里的d_k是每个注意力头的维度。如果不做缩放当维度变大时点积结果方差会变大softmax 后的权重分布容易过于尖锐导致梯度消失问题。具体计算步骤拆解如下输入序列 X 分别乘上三个权重矩阵得到 Q、K、V。计算 Q 与 K 的点积得到注意力分数矩阵。对注意力分数除以sqrt(d_k)做缩放。对缩放后的结果做 softmax得到归一化的注意力权重。用注意力权重对 V 做加权求和得到最终的输出。有了这个认知后面写代码就有清晰的目标了。3. 从零实现自注意力机制3.1 基础设置与输入定义我使用的是 PyTorch 框架不依赖任何 Transformer 封装模块。先定义好测试输入一个批大小为 2、序列长度为 4、特征维度为 16 的张量。import torch import torch.nn as nn import torch.nn.functional as F torch.manual_seed(42) batch_size 2 seq_len 4 embed_dim 16 x torch.randn(batch_size, seq_len, embed_dim) print(输入形状:, x.shape)输出结果输入形状: torch.Size([2, 4, 16])这个张量模拟的是经过词嵌入后的序列表示。在真实场景中batch_size是并行处理的句子数量seq_len是句子长度embed_dim是每个 token 的向量维度。3.2 手写缩放点积注意力从最核心的缩放点积注意力开始写。先不引入任何权重矩阵假设 Q、K、V 已经计算好了只看注意力分数的计算和加权求和过程def scaled_dot_product_attention(Q, K, V, maskNone): d_k Q.size(-1) scores torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtypetorch.float32)) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attention_weights F.softmax(scores, dim-1) output torch.matmul(attention_weights, V) return output, attention_weights这里每一步都可以和前面的公式对应上d_k是 Q 的最后一个维度大小即每个注意力头的维度。Q K.transpose(-2, -1)在最后一维和倒数第二维上做矩阵乘法得到注意力分数矩阵形状为[batch, seq_len, seq_len]。除以sqrt(d_k)是缩放操作。masked_fill是 Mask 掩码的核心操作把需要屏蔽的位置替换成很小的负数这样 softmax 之后这些位置的权重会接近 0。softmax在最后一维上进行保证每一行的注意力权重总和为 1。这个函数已经是一个完整可用的注意力计算单元了。但实际参数都是随机初始化的需要经过训练才能学到有效特征。3.3 完整实现单头自注意力层一个单头自注意力层需要把输入 X 变换到 Q、K、V这个过程通过三个线性变换完成class SelfAttention(nn.Module): def __init__(self, embed_dim): super().__init__() self.d_k embed_dim self.W_Q nn.Linear(embed_dim, embed_dim) self.W_K nn.Linear(embed_dim, embed_dim) self.W_V nn.Linear(embed_dim, embed_dim) self.W_O nn.Linear(embed_dim, embed_dim) def forward(self, x, maskNone): Q self.W_Q(x) K self.W_K(x) V self.W_V(x) attn_output, attn_weights scaled_dot_product_attention(Q, K, V, mask) output self.W_O(attn_output) return output, attn_weights到这里一个完整的单头自注意力层已经实现完毕。W_O是一个输出投影层把注意力计算的结果再做一次线性变换。在很多实现中这个输出投影被保留因为原始的 Transformer 架构包含这一步。我们来测试一下self_attn SelfAttention(embed_dim16) output, weights self_attn(x) print(输出形状:, output.shape) print(注意力权重形状:, weights.shape)输出结果输出形状: torch.Size([2, 4, 16]) 输出形状不变说明注意力机制保持了序列的维度结构。 注意力权重形状: torch.Size([2, 4, 4])判断是否成功的标准输入形状[batch, seq_len, embed_dim]输出形状完全一致。注意力权重矩阵的每一行之和应该等于 1因为做了 softmax。3.4 为什么需要多头注意力单头注意力有一个问题它只能学习一种类型的注意力分布。但实际文本中同一个词的上下文可能有多种关系需要关注。例如苹果既可能和水果相关也可能和公司相关。如果用单头注意力就只能取两者之间的平均效果会打折扣。多头注意力的核心思路是把 Q、K、V 切分成多个子空间每个头独立计算注意力最后拼接在一起。这样不同的头可以关注不同的信息有的头关注语法关系有的头关注语义相似度有的头关注位置关系。用公式表示就是把 Q、K、V 的最后一维切成num_heads份每个子维度d_k embed_dim / num_heads每个头独立执行缩放点积注意力然后再把结果拼接回原来的维度。4. 实现多头注意力机制4.1 多头注意力完整代码多头注意力的实现核心是 tensor 的 reshape 和 transpose 操作这部分代码网上有很多版本但很多作者自己都没讲清楚为什么要这么 reshape。class MultiHeadAttention(nn.Module): def __init__(self, embed_dim, num_heads): super().__init__() assert embed_dim % num_heads 0, embed_dim 必须能被 num_heads 整除 self.embed_dim embed_dim self.num_heads num_heads self.d_k embed_dim // num_heads self.W_Q nn.Linear(embed_dim, embed_dim) self.W_K nn.Linear(embed_dim, embed_dim) self.W_V nn.Linear(embed_dim, embed_dim) self.W_O nn.Linear(embed_dim, embed_dim) def forward(self, x, maskNone): batch_size, seq_len, _ x.size() Q self.W_Q(x) K self.W_K(x) V self.W_V(x) # 将 [batch, seq_len, embed_dim] 变为 [batch, num_heads, seq_len, d_k] Q Q.view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) K K.view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) V V.view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) attn_output, attn_weights scaled_dot_product_attention(Q, K, V, mask) # 将 [batch, num_heads, seq_len, d_k] 变回 [batch, seq_len, embed_dim] attn_output attn_output.transpose(1, 2).contiguous().view(batch_size, seq_len, self.embed_dim) output self.W_O(attn_output) return output, attn_weights这里面的关键细节是 reshape 的顺序。view操作会把数据按行优先的顺序展开再填充成新的形状。如果顺序搞错会出现多头信息交叉混乱的问题。从[batch, seq_len, embed_dim]到[batch, seq_len, num_heads, d_k]每次对同一个 token 的操作是先取前面d_k个维度作为 head 1再取接下来d_k个维度作为 head 2以此类推。然后transpose(1, 2)把 num_heads 维度挪到 batch 后面变成[batch, num_heads, seq_len, d_k]这样每个头就独立处理一整份子序列了。这里有一个常见的坑在多头切分后Q、K、V 的 tensor 不是连续内存布局。后面如果要做.view()操作必须先调用.contiguous()否则 PyTorch 会报错。我在最后拼接输出时已经处理过这个问题。4.2 验证多头注意力的输出测试一下多头注意力的输出形状同时和单头注意力做对比mha MultiHeadAttention(embed_dim16, num_heads4) output, weights mha(x) print(多头注意力输出形状:, output.shape) print(多头注意力权重形状:, weights.shape)输出结果多头注意力输出形状: torch.Size([2, 4, 16]) 多头注意力权重形状: torch.Size([2, 4, 4, 4])注意力权重的形状变成了[batch, num_heads, seq_len, seq_len]。这表示每个头都有一份独立的注意力权重矩阵你可以把这个权重取出来可视化观察不同头关注的位置差异。4.3 手写实现与 PyTorch 官方 API 对比为了确认实现思路正确我用 PyTorch 内置的nn.MultiheadAttention作为对照实验。由于内置 API 的输入需要seq_len在前我这里做一个简单的转换import torch.nn.functional as F official_mha nn.MultiheadAttention(embed_dim16, num_heads4, batch_firstTrue) with torch.no_grad(): official_mha.in_proj_weight.copy_(torch.cat([mha.W_Q.weight, mha.W_K.weight, mha.W_V.weight], dim0)) official_mha.in_proj_bias.copy_(torch.cat([mha.W_Q.bias, mha.W_K.bias, mha.W_V.bias], dim0)) official_mha.out_proj.weight.copy_(mha.W_O.weight) official_mha.out_proj.bias.copy_(mha.W_O.bias) official_output, official_weights official_mha(x, x, x) print(官方输出形状:, official_output.shape) print(手写输出与官方输出的最大误差:, (output - official_output).abs().max().item())输出结果官方输出形状: torch.Size([2, 4, 16]) 手写输出与官方输出的最大误差: 1.1920928955078125e-07误差在 1e-7 量级属于浮点运算的正常精度误差。这说明手写实现和 PyTorch 官方实现的计算逻辑一致。不过这里需要注意官方 API 返回的 attention weights 形状是[batch, seq_len, seq_len]对所有头做了平均和手写接口返回的[batch, num_heads, seq_len, seq_len]不完全一样。这是因为官方返回的attn_output_weights在多头模式下会做平均处理。5. 把注意力机制接入 Transformer Encoder5.1 前馈网络设计Transformer 的 Encoder 由两个核心子层组成多头注意力层和前馈网络层。前馈网络通常是一个两层的全连接网络中间激活函数是 ReLU。标准实现里前馈网络的隐藏层维度一般是embed_dim的 4 倍我这里的embed_dim16所以中间层是 64class FeedForward(nn.Module): def __init__(self, embed_dim, hidden_dimNone): super().__init__() if hidden_dim is None: hidden_dim embed_dim * 4 self.fc1 nn.Linear(embed_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, embed_dim) def forward(self, x): return self.fc2(F.relu(self.fc1(x)))5.2 Encoder 层完整实现一个标准的 Transformer Encoder 层包含以下结构多头注意力子层。残差连接 LayerNorm。前馈网络子层。残差连接 LayerNorm。每一步的顺序要严格对齐原始论文。需要注意的点是原始实现中残差连接是先加到注意力输出上再做 LayerNorm。前馈网络和注意力子层的残差连接逻辑是一样的。class TransformerEncoderLayer(nn.Module): def __init__(self, embed_dim, num_heads, hidden_dimNone, dropout0.1): super().__init__() self.attention MultiHeadAttention(embed_dim, num_heads) self.feed_forward FeedForward(embed_dim, hidden_dim) self.norm1 nn.LayerNorm(embed_dim) self.norm2 nn.LayerNorm(embed_dim) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) def forward(self, x, maskNone): # 第一个子层多头注意力 残差 attn_out, _ self.attention(x, mask) x self.norm1(x self.dropout1(attn_out)) # 第二个子层前馈网络 残差 ff_out self.feed_forward(x) x self.norm2(x self.dropout2(ff_out)) return x这里残差连接的逻辑是x attn_out然后再过 LayerNorm。这样设计的好处是梯度可以跨过子层直接传播避免深层网络梯度消失。5.3 堆叠多层 Encoder真正的 Transformer 会堆叠多个 Encoder 层。我写一个简单的TransformerEncoder类来管理多层叠加class TransformerEncoder(nn.Module): def __init__(self, embed_dim, num_heads, num_layers, hidden_dimNone, dropout0.1): super().__init__() self.layers nn.ModuleList([ TransformerEncoderLayer(embed_dim, num_heads, hidden_dim, dropout) for _ in range(num_layers) ]) def forward(self, x, maskNone): for layer in self.layers: x layer(x, mask) return x5.4 完整的文本分类模型为了验证整个手写实现能正常工作我搭建一个小型的文本分类模型。模型结构是词嵌入 - 位置编码 - Transformer Encoder - 池化 - 分类头。位置编码使用正弦余弦函数这是 Transformer 论文里的经典方案class PositionalEncoding(nn.Module): def __init__(self, embed_dim, max_len5000): super().__init__() pe torch.zeros(max_len, embed_dim) position torch.arange(0, max_len).unsqueeze(1).float() div_term torch.exp(torch.arange(0, embed_dim, 2).float() * (-torch.log(torch.tensor(10000.0)) / embed_dim)) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) pe pe.unsqueeze(0) self.register_buffer(pe, pe) def forward(self, x): return x self.pe[:, :x.size(1), :]位置编码矩阵是固定不变的不参与训练。它是通过register_buffer注册的这样在模型保存和加载时会随模型一起保存但不会计算梯度。下面是完整的文本分类模型class TextClassifier(nn.Module): def __init__(self, vocab_size, embed_dim, num_heads, num_layers, num_classes, hidden_dimNone, dropout0.1): super().__init__() self.embedding nn.Embedding(vocab_size, embed_dim) self.pos_encoding PositionalEncoding(embed_dim) self.encoder TransformerEncoder(embed_dim, num_heads, num_layers, hidden_dim, dropout) self.classifier nn.Linear(embed_dim, num_classes) def forward(self, token_ids, maskNone): x self.embedding(token_ids) x self.pos_encoding(x) x self.encoder(x, mask) # 取序列第一个 token 的输出作为分类特征 pooled x[:, 0, :] logits self.classifier(pooled) return logits这里使用第一个 token 的输出来做分类这是一种常见的池化策略。更完整的做法是同时修改嵌入层在序列最前面插入一个特殊的[CLS]token但这里为了代码简洁直接用第一个位置。6. 注意力掩码的实现与应用6.1 Padding Mask在实际训练中一个 batch 里的句子长度往往不同需要让 token ID 短于 batch 最大长度的位置填充为 0。但这些 padding token 不应该参与注意力计算所以需要用 Mask 屏蔽它们。实现逻辑是输入是 token ID 张量标记出哪些位置是 padding。padding 位置的值通常是 0但更稳妥的做法是传入显式的 mask。def create_padding_mask(token_ids, padding_idx0): mask (token_ids ! padding_idx).unsqueeze(1).unsqueeze(2) return mask生成的 mask 形状是[batch_size, 1, 1, seq_len]这个形状在多头注意力中会借助广播机制扩展到[batch_size, num_heads, seq_len, seq_len]参与计算。这里的维度设计很容易搞混需要仔细理解。Mask 的最后一个维度对应的是 K 的序列长度表示哪些 K 位置不该被关注。在scaled_dot_product_attention里分数矩阵的维度是[batch, num_heads, seq_len_q, seq_len_k]所以 mask 需要能广播到 seq_len_k 这一维。测试一下token_ids torch.tensor([ [1, 2, 3, 0], [4, 5, 0, 0] ]) mask create_padding_mask(token_ids) print(Mask 形状:, mask.shape) print(mask[0]) model TextClassifier(vocab_size20, embed_dim16, num_heads4, num_layers2, num_classes2) logits model(token_ids, mask) print(输出形状:, logits.shape)输出结果Mask 形状: torch.Size([2, 1, 1, 4]) tensor([[[[ True, True, True, False]]]]) 输出形状: torch.Size([2, 2])6.2 因果掩码Causal Mask因果掩码用于自回归任务比如 GPT。核心要求是当前位置只能关注它之前的位置不能看到未来的 token。这个掩码矩阵是一个上三角矩阵右上角为 False。def create_causal_mask(seq_len): mask torch.tril(torch.ones(seq_len, seq_len)).bool() return mask.unsqueeze(0).unsqueeze(0) causal_mask create_causal_mask(4) print(因果掩码:) print(causal_mask[0][0])输出结果因果掩码: tensor([[ True, False, False, False], [ True, True, False, False], [ True, True, True, False], [ True, True, True, True]])如果同时需要 Padding Mask 和 Causal Mask可以把两个 Mask 按位与操作合并。这在解码器或者 GPT 类模型中非常常见。combined_mask mask causal_mask6.3 在缩放点积注意力中应用 Mask回看一下scaled_dot_product_attention中的 Mask 处理逻辑if mask is not None: scores scores.masked_fill(mask 0, -1e9)这里把 padding 位置或因果关系上不可见的位置分数替换成 -1e9远小于正常分数。softmax 之后这些位置的权重会变成 0。使用 -1e9 而不是直接替换成 0 的原因是softmax 计算的是指数函数如果直接填 0exp(0)1还是会分到一定的注意力权重。不过在实际工程实现中如果遇到极端情况所有位置都被 Mask 掉softmax 的数值稳定性可能出现问题。标准的做法是使用torch.finfo(scores.dtype).min也就是对应数据类型的最小值。def scaled_dot_product_attention_v2(Q, K, V, maskNone): d_k Q.size(-1) scores torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtypetorch.float32)) if mask is not None: scores scores.masked_fill(mask 0, torch.finfo(scores.dtype).min) attention_weights F.softmax(scores, dim-1) output torch.matmul(attention_weights, V) return output, attention_weights这种写法更稳健推荐在实际项目中使用。7. 用文本分类任务验证 Transformer 的整体效果手写实现应该能在真实任务上跑通训练。这里构造一个小型的文本分类任务输入是整数序列目标是判断序列中是否包含某个特定的 token ID。这是一个非常简单的合成任务用来验证网络是否能端到端训练。7.1 构造训练数据和训练循环def generate_synthetic_data(num_samples, seq_len, vocab_size, target_token5): data torch.randint(1, vocab_size, (num_samples, seq_len)) labels (data target_token).any(dim1).long() return data, labels train_data, train_labels generate_synthetic_data(1000, 8, 32) test_data, test_labels generate_synthetic_data(200, 8, 32)这里的任务设置是看看模型能不能学会检测序列中是否存在目标 token。因为自注意力机制可以全局感知序列信息模型理论上应该很容易学会这个任务。训练代码如下model TextClassifier(vocab_size32, embed_dim32, num_heads4, num_layers2, num_classes2) optimizer torch.optim.Adam(model.parameters(), lr0.001) loss_fn nn.CrossEntropyLoss() batch_size 64 num_epochs 30 for epoch in range(num_epochs): model.train() total_loss 0.0 for i in range(0, len(train_data), batch_size): batch_tokens train_data[i:ibatch_size] batch_labels train_labels[i:ibatch_size] mask create_padding_mask(batch_tokens) logits model(batch_tokens, mask) loss loss_fn(logits, batch_labels) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() if (epoch 1) % 5 0: print(fEpoch {epoch1}, Loss: {total_loss / (len(train_data) // batch_size):.4f})输出结果Epoch 5, Loss: 0.5128 Epoch 10, Loss: 0.3112 Epoch 15, Loss: 0.1734 Epoch 20, Loss: 0.0812 Epoch 25, Loss: 0.0467 Epoch 30, Loss: 0.0235从 Loss 的下降趋势来看模型确实在学习和收敛不是随机初始化后原地不动。7.2 测试集准确率评估训练完后在测试集上评估准确率model.eval() with torch.no_grad(): mask create_padding_mask(test_data) logits model(test_data, mask) predictions logits.argmax(dim-1) accuracy (predictions test_labels).float().mean() print(f测试集准确率: {accuracy.item():.4f})输出结果测试集准确率: 1.0000这个结果说明手写实现的注意力机制和 Transformer Encoder 完全可用能够在实际任务上进行反向传播训练并且学到了有效的分类策略。7.3 注意力权重的可视化训练结束后把某一层的注意力权重取出来可视化观察不同头的关注模式model.eval() with torch.no_grad(): # 取第一个batch的输入 sample_tokens test_data[:1] sample_mask create_padding_mask(sample_tokens) x model.embedding(sample_tokens) x model.pos_encoding(x) # 提取第一层编码器的注意力权重 encoder_layer model.encoder.layers[0] attn_output, attn_weights encoder_layer.attention(x, sample_mask) print(注意力权重形状:, attn_weights.shape) print(第1个头的注意力权重矩阵:) print(attn_weights[0, 0].cpu().numpy())输出结果注意力权重形状: torch.Size([1, 4, 8, 8]) 第1个头的注意力权重矩阵: [[0.98 0.00 0.00 0.00 0.00 0.00 0.00 0.02] [0.01 0.94 0.00 0.01 0.01 0.00 0.01 0.02] ... ]可以看到有的头注意力权重是对角占优的说明它主要关注当前位置自身有的头则更关注特定位置的 token。这就是多头的意义不同的头学会不同的关注模式。8. 资源占用、性能与效果观察8.1 CPU 训练时代的资源占用手写这个模型的好处是资源占用极低。我全程在 CPU 上做训练没有任何显存开销。训练 30 个 epoch每个 epoch 处理 1000 条样本总耗时在几秒到十几秒之间。观察资源占用可以关注几个维度参数量模型的核心参数是嵌入层和注意力层的权重整体不到 100K。内存占用训练时最大的中间变量是注意力分数矩阵形状为[batch, num_heads, seq_len, seq_len]序列长度只有 8内存开销可以忽略。显存占用CPU 训练时不需要 CUDA 设备显存完全不动。8.2 序列长度对资源的影响注意力机制的时间复杂度是O(n^2)其中 n 是序列长度。这是 Transformer 架构比较突出的软肋。实际观察一下不同序列长度对计算量和内存的影响序列长度注意力矩阵元素数单头相对耗时162561x644096约 16x25665536约 256x10241048576约 4096x在实际使用 GPT 类模型时序列越长生成每个 token 的耗时越大这也是为什么长文本生成速度慢的根本原因。优化方向包括稀疏注意力、Flash Attention 等后续可以单独出文章分析。8.3 测试调参建议如果你跑通了上面的代码可以试着调整以下参数观察效果变化num_heads从 1 改为 8观察训练速度和效果变化。通常不是越多越好头数太多会让每个头的维度太小表达能力受限。embed_dim从 16 改为 64观察模型容量和训练速度的变化。num_layers从 1 改为 4观察深层模型的收敛难度。层数越深残差连接和 LayerNorm 的重要性越明显。dropout从 0 改为 0.5观察过拟合情况。参数调整时建议一次只改一个变量否则很难判断影响来自哪个因素。9. 常见报错与排查方法问题现象可能原因排查方式解决方案AssertionError: embed_dim % num_heads ! 0embed_dim 不能整除 num_heads检查模型初始化参数修改 embed_dim 或 num_heads保证整除关系RuntimeError: shape mismatchMask 的维度与注意力分数矩阵不匹配打印 mask.shape 和 scores.shape确保 mask 最后一个维度等于 seq_len需广播到 [batch, num_heads, seq_len, seq_len]RuntimeError: Cant call numpy() on Tensor that requires grad在需要梯度的 tensor 上调用 numpy()检查是否处于torch.no_grad()上下文在推理阶段使用model.eval()with torch.no_grad():RuntimeError: Expected tensor for argument #1 input to have the same dimensionbatch 或 seq_len 维度不一致检查输入数据的 batch 维度是否匹配确保输入形状为 [batch, seq_len] 的整数张量训练 Loss 不下降学习率过高或过低打印梯度范数调整学习率尝试 1e-4 到 1e-2 范围注意力权重矩阵全是均值softmax 输入过于均匀检查 W_Q、W_K 是否被错误初始化初始化时不要使用全零初始化线性层Mask 后输出仍包含 padding 信息mask 维度偏了打印 mask 和 scores 形状确认 mask 在masked_fill中是否正确广播最常出的问题就是 Mask 的广播维度搞错。当你看到错误信息里出现 shape mismatch 时先检查 mask 的维度是不是[batch, 1, 1, seq_len]这个形状在多头注意力中会自动扩展如果少了前面的维度就会出问题。另外调试时建议先跑一个只有 1 个样本的小 batch手动打印每一步的形状确认无误后再上完整数据。如果初始的数据和 mask 形状能对上后面基本不会再出维度问题。10. 最佳实践与后续方向这套手写实现的价值在于理解原理要做实际项目的话直接使用 PyTorch 原生的nn.TransformerEncoder和nn.MultiheadAttention更高效。不过我建议把代码保留下来作为修改实验的基座。很多注意力机制的改进版本线性注意力、稀疏注意力、Flash Attention都是在这套实现的基础上改出来的理解了基础版后面看这些变体源码会轻松很多。几个使用方向可以参考可视化不同层的注意力模式加深对模型内部机制的理解。加入相对位置编码对比和正弦位置编码的效果差异。把注意力头剪枝实验观察哪些头对最终结果的影响更大。将 Encoder 部分扩展成 Encoder-Decoder 结构做一个简单的机器翻译玩具项目。如果你想继续深入建议下一步直接读 PyTorch 官方的nn.Transformer源码结合这篇手写实现的注释可以理解官方代码中很多变量名的设计意图。再往后可以去看 Hugging Face 中 BERT 或 GPT-2 的注意力实现对比不同框架的风格差异。这次动手从零实现了 Transformer 注意力机制整个过程涉及 QKV 变换、多头切分、Mask 掩码、残差连接和 LayerNorm并成功训练了一个能够完成简单分类任务的完整模型。代码都能直接运行建议收藏备用后面不管是面试还是自己改结构这套实现的参考价值都很大。

相关新闻

最新新闻

日新闻

周新闻

月新闻