首页
/ Vision Transformer (ViT) 原理与 PyTorch 实现:从 Patch Embedding 到分类头的完整解析

Vision Transformer (ViT) 原理与 PyTorch 实现:从 Patch Embedding 到分类头的完整解析

2026-09-04 20:37:48作者:冯梦姬Eddie

本文围绕本仓库(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_embtorch.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=8d_model=512n_layers=6dropout=0.1encoder_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)、RandomHorizontalFlipToTensorNormalize((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 在本仓库实现中的几个要点:

  1. 零卷积结构:唯一带"卷积"痕迹的 PatchEmbeddings 仅用于一次性完成 patch 切分 + 线性投影,此后全是标准 Transformer 层(self-attention + FFN 的 pre-norm 编码器,见 labml_nn/transformers/models.py);
  2. 序列长度由 patch 数决定:32x32 图像、patch_size=4 时序列长度为 64 + 1([CLS]),位置嵌入表按 max_len 上限预留;
  3. [CLS] 单 token 分类:不聚合所有 patch 的输出,而是专门学习一个分类 token 的表达,经最终 LayerNorm 与两层 MLP 输出 logits;
  4. 性能依赖数据规模:文档明确指出 CIFAR-10 实验效果有限是因为数据集小,论文级结论来自 3 亿图像预训练 + 高分辨率推理时的位置嵌入插值,这在使用 ViT 做模型选型时是关键前提。
登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.12 K
2.72 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
903
1.82 K
docsdocs
暂无描述
Markdown
888
5.78 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
854
1.34 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
527
590
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.51 K
1.01 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.33 K
1.45 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
540
384
flutter_flutterflutter_flutter
本仓库是 Flutter SDK 与 Flutter Engine 的 OpenHarmony 适配版本,由 CPF-Flutter 团队维护。开发者可使用熟悉的 Flutter 技术栈开发 OpenHarmony 应用,3.35.7 及以后的适配版本可基于本仓库源码构建支持 OpenHarmony 的 Flutter Engine。
Dart
1.17 K
341