news 2026/8/30 4:54:07

图解Transformer:从自注意力到ViT的架构拆解与实现

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
图解Transformer:从自注意力到ViT的架构拆解与实现

Transformer 这个架构,现在几乎成了人工智能大模型的基础组件。从文本生成、机器翻译到图像分类、视频理解,凡是你能看到的大模型,底层基本都绕不开它。这篇文章就负责把 Transformer 从头到尾拆开,配合图解思路讲清楚每个模块为什么存在、怎么计算、维度怎么变、训练时容易在哪里翻车。

我默认读者已经掌握了基本的神经网络和深度学习概念,但不要求你手写过 Attention。这篇内容适合三类人:正在准备面试的算法岗同学、需要做 NLP 或视觉项目的工程师、以及想把模型架构搞明白再动手写代码的初学者。整篇按“结构认知 + 逐块图解 + 最小实现 + 排错”的顺序来写,不用按章节顺序读,卡在哪一章就回到哪一章看。

1. 先搞清楚:Transformer 到底解决了什么问题

1.1 为什么“最后是 Transformer”:从序列建模的困境说起

在 Transformer 出现之前,处理序列数据主要靠 RNN/LSTM 和 CNN。RNN 是顺序扫描的,第 t 个词要等前 t-1 个词算完,并行性很差。长距离依赖问题虽然有 LSTM 的门控机制缓解,但模型仍然是按时间步展开,序列一长,梯度传播和信息丢失依然很麻烦。CNN 可以并行处理窗口内的词,但需要靠堆层数扩大感受野,才能建立跨度很大的依赖关系。注意力机制虽然很早就被用在翻译模型里,但当时通常是配合 RNN 使用,不是单独作为主架构。

Transformer 的做法是彻底去掉循环,直接对整句话计算两两之间的注意力得分。这样有两个直接好处。第一,并行度高,一句话内部的所有 token 可以同时参与计算。第二,任意两个位置之间的信息传递只需要一步,不需要沿着时间步或层逐层传递。

所以“为什么最后是 Transformer”这个问题,答案并不复杂:它把序列建模里“长距离依赖”和“并行计算”两个核心难题同时解决掉了。当然代价也很明显,自注意力的计算复杂度是 O(n²),输入序列越长,计算量和内存占用涨得越快。这也是后来各种稀疏注意力、线性注意力、窗口注意力变体出现的原因。

1.2 一张图记住整体结构:编码器、解码器、注意力、前馈

原始 Transformer 是编码器-解码器结构。左边编码器负责把输入序列编码成一组上下文表示,右边解码器负责根据这些表示和已经生成的内容逐步输出目标序列。

先记四个关键部件:

  • 自注意力(Self-Attention):处理序列内部的关系。
  • 交叉注意力(Cross-Attention):编码器和解码器之间的桥梁。
  • 前馈网络(FFN):对每个位置做非线性变换。
  • 残差连接和层归一化:稳定训练并缓解深层网络退化。

实际项目中,现在很多人只使用编码器(如 BERT)来做理解任务,或者只使用解码器(如 GPT 系列)来做生成任务。但无论哪种变体,内部的模块基本都还是 Transformer Block,所以先把原始结构看懂,后面看任何变体都不会太慌。

图解思路是:输入句子经过 Embedding 变成向量,加上位置编码,进入编码器;编码器输出 K 和 V,给到解码器的交叉注意力;解码器输入已经生成的部分序列,经过带掩码的自注意力后,再和编码器输出做交叉注意力;最后通过线性层和 Softmax 输出下一个词的概率。

2. 图解核心组件:Self-Attention 到底是怎么算的

2.1 Q、K、V 是什么:不要死记公式

很多初学者被 Q、K、V 这三个字母劝退。其实可以把它想象成检索过程。

  • Query(查询):你想找什么。
  • Key(索引):每个候选位置的标签。
  • Value(内容):每个候选位置真正携带的信息。

比如翻译“I love you”的时候,要决定“love”和后面哪个词关系最强,就把“love”的 Query 拿去和所有词的 Key 做点积,得到一组分数,然后按分数加权所有词的 Value。分数高,说明这两个位置的相关性强,输出里就会更多地保留那个位置的 Value 信息。

具体到计算上,输入 X 的形状是 [batch, seq_len, d_model],经过三个可学习的权重矩阵 W_Q、W_K、W_V,分别得到 Q、K、V。这三个矩阵通常形状都是 [d_model, d_model],所以 Q、K、V 的形状和 X 一致。这里“可学习”是关键,模型在训练时不断调整 W 矩阵,让注意力关注到真正有用的位置。

2.2 缩放点积注意力的计算过程和维度变化

计算分四步:

  1. 算注意力分数:Q 乘以 K 的转置,形状变成 [batch, seq_len, seq_len]。
  2. 缩放:除以 sqrt(d_k),防止点积过大导致 Softmax 梯度消失。
  3. Softmax:对每一行做归一化,得到注意力权重。
  4. 加权求和:注意力权重乘以 V,输出形状仍然是 [batch, seq_len, d_model]。

这里 d_k 是每个注意力头的维度,不是 d_model。为什么要除以 sqrt(d_k)?因为当维度变高时,点积的方差会变大,Softmax 之后容易进入饱和区,梯度会非常小。除以 sqrt(d_k) 是让方差回到一个可控范围,这样训练更稳定。

下面给一个精简版 PyTorch 实现,方便对照维度:

import torch import torch.nn as nn class ScaledDotProductAttention(nn.Module): def __init__(self, d_k): super().__init__() self.d_k = d_k def forward(self, q, k, v, mask=None): scores = torch.matmul(q, k.transpose(-2, -1)) / (self.d_k ** 0.5) if mask is not None: scores = scores.masked_fill(mask == 0, float('-inf')) weights = torch.softmax(scores, dim=-1) return torch.matmul(weights, v), weights

这里的 mask 是可选参数,解码器或者 padding 位置通常会用到。用 masked_fill 把非法位置置为负无穷,Softmax 后对应位置概率变成 0,就相当于屏蔽掉了不合法位置的注意力。

2.3 多头注意力:为什么要“多头”

多头注意力不是算一次,而是把 d_model 拆成 h 份,每份维度 d_k = d_model / h,分别做注意力计算,最后拼接起来再过一次线性层。

拆成多头有什么用?可以从两个角度理解。

  • 表达能力:不同的头可以关注不同类型的模式。有的头关注句法关系,有的头关注指代关系,有的头关注局部相邻词。
  • 计算稳定性:每个头的维度变小,矩阵运算规模降低,整体计算量在可控范围。

但多头也不是越多越好。头数增加后,每个头的维度变小,能表达的信息量会下降。同时训练数据不够时,多头之间的重复度会变高。实际使用常见的头数是 8、12、16、32 等,具体要根据模型总维度和任务来决定。如果调参,不要只调头数,要把“头数 × 每个头的维度”和总维度一起看。

3. 位置编码:序列顺序信息怎么进来

3.1 为什么 Transformer 需要位置编码

Self-Attention 本身是集合操作,它不知道 token 在句子里的顺序。把“我打你”和“你打我”的 token 向量输入进去,如果不加位置信息,模型会认为两个句子完全一样。这个问题在 Transformer 里是靠显式注入位置向量解决的。

主流做法有两种:

  • 固定位置编码:正余弦函数生成,不需要学习。原版 Transformer 用的就是这个。
  • 可学习位置编码:把位置向量当作参数训练,BERT 等模型常用。

固定位置编码的好处是外推到更长序列时相对稳定。可学习编码在训练数据足够时效果通常也不错。到底哪个更好,要看任务和数据量。实际项目里,很多人直接用可学习编码,因为实现简单,而且在训练长度附近表现够用。

3.2 正余弦位置编码为什么这么设计

正余弦位置编码的公式是:

PE(pos, 2i) = sin(pos / 10000^(2i / d_model)) PE(pos, 2i+1) = cos(pos / 10000^(2i / d_model))

这里 pos 是 token 的位置,i 是维度下标。设计思路是:不同维度使用不同的频率,低维度频率高,高维度频率低。这样一来,位置向量既能区分不同位置,又会表现出一定的平滑性。对模型来说,它可以学习利用这些位置的相对关系来推断“谁在谁的附近”。

这个公式还有一个好处:不需要额外参数,位置编码是确定的,任何输入长度都能计算,不占用训练参数。

如果使用可学习位置编码,要注意最大位置长度。如果训练时只覆盖 512 个位置,下游突然输入 1024 个 token,超出部分没有对应的位置向量,处理起来会很麻烦。常见做法是预先设置一个足够大的 max_len,或者在推理时用插值方式扩展。

3.3 位置编码的现代演进:绝对位置和相对位置

正余弦编码属于绝对位置编码,它把位置信息直接加到 token 上。BERT 的可学习位置编码也属于这一类。

但很多任务更需要的是“相对位置信息”。比如判断“猫”和“椅子”的关系时,真正重要的是它们之间距离多远,而不是它们在整句话里的绝对位置在第几个。于是出现了很多相对位置编码方案,比如在注意力分数计算时直接加上一个位置偏置项。近年来大模型里常用的 RoPE(旋转位置编码)也是一种相对位置方案,它让 Q 和 K 先旋转一定角度再计算点积,从而把相对位置信息编码进注意力分数里。

如果你在阅读 Transformer 变体论文时遇到位置编码相关的问题,只要记住一条主线:从绝对位置到相对位置,再到把位置信息融合进注意力计算。这个演进方向都是为了解决同一个问题,让模型知道“顺序”和“距离”。

4. 完整拆解一个 Transformer Block:从输入到输出

4.1 编码器层:Norm、注意力、残差、前馈网络的顺序

一个标准的编码器层包含两个子层:

  • 第一子层:多头自注意力 + 残差连接 + Layer Norm。
  • 第二子层:前馈网络 + 残差连接 + Layer Norm。

很多初学者会纠结 LayerNorm 放在哪里。原版采用“先残差相加,再 LayerNorm”,也就是:

x = LayerNorm(x + Attention(x)) x = LayerNorm(x + FFN(x))

后来很多模型使用 Pre-LN,也就是先把 x 归一化再进入子层:

x = x + Attention(LayerNorm(x)) x = x + FFN(LayerNorm(x))

两种方式在训练稳定性上有区别。Pre-LN 在深层网络里更容易训练,不一定需要特别的 warmup 策略,因此很多大模型实际采用 Pre-LN。学习的时候,不要把某一派的顺序当成唯一真理,重点理解每个组件的作用。

前馈网络通常是两个线性层加一个激活函数,中间层维度一般是 d_model 的 4 倍。以 d_model=512 为例,FFN 中间层通常是 2048。这个模块对每个词位置独立计算,不做跨位置交互。真正负责“跨位置交互”的就是注意力。

4.2 解码器层:掩码注意力和交叉注意力

解码器层比编码器多了一个子层:与编码器输出做交叉注意力。解码器里有三块:

  1. 带掩码的自注意力:生成第 t 个词时,只能看到当前位置之前的内容,不能看未来。这个掩码是一个上三角矩阵,未来位置被置为负无穷。
  2. 交叉注意力:Query 来自解码器,Key 和 Value 来自编码器输出。解码器通过这一步去“读取”输入句子的信息。
  3. 前馈网络,和编码器一致。

为什么必须加掩码?因为在训练阶段,我们同时把整个目标句子交给模型,如果不加掩码,模型就能“偷看”到后面要预测的词,训练会崩。推理阶段虽然不是整体输入,但掩码保证了训练和推理的一致性。

交叉注意力的出现也解释了一个常见问题:为什么纯解码器模型(如 GPT)没有交叉注意力?因为纯解码器只做文本生成,没有单独的编码器输入来源,它的 K、V 都来自自身历史上下文。而机器翻译这类任务有“源语言”和“目标语言”两个序列,所以才需要交叉注意力做桥梁。

4.3 输出层:如何把向量变成词概率

解码器最后一层输出的向量,经过线性层映射到词表大小,然后经过 Softmax,得到每个候选词的概率分布。训练时用交叉熵损失计算预测和真实下一个词的差距;推理时每次选一个词,把它拼到已有序列中,反复执行直到遇到结束符或达到最大长度。

这部分看起来简单,但推理时有个容易被忽略的细节:解码是逐步进行的,每一步都要重新计算当前所有位置的注意力。如果句子很长,这个逐步计算的过程会越来越慢。工程上常用 KV Cache,把已经算好的 Key 和 Value 缓存起来,避免每一步都重复计算旧位置。这也是很多大模型推理性能优化的核心手段之一。

5. 图解变体:ViT 和 Swin Transformer 怎么把注意力用到图像上

5.1 ViT 的核心思路:把图像切成 Patch

Vision Transformer(ViT)把图像当作“一系列单词”来处理。假设输入是 224x224 的 RGB 图像,先切成 16x16 的 Patch,共 14x14=196 个 Patch。每个 Patch 展平成一个向量,经过一个线性投影变成 d_model 维,这就类似文本里的 token embedding。再加上位置编码和可选的 [CLS] token,最后扔进标准 Transformer 编码器。

ViT 的最大优势是:结构统一,图像和文本可以用同一套架构来建模,这对多模态模型非常友好。但它有一个明显的边界:如果训练数据不够,它不像卷积神经网络那样自带局部性先验,可能不容易收敛。所以 ViT 早期主要是在大规模数据上效果突出,在小数据集上直接训练,效果往往不如 CNN。

5.2 Swin Transformer:窗口注意力解决什么问题

ViT 把整个图像当作一个全局序列,Token 数量一多,自注意力的 O(n²) 复杂度就很难受。比如高分辨率图像,Patch 数量上万后,直接算全局注意力内存会爆炸。

Swin Transformer 的核心改动:只在窗口内部做注意力。窗口大小固定,比如 7x7 个 Patch,这样注意力复杂度只和窗口大小有关,而不是整个图像大小。为了让不同窗口之间能交换信息,Swin 还引入了“移动窗口”机制,在相邻层之间把窗口偏移一些,让原本边界处的 Patch 在新的窗口里能互相看到。

Swin Transformer 还做了类似于 CNN 金字塔的层级结构:前面层的分辨率高、尺寸大,后面层通过 Patch Merging 逐步降低分辨率、增大通道数。这个设计让它做检测、分割等密集型任务时,可以直接套用目标检测框架里的多尺度思路。

5.3 图像任务里为什么需要层级和归纳偏置

很多人问:Transformer 做图像,为什么不能像文本一样直接全局注意力就好?原因主要有两个。

  • 图像分辨率高,Token 数远多于文本,全局注意力内存和计算代价太大。
  • 图像本身有强局部性,自然图像的多尺度结构决定了模型最好能分层提取特征。

窗口注意力本质上是在没有用卷积的情况下,给模型注入了一种“先看局部、再跨窗口看全局”的偏置。这也是为什么 Swin 在检测和分割上往往比 ViT 更容易落地。

对比角度ViTSwin Transformer
核心思路全局注意力窗口注意力 + 移动窗口
计算复杂度O(n²)O(window_size²)
层级结构基本没有多尺度有 Patch Merging 多尺度
适合任务中等分辨率图像分类等检测、分割等大分辨率场景

如果你想学图像 Transformer,我建议先跑 ViT 理解“图像怎么变成 token”,再跑 Swin 理解“窗口和移动窗口”,最后把两者对比起来看。直接啃论文图容易晕,先用现成库跑一个分类任务会更靠谱。

6. 动手理解的最好方式:跑一个最小 Attention 实现和可视化

6.1 环境准备和依赖

要看懂图解,最好的验证方式是亲手跑一次。如果只是想理解 Attention 计算流程,普通电脑就够了,不需要大显卡。常见依赖:

  • Python 3.8 及以上
  • PyTorch 1.13 或 2.x 均可
  • NumPy
  • Matplotlib(可视化用)

如果做的是 ViT 或者 Swin 的完整训练,再考虑 GPU 显存问题。ViT-Base 在 224x224 分辨率下,如果用完整训练,显存需求会比较高;但如果只是前向推理一个 Batch,8G 显存通常也可以做实验。低配置机器能跑通 Demo,不代表适合批量训练,这一点要提前有预期。

6.2 用 PyTorch 手写一个精简版 Attention

先跑一个随机输入的最小样例,确认形状没问题:

import torch import torch.nn.functional as F def attention(q, k, v, mask=None): d_k = q.size(-1) scores = torch.matmul(q, k.transpose(-2, -1)) / (d_k ** 0.5) if mask is not None: scores = scores.masked_fill(mask == 0, float('-inf')) weights = F.softmax(scores, dim=-1) return torch.matmul(weights, v), weights batch, seq_len, d_model, heads = 2, 8, 64, 4 d_k = d_model // heads x = torch.randn(batch, seq_len, d_model) w_q = torch.randn(d_model, d_model) w_k = torch.randn(d_model, d_model) w_v = torch.randn(d_model, d_model) q = x @ w_q k = x @ w_k v = x @ w_v # 拆成多头 q = q.view(batch, seq_len, heads, d_k).transpose(1, 2) k = k.view(batch, seq_len, heads, d_k).transpose(1, 2) v = v.view(batch, seq_len, heads, d_k).transpose(1, 2) output, attn_weights = attention(q, k, v) print(output.shape) # [batch, heads, seq_len, d_k] print(attn_weights.shape) # [batch, heads, seq_len, seq_len]

这个例子只是为了验证维度,不是完整实现。建议跑的时候把输出形状打印出来,对照上面的说明看一遍,比单纯看文章有效很多。

6.3 检查注意力权重、输出形状和梯度

如果自己实现,容易出现三类问题:

  • 形状错乱:多头拼接后没有恢复到 d_model,最后输出维度对不上。
  • Softmax 维度选错:应该对最后一维 seq_len 方向做归一化,不是对 batch 方向。
  • 掩码位置处理错误:负无穷没有准确覆盖,导致未来位置仍被关注。

判断标准很简单:

  • 任何中间输出的最后一个维度,最后都应该能对应回 d_model。
  • 注意力权重每一行的和应该等于 1,在有效位置内。
  • 训练时梯度不为 NaN,loss 能稳定下降。

如果发现注意力权重分布非常平均,先不要急着改结构。这通常意味着训练不充分、残差和 Norm 没有正确叠加,或者输入数据本身没有区分度。

6.4 可视化注意力的思路

可视化可以用 Matplotlib 把 attn_weights 画成热力图。横轴和纵轴都是序列位置,颜色越深代表注意力权重越高。这样可以直观看到某个 token 主要关注了哪些位置。

我一般会这样检查:

  • 如果是机器翻译样例,看解码时是否关注了源句中的对应词。
  • 如果有一段长文本,看模型是否关注了关键实体而忽略无关词。
  • 多头模式下,把不同头分开画,通常会看到有的头关注局部、有的头关注远距离指代。

训练充分的模型注意力图往往比随机初始化清晰很多。这也是“图解 Transformer”最有意思的部分,把一个公式变成看得见的矩阵。

7. 从图解到实战:常见坑和排查顺序

7.1 维度对不上:先查 batch、seq_len、d_model

Transformer 代码里最常见错误就是维度。遇到报错,先打印 Q、K、V 的形状,再确认拆分多头之后的形状。遵循一个原则:形状错误时先改维度,再看逻辑,不要一上来就调学习率或换模型结构。

报错现象优先排查
维度不匹配Q/K/V 形状、多头拆分和拼接方式
loss 为 NaN学习率、输入数据、掩码
注意力权重几乎一样训练不充分、随机初始化、Norm 顺序
内存溢出序列长度、batch size、头数、d_model

7.2 训练不收敛:先看掩码、学习率和数据格式

如果 loss 不下降,或者掉到 NaN,排查顺序是:

  1. 输入数据是否包含 NaN 或异常值。先检查数据,再改模型。
  2. 损失函数和标签是否对齐。文本任务里经常是标签错位。
  3. 学习率是否过大。Transformer 对学习率比较敏感,常用 warmup + 动态衰减。
  4. 是否加入了正确掩码。padding 位置不应该参与注意力,否则模型会把无意义位置也纳入计算。
  5. 初始化策略。不同框架的默认初始化不同,不要随意叠加自定义缩放。

如果只是学习而不是生产,建议先跑小数据集、小模型、短序列。把一条样例跑通,再慢慢扩大范围。

7.3 要不要从头手写:什么时候用现成框架

“手撕 Transformer”确实能加深理解,尤其是准备面试时,自己动手写一遍和只看文章完全不一样。但从工程角度,生产项目不需要自己实现 Attention,直接用 PyTorch、Hugging Face 等现成库更稳。

原因有几个:

  • 现成库考虑了大量边界条件、数值稳定性和计算优化。
  • 手写实现可能在短序列上能跑,但长序列、批处理、量化场景下性能和稳定性都差很多。
  • 后续维护成本更高。

我的建议是:把“手写”当作学习手段,把“现成库”当作落地手段。学习阶段一定要写,生产阶段不一定要写。

7.4 学习 Transformer 的一个合理路径

如果你刚入门,可以按这个顺序走:

  1. 先理解整体结构:输入、编码器、解码器、输出。
  2. 再理解自注意力:Q/K/V、缩放点积、多头。
  3. 自己跑一个最小注意力样例,确认输出形状。
  4. 复现一个单层编码器 Block,加入残差和 LayerNorm。
  5. 看 ViT 或 Swin,理解怎么把图像转成 token。
  6. 最后再看源码、看论文,结合公式对照代码。

如果时间紧张,前四步就足以应付绝大多数面试中的架构问题。后面两步是为了视觉和多模态场景准备的。

最后说一个个人体会。Transformer 图解的价值,不在于背出某个公式,而在于你能在写代码时预判“这一步输出形状是什么”“这个掩码会影响哪些位置”“这个参数拉大后资源占用会怎么涨”。把这三个问题想清楚,比多背十篇论文都有用。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/30 4:53:16

AI生成内容的责任归属:从可追溯性到可信度治理

最近在参与一个 AI 客服项目时,我意识到一个比“AI 会不会取代我”更现实的问题:AI 生成的内容一旦进入真实交易链路,谁为错误负责?我们做了一个很常见的客服助手,模型读产品文档,按用户问题生成回复。速度…

作者头像 李华
网站建设 2026/8/30 4:52:32

Agentic Programming没凉:从工具调用到工程落地的实用指南

最近在 Hacker News 上出现了一个很有意思的提问:Agentic Programming 是不是已经变成了一个 flop。所谓 flop,可以理解成“雷声大雨点小”的失败品。这个话题能在技术社区引发讨论,本身就说明一个问题:过去两年被各种 Demo 视频和…

作者头像 李华
网站建设 2026/8/30 4:51:17

软件测试面试高频考点与实战技巧:从八股文到Offer收割

“金三银四”的春招号角已经吹响,软件测试岗位的竞争一年比一年激烈。最近很多读者私信我,问得最多的就是“2023年软件测试面试到底在考什么”、“八股文背了这么多,为什么一到面试官面前就卡壳”。作为在测试行业摸爬滚打了十年、参与过上百…

作者头像 李华
网站建设 2026/8/30 4:48:12

SR5E1E570C30F01X车规MCU实战指南:从启动到CAN-FD

做嵌入式的朋友,第一次看到SR5E1E570C30F01X这个型号时,大概率会愣一下——这串字符既不像STM32那样直观,也套不进传统MCU型号的命名规则。它是面向车身控制类场景的一款32位车规MCU,刚拿到手时,我还按老经验去点灯&am…

作者头像 李华
网站建设 2026/8/30 4:47:31

智能仓储优化:从WMS到数据驱动的仓库效率革命

简介:本资源是一套面向企业信息化开发者与物流系统学习者的智能仓库管理系统优化方案实现,聚焦于仓储布局、库存预测、AGV调度、配送路径规划等核心业务场景的代码级落地。压缩包共366个文件,含106个Java后端逻辑文件、49个HTML前端页面、42个…

作者头像 李华