news 2026/9/12 22:30:58

Swin Transformer图像分类实战:窗口注意力机制与模型微调解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Swin Transformer图像分类实战:窗口注意力机制与模型微调解析

简介: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_size4
window_size7
embed_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_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 x

PatchEmbed 把图像切成 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.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', 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。下面是我在自定义数据集上常用的组合:

超参数推荐值说明
optimizerAdamW解耦权重衰减,适合 Transformer
learning rate5e-4线性衰减到 1e-5
weight decay0.05Swin 官方默认
batch size16/32根据显存调整
epochs100配合 early stopping
drop_path0.1增强模型泛化能力
warmup epochs20先线性升温再衰减

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 输入,或者在服务端统一缩放一次,简单又稳定。

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

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

用Res2Net改进UNet实现舌头图像语义分割实战

简介&#xff1a;这套基于UNet与Res2Net模块改进的舌头图像语义分割项目&#xff0c;以PyTorch为框架&#xff0c;面向医学影像分析、深度学习入门及语义分割进阶的研究者与学生&#xff0c;提供从数据预处理到训练评估的完整流程。资源包共610个文件&#xff0c;约7.37MB&…

作者头像 李华
网站建设 2026/9/12 22:30:13

级联H桥STATCOM低频纹波抑制技术解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 22:29:38

基因疗法在罕见癫痫症治疗中的突破与应用

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/12 22:29:22

混沌时间序列预测:基于时空RBF神经网络的MATLAB实现

简介&#xff1a;这是一份面向本科、硕士阶段教研学习的Matlab实现资料&#xff0c;聚焦时空RBF神经网络&#xff08;时空RBF-NN&#xff09;在混沌时间序列预测中的应用。压缩包内共10个文件&#xff0c;包含3个.m脚本、3个.mat数据集和4个png结果图&#xff0c;整体仅1.32MB&…

作者头像 李华
网站建设 2026/9/12 22:27:25

骨折图像数据集构建:从DICOM清洗到模型训练全流程指南

简介&#xff1a;面向医学影像分析与深度学习实战&#xff0c;这份骨折X射线图像数据集源自孟加拉国三家主要医院&#xff0c;原始扫描超过1.4万张&#xff0c;其中4083张经两名放射科专家独立标注并由医疗官员复核&#xff0c;专用于骨折分类、定位与分割任务&#xff0c;适合…

作者头像 李华