前言
本站的《Transformer 详解》讲了如何用注意力机制处理文本:把句子切分成 token,让任意两个位置直接交互。本文回答一个相关的问题:同一套机制,能不能用来处理图像?
答案是能,这个架构就是 ViT(Vision Transformer,视觉 Transformer),由 Google 在 2020 年提出。它的核心思路是:把图像切成固定大小的小块(Patch),每块展平成一个向量,再把这些向量按顺序组成序列输入 Transformer。像素是连续矩阵,文本是离散符号,两者结构原本不同;这一步变形把图像转换成了 Transformer 能够处理的序列形式。
ViT 是本站 VLM 专栏的基础:《CLIP 详解》的视觉编码器是 ViT,《LLaVA 详解》复用的也是它。理解了 ViT,就能理解这两篇文章中的图像特征是如何产生的。
阅读前提:已经读过本站《Transformer 详解》(理解自注意力、多头、位置编码),并会写基础的 PyTorch 代码。本文会与 CNN 做对比,但不要求你已经精通 CNN。
为什么视觉需要 Transformer
CNN 的归纳偏置
在 ViT 之前,视觉任务的标准做法是 CNN(卷积神经网络)。CNN 有两个内在的设计假设:
- 局部性(Locality):CNN认为图上相邻的像素一定是有关系的,因此它让一个固定大小的 3*3 卷积核在图片上滑动,每个输出位置只"看"周围的一小片区域。要利用远处的信息,只能逐层堆叠卷积,逐步扩大感受野。
- 平移等变性(Translation Equivariance):同一个卷积核在整张图上共享权重,物体平移后,对应的特征也随之平移。也就是说,一个物体无论出现在一张图片中的哪一个地方,其特征都是一样的。
这两个假设统称归纳偏置(Inductive Bias),即用先验知识规定模型理解图像的方式。优点是数据效率高,用 ImageNet 规模的数据就能训练出不错的模型;缺点是这些先验限制了模型的表达能力:局部性意味着远距离信息只能通过加深网络来间接获取,而且这些先验并非在所有任务中都成立。
ViT 的选择:去掉归纳偏置
ViT 采取相反的思路:几乎完全去掉 CNN 的归纳偏置,只保留 Transformer 结构。图像被切成 patch 后,任意两个 patch 之间可以直接交互(自注意力),远距离依赖不再是问题;平移等变性也不再由结构保证,需要模型从数据中自行学习。
这种设计的代价是数据需求:归纳偏置越少,需要的数据就越多。ViT 论文的对比实验直接说明了这一点(详见"训练"一节):在 ImageNet(128 万张图)上从头训练时,ViT 的表现与同规模的 ResNet 相当;在更大的数据集(ImageNet-21k、JFT-300M)上预训练后,ViT 超过最强的 CNN 基线,而且数据规模越大,优势越明显。
以下为一张对比图,展示 CNN 的局部滑动卷积与 ViT 的全局自注意力:

如上图所示,左边的 CNN 只能看红星附近的一小圈蓝星;而右边 ViT 的红星,可以直接且瞬间与画面各个角落的蓝星建立联系。由此可见,ViT 的自注意力机制具有强大的能力。
简言之,CNN 用归纳偏置降低数据需求,ViT 用更大的数据规模补偿归纳偏置的缺失。
核心思想:把图像变成序列
我们知道 Transformer 的输入是序列,也即一连串的按顺序排列的向量,但现在需要处理的图像不是序列。因此第一步必须完成一次转换:
一张 $224 \times 224$ 的彩色图像 → 一段长度为 196 的序列,每个元素是一个 768 维的向量。
实现这个转换的组件叫 Patch Embedding(图像块嵌入),分四步:切分、展平、投影、加位置编码。下面逐一说明。
Patch Embedding:图像如何变成 token
下图展示了从图像到 token 序列的完整流程,后面的四步将逐一展开说明:

第一步:切分成 patch 网格(Reshape)
切分在代码中通过张量变形(reshape)实现,而不是真的把图像切开。假设输入图片的形状是 $[B, C, H, W] = [1, 3, 224, 224]$(B 是批大小,C 是通道数),patch 边长取 $P = 16$:
- 高方向:$224$ 拆成 $14 \times 16$(14 个网格位置 × 每个位置 16 像素);
- 宽方向同理:$224$ 拆成 $14 \times 16$。
用 reshape 把图片变成 $[1, 3, 14, 16, 14, 16]$。此时数据被组织成 14 行 × 14 列的网格,每个格子是一个 $16 \times 16 \times 3$ 的小方块。reshape 不改变数据内容,只改变排列方式,因此没有信息损失。
这里有一个容易出错的细节:不能直接把 $[1, 3, 14, 16, 14, 16]$ reshape 成最终形状,因为通道维(3)与 patch 内部的空间维(16、16)不相邻,而 reshape 只能合并相邻的维度。必须先做一次 permute(维度重排),把「通道、patch 高、patch 宽」三个维度调整到相邻位置:
$$ [1, 3, 14, 16, 14, 16] \xrightarrow{\text{permute}} [1, 14, 14, 3, 16, 16] $$如果漏掉这一步,768 维向量中各通道与空间位置的对应关系会发生错乱。为什么不能直接 reshape?原因在于底层的内存布局:
为什么不能直接 reshape?
事实上,无论张量(Tensor)有多少维,它在计算机内存里本质上都只是一条一维长链。reshape 操作本身不会改变数据在内存里的物理排布,它只是给这条一维长链重新划分子段。
1. 不加 permute 时,内存里装的是什么?
当张量形状为 $[1, 3, 14, 16, 14, 16]$ 时,维度从左到右依次代表:
[Batch, Channel(C), Grid_H, Patch_H, Grid_W, Patch_W]。在 PyTorch 的默认内存排列(C 语言行优先格式)中,最右边的维度变化最快,最左边的维度变化最慢。如果直接按这个形状去内存里顺序读取前 768 个数(也就是展平后给 Token 0 的数据),内存里的数据其实是这样的排布:
固定 $C = 0$(红色通道)、
Grid_H = 0(第 0 行网格)、Patch_H = 0(Patch 的第 0 行像素),让Grid_W在 0~13 之间变动、Patch_W在 0~15 之间变动,会连续读出 $14 \times 16 = 224$ 个像素值。物理含义:这是整张图片顶部第一行所有 14 个 Patch 的 R 通道像素。接着看
Patch_H = 1(Patch 的第 1 行像素):又读出 $14 \times 16 = 224$ 个像素值(整张图第二行所有 Patch 的 R 通道像素)。接着看
Patch_H = 2:再读出 224 个像素值。接着看
Patch_H = 3:读出剩下的 96 个像素值,凑满 768 个数。结果:这 768 个数包含了横向 14 个 Patch 混在一起的像素,而且全都是 R 通道(完全没有 G 和 B 通道),甚至还丢失了 Patch 下半部分的像素。这个 Token 已经失去单个 Patch 的语义,变成无法使用的数据。
2. 做了 permute 之后,发生了什么?
目标是让 Token 0 包含左上角这第一个 Patch 的 3 个通道、$16 \times 16$ 所有像素。通过执行
permute(0, 2, 4, 3, 5, 1)(把维度顺序调整为[Batch, Grid_H, Grid_W, Channel, Patch_H, Patch_W],形状变为 $[1, 14, 14, 3, 16, 16]$),张量维度的语义被重新排布;结合后续的 reshape(内部会先做内存重整),一维长链的数据顺序被物理重排,最右侧相邻的三个维度变成了[3, 16, 16](通道、Patch 高、Patch 宽)。此时在内存里,排在最前面的 $3 \times 16 \times 16 = 768$ 个连续数据,正好就是第 0 行第 0 列那个 Patch 自己的 RGB 完整像素。对比总结:
操作 内存里前 768 个像素的来源 Token 0 的真实语义 直接 reshape 整张图横向 14 个 Patch 的前 3.4 行 R 通道像素 ❌ 多个 Patch 像素横向杂糅,缺失 G、B 通道 先 permute 再 reshape 仅第 0 个 Patch 内部的 3 个通道 × $16 \times 16$ 像素 ✅ 完整提取单个 Patch 这就是为什么维度重排(permute)不能省略:它决定了拆分出来的数字到底是一个独立的图像块,还是全图像素打碎后的混合物。
[1, 3, 14, 16, 14, 16]"] b --> c["permute
[1, 14, 14, 3, 16, 16]"] c --> d["reshape
[1, 196, 768]"] d --> e["线性投影
[1, 196, D]"]
第二步:展平成序列(Flatten)
网格是二维的,而 Transformer 的输入是一维序列。把 $14 \times 14$ 的网格拉直,形状变成 $[1, 196, 768]$,其中:
- 196 是序列长度,对应图像上的 196 个区域;
- 768 是每个 token 的向量维度,等于 $3 \times 16 \times 16$,即一个 patch 内所有通道的所有像素值。
拉直之后,原图中相邻的 patch 在序列里不一定相邻(比如垂直相邻的 patch 会相隔 14 个位置)。这会影响后续的注意力计算吗?答案是:不会。因为自注意力只依赖 token 的向量内容,不依赖 token 之间的空间距离。以 Token 0(左上角第一个 patch)为例,它与序列中任意 token 的交互计算都是同一种形式:
Token 0 与 Token 1(原图中水平相邻的 patch):
$$ \text{Score}_{0 \to 1} = Q_0 \cdot K_1^\top $$Token 0 与 Token 14(原图中垂直相邻的 patch,即第二行第一个):
$$ \text{Score}_{0 \to 14} = Q_0 \cdot K_{14}^\top $$Token 0 与 Token 195(原图中右下角、距离最远的 patch):
$$ \text{Score}_{0 \to 195} = Q_0 \cdot K_{195}^\top $$三者的计算规则完全一致:把 Token 0 的查询向量 $Q_0$ 与对应 Token 的键向量做点积,点积结果越大,关联度越高。公式里只有向量内容的相乘,没有「计算两者坐标差 $|i - j|$」这一项。因此二维网格拉直成一维序列,不会丢失注意力计算所需的任何信息;空间位置关系由位置编码另行提供(见下一节)。
这一步是整个方法的关键:图像的空间网格从此被当作序列来处理,后续的自注意力、多头、残差等操作都作用在这段序列上。
第三步:线性投影(Patch Projection)
展平得到的 768 维向量只是原始像素值,只记录了每个位置的颜色,不包含语义信息。ViT 用一个可学习的线性层把 768 维映射到模型内部使用的维度 $D$(ViT-Base 用 $D = 768$,ViT-Large 用 $D = 1024$):
$$ z_p = x_p E $$其中 $x_p \in \mathbb{R}^{768}$ 是单个 patch 展平后的向量,$E \in \mathbb{R}^{768 \times D}$ 是可学习的投影矩阵,$z_p \in \mathbb{R}^{D}$ 是投影后的 patch 嵌入。这一步把像素值映射到模型的嵌入空间,是 ViT 中直接处理图像结构的可学习参数。
第四步:加位置编码
到这一步,图像已经变成了序列,但还缺少位置信息。展平操作丢弃了空间位置:模型无法区分哪个 token 来自左上角、哪个来自右下角。因此需要给每个 token 加上一个可学习的位置编码向量,使模型能够区分不同的空间位置。具体做法见下一节。
一个等价实现:卷积
切 patch 与线性投影这两步可以合并为一步:一个卷积核大小为 16、步长为 16 的卷积,输出恰好是 $14 \times 14$ 个位置,每个输出位置对应输入上的一个 $16 \times 16 \times 3$ 的区域。卷积窗口的滑动对应切 patch,卷积核权重对应线性投影矩阵:
nn.Conv2d(in_channels=3, out_channels=D, kernel_size=16, stride=16)把卷积核权重 $W \in \mathbb{R}^{D \times 3 \times 16 \times 16}$ 展平,就得到投影矩阵 $E$ 的转置。两者数学上等价,工程上卷积写法更简单,timm、Hugging Face 等库的 ViT 实现都采用这种方式。
一个具体的数值例子
下面用一个极小的例子演示计算过程。设输入是一张 $4 \times 4$ 的灰度图(单通道),patch 边长 $P = 2$:
$$ X = \begin{bmatrix} 1 & 2 & 3 & 4 \\ 5 & 6 & 7 & 8 \\ 9 & 10 & 11 & 12 \\ 13 & 14 & 15 & 16 \end{bmatrix} $$reshape 成 $[2, 2, 2, 2]$(行 2 组 × 列 2 组,每组 $2 \times 2$),permute 后展平,得到 4 个 patch、每个 4 维:
$$ x_p = \begin{bmatrix} 1 & 2 & 5 & 6 \\ 3 & 4 & 7 & 8 \\ 9 & 10 & 13 & 14 \\ 11 & 12 & 15 & 16 \end{bmatrix} \in \mathbb{R}^{4 \times 4} $$第一行是左上角 $2 \times 2$ 区域按行展开:$[1, 2, 5, 6]$,而不是 $[1, 2, 3, 4]$——后者会把不同 patch 的像素混在一起,这正是 permute 的作用。最后乘投影矩阵 $E \in \mathbb{R}^{4 \times D}$,得到 4 个 $D$ 维的 token。
位置编码:找回顺序信息
为什么需要位置编码
自注意力对位置不敏感:交换 token 的顺序,两两之间的相似度计算结果不变(本站《Transformer 详解》讲过这一点)。对文本,“我打你"和"你打我"含义不同;对图像,左上角的 patch 和右下角的 patch 语义不同。位置信息必须显式加入输入。
ViT 使用可学习的 1D 位置编码
文本 Transformer 常用固定的 sinusoidal 编码(公式见本站 Transformer 主文),ViT 选择可学习的 1D 位置编码:一个形状为 $(N + 1) \times D$ 的参数矩阵,$N$ 是 patch 数,加 1 是因为序列开头还有一个分类 token(下一节说明)。$N$ 个 patch 展平后是一维序列,因此位置编码也设计成一维,直接加到每个 token 上:
$$ z_0 = \left[ x_{\text{class}}; \; x_p^1 E; \; x_p^2 E; \; \dots; \; x_p^N E \right] + E_{\text{pos}} $$其中 $E_{\text{pos}} \in \mathbb{R}^{(N+1) \times D}$ 是位置编码矩阵,分号表示沿序列方向拼接。
为什么 1D 位置编码就够了
图像是二维的,位置编码却是一维的,这可能让人疑惑。ViT 论文对此做了消融实验,对比了 1D、2D 和相对位置编码三种方案,结果是 1D 与 2D 的表现几乎相同。这说明:位置编码只是提供位置线索,二维空间结构可以由注意力层从数据中自行学习,不需要显式建模。
[CLS] token 与分类头
序列准备好之后,还有一个问题:Transformer 输出 196 个向量,分类时使用哪一个?
ViT 沿用 BERT 的设计,在序列最前面拼接一个额外的 [CLS] token(形状 $1 \times D$,参数可学习)。它不对应任何 patch,不携带图像信息;经过全部 Transformer 层之后,它的输出向量作为整张图像的汇总表示,接一个线性分类头得到预测:
$$ y = \text{Linear}\left( \text{LN}\left( z_L^0 \right) \right) $$其中 $z_L^0$ 是最后一层 [CLS] 位置的输出,LN 是 LayerNorm。
为什么不直接对 196 个 patch 的输出取平均(全局平均池化,GAP)?ViT 论文对比了两种做法,结果显示两者表现接近,[CLS] 略好,因此成为默认。直观上,平均池化给所有 patch 相同的权重;[CLS] 的表示由注意力机制决定从哪些 patch 聚合信息,灵活性更高。
完整架构
把各个组件按顺序组合,ViT 的整体流程是:
切 patch + 线性投影"] pe --> seq["196 个 patch token + [CLS]
+ 位置编码"] seq --> enc["Transformer 编码器 ×L 层
(多头自注意力 + MLP + 残差 + LayerNorm)"] enc --> head["分类头(取 [CLS] 输出)"] head --> out["类别预测"]

上图左侧是 ViT 的端到端流程,右侧是编码器内部的一个处理单元(Block),各部分说明如下:
- 图像分块(Patches):最下方是输入图像,它被均匀地切割成多个固定大小的方块(Patch),类似于把一篇文章分成一个个单词。
- 线性投影(Linear Projection of Flattened Patches):每个二维的图像块被展平成一维向量,然后通过一个线性层映射到固定的维度。图中用一排彩色长方形表示这些映射后的向量。
- 位置嵌入与类别嵌入(Patch + Position Embedding):Transformer 本身不包含位置信息,因此需要为每个图像块向量加上位置编码,让模型知道每个块在原图中的位置(图中小方块上方的 0、1、2……9 代表位置);序列最前面额外添加一个特殊的 [CLS] 标记(图中对应 0 的位置),它的作用是在经过层层处理后汇聚整张图的全局信息,用来做最终的分类。
- Transformer 编码器(Transformer Encoder):带有位置信息的向量序列被送入中间的灰色长条,即 Transformer 编码器(其内部结构见右图)。
- 分类头(MLP Head):经过编码器处理后,取出对应 [CLS] 标记的输出向量,送入 MLP Head(多层感知机分类器)。
- 输出类别(Class):MLP Head 最终输出图像属于各个类别的概率,如"鸟(Bird)"、“球(Ball)"、“车(Car)“等。
- 编码器内部结构(右图虚线框):数据流向自下而上。先经 LayerNorm(Norm)归一化,帮助加快训练并稳定模型;然后是多头自注意力(Multi-Head Attention),让每一个图像块都能"关注"到其他所有图像块,捕捉图像各部分之间的全局依赖关系;注意力层的输出与最初的输入相加(残差连接,防止深层网络中的梯度消失问题);再次 LayerNorm 归一化;随后是包含两个全连接层的 MLP,对提取的特征做非线性变换;最后再将 MLP 的输出与上一步的输入相加(残差连接),完成整个编码器模块的处理。
编码器内部的每一层与本站《Transformer 详解》中讲的完全一致:多头自注意力在 patch 之间交换信息,MLP(升维到 $4D$ 再降回)对每个 token 独立加工特征,残差连接和 LayerNorm 保证深层网络稳定训练。ViT 没有引入任何新的 Transformer 组件,唯一的视觉专用设计是图像到序列的转换方式。
超参数与模型家族
ViT 论文定义了三个规格,结构相同,尺寸不同:
| 模型 | 层数 L | 隐藏维度 D | 注意力头数 | 参数量 |
|---|---|---|---|---|
| ViT-Base | 12 | 768 | 12 | 86M |
| ViT-Large | 24 | 1024 | 16 | 307M |
| ViT-Huge | 32 | 1280 | 16 | 632M |
patch 边长常见 16 或 32(记作 ViT-B/16、ViT-B/32)。patch 越小,序列越长,保留的细节越多,但注意力计算量随序列长度的平方增长,需要在精度与计算量之间权衡。
训练:为什么 ViT 需要大数据
如前所述,ViT 去掉了 CNN 的归纳偏置,代价是数据需求增大。论文中的实验对比说明了这一点:
- 在 ImageNet(128 万张)上从头训练:ViT-B/16 的准确率约为 79.9%,与同规模的 ResNet 相当,没有体现出优势;
- 在 ImageNet-21k(1400 万张)上预训练后微调:ViT 超过同规模 CNN;
- 在 JFT-300M(3 亿张)上预训练后微调:ViT 全面领先,其中 ViT-H/14 在 ImageNet 上达到 88.55%。
结论是:预训练数据的规模直接决定 ViT 能达到的性能上限。这也是 ViT 早期被认为"需要海量数据"的原因。后续的 DeiT 通过知识蒸馏和更强的数据增强,使 ViT 能够在 ImageNet-1k 单个数据集上从头训练并取得较好的效果。
ViT 的常规用法是"预训练 + 微调”:在大数据集上预训练,然后在目标数据集上微调。微调时通常更换分类头;如果序列长度不匹配,位置编码需要做插值。论文还强调了数据增强和正则化(Mixup、CutMix、DropPath 等)对 ViT 的重要性:归纳偏置少,需要更强的正则化。
局限与演进
计算复杂度
ViT 的自注意力计算复杂度是 $\mathcal{O}(N^2)$,$N$ 是 patch 数(本站 Transformer 主文有详细推导)。对 $224 \times 224$ 的输入,$N = 196$,可以接受;但高分辨率图像(例如 $1024 \times 1024$、patch 为 16 时 $N = 4096$)会使注意力矩阵急剧膨胀,显存和计算时间都难以承受。针对这一瓶颈,出现了一系列改进:
- Swin Transformer:把注意力限制在局部窗口内,并像 CNN 一样分层下采样,计算复杂度与序列长度成线性关系,同时保留多尺度信息。
- DeiT:用知识蒸馏和更强的数据增强,使 ViT 摆脱对超大数据集的依赖。
- MAE(掩码自编码):遮住大部分 patch,让模型重建被遮住的部分,用自监督方式预训练 ViT,数据效率显著提高。
- DINO:自监督对比学习方法,学到的特征可以用于语义分割等下游任务。
作为 VLM 的视觉编码器
对本站 VLM 专栏而言,ViT 的角色是视觉编码器。《CLIP 详解》用 ViT 把图像编码成特征向量,与文本特征做对比学习;《LLaVA 详解》复用 CLIP 的 ViT,把 patch 特征投影成语言模型能读的 token。CLIP 和 LLaVA 的图像理解能力都建立在 ViT 的图像到序列转换之上。理解了本文,就能回答这两篇文章中"图像特征从哪里来"的问题。
代码实战:从零实现 ViT
下面用 PyTorch 从零实现一个可训练的 ViT,不借助 torchvision.models 里的现成封装。代码分为四部分:Patch Embedding、编码器、完整模型、训练。
Patch Embedding
先写教学版(与前面的数学推导一一对应),再给工程版(用卷积一步完成):
import torch
import torch.nn as nn
# ---------- 超参数 ----------
PATCH_SIZE = 16 # patch 边长(像素)
EMBED_DIM = 192 # 嵌入维度 D
NUM_HEADS = 4 # 注意力头数
NUM_LAYERS = 6 # Transformer 块数
HIDDEN_DIM = 768 # MLP 隐藏层维度(通常 4 × D)
NUM_CLASSES = 10 # CIFAR-10 类别数
DROPOUT = 0.1
class PatchEmbed(nn.Module):
"""教学版:切 patch + 展平 + 线性投影,与公式一一对应"""
def __init__(self, in_channels=3, patch_size=16, embed_dim=192):
super().__init__()
self.patch_size = patch_size
self.proj = nn.Linear(patch_size * patch_size * in_channels, embed_dim)
def forward(self, x):
B, C, H, W = x.shape
P = self.patch_size
# 1. 切分:高、宽各拆成 (网格数, patch 边长)
x = x.reshape(B, C, H // P, P, W // P, P) # (B, 3, 14, 16, 14, 16)
# 2. 重排:把 [通道, patch 高, patch 宽] 挪到一起
x = x.permute(0, 2, 4, 3, 5, 1) # (B, 14, 14, 3, 16, 16)
# 3. 展平:网格 -> 序列,patch 内部 -> 一维向量
x = x.reshape(B, -1, P * P * C) # (B, 196, 768)
# 4. 线性投影
x = self.proj(x) # (B, 196, D)
return x
class PatchEmbedConv(nn.Module):
"""工程版:一个卷积完成"切 patch + 线性投影"(timm 等库的写法)"""
def __init__(self, in_channels=3, patch_size=16, embed_dim=192):
super().__init__()
self.proj = nn.Conv2d(in_channels, embed_dim,
kernel_size=patch_size, stride=patch_size)
def forward(self, x):
x = self.proj(x) # (B, D, 14, 14)
return x.flatten(2).transpose(1, 2) # (B, 196, D)两个版本输出一致。工程版利用的是"卷积窗口滑动对应切 patch、输出通道对应线性投影"这一等价关系。
编码器
编码器块与本站 Transformer 主文的代码同构:Pre-LN 顺序(先归一化,再做注意力或前馈,最后加残差):
class Attention(nn.Module):
"""多头自注意力(与 Transformer 主文一致,这里合并了 Q/K/V 投影)"""
def __init__(self, embed_dim, num_heads, dropout=0.1):
super().__init__()
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
self.qkv = nn.Linear(embed_dim, embed_dim * 3) # 一次算出 Q、K、V
self.proj = nn.Linear(embed_dim, embed_dim)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
B, N, C = x.shape
qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim)
qkv = qkv.permute(2, 0, 3, 1, 4) # (3, B, H, N, d)
q, k, v = qkv.unbind(0)
attn = (q @ k.transpose(-2, -1)) * (self.head_dim ** -0.5) # 缩放点积
attn = attn.softmax(dim=-1)
attn = self.dropout(attn)
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
return self.proj(x)
class MLP(nn.Module):
"""逐 token 的前馈网络:升维 -> GELU -> 降维"""
def __init__(self, embed_dim, hidden_dim, dropout=0.1):
super().__init__()
self.net = nn.Sequential(
nn.Linear(embed_dim, hidden_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_dim, embed_dim),
nn.Dropout(dropout),
)
def forward(self, x):
return self.net(x)
class EncoderBlock(nn.Module):
"""一个 Transformer 块:注意力 + MLP,各自带残差,Pre-LN 顺序"""
def __init__(self, embed_dim, num_heads, hidden_dim, dropout=0.1):
super().__init__()
self.norm1 = nn.LayerNorm(embed_dim)
self.attn = Attention(embed_dim, num_heads, dropout)
self.norm2 = nn.LayerNorm(embed_dim)
self.mlp = MLP(embed_dim, hidden_dim, dropout)
def forward(self, x):
x = x + self.attn(self.norm1(x)) # 子层 1 + 残差
x = x + self.mlp(self.norm2(x)) # 子层 2 + 残差
return x完整模型
把 Patch Embedding、[CLS] token、位置编码、编码器堆叠、分类头组合起来:
class ViT(nn.Module):
"""完整 ViT:Patch Embedding -> [CLS] + 位置编码 -> L 层编码器 -> 分类头"""
def __init__(self, in_channels=3, patch_size=16, embed_dim=192,
num_heads=4, num_layers=6, hidden_dim=768,
num_classes=10, num_patches=196, dropout=0.1):
super().__init__()
self.patch_embed = PatchEmbed(in_channels, patch_size, embed_dim)
# [CLS] token 与位置编码都是可学习参数
# 实践中常用截断正态初始化(nn.init.trunc_normal_),这里用零初始化保持简洁
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
self.pos_drop = nn.Dropout(dropout)
self.blocks = nn.ModuleList([
EncoderBlock(embed_dim, num_heads, hidden_dim, dropout)
for _ in range(num_layers)
])
self.norm = nn.LayerNorm(embed_dim)
self.head = nn.Linear(embed_dim, num_classes)
def forward(self, x):
B = x.shape[0]
x = self.patch_embed(x) # (B, N, D)
cls = self.cls_token.expand(B, -1, -1) # (B, 1, D)
x = torch.cat([cls, x], dim=1) # (B, N+1, D)
x = x + self.pos_embed # 加位置编码
x = self.pos_drop(x)
for block in self.blocks:
x = block(x)
x = self.norm(x)
cls_out = x[:, 0] # 取 [CLS] 位置的输出
return self.head(cls_out)训练与推理
训练循环与普通分类模型相同。这里用 CIFAR-10 演示:
import torch.nn.functional as F
import torchvision
import torchvision.transforms as T
# 为教学简洁,把 32×32 的 CIFAR 图放大到 224×224;
# 想跑得快可以把 PATCH_SIZE 改成 4,直接用原始 32×32
transform = T.Compose([
T.Resize((224, 224)),
T.ToTensor(),
T.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
])
trainset = torchvision.datasets.CIFAR10(root='./data', train=True,
download=True, transform=transform)
trainloader = torch.utils.data.DataLoader(trainset, batch_size=64, shuffle=True)
model = ViT()
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
criterion = nn.CrossEntropyLoss()
for epoch in range(10):
total, correct = 0, 0
for images, labels in trainloader:
logits = model(images)
loss = criterion(logits, labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()
total += labels.size(0)
correct += (logits.argmax(1) == labels).sum().item()
print(f"epoch {epoch + 1}, loss {loss.item():.4f}, acc {correct / total:.4f}")需要说明的是:在 CIFAR-10 这种小数据集上从头训练,ViT 的效果会明显不如同规模的 CNN,这正是"ViT 需要大数据"的一个实例。实际使用时通常加载预训练权重(例如 timm 中的 vit_base_patch16_224)再微调,或者采用 DeiT 的蒸馏方案。
总结
回到开头的问题:Transformer 的输入是序列,ViT 的工作就是把图像转换成序列。
- Patch Embedding 用一次张量变形(reshape → permute → reshape)把 $224 \times 224 \times 3$ 的图像变成 196 个 768 维的 token,再用线性层投影到模型维度;切分、展平、投影也可以用一步卷积等价实现。
- 位置编码 补回展平过程中丢失的空间信息;实验表明 1D 可学习编码已经足够,模型可以自行学习二维结构。
- [CLS] token 提供整张图像的汇总表示,其输出接分类头。
- 编码器部分与文本 Transformer 完全同构,没有引入新组件。
ViT 证明了自注意力可以用于视觉任务,代价是 $\mathcal{O}(N^2)$ 的计算复杂度和对大规模预训练数据的依赖,这两点催生了 Swin、DeiT、MAE 等后续工作。对本站而言,它是 CLIP 与 LLaVA 的视觉编码器,是 VLM 专栏知识链上的一环。
建议的阅读顺序:本站《Transformer 详解》(注意力机制)→ 本文(图像到序列的转换)→ 《CLIP 详解》(对比学习)→ 《LLaVA 详解》(生成式多模态)。
参考
- An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale — Dosovitskiy et al., 2020。ViT 原论文。
- Training data-efficient image transformers & distillation through attention(DeiT)— Touvron et al., 2021。
- Swin Transformer: Hierarchical Vision Transformer using Shifted Windows — Liu et al., 2021。
- pytorch-image-models (timm) — 工程版 Patch Embedding(Conv2d 实现)与预训练权重的参考实现。
- 本站《Transformer 详解:从注意力机制到代码实战》— 自注意力、多头、位置编码的前置基础。