新闻详情

新闻详情

首页 / 资讯中心 / 详情

Swin Transformer图像分类实战:窗口注意力与PyTorch微调全解析

发布时间:2026/9/12 22:31:32来源:尧图网络
Swin Transformer图像分类实战:窗口注意力与PyTorch微调全解析
简介这是一份基于 PyTorch 的 Swin Transformer 图像分类完整实现面向有一定深度学习基础、希望掌握 Transformer 架构在视觉任务中应用的开发者与研究者。配套代码覆盖模型定义、训练、验证、预测全流程包含多种预训练权重与类别索引文件可直接用于新数据集微调或推理。资源共 3691 个文件以 jpg 图像数据为主并含 7 个 Python 脚本、3 个预训练 pth 权重、json 标签映射及说明文档压缩包约 586MB。除核心模型与训练逻辑外还提供混淆矩阵生成与错误样本筛选脚本便于直观评估分类效果、定位模型薄弱点。已有 13377 人学习下载适合作为理解局部窗口自注意力机制、层次化特征提取及 Swin Transformer 工程落地的入门到进阶资料。1. Swin Transformer 不是来替换 CNN 的它只是比 ViT 更懂图像做图像分类的工程师这两年被 ViT 系模型折腾得不轻全局注意力确实强但一到高分辨率输入或者目标检测这类密集预测任务计算量就顶不住了。Swin Transformer 在 2021 年提出时最核心的贡献不是“又一个 Transformer 分类器”而是把视觉 Transformer 的全局建模能力用窗口注意力重新组织成了 CNN 式的高效层级结构。换句话说缩放位移窗口Shifted Window让 Transformer 第一次能像 ResNet 一样处理任意尺寸的输入同时保持线性计算复杂度。对做分类的团队来说它的价值在于既继承预训练大模型时代的迁移学习红利又不用在推理时担心输入尺寸一变就内存爆炸。这篇文章会从窗口注意力的原理讲起给出可直接运行的 PyTorch 推理和微调代码并把参数设置和踩坑点交代清楚。适合正在评估 Swin 做分类 baseline、或者想把 ViT 换掉的工程人员。2. Swin Transformer 的窗口注意力与层级结构为什么对分类更友好2.1 全局注意力的问题语义强但算力贵ViT 把图像切成 16×16 的 patch然后让每个 patch 和所有其他 patch 做注意力交互。这在中等分辨率下表现尚可但图像的语义信息是有局部性的猫的耳朵大概率只和附近的像素有关很少需要直接和背景里的天空做关联。全局注意力机制在数学上最优雅但它回避了视觉任务的两个现实约束。第一个约束是输入分辨率敏感。全局注意力的复杂度是序列长度的平方。对于 224×224 输入patch 16×16 时序列长度是 196这个规模 GPU 能扛住但分类任务迁移到生产环境时经常要处理 384×384 甚至 512×512 的输入序列长度直接变成 576 和 1024自注意力矩阵占用的显存和算力陡增。第二个约束是丧失了图像的二维先验。ViT 虽然加了位置编码但 patch 之间的交互模式仍然是全连接的CNN 局部感受野带来的归纳偏置被丢弃了小数据集上很容易过拟合。Swin Transformer 的处理方式非常直接把注意力限制在固定大小的窗口内。默认窗口尺寸是 7×7 的 patch每个 patch 大小为 4×4 像素因此一个窗口覆盖 28×28 像素区域。窗口内的 token 数量恒定为 49不管输入图像多大自注意力的计算量只跟窗口数量成正比而不跟图像尺寸的平方成正比。这就是 O(n) 复杂度的来源。2.2 移位窗口在局部与全局之间搭桥只做窗口内注意力会带来一个明显的问题窗口之间没有信息交流每个窗口相当于一个孤立的小模型整个 Transformer 就退化成了对图像做分块独立建模。Swin 的解法是交替使用两种窗口划分方式。在第一层使用规则窗口第二层把窗口整体向右下方向移位移位的步长是窗口尺寸的一半默认 3 个 patch然后再划分窗口。仔细观察会发现移位后的窗口数量从 4 个变成 9 个多出来的边界窗口并不完整。Swin 在实现上使用了 masked attention 技巧保证移位后的每个窗口仍然只计算有效 token 之间的注意力而不会跨越窗口边界。从效果上说一层的窗口内注意力负责提取局部模式下一层的移位窗口让相邻窗口边缘的 token 有机会交互两层叠起来等效于感受野扩大。对图像分类而言这个设计直接替代了 ViT 中昂贵的全局交互同时保留了跨区域建模能力。2.3 Patch Merging 与窗口配置参数详解Swin Transformer 采用金字塔结构堆叠了四个 Stage每经过一个 Stagetoken 数量减半通道数加倍。这个设计和 ResNet 的降采样策略几乎一致也是 Swin 能被顺利嵌入检测、分割框架的核心原因。每个 Stage 由若干 Swin Transformer Block 和末尾一个 Patch Merging 层组成Patch Merging 把相邻 2×2 的 patch 拼接后通过线性层融合实现降采样。Swin-TTiny 版本的详细配置如下表所示这也是实践中最常用的分类模型。参数项Swin-T 默认值说明patch_size4每个 patch 覆盖 4×4 像素图像先被切分为 patch 序列embed_dim96第一个 Stage 的通道维度后续 Stage 依次翻倍depths[2, 2, 6, 2]四个 Stage 中 Swin Block 的数量num_heads[3, 6, 12, 24]每个 Stage 的注意力头数头数随通道数翻倍window_size7注意力窗口内的 patch 数量7×749 个 tokenmlp_ratio4FFN 隐藏层维度是输入维度的 4 倍qkv_biasTrueQKV 投影层是否使用偏置项apeFalse是否使用绝对位置编码Swin 默认只用相对位置偏置关于 window_size 需要特别强调这个参数是和输入分辨率耦合的。Swin 官方预训练权重在 ImageNet-1K 上使用 224×224 输入、window_size 7 训练如果直接把输入分辨率提升到 448×448token 数量增加但窗口内的 token 数量不变因此窗口数量会变为原来的 4 倍自注意力计算量线性增长这个是可以接受的。但如果你想保持 64×64 的原始输入window_size 就必须改成 4否则图像会不足以撑起完整的窗口。2.4 相对位置偏置分类精度的一个隐蔽来源Swin 在每个自注意力模块里加了一个可学习的相对位置偏置表形状是 (2×window_size-1) × (2×window_size-1)。计算注意力分数时先算出 token 之间的相对位置索引再查表得到偏置值加在 QK^T 的结果上。这个设计替代了 ViT 的绝对位置编码因为图像内容本身是平移等变的绝对位置信息反而可能干扰模式识别。从工程实践角度看这个偏置表对分类任务有多重要有消融实验表明去掉相对位置偏置后Swin-T 在 ImageNet-1K 上的 top-1 精度会掉大约 1.2 到 1.5 个百分点。所以当你基于 Swin 做微调时不要随意删改位置偏置相关的代码。使用 timm 库加载预训练权重时这个偏置表已经包含在权文件中不需要额外处理。3. 用 PyTorch 快速跑通 Swin Transformer 图像分类的最小推理流程3.1 环境准备与依赖安装Swin Transformer 的官方代码仓库基于 mmdetection 体系分类部分则可以直接选用 timm 里的实现。两者都能跑但 timm 的模型封装更贴近 PyTorch 原生生态预训练权重下载、forward 接口、输入预处理都已经内置适合快速验证。创建新的 conda 环境并安装依赖建议 Python 3.9 以上版本PyTorch 2.0 以上对 Window Attention 里的算子有更充分的优化支持conda create -n swin python3.10 -y conda activate swin pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install timm0.9.16 pillow requests安装完成后验证导入是否正常顺便确认 GPU 可用性。Swin-T 的参数量约 2800 万在单张 RTX 3090 或更低的显卡上都可以流畅推理真正吃显存的是训练阶段。3.2 加载预训练权重并完成单张图片分类这里以 timm 的 Swin-T 为例给出一个可直接运行的推理脚本。重点不是跑通而是让你理解 timm 的 create_model 接口在背后做了什么事。import torch import timm from PIL import Image from timm.data import resolve_data_config from timm.data.transforms_factory import create_transform # 加载预训练权重num_classes1000 表示 ImageNet-1K 类别数 model timm.create_model( swin_tiny_patch4_window7_224, pretrainedTrue, num_classes1000 ) model.eval() # timm 会根据模型配置自动生成对应的预处理流程 # 包括 Resize、CenterCrop、Normalize无需手动设置 mean/std config resolve_data_config({}, modelmodel) transform create_transform(**config) # 读取本地图片并预处理为 224x224 的 tensor img Image.open(demo.jpg).convert(RGB) input_tensor transform(img).unsqueeze(0) # 增加 batch 维度 # 推理 with torch.no_grad(): logits model(input_tensor) # softmax 转为概率分布取 top-5 probs torch.softmax(logits, dim1) top5_probs, top5_indices torch.topk(probs, k5) print(Top-5 预测结果:) for i in range(5): print(f class_id{top5_indices[0][i].item()}, prob{top5_probs[0][i].item():.4f})这段代码里最关键的是create_model和resolve_data_config的组合。create_model里的模型名swin_tiny_patch4_window7_224包含了完整的结构信息patch 大小为 4窗口大小为 7输入分辨率是 224。resolve_data_config会从模型配置中读取 input_size、mean、std、interpolation 等参数生成与预训练一致的预处理方式。手动写 Normalize 时最常见的错误就是 mean/std 用错timm 里 Swin 的默认 mean 是[0.485, 0.456, 0.406]std 是[0.229, 0.224, 0.225]与 ResNet 相同。3.3 用 torchvision 替代 timmScaling 的差异要注意torchvision 在 0.13 版本之后也提供了 Swin Transformer 实现接口是torchvision.models.swin_t(weightsSwin_T_Weights.IMAGENET1K_V1)。两者的核心差异在于 window attention 的实现方式和预处理细节。torchvision 的 Swin 实现使用nn.Functional.pad处理移位窗口的边界timm 用了更复杂但更节省显存的 masked attention 方案。实际推理时两者输出精度几乎一致但在批量训练时 timm 的实现略快显存占用也更低。如果你只是做 demo 演示用 torchvision 足够如果要进入训练流程建议统一走 timm。预处理上torchvision 的Swin_T_Weights会自动绑定 transform 方法直接调用即可。但注意两者对 resize 的插值方式设置不同timm 使用 bicubictorchvision 使用 bilinear最终的 top-1 精度差异在 0.1% 以内可以忽略。4. 自定义数据集微调 Swin Transformer从配置文件到训练循环4.1 自建数据集的目录规范与加载方式分类场景里更多遇到的情况是用自己的业务数据做微调比如森林图像分类、花卉识别这类领域任务。Swin 的迁移学习效果依赖于数据目录的组织方式。工程上最稳妥的做法是遵循 PyTorch 的 ImageFolder 约定dataset/ ├── train/ │ ├── oak_forest/ │ │ ├── img_001.jpg │ │ └── ... │ ├── pine_forest/ │ │ └── ... │ └── wetland/ │ └── ... └── val/ ├── oak_forest/ │ └── ... └── ...每一类对应一个子文件夹文件夹名即类别名。使用torchvision.datasets.ImageFolder读取会自动完成标签到索引的映射。如果原始数据不是这种组织方式写一个简单的路径整理脚本即可完成转换。数据加载过程中有一个过去常被忽视的细节Swin 的预训练权重是在平衡数据集上训练的如果你的数据集存在严重的类别不均衡需要在训练时设置weights参数传给 CrossEntropyLoss否则小类别的准确率会被大幅压榨。4.2 微调的数据增强配置与归一化参数匹配Swin 微调时数据增强策略比 ViT 简单不需要添加 DeiT 那套 heavy augmentation常用的组合是from torchvision import transforms transform_train transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.08, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) transform_val 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]) ])注意验证集不要增加随机增强Resize(256)后CenterCrop(224)是 ImageNet 评测的标准流程。RandomResizedCrop 的scale参数从 0.08 开始而不是默认的 0.08因为 Swin 的窗口是局部的尺度变化太大会让模型难以捕捉稳定的上下文。4.3 训练循环的完整实现直接给出一段可以在单卡上运行的微调代码。需要适配的只有data_dir和num_classes两个变量。import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR from timm import create_model from timm.utils import ModelEmaV2 import time device torch.device(cuda if torch.cuda.is_available() else cpu) data_dir dataset num_classes 3 batch_size 32 epochs 30 lr 5e-5 # 加载数据transform 见上一节 train_dataset ImageFolder(f{data_dir}/train, transformtransform_train) val_dataset ImageFolder(f{data_dir}/val, transformtransform_val) train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue, num_workers8, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_sizebatch_size, shuffleFalse, num_workers8, pin_memoryTrue) # 加载预训练模型并替换分类头 model create_model( swin_tiny_patch4_window7_224, pretrainedTrue, num_classesnum_classes ) model.to(device) # 分类任务微调时只需训练分类头之外冻结前 2 个 Stage for name, param in model.named_parameters(): if stages.0 in name or stages.1 in name: param.requires_grad False optimizer AdamW(filter(lambda p: p.requires_grad, model.parameters()), lrlr, weight_decay0.05) scheduler CosineAnnealingLR(optimizer, T_maxepochs) criterion nn.CrossEntropyLoss() # EMA 参数平滑timm 内置实现 ema_model ModelEmaV2(model, decay0.9999) for epoch in range(epochs): model.train() total_loss, correct, total 0.0, 0, 0 start time.time() for images, labels in train_loader: images, labels images.to(device), labels.to(device) # 混合精度训练 with torch.amp.autocast(cuda): outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() # 梯度裁剪防止窗口注意力训练早期的梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step() ema_model.update(model) total_loss loss.item() * images.size(0) _, preds torch.max(outputs, 1) correct (preds labels).sum().item() total labels.size(0) # 用 EMA 模型做验证结果更稳定 ema_model.eval() val_correct, val_total 0, 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs ema_model.module(images) _, preds torch.max(outputs, 1) val_correct (preds labels).sum().item() val_total labels.size(0) train_acc correct / total val_acc val_correct / val_total scheduler.step() print(fEpoch {epoch1}/{epochs} | fTrain Loss: {total_loss/total:.4f} | fTrain Acc: {train_acc:.4f} | fVal Acc: {val_acc:.4f} | fTime: {time.time()-start:.1f}s)这段代码里有几个值得展开说明的工程决策。模型名中的swin_tiny_patch4_window7_224在 timm 的 registry 里是一种标准配置但 timm 也提供pretrained_cfg参数来控制权重的来源。如果你使用的是 ImageNet-22K 预训练版本可以在模型名后加_in22k后缀再配合num_classes0加载最后手动替换分类头。关于冻结策略前两个 Stage 的浅层特征主要是边缘和纹理通用性强冻结后能显著减少训练时间同时降低小数据集过拟合风险。后两个 Stage 和分类头保留可训练状态。如果你的数据集和 ImageNet 分布相差较大例如医学图像或卫星图像建议不要冻结任何层从较小的学习率如 2e-5开始全量微调。学习率和 weight decay 需要特别注意。Swin 的预训练使用的是 AdamW微调时继续用 AdamW 最稳定SGD 在 Swin 上表现较差因为 LayerNorm 对学习率的敏感度不同于 BN。weight_decay0.05是 ViT 系模型的常用值偏大时会对 QKV 投影层的权重过度惩罚。4.4 训练日志的解读与模型保存策略观察项正常表现异常情况排查Train Loss前 3 个 epoch 缓慢下降之后加速收敛若 loss 不降检查 lr 是否过大、数据预处理是否和预训练匹配Val Acc随 epoch 稳步上升可能震荡但总体向上若 val acc 突然掉到随机水平检查 DataLoader shuffle 和标签映射梯度的 grad norm稳定在 0.1 到 5 之间若持续大于 10说明 lr 过大或存在脏数据Val loss 与 Train loss 差值差值小于 0.3 属于正常差值过大代表过拟合需要增加增强强度或提前停止模型保存需要用单独的逻辑不能只存 model.state_dict()因为 EMA 模型的权重不在同一个字典里。推荐的做法是每个 epoch 保存最新的 model 和 ema_model 的权重最后根据 val_acc 选择最优的版本加载推理。Swin 微调收敛较快在 30 epoch 的设置下通常在 15 到 20 epoch 达到最优验证精度继续训练会有轻微过拟合但不严重。5. 部署前验证模型输出的正确性与鲁棒性三个实用技巧5.1 用单张图快速校验模型的输出是否合理微调完成后的第一个验证步骤不是看测试集准确率而是单张图的输出是否符合直觉。Swin 的窗口注意力对输入尺寸较敏感确认一下在实际业务图片上是否产生了合理的 top-1 预测。写一个极简的验证脚本import torch from PIL import Image from timm import create_model model create_model(swin_tiny_patch4_window7_224, pretrainedFalse, num_classesnum_classes) state_dict torch.load(best_model.pth, map_locationcpu) model.load_state_dict(state_dict) model.eval() img Image.open(test_forest.jpg).convert(RGB) tensor transform(img).unsqueeze(0) with torch.no_grad(): out torch.softmax(model(tensor), dim1) print(out) # 打印每个类别的概率注意create_model里必须显式pretrainedFalse否则会先加载 ImageNet 权重再被你覆盖虽然最终效果一样但白浪费一次下载。如果对输出分布不放心可以用多张图测试观察不确定度较高的样本长什么样。5.2 检查 batch size 变化对结果一致性的影响Swin 模型包含 LayerNorm 和 LayerScale不同 batch size 下结果应当完全一致除非数据加载过程有问题。一个快速验证方法是用相同的单张图片分别以 batch size 1 和 batch size 8 推理对比输出概率是否完全一致误差小于 1e-6。这个验证能排除隐藏的 batch 维耦合问题比如预处理的 batch 归一化没有关闭。5.3 TF32 精度对 Swin 分类结果的影响NVIDIA Ampere 架构起PyTorch 默认在 CUDA 上使用 TF32 精度计算矩阵乘法和卷积为的是提升速度。但对 Swin 这类注意力模型TF32 会让梯度计算引入额外噪声微调出来的模型精度可能比 FP32 低 0.3 到 0.5 个百分点。正式发布模型前用下面这行代码确定实际精度torch.backends.cuda.matmul.allow_tf32 False torch.backends.cudnn.allow_tf32 False如果业务对推理速度有硬性要求可以保留 TF32 并对比一次验证集准确率若差异在可接受范围则继续使用否则切回 FP32。Swin 推理时把输入尺寸从 224 提到 384 对 top-1 精度提升约为 0.8 到 1.0 个百分点代价是推理时间增加约 2.5 倍这个收益通常比调试 TF32 更明显。本文还有配套的精品资源点击获取
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

32GB GPU跑LoRA/QLoRA微调不OOM:显存优化指南 2026/9/12 23:13:43

32GB GPU跑LoRA/QLoRA微调不OOM:显存优化指南

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

阅读更多 →
六自由度弹道仿真与BTT控制技术详解 2026/9/12 23:13:43

六自由度弹道仿真与BTT控制技术详解

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

阅读更多 →
力扣SQL高频50题进阶:窗口函数与查询优化实战 2026/9/12 23:13:43

力扣SQL高频50题进阶:窗口函数与查询优化实战

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

阅读更多 →
Raspberry Pi Pico + MicroPython 入门实战:从点灯到温湿度监测器 2026/9/12 23:13:43

Raspberry Pi Pico + MicroPython 入门实战:从点灯到温湿度监测器

第一次拿到一块绿色的 Raspberry Pi Pico,我盯着板子看了好一会儿:没有常见的 USB 转串口芯片,没有复位按键,只有一个孤零零的 BOOTSEL 按钮。如果不懂烧录原理,还真不知道怎么把代码弄进去。很多做硬件开发的新手&…

阅读更多 →
Proteus 8.10中STM32F103精准PWM舵机控制实战 2026/9/12 23:13:43

Proteus 8.10中STM32F103精准PWM舵机控制实战

简介:本资源是一套基于Proteus 8.10的STM32F103舵机PWM控制仿真工程,面向嵌入式初学者、电子类课程设计学生及单片机爱好者,解决无硬件条件下验证舵机控制逻辑与PWM参数调试的实践难题。压缩包含82个文件,以33个.h头文件和32个.c源…

阅读更多 →
论文润色自己改还是用平台?先判断这处改动落在哪一层 2026/9/12 23:10:43

论文润色自己改还是用平台?先判断这处改动落在哪一层

「论文润色自己改还是用平台」——这个问题问早了一步。决定该谁动手的不是自己改和平台润色谁更强,而是**这一处毛病长在哪一层,你准备把改动落在哪一层**。同一句话被改顺,可能只是错字被修掉,两边谁动手都对;也可能…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞