简介:这份Swin-Transformer实战项目完整打通了图像识别任务从数据准备到训练推理的闭环,适合希望掌握Vision Transformer落地流程的算法初学者、研究生及竞赛选手。项目内置关键词图像采集脚本,能够按需批量下载图片,并通过代码自动排查损坏文件、划分训练集与测试集,同时生成符合模型训练要求的固定目录格式,显著降低自定义数据集的门槛;从数据获取到模型预测形成标准化管线,方便二次改造与复用。资源共1375个文件,压缩包约723MB,主体为1200个JPEG及90个PNG、35个WebP图像样本,另有14个Python脚本、类别JSON、模型权重PTH、说明文档等支撑文件,目录结构清晰。已有385人学习下载。训练阶段只需修改lr、epochs等超参数,类别文件与网络输出个数均由代码自动生成;预测脚本能批量推理inference文件夹下的全部图片,适合直接迁移到实际图像分类场景。 如果你的工作流里一直用 ResNet 做图像分类,那我建议你找个机会完整跑一遍 Swin-Transformer 的项目。我第一次在真实数据上把 ResNet 换成一个层级式 Transformer 时,最明显的感觉不是“涨了几个点”,而是整个调参逻辑、数据组织方式、显存预估思路都要重新梳理。这次我以一个宠物图片分类项目为起点,从获取关键词数据集、清洗数据、组织目录,到用 Swin-Transformer 微调、训练、评估,完整记录了一整套可复现的做法。
这篇文章不是模型原理的堆砌,而是从“我想做一个真实图像识别项目”这个需求出发,把每一步怎么决策、为什么这么选、踩过什么坑都写清楚。适合刚跑通 PyTorch 基础教程、想上手 Transformer 类视觉模型的读者,也适合已经用过 CNN、想对比感受一下 Swin 和传统卷积网络差异的从业者。
1. 为什么这个项目我选 Swin-Transformer 而不是 CNN
1.1 从 CNN 到 Transformer:图像识别的路线变化
图像识别这些年走得很快。ResNet 统治了相当长时间,靠的是卷积的局部感受野和层级化特征。ViT 出现后,大家发现把图像切成 Patch 序列,丢给 Transformer 做全局自注意力,也能在足够大的数据集上训练出非常好的效果,甚至在很多任务上超过 CNN。
但 ViT 有个现实问题:全局注意力计算量随输入分辨率平方增长。你想用 384×384 甚至更大分辨率训练,显存压力非常大,普通显卡基本吃不消。Swin-Transformer 走的是另一条路线:保留 Transformer 的表达能力,同时把注意力限制在窗口内部,并通过 Patch Merging 构建出类似 CNN 的层级特征。也就是说,它既具备 Transformer 的建模上限,又保留了卷积网络那种“浅层细节、深层语义”的优雅结构。
这直接决定了它在实际项目里的可迁移性。图像分类只是起点,很多检测和分割模型也把 Swin 当骨干网络,比如 Mask R-CNN、Cascade R-CNN、UperNet 都有基于 Swin 的版本。如果你想从分类延伸到更复杂的视觉任务,先在 Swin 上把训练流程跑熟,后面换任务不会太痛苦。
1.2 Swin 的“层级化+窗口注意力”到底解决了什么
Swin 这个名字来自 Shifted Window,核心就是两招:
第一招,局部窗口注意力。把一张图划分成多个 7×7 的小窗口,每个窗口内部做自注意力。普通 ViT 是对整张图算注意力,计算量随 H×W 增长很快;窗口注意力把每个 token 的注意力范围限制在窗口内,计算量只随图像尺寸线性增长。这是它“跑得动”的关键。
第二招,层级化特征。Swin-T 的基本配置是 C=96,四个 Stage 输出的特征图分辨率分别是 H/4、H/8、H/16、H/32,通道数逐层翻倍。这和 ResNet 的 C2 到 C5 很相似,所以它能直接替换 CNN 主干网络,配合 FPN、PAFPN 这类结构做密集预测。
窗口注意力有个明显的缺陷:窗口之间信息不流通。如果一直隔离,区域 A 的 token 永远看不到区域 B 的信息,全局建模能力就没了。Swin 的做法是交替使用 W-MSA(规则窗口)和 SW-MSA(移动窗口),移动一个窗口大小后,原本不相邻的区域会发生交互。这种交替设计就是“Shifted Window”的核心动机。
理解了这两点,你训练时看到 Loss、准确率的变化,才会知道模型内部在做什么。比如输入分辨率变化时,窗口数量变化,但窗口内计算逻辑不变。
2. 关键词数据集的获取与清洗:工作量的大头在这里
很多人以为训练模型最耗时的是训练过程,实际情况恰恰相反。我第一次做这个项目,光整理数据集就占了一半时间,而且这部分质量直接决定了模型上限。数据不好,换再强的模型也白搭。
2.1 关键词数据集到底是什么含义
所谓“关键词数据集”,通俗讲就是按分类关键词去组织图片。比如项目要做猫、狗、鸟三类识别,那数据集里就是 cat、dog、bird 三个类别目录,每类下面放对应的图片。关键词既是文件夹名,也是类别标签,后续做训练集、验证集划分都靠这个结构。
这里有个容易忽略的问题:关键词代表的是“人理解的类别”,但图片里可能存在背景干扰、多物体共现、相似物种等问题。比如“bird”类图片里,有的鸟占画面比例很小,有的图片里同时有猫和鸟,这类样本会导致模型学到错误关联。所以在获取数据前,想清楚类别定义很重要。我当时给三类各准备了 800 张左右图片,保证类别平衡。类别不平衡会带来很多麻烦,后面训练阶段还得专门处理,没必要一开始就给自己挖坑。
2.2 我用的数据获取方案
获取方案我建议按项目目的分两条路走:
- 快速验证思路:直接使用公开学术数据集。比如 CIFAR-10、Oxford 102 Flowers、Oxford-IIIT Pets、Food-101。这些数据集经过清洗,划分规范,适合先跑通模型、验证训练配置。我最初用 Oxford-IIIT Pets 做了个 37 类宠物分类,效果很直观。
- 贴近真实业务:按自己的关键词自建数据集。可以通过公开图片搜索接口、开源图片平台、或者自己拍摄采集。这里必须强调一点:自建数据集只建议用于个人学习交流,实际商用前要确认图片授权、肖像权、商标权等合规问题。公开接口有配额限制,也不建议绕过规则做大规模抓取。
我当时用公开搜索接口按关键词保存图片,每类收集了 1000 多张,其中一部分被清洗掉了。整个过程其实就是写一个脚本:给关键词,请求搜索接口,下载前 N 张图片,按类别存入目录。代码逻辑不复杂,但注意加延时、失败重试、超时跳过,避免给服务器太大压力。
2.3 清洗策略:去重、去模糊、人工抽检
采集下来的数据必须清洗,这一步别偷懒。我的清洗流程分四步:
- 去重:图片宽高相同、直方图相似、感知哈希一致的都可能重复。重复样本会让训练集和验证集信息重叠,验证结果虚高,部署后实际效果下滑。
- 去模糊:用 OpenCV 的 Laplacian 算子计算图片清晰度,方差过低的直接删除。
import cv2 import numpy as np def is_blurry(image_path, threshold=100): img = cv2.imread(image_path) gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY) laplacian_var = cv2.Laplacian(gray, cv2.CV_64F).var() return laplacian_var < threshold- 去无关图:有的搜索结果会混入文字、logo、漫画图,需要过滤。可以按图片尺寸比例筛掉明显异常的文件,再人工抽样确认。
- 人工抽检:每类随机抽 30 到 50 张图,快速翻一遍,把明显放错的样本移走。
清洗后的数据量通常会缩水 10% 到 20%,这很正常。宁可少而精,也不要多而杂。我当时实际单类保留 800 张左右,整体训练效果比原来 1000 张混杂数据更好。
3. 训练前的数据准备:目录约定、划分与增强
数据清洗完,接下来要把图片组织成 PyTorch 方便加载的结构,并设计训练用的增强策略。这一节做不好,后面训练脚本写得再漂亮也很难收敛。
3.1 ImageFolder 目录结构
PyTorch 的torchvision.datasets.ImageFolder可以直接按目录读取图片,目录名自动成为类别名。我的目录结构如下:
data/ ├── train/ │ ├── cat/ │ ├── dog/ │ └── bird/ └── val/ ├── cat/ ├── dog/ └── bird/划分时注意随机打乱,避免同一来源的图片全部落在训练集或验证集。我当时写了个小脚本,按 8:2 比例随机分配图片,同时保证每个类别内的划分比例一致。
3.2 增强 Pipeline:为什么顺序不能乱
图像增强不是随便写几个 Transform 就完事,顺序很关键。我的训练增强配置长这样:
train_transform = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.8, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.4), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) val_transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])为什么顺序不能乱?RandomResizedCrop必须在ToTensor之前,因为它是针对 PIL 图像做的几何变换;ColorJitter最好放在几何变换后面,避免先调色再裁剪导致裁剪区域和颜色变化叠加出奇怪的分布。Normalize必须放到最后,因为它是把 0 到 1 的像素值按 ImageNet 均值和标准差做标准化,顺序反了数值分布就错了。
这里用的是 ImageNet 数据集的均值和标准差,这是迁移学习默认的标准化参数。即便你的数据集不是 ImageNet,只要使用在 ImageNet 上预训练的权重,训练和验证阶段就必须用同一套标准化参数,否则模型看到的数据分布跟预训练时不一致,效果会明显下降。
3.3 数据加载器、归一化参数与 class 映射
加载器用 PyTorch 的DataLoader就够。几个参数需要实际调试:
batch_size:我建议从 32 开始,显存不够再降到 16。Swin-T 在 224×224 分辨率下,batch size 32 大概要 10GB 到 12GB 显存,显存不够可以用混合精度。num_workers:建议设成 CPU 核心数的一半,太小会拖慢数据加载,太大反而增加调度开销。pin_memory=True:当使用 GPU 训练时,这个参数能减少数据从 CPU 拷贝到 GPU 的耗时。
ImageFolder会自动按字母排序生成类别索引,比如 bird=0, cat=1, dog=2。训练结束做推理时,要保存一份class_to_idx映射关系,否则部署时不知道哪个数字对应哪个类别。
4. Swin-Transformer 核心机制拆解:看懂你正在训练的模型
用timm一行代码就能加载 Swin-T 预训练模型,但如果你不知道模型内部是怎么把图片变成特征的,遇到 Loss 发散、验证集精度异常这类问题就很难定位。我拆几个关键环节。
4.1 Patch Embedding 与 Patch Merging
Swin 先把图片切成 4×4 的 Patch,每个 Patch 展平成 16 维像素向量,再通过线性层映射到 96 维嵌入空间(Swin-T 配置)。这一步在代码里通常用卷积实现:一个 kernel_size=4, stride=4 的卷积层,直接完成切 Patch 和嵌入。
Patch Merging 类似卷积网络里的降采样。每个 Stage 结束,把 2×2 范围内的 4 个 Patch 合并成一个,分辨率减半,通道数翻倍。这就是为什么 Stage 1 输出 56×56,Stage 2 变成 28×28,通道数从 96 翻到 192。整个过程让特征图从高分辨率、低语义逐渐过渡到低分辨率、高语义。
4.2 窗口注意力(W-MSA)和滑动窗口(SW-MSA)
W-MSA 是 Swin 和 ViT 最核心的区别。ViT 直接对整张图的 token 序列做全局自注意力;Swin 把 56×56 的特征图分成 8×8 个 7×7 的窗口,每个窗口单独做注意力。
这里有个细节:窗口数量怎么算?输入 224×224,经过 4 倍下采样变成 56×56。窗口大小默认 7,那每行就是 56/7=8 个窗口,一共 8×8=64 个窗口。每个窗口内 49 个 token 互相计算注意力,计算量远小于 56×56=3136 个 token 的全两两计算。
但窗口独立计算导致信息隔绝,于是 Swin 设计了 SW-MSA。下一个 Stage 开始时,把窗口向右下方向移动 3 个位置(窗口大小的一半),重新划分窗口。原来窗口边缘的 token 现在会跑到新窗口内部,从而让相邻区域发生信息交互。为了让移动窗口后的计算仍然高效,论文里还用了 cyclic shift 和 mask,把不规则窗口通过位移拼成规则窗口,这部分在实际推理时不用你手动处理,但理解原理能帮你明白为什么 Swin 的 FLOPs 计算和参数量比较“友好”。
4.3 相对位置编码与整体组件
Swin 没有用 ViT 那种绝对位置编码,而是引入一个可学习的相对位置偏置表。窗口尺寸为 M(默认 7),偏置表大小是 (2M−1)×(2M−1) = 13×13。每个注意力头在计算输出时,会把相对位置索引对应的偏置加到注意力得分上。这样模型能学到“左边到右边”“上边到下边”这种相对位置关系,且对输入尺寸变化有一定容忍度。
整体结构还包括 LayerNorm、MLP、残差连接,以及 Stage 之间交替的 W-MSA 和 SW-MSA。你不需要手写实现,timm 和官方代码都很成熟,但理解这些东西后,调window_size、patch_size这些参数时心里才有底。比如你换了一个更大的窗口,参数量和计算量会上升,但模型对长距离依赖的建模能力也会更强,这是一个需要权衡的点。
5. 训练配置与完整流程:从环境搭建到 Loss 收敛
5.1 环境与依赖
我的训练环境是 PyTorch 2.x + CUDA 11.8 + timm 0.9.x。Python 版本 3.9。核心依赖如下:
pip install torch torchvision timm tensorboard模型加载用 timm 的 Swin-T 小型版本:
import timm model = timm.create_model( 'swin_tiny_patch4_window7_224', pretrained=True, num_classes=3 )选择swin_tiny是因为它对单卡机器比较友好,Swin-B 或 Swin-L 参数量大得多,显存和训练时间都会明显增加。第一次跑通流程,完成比什么都重要。
5.2 超参数怎么来的
超参数不是拍脑袋定的,Swin 官方在 ImageNet 上公布了一套比较可靠的配置,迁移学习时直接借用最方便。我的配置如下:
| 参数 | 数值 | 说明 |
|---|---|---|
| 输入分辨率 | 224×224 | Swin-T 默认尺寸 |
| Batch Size | 32 | 根据显存调整 |
| Optimizer | AdamW | 权重衰减独立于梯度更新 |
| Learning Rate | 5e-5 | 迁移学习偏保守 |
| Weight Decay | 0.05 | Swin 官方配置 |
| Warmup Epochs | 5 | 前 5 轮线性上升 |
| Scheduler | Cosine Annealing | 配合 warmup 使用 |
| Epochs | 50 | 迁移学习足够 |
| Label Smoothing | 0.1 | 缓解过拟合 |
| Mixup / CutMix | 0.8 / 1.0 | 数据混合增强 |
为什么用 AdamW?因为 Adam 的权重衰减实现方式会把衰减项错误地作用在动量项上,AdamW 把权重衰减从梯度更新中解耦,能有效改善正则化效果。为什么必须 warmup?Transformer 类模型在训练初期,学习率过大会导致注意力矩阵不稳定,损失函数容易暴涨,warmup 让优化器先用很小的学习率站稳脚跟,再逐步加速。
5.3 训练循环代码与日志
训练循环本身很常规,关键是每一轮都要记录 Loss 和验证集准确率,方便后面判断收敛状态。
import torch import torch.nn as nn from torch.cuda.amp import autocast, GradScaler device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') criterion = nn.CrossEntropyLoss(label_smoothing=0.1) optimizer = torch.optim.AdamW(model.parameters(), lr=5e-5, weight_decay=0.05) scaler = GradScaler() for epoch in range(epochs): model.train() running_loss = 0.0 for images, labels in train_loader: images, labels = images.to(device), labels.to(device) with autocast(): outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() running_loss += loss.item() # 验证 model.eval() correct = 0 total = 0 with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item() val_acc = 100.0 * correct / total print(f'Epoch {epoch+1}: loss={running_loss/len(train_loader):.4f}, val_acc={val_acc:.2f}%')这里用了混合精度(AMP)。Swin 这种模型在 FP16 下训练速度提升明显,显存占用也降低不少。GradScaler 是为了防止梯度下溢,如果前向计算开启了autocast,就必须配套使用GradScaler,否则半精度梯度的最小值范围不够,参数更新可能失效。
5.4 训练和推理的显存差异
这个问题我一直觉得值得单独拎出来讲,因为很多人估算显存时根本不区分训练和推理。同样一个 Swin-T 模型,224×224 输入,推理时显存占用可能只有 3GB 到 4GB,训练时却要 10GB 以上。原因在于训练要多保存两类东西:前向传播时每一层的激活值,以及反向传播时计算的梯度。激活值在反向传播用完之前不能释放,层数越深、batch size 越大,这部分显存占用越夸张。推理只需要前向,特征用完就丢,自然省显存。
所以配置 GPU 时,先想清楚是训练还是推理。训练需要按“激活值 + 梯度 + 模型参数 + 优化器状态”来估算;推理只需要按模型参数和单样本前向激活来估算。当时我用单张 16GB 显存的卡训练 Swin-T,batch size 32 加上 AMP 刚好能放下,验证集推理则毫无压力。
6. 从 Loss 曲线到混淆矩阵:我如何判断模型真的收敛了
训练结束不代表项目完成,你得能解释模型为什么饱和了、哪些类别容易混、验证集上的准确率是否可信。
6.1 验证集准确率:基本但重要
我的项目最终验证集准确率在 94% 左右,三类宠物分类任务不算特别难,但也没有简单到随便就上 90%。只看准确率不够,我还要看每类的表现。比如“猫”这一类准确率低,而“狗”高,这往往说明猫类图片里光线差异大或者背景干扰多。
6.2 混淆矩阵和分类报告
用sklearn可以快速生成混淆矩阵:
from sklearn.metrics import confusion_matrix, classification_report import numpy as np cm = confusion_matrix(all_labels, all_preds) print(classification_report(all_labels, all_preds, target_names=['bird', 'cat', 'dog']))分类报告里的 precision、recall、f1-score 比单纯准确率更有信息量。如果某类的 recall 低,说明大量该类别样本被误判成了别的类,这时候要回头查数据,看看是不是该类图片存在标注错误或者类间视觉相似度过高。
我当时发现猫和狗之间有一些混淆,检查后原因是不少猫的图片是俯拍视角,毛色和某些狗接近。这类问题不能只靠换模型解决,更好的办法是补充更多典型样本,或者对易混淆类别做更细粒度的定义。
6.3 从曲线判断是否欠拟合、过拟合
训练过程中,我同时记录了训练 Loss 和验证 Loss。如果训练 Loss 持续下降、验证 Loss 在某个 epoch 后开始上升,说明模型开始过拟合,记住训练集细节而牺牲泛化能力。这时应该提前停止,并增强正则化强度,比如增大 Weight Decay、提高 Mixup 强度。
如果训练 Loss 和验证 Loss 都高居不下,基本是欠拟合,模型表达力没发挥出来。优先检查学习率是否太低、训练轮数是否不够、模型是否太小。在 Swin-T 这种模型上,迁移学习很少出现严重的欠拟合,除非你的数据集和 ImageNet 分布差异特别大,或者增强流程配置错误。
7. 我实际踩过的坑与避坑建议
7.1 Batch Size 与显存:不要一开始就设 64
我第一次跑这个项目,想当然设了 batch size 64,结果直接 OOM。前面说过,训练显存包含激活值和梯度,batch size 翻倍,激活值显存大致翻倍。后来改成 32 并开启 AMP 就正常了。如果显存只有 8GB,可以再降到 16。另外,不要通过降低图像分辨率来硬塞大 batch,因为下游任务对分辨率很敏感,Swin 的窗口设计也以 224×224 为基准。
7.2 学习率太激进导致 Loss 发散的教训
我用自己从零训练的 Swin-T 做过一次实验,初始学习率设为 1e-3,结果第一个 epoch 的 Loss 直接飙到 20 多,后面再也降不回来。Swin 这类模型对学习率比较敏感,从头训练通常使用 5e-4 配合长 warmup,迁移学习则建议 5e-5 到 1e-4。不要和 CNN 的经验直接划等号,ResNet 用 1e-3 能正常收敛,Swin 不一定行。
7.3 数据加载成了训练瓶颈
有一次 GPU 利用率只有 60% 左右,查了半天发现是num_workers=0,数据加载完全靠主进程,GPU 一直在等数据。改成num_workers=8后,利用率立刻上来。此外,图片文件很碎很小时,从机械硬盘读取会严重拖慢训练,建议把数据集放到 SSD 上,或者先用脚本打包成tar文件再读取。
7.4 模型尺寸和输入分辨率的选择
Swin-T 是入门首选,但如果你想追求更高精度,可以换 Swin-S 或 Swin-B,代价是训练时间和显存占用成倍增加。我个人的建议是:先把 Swin-T 的完整流程跑通,记录基线准确率,再根据需求横向对比不同尺寸模型。输入分辨率也不要一开始就追求 384,Swin 在 384 分辨率下要用window12配置,如果直接把window7挪过去,窗口划分不匹配,代码会报错或者效果异常。
我自己实际跑下来还有一个体会:Swin-Transformer 没有想象中那么“重”,它的调参门槛主要在数据组织、学习率和显存规划上,而不是模型本身。把完整流程走一遍之后,你对“图像识别项目”的理解会从调用一个model.fit变成真正能掌控每个环节。后续想扩展的话,可以在这个骨架上换数据集、尝试 Swin-S、加入更多数据增强或迁移到检测任务,所有经验都能平滑迁移。
本文还有配套的精品资源,点击获取