免费获取学习方案
ARTICLE DETAIL

资讯详情

深耕编程基础知识与建站技术分享的一线实战洞察。

从零实现Vision Transformer:ViT原理与PyTorch代码详解

从零实现Vision Transformer:ViT原理与PyTorch代码详解 Transformer 目前是深度学习中最具影响力的序列建模架构之一它的核心思想来自 2017 年的论文 Attention Is All You Need用自注意力机制直接计算序列中任意两个位置之间的相关性不再依赖循环或者卷积逐步传递信息。Vision TransformerViT把这个架构搬到了图像任务上思路很直接先把图像切成一堆固定大小的 patch再把每个 patch 拉平并投影成一个向量也就是 Patch Embedding然后把这些向量当作 token 序列输入标准 Transformer Encoder最后用分类头输出结果。真正理解 ViT不能只停留在会调用现成库的层面还要能解释清楚一个输入图像从进入网络到输出 logits 的过程中每个张量的形状发生了什么变化每个模块为什么这么设计。这篇文章会从 Transformer 的核心机制讲起然后逐步实现 Patch Embedding、多头自注意力、Transformer Encoder Block 和完整 ViT 模型最后用 MNIST 分类任务把整条链路跑通并给出训练、验证、排错和工程化建议。1. Transformer 到底在解决什么问题1.1 RNN 和 CNN 的局限在 Transformer 出现之前序列建模最常用的是 RNN 和它的变体 LSTM、GRU。RNN 把序列元素按时间顺序逐个读入维护一个隐状态当前时刻的输出依赖上一时刻的隐状态。这种方式有两个明显问题第一长距离信息要经过很多次隐状态传播才能到达目标位置中间容易出现梯度消失或信息衰减第二时间步之间是串行计算无法像矩阵运算一样充分并行训练效率受到限制。CNN 在视觉任务上是主力。它通过局部感受野和卷积核共享权重天然带有平移不变性和局部性的归纳偏置。但是要建模整张图像的全局关系CNN 需要堆叠大量卷积层让感受野逐层扩大这导致建模远程依赖的成本比较高。而且远距离的两个像素是否相关并不一定与它们的空间距离成正比。CNN 的局部优先假设在 ImageNet 这类大数据上仍然有效但模型结构本身给全局建模带来了额外负担。Transformer 和它们最本质的区别是它不假设信息必须沿着时间顺序或空间邻接关系传递。对于输入序列中的任意两个元素Transformer 都计算一个注意力权重权重越高表示目标元素在更新自身表示时越依赖那个元素。这样全局关系在每一层都是显式建模的而且所有位置可以并行计算。1.2 自注意力机制的工作原理自注意力解决的问题是给定一个序列让序列中的每个元素都能根据全序列其他元素的信息来更新自己。可以把输入序列理解为 N 个 token每个 token 是一个向量例如形状是[N, D]。为了让每个 token 与其他 token 交互模型为每个 token 生成三个向量Query、Key、Value。Query 表示“我想找什么信息”Key 表示“我能提供什么信息”Value 表示“我实际携带的信息”。某个 token 的 Query 与所有 token 的 Key 做点积得到一个相关性分数分数经过 softmax 变成权重再用权重对所有 Value 做加权求和得到该 token 更新后的输出。写成公式就是Attention(Q, K, V) softmax(Q K^T / sqrt(d_k)) V其中d_k是每个注意力头的维度。除以sqrt(d_k)是为了避免Q K^T的数值随维度增大而变得太大导致 softmax 进入饱和区梯度变得非常小。多头注意力的做法是把 D 维的 Query、Key、Value 分别拆成 H 组每组维度为 D/H各自独立做注意力计算最后拼接起来再过一次线性投影。这样可以让模型在多个子空间里分别关注不同类型的依赖关系有的头可能关注相邻 patch有的头可能关注全局轮廓。1.3 为什么图像也要用 TransformerViT 的思路不是用卷积提取特征后再接 Transformer而是直接把图像变成 token 序列。它的动机是如果数据量足够大模型对“局部性”的人工预设并不是必需的。Transformer 可以通过注意力自己学到哪些 patch 需要相互关注。不过这里有一个重要的边界在中小规模数据集上直接从头训练 ViT效果通常不如同规模的 CNN。原因在于 CNN 的归纳偏置在数据不足时是一种保护而 Transformer 更依赖大规模数据来学习结构。所以个人项目和工业落地时常见做法是使用在大规模数据上预训练好的 ViT 权重做迁移学习而不是从随机初始化开始训练。2. ViT 的整体架构图像怎么变成 token 序列2.1 ViT 处理图像的完整流水线假设输入一张 RGB 图像形状是[B, 3, 224, 224]其中 B 是 batch size。ViT 的流水线如下图像切 patch把 224x224 的图像按 16x16 的 patch 划分得到 14x14196 个 patch。Patch Embedding把每个 patch 投影成一个 D 维向量得到[B, 196, D]的 token 序列。拼接 CLS token在序列开头添加一个特殊的可学习 token得到[B, 197, D]。加位置编码加上一个[B, 197, D]的 position embedding保持长度不变。过 L 层 Transformer Encoder每层都是多头自注意力 MLP LayerNorm 残差。取 CLS token 的输出[B, D]。分类头线性层映射到类别数得到[B, num_classes]。如果任务是目标检测或分割第 6 步会替换成其他解码器或特征图输出结构但前面第 1 到第 5 步基本相同。2.2 Patch Embedding把图像切块并投影成向量图像是一个高维数组不能直接当成一维 token 序列输入 Transformer。Patch Embedding 做的事情就是把“二维图像上的一个小块”编码成“一个向量”。具体来说一个大小为patch_size x patch_size、通道数为 C 的 patch拉平后长度是patch_size^2 * C。用一个线性层可以把这个向量投影到 embed_dim 维。实际操作中切块加投影可以用一个二维卷积一步完成卷积核大小和步长都等于 patch_size输入通道数等于 C输出通道数等于 embed_dim。卷积输出的特征图上每个位置对应原图的一个 patch且每个通道上的数值就是这个 patch 在该维度上的投影结果。例如输入[B, 3, 224, 224]patch_size16embed_dim768卷积输出就是[B, 768, 14, 14]。把它展平成[B, 768, 196]再转置成[B, 196, 768]就得到了 token 序列。这里每个 token 对应原图 16x16 的一个 patch。2.3 Position Embedding、CLS token 与分类头注意力机制对输入的顺序不敏感。如果把 patch 序列任意打乱自注意力计算结果只是行顺序变化数值不会改变。但图像的空间顺序是有意义的所以必须把位置信息注入。ViT 使用可学习的位置编码初始化后随训练一起更新。它的形状是[1, num_patches 1, embed_dim]因为前面还要拼接一个 CLS token。也有人会用正弦位置编码但 ViT 论文和工程实现大多是直接学习。CLS token 是拼接在 patch 序列开头的一个可学习向量。它没有对应的输入 patch只是作为全图信息的汇聚点。经过多层 Encoder 后CLS 位置的输出向量通过自注意力汇总了所有 patch 的信息因此可以用它来分类。关于 CLS token 和全局平均池化ViT 里两种做法都有人用。CLS token 的好处是输出与序列长度解耦而且和预训练结构一致平均池化则更显式地利用所有 patch 信息。在迁移学习中尽量保持与预训练一致的用法。3. 环境准备与项目结构3.1 环境依赖演示代码基于 Python 3.9 和 PyTorch 2.x 编写主要依赖如下依赖版本建议作用Python3.8运行环境PyTorch2.0张量计算与自动求导torchvision0.15数据集与图像变换numpy1.24间接依赖安装命令conda create -n vit-demo python3.9 -y conda activate vit-demo pip install torch torchvision如果没有 GPU安装 CPU 版即可。本文的最小演示模型很小CPU 上几个 epoch 也能跑完。如果需要 GPU 训练安装对应 CUDA 版本的 PyTorch。3.2 项目结构vit-from-scratch/ ├── vit.py # 模型定义PatchEmbed、Attention、Encoder、ViT ├── train.py # 数据加载与训练验证 └── data/ # 数据目录vit.py是核心所有模型组件都放在里面。train.py负责把 MNIST 数据加载进来创建模型执行训练和验证。3.3 演示数据集选择为了在普通机器上快速验证本文使用 MNIST 手写数字分类。MNIST 是 28x28 的单通道灰度图共 10 类。我们会把它 Resize 到 32x32patch_size4这样会得到 64 个 patch训练速度很快。如果读者想换 CIFAR-10只需要把in_channels改成 3并修改 Resize 和归一化参数想换 ImageNet 风格的图片需要调整img_size和patch_size保证两者可以整除。4. 代码手撕从零实现 Vision Transformer这一节是实现重点。这里的“手撕”不是把开源代码抄一遍而是用一个最小但完整的实现展示 ViT 的每个模块。模型定义全部放在vit.py中不使用 torchvision 或 timm 里现成的 ViT。4.1 PatchEmbed图像切块与线性投影import torch import torch.nn as nn class PatchEmbed(nn.Module): 将图像切成 patch并线性投影为 token 向量。 输入: [B, C, H, W] 输出: [B, num_patches, embed_dim] def __init__(self, img_size32, patch_size4, in_channels1, embed_dim128): super().__init__() if img_size % patch_size ! 0: raise ValueError( fimg_size{img_size} 必须能被 patch_size{patch_size} 整除 ) self.img_size img_size self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 # 卷积核大小和步长都等于 patch_size # 等价于先把每个 patch 拉平再过一个线性层 self.proj nn.Conv2d( in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size, ) def forward(self, x): B, C, H, W x.shape # 检查输入尺寸避免隐含错误在后面的位置编码阶段才暴露 if H ! self.img_size or W ! self.img_size: raise ValueError( f输入尺寸应为 {self.img_size}x{self.img_size}实际为 {H}x{W} ) # [B, embed_dim, H/patch_size, W/patch_size] x self.proj(x) # [B, embed_dim, num_patches] x x.flatten(2) # [B, num_patches, embed_dim] x x.transpose(1, 2) return x这个模块的核心是nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size)。卷积核在图像上以 patch_size 为步长滑动每个位置对应一个 patch输出通道数就是 embedding 维度。这样做的好处是切块和线性投影在同一个操作里完成计算效率高也不需要手动写torch.nn.functional.unfold。4.2 MultiHeadSelfAttention多头自注意力import math import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadSelfAttention(nn.Module): 多头自注意力。 输入: [B, N, D] 输出: [B, N, D] def __init__(self, embed_dim128, num_heads4, attn_dropout0.0): super().__init__
返回列表