ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch实现Vision Transformer:从Patch Embedding到训练实战

PyTorch实现Vision Transformer:从Patch Embedding到训练实战 简介这是一个将Transformer模型引入图像分类任务的PyTorch实现资源面向希望在计算机视觉中应用自注意力机制的深度学习者与开发者。资源包含4个Python源文件压缩包仅6KB涵盖模型定义、CIFAR-10数据加载、训练流程等模块结构紧凑便于阅读已有5597人学习下载。内容通过Patch Embedding将图像分割为Token并实现适用于图像的位置编码与Transformer Encoder同时展示了基于交叉熵损失和Adam优化器的完整训练及验证流程。项目对Transformer原理的落地进行了清晰拆解从数据增强、随机翻转裁剪到多头注意力计算的细节均有体现代码注释较丰富便于按模块理解与调试。借助这些代码可直观理解自注意力与多头注意力的核心机制并快速迁移到其他图像分类场景中复现与改进。 拿到 transformer_pytorch_inCV.rar 这个压缩包的时候我下意识先扫了一眼文件名——transformer、pytorch、CV 三个关键词凑在一起基本就是这两年视觉领域最热的技术栈。很多人第一次接触这个组合是想把 NLP 里的 Transformer 搬到图像任务上试试效果但真到写代码的时候就懵了PyTorch 里应该怎么实现 patch embedding多头注意力怎么写才对训练时为什么总是爆显存这篇内容就是围绕这些实际动手时躲不开的问题展开的从原理到代码再到踩坑记录适合那些已经有 PyTorch 基础、想在 CV 里跑通第一个 Transformer 模型的同学。1. Transformer 在 CV 领域为什么这么火1.1 从 NLP 跨界到图像本质是感受野之争2017 年 Attention Is All You Need 提出 Transformer 之后NLP 领域几乎被它统一但 CV 这边一开始并不买账。为什么因为 CNN 靠着局部感受野和权值共享在图像任务上已经统治了好几年卷积核天然携带了很强的归纳偏置——平移不变性、局部性这些先验让 CNN 在中小规模数据集上非常省力。而 Transformer 最初没有这些先验它全靠注意力机制自己去学全局关系所以早期在 ImageNet 上用 Transformer 做分类精度一直打不过 ResNet。真正改变格局的是 ViTVision Transformer。它的思路大胆又简单把一张图切成 16x16 的 patch拉平后当成一串 token 丢进标准 Transformer encoder 里。这种做法的核心价值是模型能直接对整张图像的全局关系建模不再像 CNN 那样依赖小卷积核一点点叠加感受野。你也可以把 Transformer 理解成一种极端版的全连接注意力每一步都在看整张图而 CNN 是先看局部、再看更大范围。在数据量大到一定程度后这种全局建模能力带来的上限明显更高。1.2 绕不开的经典模型ViT、Swin、DeiT真正动手前先把几个常听到的名字理清楚。ViT 是开山之作结构最纯粹适合拿来理解原理。Swin Transformer 则是工程上更实用的一版它提出层级式特征和窗口注意力把注意力限制在局部窗口内既保留了一些类似 CNN 的多尺度金字塔结构又大幅降低了计算量。DeiT 则是针对 ViT 训练难的问题引入知识蒸馏和一系列训练技巧让 ViT 在中等规模数据集上也能训出不错的效果。如果你只是想在自定义数据集上快速出结果我建议从 Swin Transformer 或 DeiT 入手如果是为了学习原理、看清每个模块怎么拼那先从 ViT 手写开始最合适。下面这张表是我常用的选型参考模型计算量数据集需求适合场景ViT中等大一般建议百万级原理学习、大规模预训练DeiT中等小各种蒸馏技巧加持中小数据集分类Swin较高中等检测、分割、通用主干网络2. 环境搭建PyTorch 与整个项目的依赖准备2.1 PyTorch 安装与 CUDA 版本匹配很多同学在环境这步就卡住了尤其是 PyTorch 和 CUDA 版本对不上装上之后torch.cuda.is_available()永远返回 False。我个人的习惯是用 conda 单独建一个环境避免把 base 环境搞得乱七八糟。conda create -n cv_transformer python3.10 conda activate cv_transformer pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121CUDA 版本怎么选不用非得装最新的。先看你自己显卡驱动支持的最高 CUDA 版本然后往下兼容一个稳定版本就行。比如我机器上是 RTX 3090驱动最高支持 12.x选了 cu121 的 PyTorch 包。装完之后一定要验证两件事第一import 不报错第二显卡真的能用。python -c import torch; print(torch.__version__, torch.cuda.is_available(), torch.cuda.get_device_name(0))注意如果输出里torch.cuda.is_available()是 False多半是 PyTorch 版本和 CUDA Toolkit 不一致或者安装成了 CPU 版。重新装对应 CUDA 版本的包基本能解决别急着重装系统。2.2 项目目录结构与依赖清单这个压缩包展开之后内部结构通常长这样虽然不是标准答案但很值得参考transformer_pytorch_inCV/ ├── models/ # 模型定义vit.py、swin.py 等 ├── data/ # 数据集加载与预处理 ├── utils/ # 工具函数、学习率调整、可视化 ├── train.py # 训练入口 ├── test.py # 推理入口 ├── configs/ # 配置文件yaml 或 py └── requirements.txt这种把模型、数据、工具分开的组织方式最大的好处是换数据集、换模型时不用大改代码。requirements.txt 里除了 torch 和 torchvision我常用的还有这几个timm里面有很多现成的 transformer 骨干、tensorboard 或 wandb看训练曲线、einops张量重排特别好用写 attention 时能少掉一半头发。pip install timm einops tensorboard3. 核心实现用 PyTorch 手写一个 Vision Transformer3.1 Patch Embedding把图像切成小方块ViT 的第一步是把 H×W×C 的图像切成 N 个 patch每个 patch 大小通常是 16×16。切完之后每个 patch 会通过一个线性层映射成 D 维向量这一步就叫 patch embedding。你可能会想切 patch 再加线性层能不能用卷积一步搞定能。一个 kernel_size 和 stride 都等于 patch_size 的 Conv2d输出通道等于 embed_dim效果和先切再线性映射完全等价而且实现更简洁。我第一次看 ViT 源码时愣了半天没想到一个卷积就把 embedding 做完了。import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.num_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: [B, 3, 224, 224] x self.proj(x) # [B, 768, 14, 14] x x.flatten(2).transpose(1, 2) # [B, 196, 768] return xflatten 那一步很多人容易绕晕。flatten(2) 把 H 和 W 两个维度压成一个维度得到 [B, 768, 196]transpose(1, 2) 再把维度顺序换成 [B, 196, 768]。至此一张图变成了 196 个 token每个 token 是 768 维的向量跟 NLP 里一句话变成一串词向量是同一个套路。3.2 Transformer Encoder多头注意力是核心中的核心ViT 的 encoder 和原始 Transformer 基本一致核心就是多头自注意力MSA。注意力机制的本质是让每个 token 根据自己的 Query 去所有 token 的 Key 里找相关信息再按相关性加权聚合 Value。多头就是做 H 次这种注意力每次关注不同的子空间关系最后拼起来。class Attention(nn.Module): def __init__(self, dim, num_heads8, qkv_biasFalse): super().__init__() self.num_heads num_heads self.head_dim dim // num_heads self.scale self.head_dim ** -0.5 self.qkv nn.Linear(dim, dim * 3, biasqkv_bias) self.proj nn.Linear(dim, dim) 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) q, k, v qkv[0], qkv[1], qkv[2] attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) x (attn v).transpose(1, 2).reshape(B, N, C) return self.proj(x)这段代码是 ViT 的核心值得慢慢看。reshape 和 permute 的目的是把 Q、K、V 按头数拆开方便做并行矩阵运算。注意力权重算出来之后除以 scale等于 head_dim 的平方根倒数是为了防止点积结果过大导致 softmax 梯度消失。一个完整的 encoder block 还要包含 LayerNorm、MLP 和残差连接结构是Norm → Attention → 残差 → Norm → MLP → 残差的顺序。PyTorch 官方实现里习惯叫 pre-norm也就是先归一化再做注意力这个设计和原始 Transformer 的 post-norm 略有区别但 pre-norm 在深层网络中更稳训练起来也更容易收敛。3.3 给序列加上 class token 和位置编码拿一张 224×224 的图切完 patch 后我们有 196 个 token。但分类任务需要输出一个全局特征怎么把 196 个 token 汇成一个ViT 的做法很巧妙在序列最前面额外拼接一个可学习的 class token。这个 token 的初始值靠随机初始化训练过程中它会通过自注意力不断聚合整张图的信息最后拿它的输出接分类头就行。class TokenLearner(nn.Module): def __init__(self, embed_dim): super().__init__() self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, 1 196, embed_dim)) def forward(self, x): B x.shape[0] cls_token self.cls_token.expand(B, -1, -1) x torch.cat([cls_token, x], dim1) x x self.pos_embed return x位置编码在 ViT 里是直接加在 token 向量上的一个可学习参数。因为自注意力本身没有顺序概念你不加位置编码模型就分不清左上角 patch和右下角 patch的关系。这一点和 NLP 一模一样都是靠叠加位置信号来保留空间位置信息。4. 训练与推理实操从配置到跑通全流程4.1 数据预处理Transformer 对数据更挑剔同样的数据增强策略用在 CNN 上可能没事用在 ViT 上可能就欠拟合。我在 CIFAR-10 上做实验时发现 ViT 对随机裁剪、翻转这些基本增强不太感冒但 ResNet 吃这一套。原因还是归纳偏置CNN 天然假设邻近像素相关所以少量增强就能泛化Transformer 没有这个假设必须靠大量数据和强增强来教会它视觉先验。我常用的预处理配置分两档。简单档适合快速验证train_transform T.Compose([ T.Resize((224, 224)), T.RandomHorizontalFlip(), T.ToTensor(), T.Normalize(0.5, 0.5) ])完整档适合认真调模型会加上 RandAugment、Mixup、CutMix 这些更强的增强策略。经验是数据量小于 10 万张时直接用完整档增强基本不会错。另外 Normalize 的均值方差一定要按数据集算别拿 ImageNet 的乱七八糟参数硬套小数据集。4.2 训练参数配置AdamW 是默认选择学习率别贪大ViT 这一类模型的默认优化器我建议直接用 AdamW不用 SGD。AdamW 对 Transformer 这类结构稳定性的提升很明显配合 cosine 学习率衰减和线性 warmup效果最好。warmup 的作用是让模型在刚开始训练时用较小的学习率慢慢进入状态避免前几步就把位置编码或 class token 的初始化分布冲乱。optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs)学习率从多少起步ViT 在小数据集上我一般用 1e-4 到 3e-4 之间。ResNet 能承受 0.1 这种大学习率但 Transformer 不行l太大很容易直接 loss 变成 NaN 或者训完 acc 还是随机水平。Batch size 也尽量调大一点比如 64 或 128Transformer 对 batch 大小比较敏感batch 越小BN 类归一化不稳定模型越容易震荡。4.3 混合精度训练显存不够时的救命稻草如果你跑 ViT-Base8600 万参数在 224×224 上batch size 开 32一张 12G 显存的卡勉强能跑。再想加大 batch 或者输入分辨率显存就爆了。这时候 AMP自动混合精度几乎是必选项。PyTorch 的torch.cuda.amp用起来很简单gradscaler 包一层就行。scaler torch.cuda.amp.GradScaler() for images, labels in dataloader: images, labels images.cuda(), labels.cuda() with torch.cuda.amp.autocast(): outputs model(images) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()混合精度做完显存能省下 30% 以上速度也有提升。代价是某些算子精度略微下降但对视觉任务来说影响很小。如果加了 AMP 后 loss 曲线出现明显异常检查一下 loss 是否变成 inf大概率是学习率太大或者某个自定义层不支持半精度把该层的 dtype 强制为 float32 就行。4.4 推理阶段加载权重和输出可视化推理相对简单load_state_dict 时记得先处理一下键名。如果你用的模型名是 model但权重里是 module.model 之类的键名那多半是之前用了 DataParallel 保存的需要去掉前缀state_dict torch.load(best.pth) new_state_dict {k.replace(module., ): v for k, v in state_dict.items()} model.load_state_dict(new_state_dict)很多时候我们还想看一眼注意力图验证模型到底在关注图像的哪些区域。做法是把推理时某个头的 attention 权重取出来reshape 到 14×14 的分辨率再插值回原图大小做热力图。这样做的好处是能直观发现模型是不是学到了有意义的区域比如分类猫时注意力集中在猫脸还是背景。我在调 ViT 时经常靠这个定位问题如果注意力全散在背景上基本就是数据或者训练策略出了问题。5. 常见问题与排查心得5.1 训练不收敛先查这四件事训练 loss 不降是 Transformer 新手最容易遇到的坑。我自己的排查顺序是固定的。第一确认输入归一化没问题像素值范围要么在 0-1 要么在 -1 到 1别混着用。第二确认位置编码和 class token 是否都加上了有人用预训练权重时漏了位置编码的加载模型直接崩。第三降低学习率Transformer 对学习率异常敏感从 1e-4 甚至 5e-5 重新试。第四检查 loss 是否 NaN如果 NaN 出现马上想到 AMP 梯度溢出把 GradScaler 关掉试一次。5.2 显存不足梯度累积和分辨率取舍普通用户没有多卡环境16G 甚至 8G 显存跑 ViT 挺吃力的。除了 AMP梯度累积也是一个好办法。它的思路是模拟更大的 batch梯度先攒几步再更新一次参数。accum_steps 4 optimizer.zero_grad() for i, (images, labels) in enumerate(dataloader): outputs model(images) loss criterion(outputs, labels) / accum_steps loss.backward() if (i 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()如果梯度累积还不行就得降低输入分辨率。ViT 这类模型对分辨率非常敏感输入从 224 降到 160显存占用明显下降精度损失通常也能接受。我实际测试下来160×160 的 ViT-Tiny 在不少小数据集上精度只掉 1 到 2 个点但显存能省近一半。5.3 从零训练还是用预训练权重很多人问ViT 一定要用预训练权重吗我的回答是如果数据量少于 10 万张强烈建议用。ViT 从头训练在小型数据集上效果通常不如 ResNet因为它的归纳偏置太弱。解决办法是用在 ImageNet 上预训练过的权重做微调。timm 库里一条命令就能加载import timm model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes10)微调时需要注意分类头输出维度变了所以 num_classes 要改成自己的类别数然后把位置编码插值到新分辨率。如果新输入分辨率和预训练时不一样比如从 224 改成 384需要把 pos_embed 做双线性插值。这部分我用过多次处理不好会导致精度暴跌。5.4 训练慢的优化思路如果觉得训练速度太慢先别急着买新卡。有几个纯代码层面的优化点第一用 F.bmm 或 einsum 替代手写的循环注意力矩阵运算一定比循环快得多第二开启 cudnn benchmarktorch.backends.cudnn.benchmark True输入尺寸固定时能显著提速第三数据加载瓶颈的话num_workers 设置成 4 到 8pin_memoryTrue很多时候 GPU 空等的时间比你想象的多。经验收尾踩过几次坑之后我的体会是 Transformer 在 CV 里并没有想象中那么难落地但确实不能照搬 CNN 那套经验。环境上先把 PyTorch 和 CUDA 的版本锁死模型结构优先看 ViT 的手写实现训练时严格控制学习率数据规模不够就别死磕从零训练。最后再分享一个小经验当你调试 attention 可视化时如果发现每个头关注的区域都差不多说明模型容量没被充分利用这时候可以试试增大 head 数量或者提高 dropout往往会有意外收获。本文还有配套的精品资源点击获取
返回列表