
简介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 的图送进 ViTpatch16 时 token 数是 196Swin 把 patch 缩到 4token 数变成 3136再走全局自注意力显存直接吃不消。Swin Transformer 把自注意力限制在 7×7 窗口里后续层再做窗口移位让信息跨窗口流动计算复杂度从 N² 降到近线性这是它能在图像分类任务里提升分辨率的原因。项目给出一套完整的 Swin Tiny 图像分类实现model.py 定义网络train.py 负责微调predict.py 做推理还有混淆矩阵和错误样本分析脚本。想从 CNN 花卉图像分类切到 transformer 图像分类模型的开发者或要验证森林图像分类场景可以直接用它起步。2. Swin Transformer 的窗口注意力与层次化结构2.1 窗口自注意力为什么比 ViT 省Swin 的窗口注意力并不是把图像切成互不相干的小块它用一个很聪明的设计解决了全局建模和计算量之间的矛盾。输入图 H×Wpatch size 为 P得到的 token 网格大小 NHW/P²。普通自注意力在每一层都要对所有 token 两两计算复杂度是 O(4N²C8NC²)这里的 N 一旦被 patch4 放大平方项增长非常快。Swin 采取的策略是只在一个窗口内部做自注意力窗口边长 M 通常设为 7于是复杂度变成 O(4NM²C8NC²)。以这个项目默认的 224×224 输入为例patch4 时 N56×563136。全局自注意力里 N²≈9.8×10⁶而窗口注意力里 N·M²≈3136×49≈1.54×10⁵两者相差约 64 倍。这就是为什么 Swin 敢在更高分辨率下训练而 ViT 只能依赖更小的 patch 或者更复杂的 FlashAttention。窗口大小 M7 是论文里的默认值它和 patch_size4 组合起来可以保证 224、448、896 这些常见分辨率下窗口正好整整齐齐覆盖整张图。层次化结构是 Swin 的另一个核心设计。整个网络输出 4 倍、8 倍、16 倍、32 倍下采样的特征和 ResNet 的 stage 很像这让 Swin 可以直接替换各种检测、分割模型的 backbone。项目里的 Swin Tiny 具体配置如下参数Swin-Tiny 配置输入尺寸224×224patch_size4window_size7embed_dim96各 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_size224, patch_size4, in_chans3, num_classes1000, embed_dim96, depths[2, 2, 6, 2], num_heads[3, 6, 12, 24], window_size7): 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( dimint(embed_dim * 2 ** i_layer), depthdepths[i_layer], num_headsnum_heads[i_layer], window_sizewindow_size, downsamplePatchMerging(...) 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(dim1) 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_locationcpu) 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, strictFalse)先过滤head.开头的权重再更新到模型里分类头保持随机初始化backbone 直接复用预训练参数。strictFalse允许缺失 head 键。如果你实在想用 mask_rcnn 权重可以尝试把键名前缀backbone.去掉再加载但通常只是把 backbone 部分初始化得差不多对最终分类指标并没有明显帮助不值得为了它额外写一套兼容逻辑。权重文件来源模型是否建议用于分类原因swin_tiny_patch4_window7_224.pthImageNet 分类 Swin-T是结构完全匹配mask_rcnn_swin_tiny_patch4_window7_1x.pthMask 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, encodingutf-8) as f: json.dump(idx_to_class, f, indent2, ensure_asciiFalse)这段代码生成的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。下面是我在自定义数据集上常用的组合超参数推荐值说明optimizerAdamW解耦权重衰减适合 Transformerlearning rate5e-4线性衰减到 1e-5weight decay0.05Swin 官方默认batch size16/32根据显存调整epochs100配合 early stoppingdrop_path0.1增强模型泛化能力warmup epochs20先线性升温再衰减drop_path 是 Swin 的残差分支随机丢弃操作它跟普通 dropout 不一样训练时能有效防止小数据集过拟合。构造 model.py 里的 SwinTransformer 时要传入drop_path_rate0.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_norm5.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 都存进同一个 checkpointtorch.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_locationcuda) 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_pathcm.png): model.eval() all_preds, all_labels [], [] with torch.no_grad(): for images, labels in val_loader: images images.cuda() preds model(images).argmax(dim1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) disp ConfusionMatrixDisplay(cm, display_labelsclass_names) disp.plot(cmapBlues, colorbarFalse) plt.xticks(rotation45, haright) plt.tight_layout() plt.savefig(save_path, dpi150)当类别数量超过 50 时建议在confusion_matrix里加上normalizetrue矩阵按行归一化。这样每个格子代表真实类别中被分到各类别的比例可以剔除样本量差异带来的视觉误导。4.2 select_incorrect_samples.py 的筛选逻辑错误样本的价值不一样。我通常最关注“高置信度错误样本”也就是模型非常确定、但结果依旧是错的。这些样本往往是标签噪声或类别定义重叠导致的。select_incorrect_samples.py 的筛选逻辑可以这样写def select_incorrect(model, val_loader, idx_to_class, topk30): 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, dim1) conf, preds probs.max(dim1) 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(keylambda x: x[confidence], reverseTrue) 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, k5): 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, dim1)[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_size4、window_size7 的整除约束。如果业务入口是任意分辨率比如摄像头输出 640×480直接丢给模型会多出不少麻烦。常见做法是在预处理阶段把图像缩放到最近的合法尺寸import torch.nn.functional as F def align_for_swin(img_tensor, window_size7, patch_size4): _, _, 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), modebicubic) return img_tensor这里用 interpolation 把图缩放到最近的可整除尺寸而不是 padding。padding 会在图像边缘增加大量无用像素Swin 在窗口注意力时会把它们当作正常内容参与计算容易干扰分类结果。考虑到 224、448、896 都满足整除条件实际部署时我更推荐固定 224 输入或者在服务端统一缩放一次简单又稳定。本文还有配套的精品资源点击获取