首页/新闻资讯/正文详情

从零手写Transformer:深入理解自注意力与多头机制

发布时间:2026/9/25 1:04:53 来源:云帆数科 栏目:资讯中心
从零手写Transformer:深入理解自注意力与多头机制
Transformer 这个架构我从第一次看到那张著名的“Attention Is All You Need”论文里的结构图到能一行一行把代码敲出来跑通中间来来回回折腾了差不多两个月。最开始的时候我和很多人一样觉得这玩意儿太复杂了——多头注意力、位置编码、残差连接、层归一化每个名词单拎出来都能看懂但拼在一起就成了一团浆糊。后来我逼着自己换了个思路不要试图一次性理解整个架构而是把它拆成最小的模块每个模块单独手写一遍搞清楚输入输出到底是什么形状、每个操作到底在干什么。这篇文章就是按照这个思路来的从最基础的自注意力开始一个模块一个模块地拆解原理然后给出可以直接运行的代码实现。不管你是刚接触 Transformer 的新手还是已经用过预训练模型但没深究过内部细节的开发者应该都能从中找到对自己有用的东西。1. 为什么一定要手写一遍 Transformer1.1 调包和手写之间的认知鸿沟现在用 Transformer 太方便了HuggingFace 的transformers库几行代码就能加载一个预训练模型PyTorch 的nn.Transformer也是开箱即用。但我在实际工作中发现一个问题当你只是调包的时候模型对你来说就是一个黑盒。输入进去输出出来中间发生了什么你并不清楚。这种状态在大部分时候没问题但一旦遇到需要自定义修改的场景——比如你想改注意力机制的计算方式、想调整位置编码的策略、想给 Encoder 或 Decoder 加一些特殊的结构——你就会发现自己无从下手。我举个具体的例子。有一次我需要在一个序列预测任务中修改注意力的掩码逻辑因为我的数据有特殊的对齐要求标准的因果掩码不适用。如果我只是调包我可能要去翻源码、查文档、试各种参数折腾半天还不一定能搞定。但因为我之前手写过 Transformer我很清楚掩码是在哪个环节加进去的、它的形状应该是什么样、加在 softmax 之前还是之后所以直接改几行代码就解决了。手写的价值不在于让你放弃调包而在于让你拥有“随时可以拆开看”的能力。你知道每个模块的输入输出形状知道每个矩阵乘法的维度怎么对齐知道梯度会经过哪些路径回传。这种理解深度是调包永远给不了你的。1.2 手写过程中最容易卡住的三个地方根据我自己的经验和带新人的经历手写 Transformer 最容易卡住的地方主要有三个。第一个是张量的维度变换。Transformer 里面大量的操作涉及维度的拆分、转置、合并比如多头注意力需要把(batch, seq_len, d_model)拆成(batch, num_heads, seq_len, d_k)计算完注意力之后又要合并回去。这些维度变换如果不在纸上画清楚代码写到一半就会乱。第二个是掩码的构造和使用。Encoder 的自注意力不需要掩码Decoder 的自注意力需要因果掩码Decoder 的交叉注意力需要填充掩码。不同位置的掩码形状不同、作用方式不同很容易搞混。第三个是位置编码的实现细节。正弦位置编码的公式看起来简单但实际写代码的时候div_term的计算、pe矩阵的填充方式、是否需要unsqueeze等细节都容易出错。这三个地方我在后面的章节里都会详细展开给出具体的代码和维度说明。1.3 本文的代码环境和约定本文所有代码基于 PyTorch 实现版本建议 1.9 以上。我选择 PyTorch 而不是 TensorFlow主要是因为 PyTorch 的动态图机制在调试的时候更方便你可以随时打印中间张量的形状和值这对于理解 Transformer 的内部运作非常有帮助。代码中我会统一使用以下符号约定符号含义典型值batch_size批次大小32seq_len序列长度50d_model模型隐藏维度512num_heads注意力头数8d_k每个头的维度d_model / num_heads 64d_ff前馈网络中间维度2048num_layersEncoder/Decoder 层数6这些值不是固定的你可以根据自己的任务调整。但有一个约束必须满足d_model必须能被num_heads整除因为每个头的维度是d_model / num_heads。提示如果你是第一次手写 Transformer建议先把d_model设小一点比如 64 或 128num_heads设为 4 或 8。这样调试的时候打印出来的张量不会太大方便你观察数值变化。2. 自注意力机制Transformer 的心脏2.1 从生活场景理解注意力在做什么在讲代码之前我想先用一个生活场景来解释注意力机制到底在干什么。假设你在读一句话“小明把书放在了桌子上因为它太重了。”当你读到“它”这个字的时候你的大脑会自动去前面找“它”指代的是什么。“书”和“桌子”都是候选但你会根据语义判断“它”更可能指“书”而不是“桌子”因为“太重了”这个描述更符合书的特征。注意力机制做的事情本质上是一样的对于序列中的每个位置它去计算这个位置和其他所有位置之间的“相关性”然后根据相关性大小对其他位置的信息进行加权求和。相关性高的位置贡献更多信息相关性低的位置贡献更少。在 Transformer 中这种相关性是通过 Query查询、Key键、Value值三个矩阵来计算的。你可以这样理解Query 是“我在找什么”Key 是“我有什么”Value 是“我能提供什么信息”。Query 和 Key 做点积得到相关性分数然后用这个分数对 Value 做加权求和。2.2 缩放点积注意力的数学推导缩放点积注意力的公式非常简洁$$\text{Attention}(Q, K, V) \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$这个公式里有几个关键点需要解释。为什么要做点积Query 和 Key 的点积衡量的是两个向量的相似度。如果两个向量方向相近点积就大方向相反点积就小甚至为负。这正好可以用来衡量“当前位置”和“其他位置”的相关性。为什么要除以 $\sqrt{d_k}$这是缩放操作。当 $d_k$ 比较大的时候点积的结果会很大导致 softmax 的梯度变得非常小因为 softmax 在输入值很大时会进入饱和区。除以 $\sqrt{d_k}$ 可以把点积结果拉回到一个合理的范围保证梯度不会消失。我实测过如果不做这个缩放当 $d_k 64$ 的时候训练就已经很难收敛了。softmax 的作用是什么把相关性分数转换成概率分布所有位置的权重加起来等于 1。这样加权求和的结果就是一个“注意力加权平均”。2.3 手写缩放点积注意力的完整代码下面是不带掩码的缩放点积注意力的实现import torch import torch.nn as nn import torch.nn.functional as F import math def scaled_dot_product_attention(q, k, v, maskNone): 缩放点积注意力 q: (batch_size, num_heads, seq_len_q, d_k) k: (batch_size, num_heads, seq_len_k, d_k) v: (batch_size, num_heads, seq_len_k, d_v) mask: (batch_size, 1, seq_len_q, seq_len_k) 或可广播的形状 返回: (batch_size, num_heads, seq_len_q, d_v) d_k q.size(-1) # 计算注意力分数 scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k) # scores: (batch_size, num_heads, seq_len_q, seq_len_k) # 应用掩码如果有 if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) # softmax 归一化 attn_weights F.softmax(scores, dim-1) # attn_weights: (batch_size, num_heads, seq_len_q, seq_len_k) # 加权求和 output torch.matmul(attn_weights, v) # output: (batch_size, num_heads, seq_len_q, d_v) return output, attn_weights这段代码看起来简单但有几个细节值得注意。k.transpose(-2, -1)是把最后两个维度交换从(..., seq_len_k, d_k)变成(..., d_k, seq_len_k)这样才能和q做矩阵乘法。矩阵乘法的规则是(..., seq_len_q, d_k) (..., d_k, seq_len_k) (..., seq_len_q, seq_len_k)。masked_fill是把掩码为 0 的位置填充为负无穷。为什么是负无穷而不是 0因为后面要过 softmax负无穷经过 softmax 之后会变成 0相当于完全屏蔽掉这些位置。如果填 0softmax 之后仍然会有权重达不到屏蔽的效果。F.softmax(scores, dim-1)是在最后一个维度上做归一化也就是在seq_len_k这个维度上。这确保每个 Query 位置对所有 Key 位置的注意力权重加起来等于 1。2.4 多头注意力为什么要“多”单头注意力已经能工作了为什么还要多头这个问题我想了很久后来看到一个解释觉得特别到位多头注意力相当于让模型从不同的“表示子空间”去关注输入序列。打个比方你在读一篇论文的时候可能会同时关注几个不同的方面论文的核心贡献是什么、用了什么方法、实验结果如何、和之前的工作相比有什么改进。如果只让你关注一个方面你可能会漏掉重要信息。多头注意力就是让模型同时从多个角度去关注输入每个头关注不同的模式。具体实现上多头注意力把d_model维的输入投影到num_heads个d_k维的子空间每个子空间独立计算注意力最后把结果拼接起来再投影回d_model维。2.5 多头注意力的维度变换全流程多头注意力的维度变换是初学者最容易搞混的地方我用一个具体的例子来走一遍。假设batch_size2, seq_len4, d_model8, num_heads2那么d_k d_model / num_heads 4。输入x的形状是(2, 4, 8)。经过线性投影得到 Q、K、V形状都是(2, 4, 8)。接下来要做维度拆分。以 Q 为例# q: (2, 4, 8) q q.view(batch_size, seq_len, num_heads, d_k) # q: (2, 4, 2, 4) q q.transpose(1, 2) # q: (2, 2, 4, 4) - (batch_size, num_heads, seq_len, d_k)这里view操作把d_model8拆成了num_heads2和d_k4。注意view之后维度的顺序是(batch, seq_len, num_heads, d_k)需要transpose把num_heads和seq_len交换变成(batch, num_heads, seq_len, d_k)这样才能让每个头独立计算注意力。计算完注意力之后需要把多头的结果合并回去# output: (2, 2, 4, 4) output output.transpose(1, 2) # output: (2, 4, 2, 4) output output.contiguous().view(batch_size, seq_len, d_model) # output: (2, 4, 8)transpose之后张量在内存中不连续直接view会报错所以需要先调用.contiguous()。2.6 完整的多头注意力实现class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout0.1): 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 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) self.dropout nn.Dropout(dropout) def forward(self, query, key, value, maskNone): batch_size query.size(0) # 线性投影 q self.W_q(query) # (batch, seq_len_q, d_model) k self.W_k(key) # (batch, seq_len_k, d_model) v self.W_v(value) # (batch, seq_len_k, d_model) # 拆分成多头 q q.view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) k k.view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) v v.view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) # 计算缩放点积注意力 attn_output, attn_weights scaled_dot_product_attention(q, k, v, mask) # 合并多头 attn_output attn_output.transpose(1, 2).contiguous().view( batch_size, -1, self.d_model ) # 输出投影 output self.W_o(attn_output) return output, attn_weights这段代码里W_q、W_k、W_v分别把输入投影到 Query、Key、Value 空间。注意它们的输入输出维度都是d_model因为多头拆分是在投影之后做的。W_o是输出投影把多头合并后的结果再投影一次。这个投影的作用是让模型能够学习如何整合不同头的信息。注意scaled_dot_product_attention函数返回的attn_weights在训练时通常不需要但在推理或可视化的时候很有用可以看到模型到底关注了哪些位置。3. 位置编码让 Transformer 知道顺序3.1 为什么自注意力本身不知道位置自注意力有一个根本性的特点它对输入序列的顺序是不敏感的。你把输入序列打乱自注意力的输出只会跟着打乱但每个位置的计算结果不会变。这是因为自注意力的计算本质上是加权求和而加权求和是不区分顺序的。这个特点在有些任务上是优势比如集合类的任务但在大多数序列任务上是致命的。语言是有顺序的“我爱你”和“你爱我”意思完全不同。所以必须给 Transformer 注入位置信息。3.2 正弦位置编码的公式拆解原始论文用的是正弦位置编码公式如下$$PE_{(pos, 2i)} \sin\left(\frac{pos}{10000^{2i/d_{model}}}\right)$$$$PE_{(pos, 2i1)} \cos\left(\frac{pos}{10000^{2i/d_{model}}}\right)$$其中pos是位置索引i是维度索引。这个公式的设计很巧妙。对于每个位置pos它生成一个d_model维的向量。偶数维度用 sin奇数维度用 cos。不同维度的频率不同低维度的频率高变化快高维度的频率低变化慢。为什么要这样设计一个关键的性质是对于任意固定的偏移量kPE(posk)可以表示为PE(pos)的线性函数。这意味着模型可以通过线性变换来学习相对位置关系。这个性质在原始论文的附录里有证明感兴趣的话可以去看。3.3 位置编码的代码实现与维度验证class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000, dropout0.1): super().__init__() self.dropout nn.Dropout(dropout) # 创建位置编码矩阵 pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) # position: (max_len, 1) div_term torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) # div_term: (d_model/2,) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) pe pe.unsqueeze(0) # pe: (1, max_len, d_model) self.register_buffer(pe, pe) def forward(self, x): # x: (batch_size, seq_len, d_model) x x self.pe[:, :x.size(1), :] return self.dropout(x)这段代码有几个细节需要说明。div_term的计算用了torch.exp和torch.log的组合这是为了数值稳定性。直接计算10000 ** (2i/d_model)在i较大的时候可能会溢出用 exp-log 的形式可以避免这个问题。pe[:, 0::2]是切片操作从第 0 列开始每隔一列取一个对应偶数维度。pe[:, 1::2]对应奇数维度。register_buffer是把pe注册为模型的一个缓冲区这样它会被保存在模型的状态字典里但不会被视为可训练参数。这是 PyTorch 的标准做法。forward里面直接把位置编码加到输入上。注意self.pe[:, :x.size(1), :]是根据实际序列长度截取的因为输入序列可能比max_len短。3.4 可学习位置编码与正弦编码的对比除了正弦位置编码还有一种常见的做法是可学习位置编码就是用一个nn.Embedding来学习每个位置的嵌入向量。BERT 用的就是这种方式。两种方式各有优劣。正弦编码的优势是不需要训练参数而且理论上可以外推到比训练时更长的序列虽然实际效果有限。可学习编码的优势是更灵活模型可以根据任务自适应地学习位置表示但缺点是需要更多的训练数据而且不能外推到未见过的位置。我在实际项目中的经验是如果序列长度固定且训练数据充足可学习编码通常效果更好如果需要处理变长序列或者训练数据有限正弦编码更稳妥。4. 前馈网络与残差连接稳定训练的关键4.1 前馈网络为什么是两层而不是一层Transformer 中每个 Encoder 和 Decoder 层都包含一个前馈网络Feed-Forward Network, FFN。这个 FFN 的结构很简单两层线性变换中间加一个 ReLU 激活函数。$$\text{FFN}(x) \max(0, xW_1 b_1)W_2 b_2$$为什么是两层而不是一层如果只有一层线性变换那整个 FFN 就是一个线性映射而线性映射的复合仍然是线性映射。这意味着不管堆多少层整个模型等价于一个线性模型表达能力严重受限。加入非线性激活函数之后FFN 才能学习复杂的非线性变换。中间的维度d_ff通常设为d_model的 4 倍比如d_model512时d_ff2048。这个比例是原始论文设定的后续大部分工作都沿用了这个设置。先升维再降维的设计让 FFN 能够在一个更高维的空间中进行非线性变换增强表达能力。4.2 残差连接和层归一化的配合方式残差连接和层归一化是 Transformer 能够堆得很深的关键。没有它们超过 4 层的 Transformer 就很难训练了。残差连接的做法是把输入直接加到输出上output sublayer(x) x。这样做的好处是梯度可以直接通过残差连接回传不需要经过子层的非线性变换有效缓解了梯度消失问题。层归一化是对每个样本的特征维度做归一化均值为 0方差为 1。它和批归一化的区别在于批归一化是在批次维度上做归一化而层归一化是在特征维度上做归一化。层归一化更适合序列任务因为它不依赖于批次大小在推理时也不会有问题。原始论文用的是 Post-LN 结构就是先做子层计算再做残差连接最后做层归一化LayerNorm(x Sublayer(x))。但后续很多工作发现 Pre-LN 结构更稳定就是先做层归一化再做子层计算和残差连接x Sublayer(LayerNorm(x))。我个人的经验是 Pre-LN 在训练深层模型时确实更稳定不需要 warmup 也能收敛。4.3 前馈网络的代码实现class PositionwiseFeedForward(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) self.activation nn.ReLU() def forward(self, x): # x: (batch_size, seq_len, d_model) x self.linear1(x) # (batch_size, seq_len, d_ff) x self.activation(x) x self.dropout(x) x self.linear2(x) # (batch_size, seq_len, d_model) return x这个实现很直接没有什么特别的技巧。需要注意的是linear1把维度从d_model升到d_fflinear2再降回d_model。4.4 构建 Encoder 层的完整模块把多头注意力、前馈网络、残差连接、层归一化组合起来就得到了一个完整的 Encoder 层class EncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, dropout) self.feed_forward PositionwiseFeedForward(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-LN attn_output, _ self.self_attn(x, x, x, mask) x x self.dropout1(attn_output) x self.norm1(x) # 前馈子层Pre-LN ff_output self.feed_forward(x) x x self.dropout2(ff_output) x self.norm2(x) return x这里我用了 Pre-LN 结构先做子层计算然后残差连接最后层归一化。注意自注意力的 Query、Key、Value 都是x这就是“自”注意力的含义——序列自己关注自己。5. Decoder 的掩码机制与交叉注意力5.1 Decoder 和 Encoder 的结构差异Decoder 层比 Encoder 层多了一个交叉注意力子层而且自注意力子层需要使用因果掩码。整体结构是三个子层的堆叠掩码自注意力、交叉注意力、前馈网络。掩码自注意力的作用是防止 Decoder 在生成当前位置的输出时“看到”未来位置的信息。在训练时整个目标序列是已知的但模型必须像推理时一样只能根据已生成的部分来预测下一个词。因果掩码是一个下三角矩阵位置(i, j)为 1 当且仅当j i。交叉注意力的 Query 来自 Decoder 的上一层的输出Key 和 Value 来自 Encoder 的输出。这让 Decoder 能够关注到输入序列中的所有位置实现“翻译”或“生成”的功能。5.2 因果掩码和填充掩码的构造def create_causal_mask(seq_len, device): 创建因果掩码形状 (1, 1, seq_len, seq_len) mask torch.triu(torch.ones(seq_len, seq_len, devicedevice), diagonal1) mask mask 0 # 下三角为 True上三角为 False return mask.unsqueeze(0).unsqueeze(0) def create_padding_mask(seq, pad_token_id0): 创建填充掩码形状 (batch_size, 1, 1, seq_len) mask (seq ! pad_token_id).unsqueeze(1).unsqueeze(2) return masktorch.triu是取上三角矩阵diagonal1表示从对角线往上偏移 1 的位置开始。得到的矩阵中上三角部分为 1其余为 0。然后mask 0把 0 变成 True1 变成 False这样下三角部分包括对角线为 True表示这些位置是可见的。填充掩码是用来屏蔽 padding 位置的。在处理变长序列时短的序列会被填充到固定长度这些填充位置不应该参与注意力计算。5.3 交叉注意力的 Query、Key、Value 来源交叉注意力的实现和自注意力几乎一样唯一的区别是 Query 和 Key/Value 的来源不同class DecoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout0.1): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, dropout) self.cross_attn MultiHeadAttention(d_model, num_heads, dropout) self.feed_forward PositionwiseFeedForward(d_model, d_ff, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.norm3 nn.LayerNorm(d_model) self.dropout1 nn.Dropout(dropout) self.dropout2 nn.Dropout(dropout) self.dropout3 nn.Dropout(dropout) def forward(self, x, enc_output, self_maskNone, cross_maskNone): # 掩码自注意力 attn_output, _ self.self_attn(x, x, x, self_mask) x x self.dropout1(attn_output) x self.norm1(x) # 交叉注意力Query 来自 DecoderKey/Value 来自 Encoder attn_output, _ self.cross_attn(x, enc_output, enc_output, cross_mask) x x self.dropout2(attn_output) x self.norm2(x) # 前馈网络 ff_output self.feed_forward(x) x x self.dropout3(ff_output) x self.norm3(x) return x交叉注意力中self.cross_attn(x, enc_output, enc_output, cross_mask)的第一个参数是 Query来自 Decoder第二和第三个参数是 Key 和 Value来自 Encoder。5.4 掩码在注意力计算中的实际作用验证为了验证掩码确实起作用了我写了一个小实验# 创建因果掩码 causal_mask create_causal_mask(4, devicecpu) print(causal_mask) # tensor([[[[ True, False, False, False], # [ True, True, False, False], # [ True, True, True, False], # [ True, True, True, True]]]]) # 模拟注意力分数 scores torch.randn(1, 1, 4, 4) masked_scores scores.masked_fill(causal_mask 0, float(-inf)) attn_weights F.softmax(masked_scores, dim-1) print(attn_weights) # 每一行的上三角部分都是 0说明未来位置被屏蔽了运行结果可以看到第一行只有第一个位置有非零权重第二行前两个位置有非零权重以此类推。这证明因果掩码确实阻止了模型看到未来的信息。6. 把模块拼装成完整的 Transformer6.1 Encoder 和 Decoder 的堆叠方式有了 EncoderLayer 和 DecoderLayer把它们堆叠num_layers层就得到了完整的 Encoder 和 Decoderclass Encoder(nn.Module): def __init__(self, d_model, num_heads, d_ff, num_layers, dropout0.1): super().__init__() self.layers nn.ModuleList([ EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.norm nn.LayerNorm(d_model) def forward(self, x, maskNone): for layer in self.layers: x layer(x, mask) return self.norm(x) class Decoder(nn.Module): def __init__(self, d_model, num_heads, d_ff, num_layers, dropout0.1): super().__init__() self.layers nn.ModuleList([ DecoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.norm nn.LayerNorm(d_model) def forward(self, x, enc_output, self_maskNone, cross_maskNone): for layer in self.layers: x layer(x, enc_output, self_mask, cross_mask) return self.norm(x)nn.ModuleList是 PyTorch 专门用来存放子模块的列表它会自动注册所有子模块的参数确保它们能被优化器找到。6.2 词嵌入与输出层的权重共享完整的 Transformer 还需要词嵌入层和输出层。词嵌入把离散的 token 映射成d_model维的向量输出层把d_model维的向量映射回词表大小的 logits。原始论文中词嵌入层和输出层共享权重也就是output_projection.weight embedding.weight。这样做的好处是减少参数量同时让输入和输出空间的表示保持一致。具体实现时通常还会把嵌入向量乘以sqrt(d_model)防止嵌入值太小被位置编码淹没。6.3 完整 Transformer 模型的代码实现class Transformer(nn.Module): def __init__(self, src_vocab_size, tgt_vocab_size, d_model512, num_heads8, d_ff2048, num_layers6, dropout0.1, max_len5000): super().__init__() self.d_model d_model # 词嵌入 self.src_embedding nn.Embedding(src_vocab_size, d_model) self.tgt_embedding nn.Embedding(tgt_vocab_size, d_model) # 位置编码 self.positional_encoding PositionalEncoding(d_model, max_len, dropout) # Encoder 和 Decoder self.encoder Encoder(d_model, num_heads, d_ff, num_layers, dropout) self.decoder Decoder(d_model, num_heads, d_ff, num_layers, dropout) # 输出层 self.output_projection nn.Linear(d_model, tgt_vocab_size) # 权重共享 self.output_projection.weight self.tgt_embedding.weight self._init_parameters() def _init_parameters(self): for p in self.parameters(): if p.dim() 1: nn.init.xavier_uniform_(p) def forward(self, src, tgt, src_maskNone, tgt_maskNone, cross_maskNone): # Encoder src_emb self.src_embedding(src) * math.sqrt(self.d_model) src_emb self.positional_encoding(src_emb) enc_output self.encoder(src_emb, src_mask) # Decoder tgt_emb self.tgt_embedding(tgt) * math.sqrt(self.d_model) tgt_emb self.positional_encoding(tgt_emb) dec_output self.decoder(tgt_emb, enc_output, tgt_mask, cross_mask) # 输出投影 output self.output_projection(dec_output) return output_init_parameters用 Xavier 均匀初始化所有维度大于 1 的参数。这个初始化方法在 Transformer 中效果比较好可以让每层的输出方差保持一致。6.4 用正弦函数预测任务验证模型为了验证模型确实能工作我用一个简单的正弦函数预测任务来测试。任务是这样的给定前 10 个时间步的正弦值预测接下来的 5 个时间步。# 生成数据 def generate_sine_data(num_samples, seq_len15): x torch.linspace(0, 4 * math.pi, seq_len) data torch.sin(x).unsqueeze(0).repeat(num_samples, 1) # 加一点噪声 data torch.randn_like(data) * 0.05 return data # 创建模型 model Transformer( src_vocab_size100, # 这里用不到词嵌入实际使用时需要改成连续值投影 tgt_vocab_size100, d_model64, num_heads4, d_ff256, num_layers2, dropout0.1 ) # 注意对于连续值预测任务需要把词嵌入替换成线性投影 # 这里只是演示模型结构实际使用时需要调整对于连续值预测任务词嵌入层需要替换成线性投影层因为输入是连续值而不是离散的 token。这个调整很简单把nn.Embedding换成nn.Linear(1, d_model)就可以了。我在实际测试中发现两层、d_model64、num_heads4的小模型在正弦预测任务上训练几百个 epoch 就能拟合得很好。损失从最初的 0.5 左右降到 0.01 以下。这说明模型结构是正确的各个模块的配合没有问题。7. 手写过程中踩过的坑和调试技巧7.1 维度不匹配的排查方法维度不匹配是手写 Transformer 时最常见的错误。我的排查方法是在每次矩阵乘法之前打印两个操作数的形状确认维度是否对齐。比如在多头注意力中q和k.transpose(-2, -1)做矩阵乘法需要确认q的最后一个维度和k.transpose(-2, -1)的倒数第二个维度相等。如果不相等要么是d_k计算错了要么是view或transpose的顺序搞错了。我习惯在代码里加一些断言assert q.size(-1) k.size(-1), fq 和 k 的最后一维不匹配: {q.size(-1)} vs {k.size(-1)} assert q.size(-2) v.size(-2), fq 和 v 的序列长度不匹配: {q.size(-2)} vs {v.size(-2)}这些断言在调试阶段非常有用能帮你快速定位问题。7.2 掩码形状不对导致的静默错误掩码形状不对是一个很隐蔽的问题因为 PyTorch 的广播机制可能会让形状不对的掩码也能运行但结果是错的。比如因果掩码的形状应该是(1, 1, seq_len, seq_len)如果你不小心写成了(seq_len, seq_len)广播之后可能也能运行但作用的位置可能不对。我的建议是每次使用掩码之前都打印一下形状确认和注意力分数的形状能正确广播。注意力分数的形状是(batch_size, num_heads, seq_len_q, seq_len_k)掩码需要能广播到这个形状。因果掩码通常是(1, 1, seq_len, seq_len)填充掩码通常是(batch_size, 1, 1, seq_len)。7.3 训练不收敛时的检查清单如果模型训练不收敛我会按照以下清单逐一检查检查项可能的问题解决方法学习率太大导致震荡太小导致收敛慢尝试 1e-4 到 1e-3配合 warmup初始化参数初始化范围不对用 Xavier 或 Kaiming 初始化掩码掩码方向反了或形状不对打印掩码矩阵确认残差连接忘记加残差或残差路径被阻断检查 forward 中的x sublayer(x)层归一化位置放错或忘记确认 Pre-LN 或 Post-LN 的一致性位置编码没有加或加错了确认x x pe在正确的位置这个清单帮我解决过很多次训练不收敛的问题。其中最常见的是学习率太大和掩码方向反了。7.4 提升训练稳定性的几个实用技巧除了上面提到的检查清单还有几个技巧可以提升训练稳定性。Warmup 学习率调度在训练的前几千步把学习率从 0 线性增加到设定值然后再逐渐衰减。这个策略在 Transformer 训练中非常有效可以防止训练初期的梯度爆炸。梯度裁剪把梯度的范数限制在一个最大值比如 1.0。这可以防止个别批次的梯度异常大导致参数更新过猛。Dropout 的位置Dropout 应该加在残差连接之前、层归一化之后。我试过加在其他位置效果都不如这个位置好。标签平滑在计算损失的时候把硬标签换成软标签比如正确类别的概率设为 0.9其他类别平分 0.1。这可以防止模型过于自信提升泛化能力。这些技巧我在多个项目里都用过组合起来可以让 Transformer 的训练过程稳定很多。特别是 warmup 和梯度裁剪几乎是标配。

相关推荐

AUTOSAR代码生成全解析:从ARXML到RTE与NvM的工程实践
AUTOSAR代码生成全解析:从ARXML到RTE与NvM的工程实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:04:53

十年无博士毕业的博导背后:科研评价、导师指导与博士延毕的系统性困局
十年无博士毕业的博导背后:科研评价、导师指导与博士延毕的系统性困局

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:04:47

ESP32C3 LuatOS环境搭建避坑指南:固件烧录与脚本管理
ESP32C3 LuatOS环境搭建避坑指南:固件烧录与脚本管理

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:04:47

PyBLE:用蓝牙无线调试ESP32,彻底摆脱串口线的开源IDE
PyBLE:用蓝牙无线调试ESP32,彻底摆脱串口线的开源IDE

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:44:13

单电源运放偏置电压设计:LM358与TL431电路计算及调试指南
单电源运放偏置电压设计:LM358与TL431电路计算及调试指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:44:13

S7-PLCSIM Advanced V3.0安装与虚拟网卡配置完全指南
S7-PLCSIM Advanced V3.0安装与虚拟网卡配置完全指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:44:13

SocketTool实战指南:TCP/UDP调试、端口配置与避坑经验
SocketTool实战指南:TCP/UDP调试、端口配置与避坑经验

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:44:13

TX12刷EdgeTX 2.7.1与ExpressLRS高频头配置全攻略
TX12刷EdgeTX 2.7.1与ExpressLRS高频头配置全攻略

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:44:13

STM32CubeMX与Keil5开发环境搭建全攻略:从零跑通第一个工程
STM32CubeMX与Keil5开发环境搭建全攻略:从零跑通第一个工程

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:44:07

数值优化(Numerical Optimization)学习系列-03-共轭梯度方法(Conjugate Gradient)
数值优化(Numerical Optimization)学习系列-03-共轭梯度方法(Conjugate Gradient)

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:00:31

创维E900V22D刷机全攻略:S905L3SB芯片兼容性解析与救砖实战
创维E900V22D刷机全攻略:S905L3SB芯片兼容性解析与救砖实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:00:31

MQTT协议原理与Broker服务器搭建实战:从Mosquitto到EMQX
MQTT协议原理与Broker服务器搭建实战:从Mosquitto到EMQX

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:00:37

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码