资讯动态

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

发布时间:2026/9/5 21:38:38 来源:尧图企业网站定制
Vision Transformer (ViT) 原理与 PyTorch 实现从 Patch Embedding 到分类头的完整解析【免费下载链接】annotated_deep_learning_paper_implementations‍ 60 Implementations/tutorials of deep learning papers with side-by-side notes ; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), reinforcement learning (ppo, dqn), capsnet, distillation, ... 项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations本文围绕本仓库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]分类 tokenpatch 嵌入序列前拼接一个分类 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, stridepatch_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_modelTransformer 嵌入维度、patch_sizepatch 边长、in_channels输入图像通道数RGB 为 3。可学习位置嵌入对应文档中位置嵌入是一组按 patch 位置区分的向量随其他参数一起训练的描述源码见 labml_nn/transformers/vit/init.pyclass 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_gradTrue) def forward(self, x: torch.Tensor): pe self.positional_encodings[:x.shape[0]] # 取出前 patches1 个位置 return x pepositional_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.pyclass 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为隐层维度实验配置中默认 2048n_classes为类别数CIFAR-10 为 10。VisionTransformer完整前向流程VisionTransformer把上述组件串联起来labml_nn/transformers/vit/init.pyclass 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_gradTrue) # 最终 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(xx, maskNone) # 无 attention mask x x[0] # 取 [CLS] 的输出 x self.ln(x) return self.classification(x) # 输出类别 logits值得注意的实现细节编码器层复用transformer_layer是仓库通用的 TransformerLayerpre-norm 结构先 LayerNorm 再 self-attention / FFNVisionTransformer只创建一份再借助 clone_module_list 复制出n_layers份独立层无掩码图像分类任务不需要因果掩码所有层均以maskNone前向[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_size4patch 边长CIFAR-10 图像为 32x32即 8x8 64 个 patchn_hidden_classification2048分类头隐层维度n_classes10CIFAR-10 类别数模型由option(Configs.model)工厂函数_vit构建labml_nn/transformers/vit/experiment.pyreturn 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_heads8、d_model512、n_layers6、dropout0.1encoder_layer会据此自动组装出 TransformerLayer。main()中通过experiment.configs覆盖的部分默认配置labml_nn/transformers/vit/experiment.pyexperiment.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, padding4)先 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_size4时序列长度为 64 1[CLS]位置嵌入表按max_len上限预留[CLS]单 token 分类不聚合所有 patch 的输出而是专门学习一个分类 token 的表达经最终 LayerNorm 与两层 MLP 输出 logits性能依赖数据规模文档明确指出 CIFAR-10 实验效果有限是因为数据集小论文级结论来自 3 亿图像预训练 高分辨率推理时的位置嵌入插值这在使用 ViT 做模型选型时是关键前提。【免费下载链接】annotated_deep_learning_paper_implementations‍ 60 Implementations/tutorials of deep learning papers with side-by-side notes ; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), reinforcement learning (ppo, dqn), capsnet, distillation, ... 项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价