简介:Swin Transformer图像分类项目完整实现,面向具备Python与PyTorch基础、希望掌握Transformer架构在视觉任务中应用的开发者与研究人员。资源围绕图像分类全流程组织,包含模型定义、数据加载、训练验证、预测推理及混淆矩阵分析等脚本,并附不同配置的预训练权重,可快速用于模型微调或实际部署。压缩包共3691个文件,以3675张JPG图像样本为主,同时包含Python源码、PyTorch权重、JSON类别映射及说明文档,整体约586MB。项目代码中class_indices.json用于类别ID与名称映射,model.py展示窗口自注意力与层次化特征提取结构,utils.py封装常用辅助函数,train.py与predict.py覆盖训练到推理的完整链路。目前已有13377人学习下载。通过该资源可深入理解Swin Transformer核心机制,并借助错误样本筛选和混淆矩阵脚本定位模型不足,适合课程设计、算法实验及工程参考等场景。
1. 为什么 Swin Transformer 能扛住高分辨率图像分类
把一张 224×224 的图送进 ViT,patch=16 时 token 数是 196;Swin 把 patch 缩到 4,token 数变成 3136,再走全局自注意力,显存直接吃不消。Swin Transformer 把自注意力限制在 7×7 窗口里,后续层再做窗口移位,让信息跨窗口流动,计算复杂度从 N² 降到近线性,这是它能在图像分类任务里提升分辨率的原因。项目给出一套完整的 Swin Tiny 图像分类实现,model.py 定义网络,train.py 负责微调,predict.py 做推理,还有混淆矩阵和错误样本分析脚本。想从 CNN 花卉图像分类切到 transformer 图像分类模型的开发者,或要验证森林图像分类场景,可以直接用它起步。
2. Swin Transformer 的窗口注意力与层次化结构
2.1 窗口自注意力为什么比 ViT 省
Swin 的窗口注意力并不是把图像切成互不相干的小块,它用一个很聪明的设计解决了全局建模和计算量之间的矛盾。输入图 H×W,patch size 为 P,得到的 token 网格大小 N=HW/P²。普通自注意力在每一层都要对所有 token 两两计算,复杂度是 O(4N²C+8NC²),这里的 N 一旦被 patch=4 放大,平方项增长非常快。Swin 采取的策略是只在一个窗口内部做自注意力,窗口边长 M 通常设为 7,于是复杂度变成 O(4NM²C+8NC²)。
以这个项目默认的 224×224 输入为例,patch=4 时 N=56×56=3136。全局自注意力里 N²≈9.8×10⁶,而窗口注意力里 N·M²≈3136×49≈1.54×10⁵,两者相差约 64 倍。这就是为什么 Swin 敢在更高分辨率下训练,而 ViT 只能依赖更小的 patch 或者更复杂的 FlashAttention。窗口大小 M=7 是论文里的默认值,它和 patch_size=4 组合起来,可以保证 224、448、896 这些常见分辨率下窗口正好整整齐齐覆盖整张图。
层次化结构是 Swin 的另一个核心设计。整个网络输出 4 倍、8 倍、16 倍、32 倍下采样的特征,和 ResNet 的 stage 很像,这让 Swin 可以直接替换各种检测、分割模型的 backbone。项目里的 Swin Tiny 具体配置如下:
| 参数 | Swin-Tiny 配置 |
|---|---|
| 输入尺寸 | 224×224 |
| patch_size | 4 |
| window_size | 7 |
| embed_dim | 96 |
| 各 stage 深度 | [2, 2, 6, 2] |
| 各 stage 注意力头数 | [3, 6, 12, 24] |
| 分类头输出 | 数据集类别数 |
理解这张表很重要,因为后面用swin_tiny_patch4_window7_224.pth做微调时,分类头会被替换,只有 backbone 部分的参数能保留。
2.2 从 model.py 看 SwinTransformer 类组装
model.py 里并没有把每个算子都堆在一个大 forward 里,而是按模块拆开。核心类大概是这样的结构:
class SwinTransformer(nn.Module): def __init__(self, img_size=224, patch_size=4, in_chans=3, num_classes=1000, embed_dim=96, depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24], window_size=7): super().__init__() self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim) self.layers = nn.ModuleList() for i_layer in range(len(depths)): layer = BasicLayer( dim=int(embed_dim * 2 ** i_layer), depth=depths[i_layer], num_heads=num_heads[i_layer], window_size=window_size, downsample=PatchMerging(...) if i_layer < len(depths) - 1 else None ) self.layers.append(layer) self.norm = nn.LayerNorm(int(embed_dim * 2 ** (len(depths) - 1))) self.head = nn.Linear(int(embed_dim * 2 ** (len(depths) - 1)), num_classes) def forward(self, x): x = self.patch_embed(x) for layer in self.layers: x = layer(x) x = self.norm(x) x = x.mean(dim=1) x = self.head(x) return xPatchEmbed 把图像切成 4×4 的小 patch,并映射成一个 token 序列。BasicLayer 是每一阶段的容器,内部包含多个 SwinTransformerBlock,每个完整 block 由 W-MSA 和 SW-MSA 组成。W-MSA 是普通窗口注意力,SW-MSA 会先把窗口平移一半,让信息能穿过窗口边界流动。两个 block 成对出现,正是 Swin 能在保持低复杂度的同时建立全局依赖的关键。
窗口划分在代码里通常直接用 reshape 实现:
def window_partition(x, window_size): B, H, W, C = x.shape x = x.view(B, H // window_size, window_size, W // window_size, window_size, C) windows = x.permute(0, 1, 3, 2, 4, 5).contiguous() return windows.view(-1, window_size, window_size, C)这里把高、宽分别拆成“窗口数量 × 窗口大小”两个维度,再把窗口位置调换到同一维度,最后展平成所有窗口的序列。这段代码决定了为什么输入尺寸必须满足 H/4 和 W/4 能被 window_size 整除,否则窗口划分会丢掉边缘像素,导致特征缺失。
2.3 预训练权重选择:swin_tiny_patch4_window7_224.pth 与 mask_rcnn 权重
很多人在这个项目里看到两个 .pth 文件后会直接困惑:到底该加载哪一个?swin_tiny_patch4_window7_224.pth是官方针对分类任务预训练的 Swin-Tiny 权重,加载到model.py里的 SwinTransformer 上非常顺。mask_rcnn_swin_tiny_patch4_window7_1x.pth则是从 Mask R-CNN 模型导出的权重,里面除了 backbone,还有很多检测头参数,直接 load 会报一堆 unexpected keys,分类任务根本用不上。
加载分类权重的常见做法如下:
checkpoint = torch.load('swin_tiny_patch4_window7_224.pth', map_location='cpu') if 'model' in checkpoint: state_dict = checkpoint['model'] else: state_dict = checkpoint # 去掉分类头相关参数,避免 num_classes 不一致时报错 state_dict = {k: v for k, v in state_dict.items() if not k.startswith('head.')} model_dict = model.state_dict() model_dict.update(state_dict) model.load_state_dict(model_dict, strict=False)先过滤head.开头的权重,再更新到模型里,分类头保持随机初始化,backbone 直接复用预训练参数。strict=False允许缺失 head 键。如果你实在想用 mask_rcnn 权重,可以尝试把键名前缀backbone.去掉再加载,但通常只是把 backbone 部分初始化得差不多,对最终分类指标并没有明显帮助,不值得为了它额外写一套兼容逻辑。
| 权重文件 | 来源模型 | 是否建议用于分类 | 原因 |
|---|---|---|---|
| swin_tiny_patch4_window7_224.pth | ImageNet 分类 Swin-T | 是 | 结构完全匹配 |
| mask_rcnn_swin_tiny_patch4_window7_1x.pth | Mask R-CNN 检测模型 | 不建议 | 参数字典复杂,匹配困难 |
3. 训练与微调:从数据目录到 train.py 参数
3.1 数据目录与类别映射
用 Swin 做分类,第一步是把数据整理成 PyTorch ImageFolder 能识别的目录格式。比如数据集根目录下分 train 和 val,每个子目录内部再按类名建文件夹:
data/ train/ 0_dog/ img_0001.jpg 1_cat/ img_0002.jpg val/ 0_dog/ img_0003.jpg 1_cat/ img_0004.jpg读取目录并生成class_indices.json的代码如下:
from torchvision import datasets import json train_set = datasets.ImageFolder('data/train') class_to_idx = train_set.class_to_idx idx_to_class = {v: k for k, v in class_to_idx.items()} with open('class_indices.json', 'w', encoding='utf-8') as f: json.dump(idx_to_class, f, indent=2, ensure_ascii=False)这段代码生成的idx_to_class是纯 Python 的int->str字典,写入 json 后 key 会自动变成字符串,也就是类似{"0": "dog"}。后面 predict.py 从 json 读回来时,索引要先用str()转换,否则会导致 KeyError。
数据增强方面,Swin 对输入分辨率比较挑剔,不是因为算力不够,而是因为 window_size 和 patch_size 有整除关系。训练时随机裁剪到 224 就好,验证集不要做随机增强:
from torchvision import transforms train_transform = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.2, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.4, 0.4, 0.4), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ])RandomResizedCrop 的 scale 下限设为 0.2 是 Swin 官方训练里的常见配置,如果数据集很小,可以把它放宽到 0.08,让裁剪覆盖更多目标比例,增强效果会更明显。
3.2 train.py 里的训练循环与超参设置
Swin 微调时最关键的几个参数是优化器、学习率、weight decay 和 drop path。下面是我在自定义数据集上常用的组合:
| 超参数 | 推荐值 | 说明 |
|---|---|---|
| optimizer | AdamW | 解耦权重衰减,适合 Transformer |
| learning rate | 5e-4 | 线性衰减到 1e-5 |
| weight decay | 0.05 | Swin 官方默认 |
| batch size | 16/32 | 根据显存调整 |
| epochs | 100 | 配合 early stopping |
| drop_path | 0.1 | 增强模型泛化能力 |
| warmup epochs | 20 | 先线性升温再衰减 |
drop_path 是 Swin 的残差分支随机丢弃操作,它跟普通 dropout 不一样,训练时能有效防止小数据集过拟合。构造 model.py 里的 SwinTransformer 时要传入drop_path_rate=0.1,而不是写在nn.Dropout里。
训练循环建议使用混合精度和梯度裁剪。Swin 深层的梯度偶尔会剧烈抖动,clip 一下省心很多:
scaler = torch.cuda.amp.GradScaler() for images, labels in train_loader: images, labels = images.cuda(), labels.cuda() optimizer.zero_grad() with torch.cuda.amp.autocast(): outputs = model(images) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) scaler.step(optimizer) scaler.update()clip_grad_norm_的 max_norm 设为 5.0 是一个比较稳妥的值,太小的裁剪会拖慢收敛速度,太大会失去保护作用。AMP 下打日志时直接取loss.item()就行,不要拿scaler.scale(loss)的值去打印,那个是放大后的数值。
3.3 断点续训与权重保存
训练到一半断掉是很常见的事。只存 best model 不够,我还习惯把 optimizer、scheduler、epoch 都存进同一个 checkpoint:
torch.save({ 'epoch': epoch, 'model': model.state_dict(), 'optimizer': optimizer.state_dict(), 'scheduler': scheduler.state_dict(), 'best_acc': best_acc, }, 'checkpoint.pth')恢复训练的代码要按相同顺序重建对象:
ckpt = torch.load('checkpoint.pth', map_location='cuda') model.load_state_dict(ckpt['model']) optimizer.load_state_dict(ckpt['optimizer']) scheduler.load_state_dict(ckpt['scheduler']) start_epoch = ckpt['epoch'] + 1 best_acc = ckpt['best_acc']这里有个容易踩的坑:如果你改了分类头,输出类别数变了,旧 checkpoint 里的优化器状态会和当前模型参数不同。解决方法是改分类头时只加载模型权重,不要恢复优化器状态,或者重建 optimizer 后再训练几个 epoch。
4. 评估与错误分析:混淆矩阵和错误样本选择
4.1 用混淆矩阵定位类别混淆
准确率只能说明模型整体水平,但看不出是哪个类拖了后腿。项目里的 create_confusion_matrix.py 专门干这个。比如森林图像分类里,杉树、松树、桦树外观接近,模型可能经常把杉树猜成松树,混淆矩阵里就会有一块明显的横向亮色带。
我自己生成混淆矩阵时习惯使用 sklearn 的交互组件:
import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay def make_confusion_matrix(model, val_loader, class_names, save_path='cm.png'): model.eval() all_preds, all_labels = [], [] with torch.no_grad(): for images, labels in val_loader: images = images.cuda() preds = model(images).argmax(dim=1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) cm = confusion_matrix(all_labels, all_preds) disp = ConfusionMatrixDisplay(cm, display_labels=class_names) disp.plot(cmap='Blues', colorbar=False) plt.xticks(rotation=45, ha='right') plt.tight_layout() plt.savefig(save_path, dpi=150)当类别数量超过 50 时,建议在confusion_matrix里加上normalize='true',矩阵按行归一化。这样每个格子代表真实类别中被分到各类别的比例,可以剔除样本量差异带来的视觉误导。
4.2 select_incorrect_samples.py 的筛选逻辑
错误样本的价值不一样。我通常最关注“高置信度错误样本”,也就是模型非常确定、但结果依旧是错的。这些样本往往是标签噪声或类别定义重叠导致的。
select_incorrect_samples.py 的筛选逻辑可以这样写:
def select_incorrect(model, val_loader, idx_to_class, topk=30): results = [] model.eval() with torch.no_grad(): for batch_idx, (images, labels) in enumerate(val_loader): images = images.cuda() logits = model(images) probs = torch.softmax(logits, dim=1) conf, preds = probs.max(dim=1) for i in range(images.size(0)): if preds[i] != labels[i]: results.append({ 'sample_id': batch_idx * val_loader.batch_size + i, 'true': idx_to_class[str(labels[i].item())], 'pred': idx_to_class[str(preds[i].item())], 'confidence': conf[i].item() }) results.sort(key=lambda x: x['confidence'], reverse=True) return results[:topk]返回列表里已经按置信度从高到低排列。拿到结果后,我会对照下面的模式来分析:
| 错误模式 | 可能原因 | 处理方向 |
|---|---|---|
| 高置信度集中错误 | 标注错误或类别边界污染 | 检查原图,修正标签 |
| 低置信度错误 | 目标过小或遮挡严重 | 增加多尺度训练 |
| 固定两个类别之间互混 | 类别外观高度重叠 | 合并类或细分标注 |
| 某个环境下集中出错 | 背景特征过强 | 加入随机擦除和背景扰动 |
如果错误样本大多来自光线很暗的图片,说明训练集缺少暗光数据。这时候与其继续调参,不如去补充一个月的实地拍摄样本,效果比堆网络层数更明显。
5. 预测提速与推理细节
5.1 predict.py 的类别反查流程
训练完成后真正要用的其实是 predict.py。这里最容易出错的是 class_indices.json 的反查。模型输出是一个索引张量,必须通过idx_to_class转成类别名。一个可用的 top-k 预测函数大概是这样的:
def predict_topk(model, image_path, transform, idx_to_class, k=5): img = Image.open(image_path).convert('RGB') img_tensor = transform(img).unsqueeze(0).cuda() model.eval() with torch.no_grad(): logits = model(img_tensor) probs = torch.softmax(logits, dim=1)[0] topk_conf, topk_idx = torch.topk(probs, k) for conf, idx in zip(topk_conf.tolist(), topk_idx.tolist()): print(idx_to_class[str(idx)], f'{conf:.4f}')这里有一处细节:json 的 key 是字符串,所以索引必须写成str(idx)。如果直接写idx_to_class[idx],int 类型是无法命中字符串 key 的。另外,model.eval()和torch.no_grad()都要写,尤其是模型里有 DropPath,运行时态不固定会导致预测结果抖动。
推理阶段如果希望提速,可以把模型和输入都转成半精度:
model = model.half().cuda() img_tensor = img_tensor.half()在同等显存下,半精度推理通常能比 FP32 快 20% 到 35%。如果后续要接入 Triton 或 ONNX Runtime,固定 224 输入尺寸的效果反而比动态尺寸更好,因为窗口划分逻辑在静态 shape 下更容易被优化。
5.2 多尺寸推理的窗口对齐与导出坑
Swin 对输入尺寸不灵活,根因是 patch_size=4、window_size=7 的整除约束。如果业务入口是任意分辨率,比如摄像头输出 640×480,直接丢给模型会多出不少麻烦。常见做法是在预处理阶段把图像缩放到最近的合法尺寸:
import torch.nn.functional as F def align_for_swin(img_tensor, window_size=7, patch_size=4): _, _, H, W = img_tensor.shape unit = window_size * patch_size nH, nW = max(1, round(H / unit)), max(1, round(W / unit)) target_h, target_w = nH * unit, nW * unit if (H, W) != (target_h, target_w): img_tensor = F.interpolate(img_tensor, size=(target_h, target_w), mode='bicubic') return img_tensor这里用 interpolation 把图缩放到最近的可整除尺寸,而不是 padding。padding 会在图像边缘增加大量无用像素,Swin 在窗口注意力时会把它们当作正常内容参与计算,容易干扰分类结果。考虑到 224、448、896 都满足整除条件,实际部署时我更推荐固定 224 输入,或者在服务端统一缩放一次,简单又稳定。
本文还有配套的精品资源,点击获取