Transformer计算机视觉精讲:自注意力到ViT与Swin
2026/9/8 13:30:20 网站建设 项目流程

Transformer 进入计算机视觉已经不再是“要不要用”的问题,而是“怎么理解、怎么选型、怎么落地”的问题。这次我们来看一套覆盖全套知识点的内容:从注意力机制讲起,逐步推进到 ViT、Swin Transformer 两篇核心论文的精读。与很多只给概念截图的教程不同,这套内容更偏向“能把论文读透、能把代码写出来”的干货路线,适合准备面试、做课题、复现论文或者使用 Transformer 改进视觉任务的读者。

文章不绕弯子,直接拆解三层内容:第一层是 QKV、自注意力、多头注意力的数学原理与实现细节;第二层是 ViT 的 Patch Embedding、位置编码、Class Token 和整体训练策略;第三层是 Swin Transformer 的层次化设计、窗口注意力、移位窗口和相对位置偏置。最后会给出从论文到代码的验证路线,以及学习中最容易踩的坑。

如果你正在为“Transformer 计算机视觉”整理知识体系,或者准备面试时被问到“ViT 和 Swin 到底差在哪”,这篇可以直接收藏备用。

1. 核心能力速览

先给这套内容做一个整体画像,方便判断是否适合自己。

能力项说明
内容定位计算机视觉中的 Transformer 全套精讲,覆盖注意力机制、ViT、Swin Transformer
面向读者准备面试的算法工程师、做视觉大作业/课题的学生、准备复现论文的研究者
前置知识Python 基础、PyTorch 基础、CNN 基本概念(卷积、池化、全连接)
核心模块自注意力机制、多头注意力、位置编码、ViT 论文、Swin Transformer 论文
代码落地可对照 PyTorch 手写注意力模块、Patch Embedding、窗口注意力
面试价值高频考点全覆盖:QKV、softmax、复杂度、归纳偏置、层次化特征
配套建议需要一面读论文原文、一面跑代码验证,不建议只看概念图

这套内容最直接的收益是:读完以后你能回答清楚几个关键问题——为什么 NLP 的 Transformer 能搬到图像上?ViT 为什么需要大数据集预训练?Swin 的窗口注意力到底省了多少计算量?这也是面试和论文复现里最容易卡住的地方。

2. 为什么计算机视觉需要 Transformer

先回顾背景。CNN 统治视觉任务很多年,靠的是局部感受野和权值共享,卷积核一层层堆叠,从边缘、纹理逐步抽象到语义。这个设计很好,但有两个天然限制。

第一个限制是感受野扩展得慢。高层特征虽然能覆盖更大区域,但依然依赖层层堆叠,对于长距离依赖关系,比如一张图里“远处的行人”和“近处的汽车”之间的语义关系,CNN 需要很深的网络才能建模。

第二个限制是归纳偏置强。卷积假设了 locality(局部性)和 translation invariance(平移不变性),这在自然图像里通常成立,但遇到需要全局建模的任务时,反而成了限制。

Transformer 的思路完全不同。它不预设局部性,而是通过自注意力机制,让序列中任意两个位置直接交互。一张图切成 patch 序列后,每个 patch 都能和所有其他 patch 计算相关性。这种全局建模能力,是 Transformer 进入视觉领域的根本动力。

但要注意,这种能力的代价是计算复杂度。标准自注意力是 O(n²) 的复杂度,n 是序列长度。对 NLP 来说,一个句子不过几十到几百个 token,问题不大;对图像来说,如果直接把每个像素当 token,一张 224×224 的图就有 5 万个 token,计算量直接爆炸。这也解释了为什么 ViT 要把图像切成 patch,也解释了为什么 Swin Transformer 要设计窗口注意力。理解这条主线,后面读论文就顺了。

3. 注意力机制核心:Q、K、V 与自注意力

注意力机制最早来自 NLP 的机器翻译场景,本质是做“信息选择”:给定一个查询(Query),在一组键(Key)上计算相关性,再按相关性加权求和对应的值(Value)。

公式长这样:

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

其中:

  • Q 是查询向量,代表“我要找什么”。
  • K 是键向量,代表“我这里有什么”。
  • V 是值向量,代表“我实际提供的信息”。
  • d_k 是键向量的维度,除以 sqrt(d_k) 是为了防止点积过大导致 softmax 梯度消失。

为什么除以 sqrt(d_k)?因为 Q 和 K 的点积结果会随维度增大而增大,如果 d_k 很大,点积进入 softmax 的饱和区,梯度会非常小。缩放一下,让点积分布保持在合适的区间,训练更稳定。

自注意力是注意力机制的一个特例:Q、K、V 都来自同一个输入序列。也就是说,序列内部自己做“注意力分配”,每个 token 在计算输出时,都会参考序列里所有 token 的信息。

用 PyTorch 实现一个基础的自注意力模块,代码很短:

import torch import torch.nn as nn import torch.nn.functional as F class SelfAttention(nn.Module): def __init__(self, embed_dim, dropout=0.0): super().__init__() self.embed_dim = embed_dim self.q_proj = nn.Linear(embed_dim, embed_dim) self.k_proj = nn.Linear(embed_dim, embed_dim) self.v_proj = nn.Linear(embed_dim, embed_dim) self.dropout = nn.Dropout(dropout) def forward(self, x): # x: [batch_size, seq_len, embed_dim] B, N, C = x.shape q = self.q_proj(x) # [B, N, C] k = self.k_proj(x) # [B, N, C] v = self.v_proj(x) # [B, N, C] attn = torch.matmul(q, k.transpose(-2, -1)) / (C ** 0.5) attn = F.softmax(attn, dim=-1) attn = self.dropout(attn) out = torch.matmul(attn, v) # [B, N, C] return out

这段代码是理解 ViT 和 Swin 的基础。ViT 里的 Encoder Block 就是在这个模块外面加上 LayerNorm、MLP 和残差连接;Swin 里的窗口注意力则是在这个模块基础上加了 mask 和相对位置偏置。把这段代码跑通,再去看论文里的公式,基本没有理解障碍。

4. 多头注意力机制

单个注意力头只能学到一种“查询-键”相关性,表达能力有限。多头注意力把 embedding 分成 h 个头,每个头独立计算注意力,最后拼接在一起再做一次线性变换。这样模型可以同时关注不同的关系:某个头关注局部纹理,另一个头关注全局轮廓,各司其职。

多头注意力的公式:

MultiHead(Q, K, V) = Concat(head_1, ..., head_h) W_O head_i = Attention(QW_Q_i, KW_K_i, VW_V_i)

在实现层面,常见的做法不是真的建立多组独立 Linear,而是把 Q/K/V 的线性层输出拆成 h 份。比如 embed_dim 是 768,头数是 12,每个头的维度就是 64。

PyTorch 里手写多头注意力:

class MultiHeadAttention(nn.Module): def __init__(self, embed_dim, num_heads, dropout=0.0): super().__init__() assert embed_dim % num_heads == 0 self.embed_dim = embed_dim self.num_heads = num_heads self.head_dim = embed_dim // num_heads self.q_proj = nn.Linear(embed_dim, embed_dim) self.k_proj = nn.Linear(embed_dim, embed_dim) self.v_proj = nn.Linear(embed_dim, embed_dim) self.out_proj = nn.Linear(embed_dim, embed_dim) self.dropout = nn.Dropout(dropout) def forward(self, x): B, N, C = x.shape q = self.q_proj(x).view(B, N, self.num_heads, self.head_dim).transpose(1, 2) k = self.k_proj(x).view(B, N, self.num_heads, self.head_dim).transpose(1, 2) v = self.v_proj(x).view(B, N, self.num_heads, self.head_dim).transpose(1, 2) attn = torch.matmul(q, k.transpose(-2, -1)) / (self.head_dim ** 0.5) attn = F.softmax(attn, dim=-1) attn = self.dropout(attn) out = torch.matmul(attn, v) # [B, num_heads, N, head_dim] out = out.transpose(1, 2).contiguous().view(B, N, C) out = self.out_proj(out) return out

多头注意力里的几个细节值得单独记住:

  • 每个头的维度是 embed_dim / num_heads,不是 embed_dim。
  • 缩放因子用的是 head_dim,不是 embed_dim。
  • 拼接完要经过 out_proj 做一次融合线性变换。
  • 头数越多,模型越能捕捉不同类型的依赖关系,但计算量也线性增加。

面试里常问“为什么 softmax 之前要缩放”,答案就是上面说的梯度问题。另一个常问的是“多头注意力的复杂度”,答案和单头一致,因为头的数量不改变序列长度的平方项。

5. 位置编码

注意力机制本身没有顺序概念。对于一组 token,无论它们怎么排列,注意力计算的结果都一样。NLP 里 token 是词,顺序决定语义;视觉里 patch 是图像块,排列决定空间结构。模型必须知道每个 patch 在原始图像中的位置,否则一张图的所有 patch 顺序打乱后,模型会得到完全一样的输出。位置编码就是解决这个问题的。

ViT 使用可学习的位置编码。它的做法是初始化一个形状为 [1, num_patches + 1, embed_dim] 的向量表,作为可学习参数参与训练。推理时,直接和 patch embedding 相加。

ViT 的 Patch Embedding 和位置编码叠加逻辑:

import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768): super().__init__() self.num_patches = (img_size // patch_size) ** 2 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) def forward(self, x): # x: [B, C, H, W] -> [B, embed_dim, H/patch, W/patch] x = self.proj(x) # flatten spatial dims -> [B, num_patches, embed_dim] x = x.flatten(2).transpose(1, 2) return x class ViTEmbedding(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768): super().__init__() self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim) self.num_patches = self.patch_embed.num_patches self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches + 1, embed_dim)) def forward(self, x): x = self.patch_embed(x) # [B, N, C] cls_token = self.cls_token.expand(x.shape[0], -1, -1) x = torch.cat([cls_token, x], dim=1) # [B, N+1, C] x = x + self.pos_embed return x

Swin Transformer 则使用相对位置偏置(relative position bias),而不是绝对位置编码。它的思路是:不直接给每个 patch 一个绝对位置向量,而是在计算注意力分数时,根据两个 patch 之间的相对坐标差加上一个可学习的偏置项。

这样做有两个好处。第一,相对位置信息更符合视觉任务的特点——某个结构出现在图像左上角还是右下角,有时并不重要,重要的是结构内部各部件之间的相对空间关系。第二,Swin 的窗口注意力需要在不同窗口间共享参数,相对位置偏置天然具备平移不变性,更容易做到这一点。

相对位置偏置的实现细节是:对每个窗口内的 patch,先计算相对坐标偏移,映射到一组可学习参数表中,得到形状为 [num_heads, window_size², window_size²] 的偏置矩阵,然后加到注意力分数上。这个过程比较绕,读论文时建议配合代码一起看,单独看公式容易晕。

6. ViT 论文精讲

ViT 全称 Vision Transformer,是 2020 年 Google 团队发表在 ICLR 的工作。它的核心主张非常激进:既然 Transformer 在 NLP 上已经证明了强大的建模能力,那么把图像切成 patch 序列,直接输入标准 Transformer Encoder,不依赖任何卷积结构,也可以取得很好的图像分类效果。

6.1 网络结构

ViT 的结构可以拆成四段:

第一段是 Patch Embedding。把图像分成固定大小的 patch,常用 16×16,通过一个 stride 等于 patch_size 的卷积,把每个 patch 映射为一个向量。一张 224×224 的图,切成 16×16 patch,得到 14×14=196 个 patch,每个 patch 展平后经过线性映射成 768 维向量。

第二段是 Class Token 和位置编码。ViT 参考 BERT 的做法,在 patch 序列最前面拼接一个可学习的 class token。最终分类时,只取这个 class token 对应的输出过 MLP Head。位置编码使用可学习参数,形状是 [1, 197, 768],其中 197 是 196 个 patch 加上 1 个 class token。

第三段是 Transformer Encoder。这是标准结构:LayerNorm -> Multi-Head Attention -> 残差连接 -> LayerNorm -> MLP -> 残差连接。MLP 通常是两层,中间维度是 embed_dim 的 4 倍,激活函数用 GELU。这个 block 会重复 L 次。

第四段是分类头。取最后一层 Encoder 输出中的 class token 向量,经过 LayerNorm 后送入 MLP Head。预训练时 MLP Head 是一个线性层,微调时换成新的线性层即可。

ViT 各阶段的维度配置有一个通用命名规则,理解它对读论文很重要:

模型Patch 大小embed_dim层数头数MLP 维度参数量
ViT-Base16×167681212307286M
ViT-Large16×16102424164096307M
ViT-Huge14×14128032165120632M

6.2 关键结论

ViT 论文里最重要的实验结论是:在中等规模数据集(比如 ImageNet-1k)上直接训练,ViT 打不过同量级的 ResNet;但在大规模数据集(比如 JFT-300M,约 3 亿张图)上预训练后,再迁移到下游任务,ViT 可以超过 ResNet 和其他 CNN 方法。

这个结论引出了 ViT 最著名的特点:缺少 CNN 的归纳偏置,数据量不够时学不好。CNN 天生带 locality 和 translation invariance,小数据也能学出不错的结果;ViT 把这些先验去掉了,换来更强的表达能力,代价是大数据需求。

后续研究中,DeiT 通过知识蒸馏和更精细的训练策略,让 ViT 在 ImageNet-1k 上也能直接训练出不错的效果,说明数据不足的问题可以通过训练策略缓解,但这一点在 ViT 原始论文里确实存在。

6.3 需要注意的细节

  • 位置编码在 ResNet-50 特征图上做的话,需要调整编码长度以匹配特征图序列长度。
  • 微调时如果要提高输入分辨率,位置编码需要用二维插值放大,这是一个常见工程坑。
  • Class Token 不是唯一选择。后续工作证明,直接对所有 patch 输出做全局平均池化也可以,效果差别不大。ViT 论文选择 class token 是沿袭 BERT 的设计习惯。

7. Swin Transformer 论文精讲

Swin Transformer 是 2021 年微软研究院发表在 ICCV 的工作,拿了 Best Paper。它的动机直接针对 ViT 的两个痛点:计算复杂度高和缺乏多尺度特征。

7.1 两个核心创新

第一个创新是层次化架构。Swin 不是把整张图一次性切成固定数量的 patch,而是像 CNN 一样,通过 Patch Merging 逐步下采样。最开始每个 patch 是 4×4 像素,经过第一个 stage 后特征图分辨率减半、维度翻倍;总共 4 个 stage,输出的特征图分别是原图的 1/4、1/8、1/16、1/32。这种多尺度设计让 Swin 可以直接用于检测和分割等密集预测任务,而不像 ViT 那样需要在最后一层上做复杂的特征恢复。

第二个创新是窗口注意力。标准自注意力要在所有 patch 之间两两计算,复杂度是 O(n²)。Swin 把特征图分成不重叠的窗口,比如 7×7 个 patch 一个窗口,注意力只在窗口内部计算。这样复杂度变成 O(n),从平方降为线性,这是 Swin 能处理高分辨率图像的关键。

但窗口注意力有一个问题:窗口之间没有信息交互。如果一直用固定窗口,每个窗口只能看到自己内部的信息,全局建模能力被锁死了。Swin 的解法是移位窗口(Shifted Window Attention):在连续的两层 Transformer Block 之间,把窗口偏移一半窗口大小,再重新划分窗口。这样,上一层窗口边界处的 patch 在下一层就会和相邻窗口的 patch 处于同一个窗口内,实现了跨窗口信息流动。

移位窗口在工程上有实现难点:直接移动窗口会导致窗口数量变化,并且部分窗口会超出特征图边界。Swin 引入了 cyclic shift(循环移位)把越界的窗口搬回另一侧,再通过 attention mask 屏蔽掉不该计算的位置。这一步是 Swin 代码里最复杂的部分,建议读代码时单独花时间。

7.2 网络结构配置

Swin Transformer 有四个常用变体:

模型embed_dim层数窗口大小参数量
Swin-T962,2,6,27×728M
Swin-S962,2,18,27×750M
Swin-B1282,2,18,27×788M
Swin-L1922,2,18,27×7197M

每个 stage 里是偶数个 Transformer Block,因为需要“标准窗口注意力 + 移位窗口注意力”成对出现。

7.3 相对位置偏置的作用

Swin 在注意力分数上加了相对位置偏置,这个偏置来自一个可学习的参数表,表的索引是两个 patch 的相对坐标偏移。论文里有专门的消融实验:去掉相对位置偏置,Swin-T 在 ImageNet 上的 top-1 准确率明显下降;换成绝对位置编码,效果也不如相对位置偏置。这说明相对位置信息对窗口注意力非常重要。

7.4 关键结果

Swin 在 ImageNet-1k 上,Swin-B 用 ImageNet-1k 预训练可以达到 83.5% top-1 准确率,用 ImageNet-22k 预训练微调后可以达到 85.2%。在 COCO 检测和 ADE20K 分割任务上,Swin 超过了当时所有基于 CNN 和 ViT 的方法,成为通用骨干网络的新选择。

Swin 的意义不仅是分类精度,更在于它证明了 Transformer 通过合理的计算约束和多尺度设计,可以成为 CNN 的完整替代方案,覆盖分类、检测、分割全场景。

8. ViT 与 Swin Transformer 对比

把两篇论文放在一起对比,能更清楚地看到设计取舍:

对比维度ViTSwin Transformer
输入处理16×16 patch 线性映射4×4 patch,逐 stage 合并下采样
特征图尺度单一尺度多尺度(1/4、1/8、1/16、1/32)
注意力计算范围全局自注意力窗口内自注意力 + 移位窗口
计算复杂度O(n²)O(n)(窗口大小固定时)
位置编码可学习绝对位置编码相对位置偏置
归纳偏置弱,依赖大数据相对较强,适合中小规模数据集
适用任务图像分类为主,配合改动可用于分割检测分类、检测、分割通用
典型参数量86M(Base)88M(Base)

面试时如果被问“ViT 和 Swin 哪个更好”,不能一句话回答。应该说清楚:ViT 结构更简洁,全局建模能力强,但计算复杂度高、需要大数据;Swin 用窗口注意力换来了线性的计算复杂度,并且天然支持多尺度特征,更适合迁移到检测分割等密集预测任务。选择哪个取决于数据规模、任务类型和计算资源。

9. 常见问题与易错点

对照实际学习和复现经历,列出最常见的几个问题:

问题现象可能原因排查方式解决方案
自注意力实现后 loss 不下降缩放因子用错或 mask 缺失检查是否除以 sqrt(d_k)修正缩放,确认 padding mask
ViT 位置编码与输入长度不匹配修改输入分辨率后未插值位置编码打印 pos_embed 形状用双线性插值放大 pos_embed
Swin 窗口注意力结果异常移位窗口的 attention mask 写错单测一个 8×8 特征图,手算对比参考官方实现逐行对照 mask 逻辑
显存不足patch 数过多或 batch 过大降低 batch、减小输入分辨率使用渐变学习率配合梯度累积
小数据集上 ViT 效果差ViT 缺少 CNN 归纳偏置对比 ResNet 在同等数据下的表现换用 Swin/DeiT 或增加数据增强
训练速度慢全局注意力 O(n²) 开销大观察每个 step 的耗时切 patch 更大,或换成窗口注意力
分类结果波动大未设置合理的 warmup 和 weight decay检查训练配置ViT 训练对优化器参数敏感,按论文默认配置设置

这里面特别要提的是 ViT 的优化敏感性。ViT 对学习率、weight decay、warmup epochs 非常敏感,直接拿 ResNet 的训练配置去训 ViT,往往效果很差。DAT(Data efficient training)相关工作中也专门研究过这个,核心经验是:使用 AdamW、较大的 weight decay、较长的 warmup、以及适当的增强策略。

10. 学习路线与代码实践建议

如果你打算把这套内容完整学透,建议按下面五步走。

第一步,先不看任何 Transformer 结构,用 PyTorch 把一个自注意力模块从零写出来,跑通前向和反向。这个模块是后面所有内容的地基,建议写到能不看资料独立写出的程度。

第二步,手写多头注意力,然后自己验证一件事:把输入序列的顺序打乱,输出是否变化。如果不变,说明位置编码确实必要;这一步能帮你真正理解位置编码的作用。

第三步,读 ViT 原文,同时保持原文和代码对齐。重点看 Patch Embedding 的实现、Class Token 的拼接、Transformer Encoder 的 Block 结构。可以跑一个简单的实验:在 CIFAR-10 上用 ViT-Tiny 配置从零训练,观察收敛情况。

第四步,读 Swin Transformer 原文,重点放在窗口注意力和移位窗口的实现。建议做两个验证实验:第一个是固定窗口注意力下,两个不同窗口的 patch 之间注意力分数是否为零;第二个是经过一层移位窗口后,跨窗口信息是否开始流动。

第五步,把两篇论文做系统性对比,整理出自己的笔记。建议按“结构设计 -> 计算复杂度 -> 归纳偏置 -> 实验结论 -> 适用场景”五个维度写,这份笔记可以直接当面试复习材料用。

代码实践上,有几个值得长期保留的最小验证脚本:

# 验证自注意力输出形状 python -c " import torch from model import SelfAttention x = torch.randn(2, 197, 768) attn = SelfAttention(embed_dim=768) out = attn(x) print(out.shape) # 期望输出 torch.Size([2, 197, 768]) "
# 验证 ViT Patch Embedding python -c " import torch from model import PatchEmbed x = torch.randn(2, 3, 224, 224) pe = PatchEmbed(img_size=224, patch_size=16, embed_dim=768) out = pe(x) print(out.shape) # 期望输出 torch.Size([2, 196, 768]) "

如果你计划把 Transformer 用在自己的视觉任务上,比如改进 YOLO 检测器、做高光谱图像分类或者视频理解,建议从 Swin 的层次化结构入手,而不是直接堆 ViT。绝大多数实际视觉任务都需要多尺度特征,Swin 的架构和现有检测分割框架的兼容性更好。当然,如果资源充足且任务本身需要极强的全局建模能力,ViT 仍然值得尝试。

关于显存占用,可按实际环境测试,这里有一个通用观察方式:训练时用nvidia-smi实时监控显存,对比不同 patch size、不同分辨率、不同 batch size 下的占用差异。ViT 在 224×224 分辨率下 token 数只有 197,显存压力不大;一旦提升到 384×384 或 512×512,全局注意力的内存会快速上涨。Swin 因为窗口注意力是局部计算,高分辨率下显存增长会平缓很多,这在工程选型里有实际意义。

11. 总结与下一步

Transformer 计算机视觉这条技术线,从注意力机制到 ViT 再到 Swin Transformer,是一条逻辑非常清晰的演进路线。抓住三条主线就能串联全部内容:注意力机制解决“如何建模全局关系”,ViT 回答“如何用纯 Transformer 做视觉”,Swin 解决“如何让 Transformer 在视觉任务上更高效、更通用”。

首次学习时,建议优先把自注意力和多头注意力写熟,然后精读 ViT 原文,最后攻克 Swin 的移位窗口。最容易踩的坑是:只看概念图、不读源码、不做对比实验。建议至少完整复现一个最小自注意力模块,并跑通 ViT 和 Swin 的官方预训练模型推理。

下一步可以根据自己的方向继续扩展:如果做图像生成,可以研究 ViT 和 Swin 之后的 Diffusion Transformer(DiT);如果做多模态,可以看 CLIP 和 BLIP 如何借用视觉 Transformer;如果做轻量化,可以调研 MobileViT 和 EdgeViT 这类轻量级视觉 Transformer。基础打牢之后,后面的路会顺很多。

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

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

立即咨询