news 2026/9/9 1:49:58

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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。

真正改变格局的是 ViT(Vision 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 python=3.10 conda activate cv_transformer pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121

CUDA 版本怎么选?不用非得装最新的。先看你自己显卡驱动支持的最高 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 tensorboard

3. 核心实现:用 PyTorch 手写一个 Vision Transformer

3.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_size=224, patch_size=16, in_chans=3, embed_dim=768): super().__init__() self.num_patches = (img_size // patch_size) ** 2 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_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 x

flatten 那一步很多人容易绕晕。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_heads=8, qkv_bias=False): 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, bias=qkv_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], dim=1) 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(), lr=1e-4, weight_decay=0.05) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)

学习率从多少起步?ViT 在小数据集上我一般用 1e-4 到 3e-4 之间。ResNet 能承受 0.1 这种大学习率,但 Transformer 不行,l太大很容易直接 loss 变成 NaN 或者训完 acc 还是随机水平。Batch size 也尽量调大一点,比如 64 或 128,Transformer 对 batch 大小比较敏感,batch 越小,BN 类归一化不稳定,模型越容易震荡。

4.3 混合精度训练:显存不够时的救命稻草

如果你跑 ViT-Base(8600 万参数)在 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", pretrained=True, num_classes=10)

微调时需要注意,分类头输出维度变了,所以 num_classes 要改成自己的类别数,然后把位置编码插值到新分辨率。如果新输入分辨率和预训练时不一样,比如从 224 改成 384,需要把 pos_embed 做双线性插值。这部分我用过多次,处理不好会导致精度暴跌。

5.4 训练慢的优化思路

如果觉得训练速度太慢,先别急着买新卡。有几个纯代码层面的优化点:第一,用 F.bmm 或 einsum 替代手写的循环注意力,矩阵运算一定比循环快得多;第二,开启 cudnn benchmark,torch.backends.cudnn.benchmark = True,输入尺寸固定时能显著提速;第三,数据加载瓶颈的话,num_workers 设置成 4 到 8,pin_memory=True,很多时候 GPU 空等的时间比你想象的多。

经验收尾

踩过几次坑之后,我的体会是 Transformer 在 CV 里并没有想象中那么难落地,但确实不能照搬 CNN 那套经验。环境上先把 PyTorch 和 CUDA 的版本锁死,模型结构优先看 ViT 的手写实现,训练时严格控制学习率,数据规模不够就别死磕从零训练。最后再分享一个小经验:当你调试 attention 可视化时,如果发现每个头关注的区域都差不多,说明模型容量没被充分利用,这时候可以试试增大 head 数量或者提高 dropout,往往会有意外收获。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/9 1:49:53

WPF仿Word富文本编辑器:基于RichTextBox与FlowDocument的完整实现

简介:面向WPF开发者的富文本编辑器开源项目,仿Word风格,适合需要构建行业专用文档编辑工具或学习自定义控件封装的中高级开发者。压缩包共289个文件,主体为76个C#源代码、10个XAML界面标记以及15个BAML编译资源,同时搭…

作者头像 李华
网站建设 2026/9/9 1:49:29

多物理场耦合仿真中的电磁场理论核心与建模要点

做多物理场耦合仿真这些年,我发现自己被问得最多的一个问题不是“怎么设置求解器”,也不是“网格怎么划分”,而是“电磁场理论到底要学到什么程度才能不做错模型”。这问题很实在,因为多物理场耦合仿真里的电磁场模块,…

作者头像 李华
网站建设 2026/9/9 1:49:15

TensorFlow实现SRCNN:图像超分入门实战与踩坑指南

简介:这是使用TensorFlow实现经典图像超分辨率算法SRCNN的完整工程代码,适合正在学习深度学习图像复原、或需要在TensorFlow环境中复现论文实验的研究者与开发者。工程共包含308个文件,压缩包约27.72MB;其中302张BMP格式图像构成训…

作者头像 李华
网站建设 2026/9/9 1:46:03

Dify Chatflow vs Workflow:选型逻辑与实战搭建指南

在Dify里新建应用的时候,平台会让你在Chatflow和Workflow之间做一个看起来简单、实际上很关键的选择。我在带团队做智能客服和知识库问答项目时,几乎每次都要跟新同事解释一遍这两个东西到底差在哪里:为什么客服机器人必须用Chatflow&#xf…

作者头像 李华
网站建设 2026/9/9 1:44:56

UTF-8与GBK编码转换工具:乱码问题排查与批量处理指南

简介:这是一款面向开发者的UTF-8编码转换小工具,支持.c、.h、.cpp、.hpp、.bat、.java等常见源码与脚本文件格式,可批量统一文件编码,并预留扩展接口,只需调整suffix判断条件即可覆盖更多类型。压缩包内共2个文件&…

作者头像 李华