手写 Transformer:从零实现注意力机制,这次终于把 QKV 搞明白了
如果你和我一样,第一次接触 Transformer 时就被“注意力机制”四个字绕得云里雾里,网上的代码一搜一大把,但真正敢说自己“手写过”的并不多。大多数项目里我们直接调用nn.MultiheadAttention,或者model.forward()一把梭,等面试官问起“Q、K、V 到底怎么来的”“为什么要除以根号 d_k”时,还是会卡壳。
这篇文章就用一个完整的实战过程,带你在 PyTorch 里手写注意力机制,并在此基础上搭建一个可运行的 Transformer 编码器。整个过程不依赖高级封装,你只需要了解基本的张量操作,跟着代码一步步走,就能彻底看清注意力机制从输入到输出的完整计算链路。
文章适合谁?适合刚学完 PyTorch 基础、想深入理解 Transformer 原理的初学者;也适合已经在用nn.Transformer但想搞清楚内部机制的开发者。全文包含核心公式推导、完整代码、可视化思路、常见报错排查和工程建议,建议收藏后边看边敲。
1. 背景与核心概念:注意力机制到底在做什么
1.1 从 RNN 和 CNN 的痛点说起
在 Transformer 之前,处理序列数据的主流方案是 RNN 和 LSTM。RNN 的核心思路是“按顺序处理”,也就是把序列看成一个时间步一个时间步地往后传递:
h_1 = f(x_1, h_0) h_2 = f(x_2, h_1) h_3 = f(x_3, h_2) ...这种方式有两个明显问题:
- 长距离依赖难捕捉:如果句子很长,最前面的信息要经过很多时间步才能传到后面,梯度容易消失,位置靠前的单词对最终结果的影响会越来越弱。
- 无法并行:第 t 步必须等第 t-1 步算完,计算效率很低,这让大模型训练变得非常吃力。
CNN 虽然可以并行,但它更擅长捕捉局部特征,需要通过堆叠很多层才能扩大感受野,对序列中“跨位置的全局依赖”建模能力有限。
注意力机制的出现就是为了解决“让模型直接关注序列中任意两个位置之间的关系”,不需要按序传递,也不需要一层层堆叠就能看到全局。
1.2 注意力机制的本质:加权求和
从直觉上讲,注意力机制就是一句话:对信息做加权求和,权重代表了当前关注点与不同信息之间的相关程度。
举个例子,翻译句子 “The animal didn't cross the street because it was too tired”,我们要确定 “it” 指代的是 “animal” 还是 “street”,就需要让模型在计算 “it” 这个位置的时候,同时“看到”句子中其他单词,并给 “animal” 一个更高的重要性权重。这个“重要性权重”就是注意力分数。
从数学上看,给定一个查询向量(Query)和一组键值对(Key-Value),注意力输出可以表示为:
Attention(Q, K, V) = softmax(Q * K^T / sqrt(d_k)) * V其中:
Q(Query):当前要查询的内容,代表“我想找什么”。K(Key):被查询内容的索引,代表“我有什么”。V(Value):实际被提取的信息,代表“我能提供什么”。d_k:Q和K的向量维度,缩放因子,防止点积数值过大。
Q、K、V 这三个词很抽象,下面我们用一张清晰的流程串起来。
1.3 从 QKV 到自注意力
自注意力(Self-Attention)是 Transformer 中最核心的注意力类型。所谓“自”,指的是 Q、K、V 都来自同一个输入序列。每个 token 既当“查询者”,也当“被查者”,这样计算出来的注意力就表达了序列内部每个位置与其他所有位置之间的关系。
具体流程如下:
- 输入序列
X(形状为[batch_size, seq_len, d_model])分别乘以三个权重矩阵W_Q、W_K、W_V,得到Q、K、V。 - 计算
Q与所有K的点积,得到原始注意力分数。 - 将分数除以
sqrt(d_k)进行缩放,防止数值过大进入 softmax 的饱和区。 - 对最后维度做 softmax,归一化成概率分布。
- 将归一化后的注意力权重与
V相乘并求和,得到加权后的输出。
我把这个过程画成文字示意:
输入 X │ ├──乘 W_Q──> Q ├──乘 W_K──> K └──乘 W_V──> V │ Q·K^T ──缩放──> softmax ──× V──> 输出理解了这一步,Transformer 的“心脏”你就已经掌握了。
2. 环境准备与项目结构
2.1 环境版本说明
本文示例使用 PyTorch 实现,核心只需要torch和numpy。版本方面,PyTorch 1.13 以上都可以正常运行,示例代码不依赖最新 API。
如果还没有安装,可以用以下命令创建虚拟环境并安装依赖:
conda create -n transformer_diy python=3.9 -y conda activate transformer_diy pip install torch numpy说明:如果你的机器有 NVIDIA GPU,可以按官网命令安装对应 CUDA 版本的 PyTorch。本文示例数据量极小,CPU 上即可运行。
2.2 项目文件结构
为了让代码便于阅读和复现,本次实战按以下结构组织文件:
transformer_diy/ ├── attention.py # 自注意力、多头注意力实现 ├── transformer_encoder.py # 位置编码、前馈网络、编码器层 ├── main.py # 最小可运行示例下面的内容中,每个完整代码块上方都会标注对应的文件路径。
3. 自注意力机制的从零实现
3.1 不带掩码的缩放点积注意力
我们先写最核心的缩放点积注意力函数。它接收 Q、K、V,输出注意力权重和加权后的结果。
# 文件路径:attention.py import torch import torch.nn as nn import torch.nn.functional as F def scaled_dot_product_attention(query, key, value, mask=None): """ 缩放点积注意力。 参数: query: [batch_size, ..., seq_len_q, d_k] key: [batch_size, ..., seq_len_k, d_k] value: [batch_size, ..., seq_len_v, d_v] mask: [batch_size, 1, seq_len_q, seq_len_k] 或广播兼容形状 返回: output: [batch_size, ..., seq_len_q, d_v] attention_weights: [batch_size, ..., seq_len_q, seq_len_k] """ d_k = query.size(-1) # 1. Q 和 K 做点积,得到原始注意力分数 scores = torch.matmul(query, key.transpose(-2, -1)) # 2. 缩放 scores = scores / torch.sqrt(torch.tensor(d_k, dtype=torch.float32)) # 3. 掩码:将被掩码位置的分数设为极小的负数 if mask is not None: scores = scores.masked_fill(mask == 0, float('-inf')) # 4. softmax 归一化 attention_weights = F.softmax(scores, dim=-1) # 5. 与 V 加权求和 output = torch.matmul(attention_weights, value) return output, attention_weights这里的mask参数很关键。在实际应用中,有两个典型场景需要掩码:
- Padding Mask:把填充位置对应的注意力分数设为
-inf,让模型忽略填充符。 - Causal Mask(因果掩码):在解码器里,禁止当前位置看到后面位置的信息。
为什么用-inf而不是0?因为softmax中-inf经过指数运算会变成0,这样对应位置的注意力权重就是 0,等价于完全忽略该位置。
3.2 为什么除以 sqrt(d_k)
这是一个高频面试点。在Q·K^T中,如果d_k很大,点积的数值会变得很大。假设Q和K的元素是独立随机变量且均值为 0、方差为 1,那么点积的均值是 0,方差是d_k。
方差大意味着什么?意味着某些维度的点积结果会异常大,进入 softmax 后,这些异常大的值对应位置的权重会接近 1,而其他位置的权重会趋近于 0。这样一来,梯度会变得非常小,模型难以学习。
除以sqrt(d_k)后,点积的方差被控制在 1 左右,softmax 的输入不至于过大,梯度更稳定。这也是论文原文《Attention Is All You Need》里的标准做法。
3.3 自注意力的完整封装
上面的函数是通用计算核心,我们需要把它封装成带可学习参数的模块。自注意力模块中,Q、K、V 来自同一个输入X,所以需要三组可学习的线性变换。
# 文件路径:attention.py(续上方代码) class SelfAttention(nn.Module): """ 单头自注意力模块。 """ def __init__(self, d_model, d_k=None, d_v=None, dropout=0.1): super().__init__() # 默认 Q、K、V 的维度都是 d_model self.d_k = d_k if d_k is not None else d_model self.d_v = d_v if d_v is not None else d_model self.w_q = nn.Linear(d_model, self.d_k, bias=False) self.w_k = nn.Linear(d_model, self.d_k, bias=False) self.w_v = nn.Linear(d_model, self.d_v, bias=False) self.dropout = nn.Dropout(dropout) def forward(self, x, mask=None): """ 参数: x: [batch_size, seq_len, d_model] mask: 广播兼容形状 返回: output: [batch_size, seq_len, d_v] attention_weights: [batch_size, seq_len, seq_len] """ query = self.w_q(x) key = self.w_k(x) value = self.w_v(x) output, attention_weights = scaled_dot_product_attention( query, key, value, mask ) output = self.dropout(output) return output, attention_weights这样我们就有了一个完整的单头自注意力模块。可以试一下前向传播:
import torch from attention import SelfAttention x = torch.randn(2, 10, 512) # batch_size=2, seq_len=10, d_model=512 attn = SelfAttention(d_model=512) out, weights = attn(x) print(out.shape) # torch.Size([2, 10, 512]) print(weights.shape) # torch.Size([2, 10, 10])输出形状符合预期:每个位置的输出融合了所有位置的信息,注意力矩阵是 10×10。
4. 多头注意力机制的实现
4.1 为什么要多头
单头注意力虽然能建模全局依赖,但它只有一个“注意力模式”。例如某个头关注语法依赖,另一个头关注语义相似,还有一个头关注位置邻近。多个头各司其职,模型表达能力更强。
用专业的话说:多头注意力把 Q、K、V 投影到多个低维子空间中,在每个子空间独立计算注意力,最后拼接起来再线性变换。这样模型可以在不同表示子空间上关注不同维度的信息。
需要注意,多头注意力并不是让每个头算出完整结果再平均,而是把d_model维度切分成h份,每个头处理d_k = d_model / h维度的子向量。
4.2 多头注意力完整实现
实现多头注意力的最优雅方式是利用 PyTorch 张量 reshape + transpose,把[batch_size, seq_len, d_model]变成[batch_size, seq_len, h, d_k],再转置成[batch_size, h, seq_len, d_k],这样就可以把“多头”当成批次维度并行计算。
# 文件路径:attention.py(续上方代码) class MultiHeadAttention(nn.Module): """ 多头注意力模块。 """ def __init__(self, d_model, num_heads, dropout=0.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, bias=False) self.w_k = nn.Linear(d_model, d_model, bias=False) self.w_v = nn.Linear(d_model, d_model, bias=False) self.out_proj = nn.Linear(d_model, d_model, bias=False) self.dropout = nn.Dropout(dropout) def split_heads(self, x): """ 把最后一个维度拆成 num_heads 份并转置。 输入: [batch_size, seq_len, d_model] 输出: [batch_size, num_heads, seq_len, d_k] """ batch_size, seq_len, _ = x.size() x = x.view(batch_size, seq_len, self.num_heads, self.d_k) x = x.transpose(1, 2) # [batch_size, num_heads, seq_len, d_k] return x def combine_heads(self, x): """ 将多头结果拼接回完整维度。 输入: [batch_size, num_heads, seq_len, d_k] 输出: [batch_size, seq_len, d_model] """ batch_size, _, seq_len, _ = x.size() x = x.transpose(1, 2).contiguous() x = x.view(batch_size, seq_len, self.d_model) return x def forward(self, x, mask=None): # 1. 生成 Q、K、V 并拆成多头 query = self.split_heads(self.w_q(x)) # [B, H, T, d_k] key = self.split_heads(self.w_k(x)) value = self.split_heads(self.w_v(x)) # 2. 调用缩放点积注意力 # mask 需要扩展成能和 [B, H, T, T] 广播的形状 if mask is not None: # 假设传入 mask 是 [B, T, T],需要变成 [B, 1, T, T] mask = mask.unsqueeze(1) output, attention_weights = scaled_dot_product_attention( query, key, value, mask ) # 3. 合并多头结果,做输出投影 output = self.combine_heads(output) output = self.out_proj(output) output = self.dropout(output) return output, attention_weights这里的核心是split_heads和combine_heads。很多初学者第一次接触多头注意力时会写循环去遍历每个头,虽然也能实现,但效率很低。用 reshape 的方式可以把所有头的计算一次性交给矩阵乘法完成,这也体现出了 PyTorch 张量操作的优势。
4.3 自测与对比
写完后可以做个简单测试:
mha = MultiHeadAttention(d_model=512, num_heads=8) out, weights = mha(torch.randn(2, 10, 512)) print(out.shape) # torch.Size([2, 10, 512]) print(weights.shape) # torch.Size([2, 8, 10, 10])注意这里weights的形状比单头多了一维,第 2 维是num_heads,说明每个头都有独立的注意力矩阵。
为了验证实现正确性,我们可以和 PyTorch 官方的nn.MultiheadAttention做一次输出维度对比,需要注意的是官方模块默认batch_first=False,且输出顺序不同,这里仅验证形状:
official = nn.MultiheadAttention(embed_dim=512, num_heads=8, batch_first=True) official_out, official_weights = official( torch.randn(2, 10, 512), torch.randn(2, 10, 512), torch.randn(2, 10, 512), ) print(official_out.shape) # torch.Size([2, 10, 512])从形状上看,我们的实现与官方一致,说明整体流程没有偏差。
5. 从零搭建 Transformer 编码器
有了多头注意力,我们距离完整的 Transformer 只差三个组件:位置编码、前馈神经网络、残差连接与层归一化。
5.1 位置编码:给序列引入顺序信息
注意力机制本身对 token 的位置不敏感。如果你把两个 token 交换顺序,计算出的 Q、K、V 点积结果是一样的。也就是说,纯注意力模型根本不知道“谁在前、谁在后”,这显然不行。
Transformer 原文使用了正弦位置编码(Sinusoidal Positional Encoding),公式如下:
PE(pos, 2i) = sin(pos / 10000^(2i / d_model)) PE(pos, 2i+1) = cos(pos / 10000^(2i / d_model))其中pos是位置下标,i是维度下标。这个设计的巧妙之处在于:
- 不同位置的编码向量不同;
- 编码值在
[-1, 1]之间,利于训练稳定; - 可以用正弦/余弦的加法性质表达相对位置关系。
代码实现如下:
# 文件路径:transformer_encoder.py import math import torch import torch.nn as nn class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=5000, dropout=0.1): super().__init__() self.dropout = nn.Dropout(p=dropout) # 生成 [max_len, d_model] 的位置编码矩阵 pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len, dtype=torch.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) # 注册为 buffer,不参与梯度更新,但会随模型保存 pe = pe.unsqueeze(0) # [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)使用位置编码时,直接将它加在词向量/输入向量的上方。这个加法本身就是一个“注入位置信息”的操作,位置编码的数值相对于输入嵌入来说不能太大,否则会破坏原有的语义信息。
5.2 前馈神经网络
Transformer 编码器中的注意力输出会进入一个前馈网络(Feed-Forward Network, FFN)。它包含两个线性层和一个 ReLU 激活函数,常被称为 Position-wise FFN,因为它对每个位置独立操作:
FFN(x) = max(0, x·W1 + b1)·W2 + b2中间隐藏层维度一般扩大 4 倍,例如d_model=512时,d_ff=2048。
# 文件路径:transformer_encoder.py(续上方代码) class FeedForward(nn.Module): def __init__(self, d_model, d_ff=2048, dropout=0.1): super().__init__() self.linear1 = nn.Linear(d_model, d_ff) self.relu = nn.ReLU() self.linear2 = nn.Linear(d_ff, d_model) self.dropout = nn.Dropout(dropout) def forward(self, x): x = self.linear1(x) x = self.relu(x) x = self.dropout(x) x = self.linear2(x) return x为什么每个注意力层后面要接 FFN?一种解释是:注意力机制本质上是“加权求和”,它的变换是线性的(准确说是凸组合)。如果没有非线性激活层,多个注意力层堆叠起来仍然近似线性变换,模型的表达能力会受到很大限制。FFN 中的 ReLU 引入了非线性,让模型能学到更复杂的特征。
5.3 残差连接与层归一化
Transformer 编码器中的每个子层都包含“残差连接 + 层归一化”。公式如下:
x = LayerNorm(x + Sublayer(x))残差连接让梯度可以跨层直接传播,训练深层的 Transformer 时更加稳定。层归一化(LayerNorm)则作用于每个 token 的所有特征维度,把数据分布拉回均值为 0、方差为 1 的状态,加速收敛。
5.4 完整 Transformer 编码器层
把以上组件组合起来,就是一个 Transformer Encoder Block:
# 文件路径:transformer_encoder.py(续上方代码) from attention import MultiHeadAttention class TransformerEncoderLayer(nn.Module): """ 单个 Transformer 编码器层。 """ def __init__(self, d_model, num_heads, d_ff=2048, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, num_heads, dropout) 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, mask=None): # 1. 子层1:多头自注意力 + 残差 attn_output, _ = self.self_attn(x, mask) x = self.norm1(x + self.dropout1(attn_output)) # 2. 子层2:前馈网络 + 残差 ffn_output = self.ffn(x) x = self.norm2(x + self.dropout2(ffn_output)) return x注意,这里使用的是 Post-Norm 结构,也就是“残差之后再做 LayerNorm”,这是原始 Transformer 论文的做法。现在很多新模型用的是 Pre-Norm,即“先归一化再进子层”,训练更稳定。这个差异在后来被大量实验验证过,不过作为入门,我们先以原版为准。
5.5 组合成完整的编码器并运行
最后,把位置编码和多个编码器层堆叠成完整编码器:
# 文件路径:transformer_encoder.py(续上方代码) class TransformerEncoder(nn.Module): def __init__(self, vocab_size, d_model, num_heads, num_layers, d_ff=2048, dropout=0.1, max_len=5000): super().__init__() self.embedding = nn.Embedding(vocab_size, d_model) self.pos_encoding = PositionalEncoding(d_model, max_len, dropout) self.layers = nn.ModuleList([ TransformerEncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers) ]) self.d_model = d_model def forward(self, token_ids, mask=None): """ 参数: token_ids: [batch_size, seq_len] mask: 扩展后可用于注意力的掩码 返回: x: [batch_size, seq_len, d_model] """ x = self.embedding(token_ids) x = x * math.sqrt(self.d_model) # 论文中的缩放操作 x = self.pos_encoding(x) for layer in self.layers: x = layer(x, mask) return x写一个最小 demo:
# 文件路径:main.py import torch from transformer_encoder import TransformerEncoder # 构造一个简单的词表 vocab_size = 100 d_model = 64 num_heads = 4 num_layers = 2 model = TransformerEncoder( vocab_size=vocab_size, d_model=d_model, num_heads=num_heads, num_layers=num_layers, d_ff=128, dropout=0.1, ) # 模拟一个 batch:batch_size=2,seq_len=8 token_ids = torch.randint(0, vocab_size, (2, 8)) output = model(token_ids) print("输入形状:", token_ids.shape) # [2, 8] print("输出形状:", output.shape) # [2, 8, 64]运行后输出:
输入形状: torch.Size([2, 8]) 输出形状: torch.Size([2, 8, 64])到这里,我们其实已经实现了一个可用的 Transformer 编码器。如果你加上一个线性分类头,就可以用它做文本分类;如果接一个 LM Head,就能训练一个小型的语言模型。
6. 注意力权重的可视化
手写注意力机制还有一个很大的好处:可以非常方便地拿到每个头的注意力权重。可视化注意力能帮我们直观理解“模型在看什么”。
下面给出一段简单的可视化代码,随机选一个头画出注意力矩阵热力图。假设你已经跑通了main.py,我们可以把编码器第一层的注意力权重导出来。
# 文件路径:main.py(续上方代码) import matplotlib.pyplot as plt def extract_attention(model, token_ids, layer_idx=0): """ 提取指定层的注意力权重。 """ model.eval() with torch.no_grad(): x = model.embedding(token_ids) x = x * (model.d_model ** 0.5) x = model.pos_encoding(x) attention_weights = None for idx, layer in enumerate(model.layers): if idx == layer_idx: _, attention_weights = layer.self_attn(x) return attention_weights x = layer(x) return attention_weights token_ids = torch.randint(0, vocab_size, (1, 8)) attn_weights = extract_attention(model, token_ids, layer_idx=0) # 形状 [batch_size, num_heads, seq_len, seq_len] # 展示第 0 个样本、第 0 个头 head_idx = 0 plt.figure(figsize=(6, 6)) plt.imshow(attn_weights[0, head_idx].cpu().numpy(), cmap='Blues') plt.colorbar() plt.title(f"Layer 0, Head {head_idx} Attention") plt.xlabel("Key Position") plt.ylabel("Query Position") plt.show()可视化后你会看到,不同的头关注模式差异很大。有的头呈明显的对角分布(当前位置主要关注邻近位置),有的头分散到整句的少数几个关键位置。这就是多头机制带来“多种关注模式”的直观体现。
7. 常见问题与排查思路
手写 Transformer 过程中,最常踩的坑大部分集中在维度不匹配和 mask 上。下面整理了一份排查清单。
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
mat1 and mat2 shapes cannot be multiplied | 线性层输入维度与d_model不一致 | 检查nn.Embedding输出维度、输入 token 的最后一个维度是否为d_model |
The size of tensor a must match the size of tensor b | 多头注意力中 head 拆分后维度不对,或者位置编码和输入 seq_len 不匹配 | 检查split_heads和combine_heads中的 reshape/transpose 顺序 |
| 注意力权重全变成 0 或全变成 1 | mask 使用不当,或-inf的位置设置错误 | 检查 mask 的每一维是否和[B, H, T, T]匹配;padding 位置应为False/0 |
| 训练 loss 下降极慢或直接 NaN | 学习率过大、未做梯度裁剪、LayerNorm 缺失 | 调小学习率,加 warmup,检查是否所有输出都经过了 LayerNorm |
多头结果和官方nn.MultiheadAttention不一致 | QKV 投影方式、是否 bias、输出投影顺序不同 | 先对比形状,再关掉 dropout、初始化相同权重后对比数值 |
排查维度问题时,最实用的方法是print(shape)大法:
print("x:", x.shape) print("query:", query.shape) print("key:", key.shape) print("score:", scores.shape) print("attention_weights:", attention_weights.shape) print("output:", output.shape)不要觉得这很初级,实际调模型时,这类逐步打印往往比看报错信息更高效,尤其是面对多层结构时。
还有一个常见问题是mask的维度。在scaled_dot_product_attention中,scores的形状是[B, H, T, T],因此mask需要能广播成这个形状。很多人在做多头时忘记给 mask 扩展维度,导致什么报错都出现了还是找不到原因。建议统一写成:
if mask is not None: mask = mask.unsqueeze(1) # [B, T, T] -> [B, 1, T, T]或者传入前就确保 mask 是[B, 1, T, T]。
8. 最佳实践与工程建议
写完一个可运行的注意力机制只是第一步。如果要把它真正用到项目里,下面这些建议值得提前了解。
8.1 模块设计建议
代码结构上,把scaled_dot_product_attention单独抽出来是一个好习惯。这个函数是纯计算逻辑,与模型参数无关,既能被单头注意力调用,也能被多头注意力调用,单元测试起来非常方便。而且后面如果你想尝试 flash-attention 等高效实现,只需要替换这个函数,不需要改动整个模块。
8.2 训练稳定性
Transformer 对训练超参数比较敏感。实际项目里建议注意以下几点:
- 学习率调度:先用 warmup,把学习率从 0 缓慢升到峰值,再按余弦退火降低。Transformer 原论文对 512 维模型使用 4000 步 warmup。
- Dropout 设置:小模型建议
0.1左右;数据量很大时 dropout 可以适当调低。 - 梯度裁剪:把梯度的范数裁剪到
1.0或者5.0,可以有效避免训练中期出现的 loss 突增。 - Pre-Norm 更稳:如果训练深层编码器经常不收敛,可以考虑把 LayerNorm 移到子层之前,也就是 Pre-LN 结构。
8.3 性能优化
自注意力的复杂度是O(n^2),n是序列长度。序列一长,显存和耗时都会迅速上升。如果项目里需要处理长文本,可以这样优化:
- 使用 PyTorch 2.0 的
F.scaled_dot_product_attention,它会自动选择最高效的内存融合实现。 - 尝试稀疏注意力或窗口注意力,只让局部 token 互相注意力。
- 在训练和推理时使用
torch.compile加速模型执行。
8.4 对比官方实现与扩展方向
手写实现最大的价值是“消除黑盒感”。等你能把代码跑通、能画出注意力热力图之后,我推荐你再打开 PyTorch 官方的nn.TransformerEncoderLayer源码做一次对比,看看官方在细节上做了哪些差异处理,比如:
- 官方默认
batch_first=False,第一维是序列长度。 - 官方支持
src_key_padding_mask和attention_mask两类掩码。 - 官方还提供了解码器
nn.TransformerDecoderLayer,支持交叉注意力(Cross-Attention),即 Q 来自解码器输入、K/V 来自编码器输出。
交叉注意力是多头注意力的重要变体,在机器翻译、图像描述、语音识别等场景中非常关键。看懂了本文的多头注意力,交叉注意力只需要把SelfAttention的 Q 改成外部输入、K/V 保持从编码器输出生成即可,原理完全一致。
8.5 不要盲目手写
最后说句实在的:如果你只是在业务项目里用 Transformer,完全不需要手写,直接使用 PyTorch 官方封装更高效、更稳定。手写价值体现在下面几个地方:
- 面试前理解原理,避免用“调包侠”人设被问倒。
- 做研究/发论文时,需要对注意力机制做自定义改造。
- 学习阶段,用一个小项目彻底打通原理到实现的路径。
把本文代码跑通后,你的下一步可以尝试:
- 在文本分类任务上训练一个完整的 Transformer 分类器;
- 给代码加上因果掩码,实现一个自回归语言模型;
- 用自己实现的编码器替换
nn.TransformerEncoder做对比实验; - 尝试加入相对位置编码(如 RoPE、ALiBi),感受位置编码的演化。
手写一次注意力机制,远比调包十次学到的多。希望这篇文章能帮你真正跨过 Transformer 入门的门槛。如果过程中遇到报错,可以照着上面的排查表格一步步找原因,也欢迎在评论区交流你踩到的坑。