亲手实现Transformer核心:注意力机制从原理到PyTorch代码
2026/9/1 7:45:59 网站建设 项目流程

为什么我把"看懂的"Transformer又亲手写了一遍

Transformer 这个词,在深度学习圈子里几乎被说烂了。从 NLP 到 CV,从推荐系统到多模态,随处可见它的影子。但很多时候,我们看完了各种图解、听完了各种原理讲解,心里却始终有一个模糊地带:注意力机制内部到底是怎么算的?为什么是除以根号 d_k?Q、K、V 到底是哪来的?

如果你只是停留在调用nn.TransformerEncoderLayer或者model = BertModel.from_pretrained(...)这种层面,那么注意力机制对你来说本质上还是一个"黑盒"。

更好的方式是哪一种?亲手用 PyTorch 把注意力机制从零实现一遍。这个过程不需要你从底层写 CUDA 核函数,不需要你手写反向传播,只需要把矩阵乘法、softmax、mask 这些基本操作组合起来,就能揭开 Transformer 最核心的一层窗户纸。

这篇文章记录我从"会调用"到"能实现"的过程,重点不是完整复刻一份官方源码,而是带你走过一条务实的实现路径。我会先讲注意力机制到底在解决什么问题,再逐个实现缩放点积注意力(Scaled Dot-Product Attention)、多头注意力、因果掩码机制,最后聊聊实现过程中最大的几个认知转折。看完之后,你可以对照自己的项目,判断哪些场景真的需要"手写"注意力机制,哪些场景用现成封装就足够了。

1. 注意力机制真正解决的问题是什么

很多人在理解 Transformer 时都犯了一个误区:把注意力机制当成一种"高级特征提取器"来背公式。其实,一个机制的价值不在公式漂亮,而在于它解决了什么现实问题。

在注意力机制出现之前,处理序列数据的主流方案是 RNN(循环神经网络)及其变体 LSTM、GRU。RNN 的思路是"逐个处理,状态传递":每读取一个 token,就更新一次隐藏状态,下一个 token 的处理依赖上一个状态。这种串行机制导致两个严重问题:

  1. 长距离依赖难以捕捉。当一个句子很长时,前面的信息经过多步传递后,要么被遗忘,要么被后面的信息淹没。LSTM 通过门控机制缓解了这个问题,但没有根本解决。
  2. 无法并行计算。由于每个时间步都依赖前一个时间步的状态,RNN 的天然结构就是串行的,这在 GPU 时代几乎是致命的性能瓶颈。

注意力机制的思路完全不同:直接计算任意两个位置之间的关联权重,让每个 token 都能"看到"整个序列中的所有其他 token。它不依赖中间状态的逐步传递,一步到位建立全局依赖关系。

通俗地理解,注意力机制在一个句子里做的事情,就是为每个词分配一组权重,表示"我在理解这个词的时候,应该重点参考哪些词"。比如"苹果"这个词在"我爱吃苹果"和"苹果发布了新手机"中,模型关注的重点应该完全不同。注意力机制就是让模型"按需关注"。

也正是因为这种"全局直接连接"的特性,Transformer 可以把序列中的所有 token 一次性并行输入,训练速度显著提升。这为后来 GPT、BERT 等大规模预训练模型的出现扫清了算力障碍。

2. Q、K、V 到底在做什么

第一次接触 Q、K、V 时,最常见的反应是:这三个矩阵是怎么来的?为什么要分成三份?

回答这个问题的关键,是理解查询、键、值的类比。

  • Query(查询):代表"我要找什么"。一个 token 作为查询者时,它想知道自己在整个序列中应该关注哪些信息。
  • Key(键):代表"我有什么可供匹配"。每个 token 都作为被查询的目标,把自己的"特征标签"暴露出来。
  • Value(值):代表"真正的内容信息"。一旦 Query 和 Key 匹配上了,对应 Value 的内容才会被加权提取出来。

可以把它类比成文件检索系统:Query 是你输入的关键词,Key 是文件系统里的索引标签,Value 是文件正文。检索时先匹配索引,再提取正文。

在 Transformer 中,这三个矩阵不是独立训练的,而是从同一个输入通过三个不同的线性变换得到:

Q = x @ W_q K = x @ W_k V = x @ W_v

也就是每一个 token 的输入向量,分别乘以三个权重矩阵,映射到三个不同的语义空间。权重矩阵W_qW_kW_v就是模型学习到的"如何提问、如何匹配、如何提取"的参数。

这里有一个新手经常踩的坑:Q、K、V 在自注意力(Self-Attention)里都来自同一个输入 x,在交叉注意力(Cross-Attention)里则来自不同的输入。比如在 Encoder-Decoder 架构中,Decoder 的 Query 来自 decoder 侧,而 Key 和 Value 来自 encoder 侧的输出。这一点理解错了,后面看模型结构图就会一直觉得别扭。

3. 手写缩放点积注意力:从公式到 PyTorch 实现

实现注意力机制之前,先把目标公式写清楚。缩放点积注意力的计算过程可以拆成四个步骤:

  1. 计算 Q 和 K 的点积,得到注意力分数矩阵(Attention Scores)。
  2. 对分数进行缩放,除以sqrt(d_k),其中d_k是 Key 的维度。
  3. 对分数进行 softmax 归一化,得到注意力权重矩阵。
  4. 用权重矩阵加权 V,得到最终输出。

公式表达如下:

Attention(Q, K, V) = softmax(Q K^T / sqrt(d_k)) V

这里有个值得单独说的细节:为什么要除以sqrt(d_k)

这一步很多教程直接当作"工程技巧"略过了,但它其实是保证训练稳定的关键。当d_k较大时,Q 和 K 的点积结果方差也会增大,导致 softmax 输入的数值落在梯度很小的饱和区,反向传播时梯度几乎消失。除以sqrt(d_k)相当于把点积结果的方差重新拉回到 1 附近,让 softmax 的梯度保持在一个合理的区间。

下面是最基础版本的 PyTorch 实现,不包含任何封装,用于演示核心计算流程:

import torch import torch.nn.functional as F def scaled_dot_product_attention(query, key, value, mask=None): """ query: (batch, seq_len, d_k) key: (batch, seq_len, d_k) value: (batch, seq_len, d_v) """ d_k = query.size(-1) # 1. 计算 Q 和 K 的点积 # scores shape: (batch, seq_len, seq_len) scores = torch.matmul(query, key.transpose(-2, -1)) # 2. 缩放 scores = scores / torch.sqrt(torch.tensor(d_k, dtype=torch.float32)) # 3. 可选:应用 mask(后面会详细讲) 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

需要说明的是,上述代码只是演示核心逻辑。在实际的 Transformer 实现中,Q、K、V 的形状可能会包含 heads 维度,变成(batch, heads, seq_len, d_k),那是在多头注意力阶段处理的问题。

验证一下这个基本实现:

# 构造一个小例子 batch_size = 2 seq_len = 4 d_k = 8 d_v = 8 query = torch.randn(batch_size, seq_len, d_k) key = torch.randn(batch_size, seq_len, d_k) value = torch.randn(batch_size, seq_len, d_v) output, attention_weights = scaled_dot_product_attention(query, key, value) print("输出维度:", output.shape) # (2, 4, 8) print("注意力权重维度:", attention_weights.shape) # (2, 4, 4) print("每行注意力权重之和:", attention_weights.sum(dim=-1))

预期输出是每行注意力权重之和接近 1,因为每一行都经过了 softmax 归一化。

这一版实现已经能让你直观地理解"注意力机制到底是什么":它不过就是计算一个加权平均,只不过权重不是外部指定的,而是通过查询和键的匹配程度计算出来的。它的所有魔力,都在于权重由模型自己学习如何衡量"匹配"

4. 把多头注意力补完:为什么一个头不够用

单个注意力头有一个明显的局限:一次只能学习一种距离度量方式。在实际语言中,"关注什么"是一个多维度的概念。一个词可能因为语法关系关注主语,因为指代关系关注前文的名词,因为修辞关系关注形容词。单头注意力只能在这几种关系里折中。

多头注意力(Multi-Head Attention)的解决思路是:并行运行多套 Q、K、V 变换,每一套学习不同的注意力模式,最后把结果拼接起来。通俗地说,就是用多组不同的"透镜"去观察同一个序列,每一组看到的东西侧重不同。

实现多头注意力时,最常见的问题不是公式本身,而是维度变换。在实际代码中,为了效率,我们不会真的拆成多个独立的线性层,而是通过一次大的矩阵乘法,再通过viewtranspose把数据拆成多个头。

下面是一个完整的多头注意力实现:

import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads, dropout=0.1): super().__init__() assert d_model % n_heads == 0, "d_model 必须能被 n_heads 整除" self.d_model = d_model self.n_heads = n_heads self.d_k = d_model // n_heads # 通过一次线性变换同时得到 Q、K、V,减少参数和计算 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_out = nn.Linear(d_model, d_model) self.dropout = nn.Dropout(dropout) def forward(self, query, key, value, mask=None): batch_size = query.size(0) # 1. 线性变换 Q = self.w_q(query) # (batch, seq_len, d_model) K = self.w_k(key) V = self.w_v(value) # 2. 拆分成多头 # (batch, seq_len, n_heads, d_k) -> (batch, n_heads, seq_len, d_k) Q = Q.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2) K = K.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2) V = V.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2) # 3. 计算注意力分数 scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_k, dtype=torch.float32, device=Q.device)) # 4. 应用掩码 if mask is not None: scores = scores.masked_fill(mask.unsqueeze(1) == 0, float('-inf')) # 5. softmax + 加权 attn_weights = F.softmax(scores, dim=-1) attn_weights = self.dropout(attn_weights) output = torch.matmul(attn_weights, V) # 6. 拼接多头结果并线性变换 # (batch, n_heads, seq_len, d_k) -> (batch, seq_len, n_heads, d_k) -> (batch, seq_len, d_model) output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) output = self.w_out(output) return output, attn_weights

关于这段代码,有三个地方值得单独提一下。

第一,viewtranspose的组合是拆多头时的标准写法。先view变成四维,再用transpose把头的维度换到第二位。这里要特别注意:执行完transpose后,张量的内存布局是不连续的,所以后面拼接多头结果时,必须先调用contiguous()view,否则会报错。

第二,mask.unsqueeze(1)是为了把原始形状(batch, seq_len, seq_len)扩展成(batch, 1, seq_len, seq_len),这样才能和(batch, n_heads, seq_len, seq_len)的 scores 广播对齐。

第三,输出投影w_out很容易被遗漏。多头注意力最后必须把所有头的拼接结果再经过一个线性层,否则各头之间的信息无法融合。

5. 因果掩码:让模型"看不见未来"

如果你只实现到上面这一版多头注意力,那么你实现的只是"普通的自注意力"。但如果你想理解 GPT 这类自回归语言模型,就绕不开因果掩码(Causal Mask)

自回归语言模型的训练目标很简单:给定前面的词,预测下一个词。因此,当模型处理第 i 个 token 时,它只能看到前 i 个 token,绝对不能看到第 i 个 token 之后的内容。否则,模型就相当于"作弊"了——预测时答案已经摆在眼前。

实现因果掩码的思路是构造一个上三角掩码矩阵。矩阵中第 i 行、第 j 列的值表示:当计算第 i 个位置时,是否允许看到第 j 个位置。如果是允许,就保留原来的注意力分数;如果不允许,就把分数置为负无穷,这样 softmax 之后权重几乎为 0。

生成上三角掩码的经典代码:

import torch def create_causal_mask(seq_len): """ 生成因果掩码矩阵,用于自回归模型 mask[i][j] == 1 表示位置 i 可以关注位置 j """ mask = torch.tril(torch.ones(seq_len, seq_len)).bool() return mask mask = create_causal_mask(5) print(mask)

输出结果:

tensor([[ True, False, False, False, False], [ True, True, False, False, False], [ True, True, True, False, False], [ True, True, True, True, False], [ True, True, True, True, True]])

可以看到,第一行只有第一个位置为 True,第二行只有前两个位置为 True,以此类推。这就是"每个位置只能看到自己及之前位置"的量化表达。

把因果掩码接入之前的MultiHeadAttention中:

# 假设 seq_len = 10, batch_size = 2, d_model = 512, n_heads = 8 mha = MultiHeadAttention(d_model=512, n_heads=8) query = torch.randn(2, 10, 512) key = torch.randn(2, 10, 512) value = torch.randn(2, 10, 512) causal_mask = create_causal_mask(10) output, attn_weights = mha(query, key, value, mask=causal_mask) print("输出维度:", output.shape) # (2, 10, 512) print("注意力权重维度:", attn_weights.shape) # (2, 8, 10, 10)

在实现时,需要注意掩码的一致性检查。确认一下:经过 softmax 之后,被掩码的位置权重是否为 0?看一眼前几个位置的注意力权重:

# 查看第一个样本、第一个头、第 3 个位置(下标 2)的注意力权重 print(attn_weights[0, 0, 2])

预期结果中,最后一个值(对应下标 3 的位置,也就是"未来")应该非常接近 0。因为 3 > 2,该位置被掩码了。

因果掩码实现看起来不难,但它的含义非常深刻:它决定了 Transformer 的序列建模边界。带因果掩码的 Transformer 不能利用未来信息,这既是它的限制,也是它能做生成模型的原因。

6. 实现过程中最容易踩中的四个认知误区

把注意力机制真正实现了一遍之后,有几个"看起来简单、实际容易错"的地方给我留下了深刻印象,这里单独列出来分析,帮你少走弯路。

误区一:只关注矩阵乘法,没有注意维度匹配。在实现多头注意力时,最容易犯的错误是 Q、K、V 在拆多头之后的维度对不上。Q 和 K 在最后两维(seq_len 和 d_k)上要能进行矩阵乘法,V 的最后一维(d_v)可以不同于 d_k,但实际实现中通常保持一致。如果你的代码报mat1 and mat2 shapes cannot be multiplied错误,首先检查transpose之后各张量的最后一维是否匹配。

误区二:把 drop mask 直接加在原始 mask 上。有一种常见写法是:

scores = scores.masked_fill(mask == 0, -1e9)

-1e9代替float('-inf')在大多数情况下也能工作,因为 softmax 之后接近 0,但在数值极端情况下可能会导致梯度问题。更稳妥的写法是用float('-inf')。同时要注意,masked_fill的 mask 类型必须是 bool 张量。

误区三:没有对 attention 权重做 dropout。很多早期实现会把 dropout 加在 output 之后,却忘了对 attention 权重本身做 dropout。在原始 Transformer 论文的实现中,dropout 是加在注意力权重上的,这样做的目的是让模型不过度依赖某几个位置的组合,有正则化效果。实际可以自己对比一下,加在权重上效果更稳定。

误区四:不能区分自注意力和交叉注意力的 mask 差异。在 Encoder 的自注意力中通常不需要 mask(或者只需要 padding mask);在 Decoder 的自注意力中要同时使用因果 mask 和 padding mask;在 Encoder-Decoder 交叉注意力中,Decoder 侧 Query 可以关注 Encoder 输出的所有位置,不需要因果 mask,但要考虑 padding mask。

7. 一个通俗的应用示例:机器翻译场景

把实现的注意力机制放到一个简单的应用场景中,能更好地理解每个模块的作用。这里用一个经典的「英译中」的 Encoder-Decoder 场景来展示,不写出完整训练代码,只描述数据流动过程。

  • Encoder 侧:输入英文句子 "I love deep learning",经过词嵌入和位置编码后,进入多层自注意力。在 Encoder 自注意力中,每个英文单词都能看到整个句子的所有单词,因此可以建立"love"和"deep learning"之间的语义关联。
  • Decoder 侧:输入中文目标句子的前缀,比如"我"、然后"我爱"、然后"我爱深度",每走一步,Decoder 都用因果掩码保证自己看不见未来。这一步的 Q 来自 Decoder 当前已经生成的所有词,K、V 则来自 Encoder 的输出,实现"根据英文原句的内容,生成中文译文"。

把这段流程用代码来表达,就是:

# Encoder 输出是英文句子的语义表示 encoder_output, _ = multihead_attn_enc(query=enc_input, key=enc_input, value=enc_input) # Decoder 交叉注意力:Q 来自 decoder,K、V 来自 encoder decoder_output, _ = multihead_attn_dec_cross( query=dec_input, key=encoder_output, value=encoder_output )

可以看到,同样的MultiHeadAttention模块,只是输入来源不同,就实现了两种完全不同的功能。这也解释了为什么 Transformer 的抽象能力如此强:它以极少的"原语"(attention + FFN + LayerNorm),通过不同的组合方式,构建出了不同用途的模型架构

8. 注意力机制家族的坐标梳理

在热词中出现了大量注意力变体:自注意力、交叉注意力、多头注意力、通道注意力、SE注意力、CBAM、ECA、CA注意力等。它们并不完全属于同一个应用层级,如果放在一张坐标图上可以更清晰地理解。

第一类:按输入来源划分。自注意力(Self-Attention)是 Q、K、V 均来自同一输入;交叉注意力(Cross-Attention)是 Q 来自一个输入,K、V 来自另一个输入。这是架构层面的区分。

第二类:按"分头"划分。单头注意力一次只学习一种关联模式;多头注意力(Multi-Head Attention)并行学习多种关联模式,然后拼接融合。这是特征表达能力层面的区分。

第三类:按"作用维度"划分。上述注意力机制主要作用在 token 与 token 之间(序列维度),而 CV 领域常说的 SE 注意力、CBAM、ECA 等,主要作用在通道维度空间维度上。SE 注意力先全局平均池化,再通过两个全连接层学习每个通道的重要性权重,本质上是对特征图的"通道重标定"。ECA 是 SE 的轻量化改进,把两个全连接层替换为 1D 卷积,减少了参数。CBAM 则同时串行通过通道注意力模块和空间注意力模块,让模型既关注"什么样的特征重要",也关注"哪个位置的重要"。

CATEGORICAL. 对于做 NLP 或大模型的开发者来说,首先要掌握的是前两类;对于做图像分类、目标检测的开发者来说,第三类才是重点。这一点做方向选择时值得明确。

9. 手写到底值不值:实践路线建议

回到一个很实际的问题:官方早就提供了封装好的nn.MultiheadAttention,那么自己手写一遍到底有没有价值?

我的判断是:如果你是第一次接触 Transformer,手写一遍的价值非常大;如果你是做应用开发的,不必在业务代码里重复造轮子,但至少要把手写过程作为学习路径走一遍

推荐的实践路线如下:

第一步,按本文的路径,实现一个不带任何封装的scaled_dot_product_attention,确认注意力机制的基本结构。这个函数不超过 20 行,目的是理解机制。

第二步,实现MultiHeadAttention模块,重点理解维度变换和transposecontiguous()view的关系。这一步能让你真正理解"多头的本质是分组"。

第三步,把实现的MultiHeadAttention放进一个简化版的 Transformer Encoder 中,在小规模数据上做一个文本分类任务(比如句子情感分类),验证前向传播和反向传播是否正常。确定能跑通后,再对照源码逐行看差异,这是一个渐进的过程。

第四步,如果你的方向是 NLP/大模型,建议继续实现因果掩码、位置编码、前馈网络和残差连接,构成一个完整的 Decoder Block,然后跑一个玩具规模的"字符级语言生成"任务,也就是给定几个字符,预测下一个字符。

实现层面有一个很明确的判断标准:当你能从零写出一个能跑通正向传播和反向传播的 Transformer Block 时,你对"注意力机制"的理解就已经超过了大量只会调包的人

10. 工程环境与常见问题排查

环境方面,建议使用 Python 3.8+ 和 PyTorch 2.0+。不要求 GPU,但如果有 GPU,训练效率会更高。完整的环境依赖如下:

pip install torch numpy

也可以直接用一个脚本来验证自己的环境:

import torch print("PyTorch 版本:", torch.__version__) print("CUDA 是否可用:", torch.cuda.is_available())

下面是一份常见问题排查清单。

问题现象可能原因排查方式解决方案
前向传播报 shape 不匹配错误拆多头后 Q、K、V 维度不匹配打印每一步张量的 shape,重点检查 transpose 后各维度的顺序确保d_model = n_heads * d_k,并在view之前确认张量的形状
view操作报错transpose 后张量内存不连续检查报错信息中是否有 "not contiguous"view之前调用.contiguous()
训练损失不下降没有做缩放或缩放值不对检查 scores 是否除以了sqrt(d_k)确保除以的根号里是d_k,而不是d_model
输出结果全为同一个向量注意力权重过度平均检查是否为 transformer 初始化问题,或学习率设置过大调整学习率,或在初始化时适当调整参数范围
模型推理时生成重复内容因果掩码可能未生效打印注意力权重矩阵,观察是否有未来位置的权重大于 0确认 mask 传入的是三角矩阵,且在下采样到(batch, 1, seq_len, seq_len)时没有出错

11. 工程化与性能优化的进阶建议

如果你已经能完整手写并运行注意力机制,下一步需要关注工程化层面的问题。这些经验在真实项目中非常有用。

第一,能调用现成算子的地方不要手写。PyTorch 从 2.0 开始提供了F.scaled_dot_product_attention,不仅支持缩放点积注意力、因果掩码等多种配置,还会自动选择高效的 Flash Attention 融合算子。在 GPU 上,Flash Attention 能显著节省显存并提升计算速度。学习手写是为了理解原理,生产环境则应该优先使用经过优化、语义清晰的原生算子

第二,注意力计算的复杂度是序列长度的平方。如果序列长度从 512 增加到 1024,计算量会扩大 4 倍。当你要处理超长序列时,就需要考虑稀疏注意力、线性注意力等近似方法,或者考虑长文本分块等工程手段。理解了这一点,你就能明白为什么大模型对输入长度有硬性限制,以及为什么近期各种"长上下文"技术会成为热门话题。

第三,FP16 混合精度训练在现代大模型训练中已是标配,但由于 softmax 和注意力计算涉及指数和对数操作,精度敏感度高。从实践的角度看,如果你手动实现注意力机制,在 FP16 下可能出现训练不稳定的问题。建议在用 FP16 训练时,优先使用 PyTorch 官方封装好的注意力算子,因为它已经在数值稳定性上做过专门处理。

12. 最后再分享一点实践心得

从最初看到Q K^T / sqrt(d_k)这个公式时的茫然,到最终亲手把每个维度变换理清并跑通前向传播,这条路并不长。真正难的不是代码,而是愿意抛开"调用者心态",深入到组件内部去理解它的边界、它的细节、它的变换逻辑

如果你手头正有一个序列问题,建议直接动手:取一段文本,实现多头自注意力,观察不同 token 之间的注意力权重分布,你会发现一些直觉上"该被关注"的词确实获得了较高的权重。这种"模型内部"的观察比看任何论文插图都更直观。

当然,手动实现注意力机制不是终点。下一步值得继续深入的内容包括:

  • 位置编码的演进:从正弦编码到 RoPE(旋转位置编码),理解位置信息如何融入、以及为什么 RoPE 在长上下文场景有明显优势;
  • Flash Attention 的核心优化思路,以及它对训练效率和显存的影响;
  • GQA(分组查询注意力)和 MQA(多查询注意力)在不同大模型中的应用,也就是每组 Query 共享多少 Key 和 Value 的取舍逻辑;
  • 稀疏注意力与线性注意力在大规模长序列任务中的取舍。

读完这篇文章,我建议你先做一件事:打开代码编辑器,把本文中的scaled_dot_product_attention亲手敲一遍,然后给输入 X 打上梯度,试着优化它对某个监督信号的预测效果。当你能用自己实现的注意力模块搭建出一个能训练的小模型时,你就不再是"听过 Transformer"的人,而是"做过 Transformer"的人。这条从"看"到"做"的分界线,价值非常大。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询