Vision Transformer (ViT) 原理与 PyTorch 实现:从 Patch Embedding 到分类头的完整解析
本文围绕本仓库(annotated_deep_learning_paper_implementations)中 ViT 教程文档 labml_nn/transformers/vit/readme.md 展开,系统讲解 ViT 如何用纯 Transformer 处理图像(无任何卷积结构):包括 patch 切分与线性变换生成嵌入、[CLS] 分类 token、可学习位置嵌入、MLP 分类头,以及配套的 CIFAR-10 实验配置与运行方式。读完后可完整理解 labml_nn/transformers/vit/init.py 的实现细节,并能独立运行 labml_nn/transformers/vit/experiment.py 复现实验。
ViT 概览:把图像当作"句子"
该实现出自论文《An Image Is Worth 16x16 Words: Transformers For Image Recognition At Scale》(论文原文 PDF 收录于 papers/vit.pdf)。其核心思想是:
- 纯 Transformer 处理图像:ViT 不含任何卷积层,直接将标准 Transformer 编码器应用到图像上;
- Patch 切分:图像被切成固定大小的 patch,对每个 patch 展平后的像素值做一次线性变换,得到 patch 嵌入;
[CLS]分类 token:patch 嵌入序列前拼接一个分类 token[CLS],将其最终编码经 MLP 即可得到图像类别的 logits;- 可学习位置嵌入:patch 嵌入本身不携带"这块 patch 来自图像哪个位置"的信息,因此要加上按 patch 位置区分的可学习位置嵌入,并通过梯度下降与其他参数一起训练;
- 大数据集预训练:ViT 在大数据集上预训练时表现良好。论文建议先用 MLP 分类头预训练,微调时只保留单个线性层。在 3 亿图像规模的数据集上预训练后,ViT 超过了当时的 SOTA;推理时还可以使用更高分辨率的图像(patch 大小不变),新 patch 位置的位置嵌入通过对已有位置嵌入插值得到。
仓库同时提供了一个在 CIFAR-10 上训练 ViT 的简单实验 labml_nn/transformers/vit/experiment.py——由于 CIFAR-10 数据规模小,它并不能很好地发挥 ViT 的能力,但胜在简单、任何人都可以直接运行来把玩 ViT。
Patch 嵌入:用卷积实现逐 patch 线性变换
论文中的做法是"把图像切成相同大小的 patch,再对每个 patch 的展平像素做线性变换"。源码在 labml_nn/transformers/vit/init.py 中用了一个等价的技巧来实现:
class PatchEmbeddings(nn.Module):
def __init__(self, d_model: int, patch_size: int, in_channels: int):
# 卷积核大小和步长都等于 patch_size
self.conv = nn.Conv2d(in_channels, d_model, patch_size, stride=patch_size)
def forward(self, x: torch.Tensor):
x = self.conv(x) # [bs, d_model, h, w]
bs, c, h, w = x.shape
x = x.permute(2, 3, 0, 1) # -> [h, w, bs, d_model]
x = x.view(h * w, bs, c) # -> [patches, bs, d_model]
return x
关键点(见 labml_nn/transformers/vit/init.py 注释):
- 一个 kernel size = stride = patch_size 的卷积,在数学上等价于"把图像切成不重叠的 patch,再对每个 patch 做线性变换",因此源码选择卷积这一更简洁的实现方式;
- 输入
x形状为[batch_size, channels, height, width],输出形状为[patches, batch_size, d_model],即 patch 数在第一维(注意这与 NLP 中 token 在 batch 维的常规布局不同,这是为了便于与可学习位置嵌入直接相加); - 三个构造参数:
d_model(Transformer 嵌入维度)、patch_size(patch 边长)、in_channels(输入图像通道数,RGB 为 3)。
可学习位置嵌入
对应文档中"位置嵌入是一组按 patch 位置区分的向量,随其他参数一起训练"的描述,源码见 labml_nn/transformers/vit/init.py:
class LearnedPositionalEmbeddings(nn.Module):
def __init__(self, d_model: int, max_len: int = 5_000):
# 每个位置对应一个可学习向量
self.positional_encodings = nn.Parameter(torch.zeros(max_len, 1, d_model), requires_grad=True)
def forward(self, x: torch.Tensor):
pe = self.positional_encodings[:x.shape[0]] # 取出前 patches+1 个位置
return x + pe
positional_encodings是形状[max_len, 1, d_model]的可训练参数(max_len默认 5000,即最多支持 5000 个位置);- 前向时按序列长度切片取出对应位置向量,直接与 patch 嵌入相加。注意此处序列已经包含了
[CLS]token(见下文),因此实际上[CLS]也会占据"位置 0"的嵌入; - 与仓库中 Transformer 文本模型使用的固定正弦位置编码(见 labml_nn/transformers/positional_encoding.py)不同,ViT 采用完全可学习的位置嵌入——这也是 ViT 支持"高分辨率推理时插值位置嵌入"做法的基础。
[CLS] Token 与 MLP 分类头
分类头是一个两层 MLP,输入为 [CLS] token 的 Transformer 编码,见 labml_nn/transformers/vit/init.py:
class ClassificationHead(nn.Module):
def __init__(self, d_model: int, n_hidden: int, n_classes: int):
self.linear1 = nn.Linear(d_model, n_hidden)
self.act = nn.ReLU()
self.linear2 = nn.Linear(n_hidden, n_classes)
def forward(self, x: torch.Tensor):
x = self.act(self.linear1(x))
return self.linear2(x)
参数含义:d_model 为嵌入维度,n_hidden 为隐层维度(实验配置中默认 2048),n_classes 为类别数(CIFAR-10 为 10)。
VisionTransformer:完整前向流程
VisionTransformer 把上述组件串联起来(labml_nn/transformers/vit/init.py):
class VisionTransformer(nn.Module):
def __init__(self, transformer_layer: TransformerLayer, n_layers: int,
patch_emb: PatchEmbeddings, pos_emb: LearnedPositionalEmbeddings,
classification: ClassificationHead):
self.patch_emb = patch_emb
self.pos_emb = pos_emb
self.classification = classification
# 复制 n_layers 份 transformer 层
self.transformer_layers = clone_module_list(transformer_layer, n_layers)
# [CLS] token 嵌入(1 个位置,广播到整个 batch)
self.cls_token_emb = nn.Parameter(torch.randn(1, 1, transformer_layer.size), requires_grad=True)
# 最终 LayerNorm
self.ln = nn.LayerNorm([transformer_layer.size])
def forward(self, x: torch.Tensor):
x = self.patch_emb(x) # [patches, bs, d_model]
cls_token_emb = self.cls_token_emb.expand(-1, x.shape[1], -1)
x = torch.cat([cls_token_emb, x]) # [CLS] 放在序列最前面
x = self.pos_emb(x) # 加可学习位置嵌入
for layer in self.transformer_layers:
x = layer(x=x, mask=None) # 无 attention mask
x = x[0] # 取 [CLS] 的输出
x = self.ln(x)
return self.classification(x) # 输出类别 logits
值得注意的实现细节:
- 编码器层复用:
transformer_layer是仓库通用的 TransformerLayer(pre-norm 结构:先 LayerNorm 再 self-attention / FFN),VisionTransformer只创建一份,再借助 clone_module_list 复制出n_layers份独立层; - 无掩码:图像分类任务不需要因果掩码,所有层均以
mask=None前向; [CLS]初始化:cls_token_emb用torch.randn初始化(而位置嵌入用全零初始化),形状[1, 1, d_model]通过expand广播到整个 batch;- 取第一个位置的输出:因为
[CLS]被拼接在序列最前面,x[0]即其编码,再经 LayerNorm 和 MLP 头输出n_classes维 logits。
CIFAR-10 实验:配置与运行
实验脚本 labml_nn/transformers/vit/experiment.py 基于 labml 实验框架,配置文件 Configs 继承自 CIFAR10Configs,后者又继承自 MNISTConfigs 提供的通用训练循环(模型前向 → CrossEntropyLoss → 记录 loss/accuracy → loss.backward() → optimizer.step(),见 labml_nn/experiments/mnist.py)。
ViT 特有的三个配置项(labml_nn/transformers/vit/experiment.py):
| 配置项 | 默认值 | 含义 |
|---|---|---|
patch_size |
4 | patch 边长(CIFAR-10 图像为 32x32,即 8x8 = 64 个 patch) |
n_hidden_classification |
2048 | 分类头隐层维度 |
n_classes |
10 | CIFAR-10 类别数 |
模型由 @option(Configs.model) 工厂函数 _vit 构建(labml_nn/transformers/vit/experiment.py):
return VisionTransformer(c.transformer.encoder_layer, c.transformer.n_layers,
PatchEmbeddings(d_model, c.patch_size, 3),
LearnedPositionalEmbeddings(d_model),
ClassificationHead(d_model, c.n_hidden_classification, c.n_classes)).to(c.device)
其中 c.transformer 是仓库通用的 TransformerConfigs,关键默认值为 n_heads=8、d_model=512、n_layers=6、dropout=0.1,encoder_layer 会据此自动组装出 TransformerLayer。
main() 中通过 experiment.configs 覆盖的部分默认配置(labml_nn/transformers/vit/experiment.py):
experiment.configs(conf, {
# 优化器
'optimizer.optimizer': 'Adam',
'optimizer.learning_rate': 2.5e-4,
# Transformer 嵌入维度
'transformer.d_model': 512,
# 训练轮数与 batch size
'epochs': 32,
'train_batch_size': 64,
# 训练集使用增强,验证集不增强
'train_dataset': 'cifar10_train_augmented',
'valid_dataset': 'cifar10_valid_no_augment',
})
数据增强的具体定义在 labml_nn/experiments/cifar10.py:训练集依次执行 RandomCrop(32, padding=4)(先 4 像素 padding 再随机裁回 32x32)、RandomHorizontalFlip、ToTensor、Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5));验证集只做 ToTensor + 同样的归一化,不做增强。模型通过 experiment.add_pytorch_models({'model': conf.model}) 注册以支持保存/加载,with experiment.start(): conf.run() 启动训练循环。
运行方式
前提:已安装 labml 实验框架及仓库依赖(见 requirements.txt)。由于实验使用了 labml 的 lab.get_data_path(),数据集(CIFAR-10)会自动下载。典型运行方式:
# 本地运行
python -m labml_nn.transformers.vit.experiment
# 或指定数据目录
labml --data /path/to/data run -m labml_nn.transformers.vit.experiment
训练产物(实验配置、模型 checkpoint、指标曲线)会写入当前 labml 项目的 experiments/ 目录。
小结:ViT 与 CNN 的关键区别
结合本文档与源码,可以总结 ViT 在本仓库实现中的几个要点:
- 零卷积结构:唯一带"卷积"痕迹的
PatchEmbeddings仅用于一次性完成 patch 切分 + 线性投影,此后全是标准 Transformer 层(self-attention + FFN 的 pre-norm 编码器,见 labml_nn/transformers/models.py); - 序列长度由 patch 数决定:32x32 图像、
patch_size=4时序列长度为 64 + 1([CLS]),位置嵌入表按max_len上限预留; [CLS]单 token 分类:不聚合所有 patch 的输出,而是专门学习一个分类 token 的表达,经最终 LayerNorm 与两层 MLP 输出 logits;- 性能依赖数据规模:文档明确指出 CIFAR-10 实验效果有限是因为数据集小,论文级结论来自 3 亿图像预训练 + 高分辨率推理时的位置嵌入插值,这在使用 ViT 做模型选型时是关键前提。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0622
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00