PyTorch猫狗识别实战:CNN、ResNet与Swin Transformer源码解析

发布时间:2026/10/1 10:36:34
PyTorch猫狗识别实战:CNN、ResNet与Swin Transformer源码解析 简介这份资源是面向计算机、人工智能及相关专业学生与开发者的猫狗识别分类项目源码包可作为毕业设计、课程设计、作业或项目立项演示的参考方案也适合具备一定基础的学习者进阶练手。压缩包共16个文件约1.67MB以7个Python脚本为核心涵盖数据读取、CNN与ResNet、Swin Transformer等模型训练与测试代码另含2个Markdown说明文档、1个docx论文、1个txt数据说明及1个pth模型权重文件结构清晰、便于按模块查阅。目前已有68人学习下载。项目代码完整、资料齐全包含设计文档与训练好的模型读者可据此理解从数据预处理、模型搭建到训练评估的完整流程并在此基础上修改网络结构或更换数据集以实现其他分类功能遇到配置与运行问题还可获得远程指导适合作为机器学习入门与实战的参考素材。1. 从一份能跑通的猫狗识别源码说起如果你正在找一份能直接跑起来的 PyTorch 图像二分类项目这份基于 Python 机器学习的猫狗识别分类源码包值得先看一眼。它把训练、测试、推理三条链路都拆成了独立脚本get_data.py负责数据加载与增强cnn.py和resnet.py分别给出自定义卷积网络与残差网络的实现swin_transformer.py补上了 Transformer 路线的对照train目录下还留了 TensorBoard 的 events 日志model目录里直接放了训练好的cnn_epoch400.pth权重。配套的说明文档、Swin-trans 笔记和一份 docx 论文基本覆盖了课程设计或毕业设计从开题到答辩的材料需求。适合谁计算机、人工智能、通信、自动化方向的学生以及想拿一个干净二分类基线做迁移实验的从业者。下面按「资源结构 → 环境与数据 → 三条模型路线 → 训练与验证 → 避坑 → 进阶技巧」的顺序拆开讲。2. 资源结构与运行环境先看清目录再动手2.1 目录里每个文件到底干什么拿到压缩包先别急着python train.py这个项目没有统一的入口脚本而是按功能散落成多个文件。先把结构理清楚后面调参和排错才不会迷路。文件/目录作用是否可直接运行get_data.py数据集加载、划分、图像增强是需先配好数据路径cnn.py自定义 CNN 网络定义否被训练脚本调用resnet.pyResNet 迁移学习网络定义否被训练脚本调用swin_transformer.pySwin Transformer 网络定义否被训练脚本调用test_cnn.py/test_resnet.py单张或批量图片推理测试是需指定权重路径show.py结果可视化损失曲线、预测展示是train/训练脚本与 TensorBoard 日志是model/cnn_epoch400.pth已训练好的 CNN 权重直接加载data.txt数据清单或路径配置视内容而定说明文档.md/Swin-trans.md环境配置与模型说明阅读用*.docx论文/设计文档阅读用这里有个血泪经验data.txt在不同项目里含义差别很大有的是图片路径列表有的是类别映射有的干脆是超参配置。打开它之前不要假设格式先head -n 20 data.txt看一眼否则后面get_data.py报的错会让你怀疑人生。2.2 环境搭建PyTorch 版本对应是第一个坎项目基于 PyTorch环境搭建是新手最容易翻车的地方。核心原则是Python 版本、PyTorch 版本、CUDA 版本三者必须对应不能随便pip install torch了事。# 建议用 conda 建独立环境避免污染系统 Python conda create -n catdog python3.8 -y conda activate catdog # 先查显卡驱动支持的 CUDA 版本 nvidia-smi # 按官方对应关系安装以 CUDA 11.3 为例 pip install torch1.10.0cu113 torchvision0.11.1cu113 \ -f https://download.pytorch.org/whl/torch_stable.html # 其余依赖 pip install numpy pillow matplotlib tensorboard tqdm逻辑说明nvidia-smi右上角显示的 CUDA Version 是驱动支持的上限不是你要装的版本装等于或低于它的都行。参数上torch1.10.0cu113里的cu113表示编译时链接的 CUDA 版本写错会装成 CPU 版训练时torch.cuda.is_available()返回 False速度差几十倍。没有独显的机器直接装 CPU 版即可pip install torch torchvision只是 400 个 epoch 会跑到你怀疑项目是不是卡死了。提示装完先跑一句python -c import torch; print(torch.__version__, torch.cuda.is_available())确认版本和 GPU 可用性再往下走。2.3 数据组织ImageFolder 的目录约定PyTorch 的ImageFolder对目录结构有硬性要求猫狗识别这类二分类项目通常按类别建文件夹。常见做法是dataset/ ├── train/ │ ├── cat/ │ └── dog/ └── val/ ├── cat/ └── dog/get_data.py里一般会用transforms.Compose做增强典型配置是训练集随机裁剪加翻转、验证集只做缩放和归一化。归一化的mean和std用 ImageNet 的[0.485, 0.456, 0.406]和[0.229, 0.224, 0.225]因为后面 ResNet 和 Swin 都是 ImageNet 预训练权重输入分布对齐才能发挥迁移学习的效果。这一步如果偷懒不做或者训练验证用了不同的归一化参数模型收敛会变得非常玄学。3. 三条模型路线CNN、ResNet、Swin 怎么选3.1 自定义 CNN理解卷积堆叠的最小闭环cnn.py里的网络是理解整个项目的地基。它通常由若干「卷积 → 激活 → 池化」块堆叠最后接全连接层输出 2 类。这种结构参数量小、训练快适合先跑通流程。import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self, num_classes2): super().__init__() self.features nn.Sequential( nn.Conv2d(3, 32, kernel_size3, padding1), # 输入3通道RGB nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # 尺寸减半 nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.AdaptiveAvgPool2d(1) # 自适应池化免去尺寸计算 ) self.classifier nn.Linear(128, num_classes) def forward(self, x): x self.features(x) x x.flatten(1) return self.classifier(x)逻辑说明padding1配合kernel_size3保证卷积后空间尺寸不变尺寸缩减全靠MaxPool2d这样每经过一个池化层特征图边长减半。AdaptiveAvgPool2d(1)是省心设计无论输入图片多大输出都是 1×1避免全连接层输入维度写死导致换分辨率就报错。参数上通道数 32→64→128 是常见的翻倍策略显存不够就整体减半。这个网络在猫狗这种类间差异明显的任务上几百个 epoch 能到 90% 以上准确率但泛化能力不如预训练模型。3.2 ResNet 迁移学习小数据集的主力方案猫狗数据集通常只有两万多张图从零训练深层网络容易过拟合resnet.py走的是迁移学习路线加载 ImageNet 预训练权重替换最后的全连接层。import torch.nn as nn from torchvision import models def build_resnet(num_classes2, freeze_backboneTrue): model models.resnet18(pretrainedTrue) # 加载ImageNet预训练权重 if freeze_backbone: for param in model.parameters(): param.requires_grad False # 冻结特征提取层 in_features model.fc.in_features model.fc nn.Linear(in_features, num_classes) # 只训练新的分类头 return model逻辑说明pretrainedTrue会下载 ImageNet 权重第一次运行需要联网下不到就手动放到~/.cache/torch/hub/checkpoints/。freeze_backboneTrue时只更新最后的fc层训练快、显存省适合数据量小的情况数据量上万后可以解冻后面几个 stage 做微调准确率还能再涨一两个点。参数上resnet18是最轻的选择显存够可以换resnet50但要注意model.fc这个名字在 resnet50 上是一样的不用改代码。学习率方面冻结时用 1e-3解冻微调时降到 1e-4否则预训练权重会被大学习率冲垮。3.3 Swin Transformer想冲高准确率的对照实验swin_transformer.py和Swin-trans.md说明项目作者把 Transformer 路线也纳入了对比。Swin 通过窗口注意力和层级结构在图像分类上常比同量级 CNN 更强但代价是显存和训练时间。# 常见做法是借助 timm 库加载 Swin # pip install timm import timm import torch.nn as nn def build_swin(num_classes2): model timm.create_model( swin_tiny_patch4_window7_224, pretrainedTrue, num_classesnum_classes ) return model逻辑说明swin_tiny_patch4_window7_224里的224是输入分辨率Swin 对输入尺寸比 CNN 敏感最好统一到 224×224。num_classes2直接替换分类头timm 会自动处理。参数上Swin 训练需要更小的学习率和更长的 warmup常见配置是lr1e-4、weight_decay0.05batch size 受显存限制往往只能开到 16 或 32。如果你的显卡只有 6G 显存跑 Swin 大概率 OOM这时候老老实实用 ResNet 更实际。三条路线的取舍可以概括为CNN 用来理解原理ResNet 用来交作业和做基线Swin 用来写论文里的对比实验。4. 训练、验证与推理把流程跑成闭环4.1 训练循环的关键参数训练脚本在train/目录下核心是标准的 PyTorch 训练循环。下面这段是骨架重点看参数怎么设。import torch from torch import nn, optim device torch.device(cuda if torch.cuda.is_available() else cpu) model build_resnet(num_classes2).to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam( filter(lambda p: p.requires_grad, model.parameters()), lr1e-3, weight_decay1e-4 ) scheduler optim.lr_scheduler.StepLR(optimizer, step_size20, gamma0.1) for epoch in range(400): model.train() for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() scheduler.step() # 每个epoch后在验证集上评估保存最优权重逻辑说明filter(lambda p: p.requires_grad, ...)只把需要更新的参数交给优化器冻结 backbone 时这一步很关键否则优化器会带着一堆requires_gradFalse的参数空转。weight_decay1e-4是 L2 正则抑制过拟合。StepLR每 20 个 epoch 把学习率乘 0.1让后期收敛更稳。400 个 epoch 是项目里cnn_epoch400.pth的由来但用预训练 ResNet 通常 30 到 50 个 epoch 就收敛了没必要照搬这个数字。验证阶段记得model.eval()加torch.no_grad()否则 BatchNorm 和 Dropout 行为不对评估结果会偏低。4.2 用 TensorBoard 看训练曲线项目Logger目录下的events.out.tfevents.*就是 TensorBoard 日志这是排查训练问题的黑匣子。tensorboard --logdir ./Logger --port 6006启动后浏览器打开localhost:6006重点看两条曲线训练 loss 是否稳定下降验证准确率是否在某个 epoch 后停滞甚至下降。如果训练 loss 一直降但验证准确率不涨就是过拟合该加数据增强或提前停止如果两条都不动多半是学习率太小或数据归一化有问题。参数上--logdir指向日志目录--port换端口可以同时开多个实验对比。4.3 推理测试加载权重验证单张图片test_cnn.py和test_resnet.py负责推理核心是加载权重后走一遍前向。import torch from PIL import Image from torchvision import transforms model build_resnet(num_classes2) model.load_state_dict(torch.load(model/cnn_epoch400.pth, map_locationcpu)) model.eval() tf transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) img tf(Image.open(test.jpg).convert(RGB)).unsqueeze(0) with torch.no_grad(): pred model(img).argmax(dim1).item() print(cat if pred 0 else dog)逻辑说明map_locationcpu让权重能在没有 GPU 的机器上加载避免CUDA error。unsqueeze(0)是给单张图补上 batch 维度模型永远按 batch 处理。convert(RGB)防止灰度图或带 alpha 通道的 PNG 导致通道数不匹配。注意推理时的预处理必须和验证集完全一致少一个归一化预测结果就可能反过来。5. 避坑与常见问题排查5.1 报错 CUDA out of memory现象训练刚开始或跑到某个 batch 时抛RuntimeError: CUDA out of memory。原因batch size 太大、模型太大或显存被其他进程占用。解决先把 batch size 减半用nvidia-smi看是否有残留进程必要时torch.cuda.empty_cache()Swin 路线显存吃紧就换 ResNet。5.2 准确率卡在 50% 不动现象训练 loss 缓慢下降但验证准确率始终在 0.5 附近等于随机猜。原因最常见是标签和输出维度对不上或者归一化参数用错。解决检查num_classes是否为 2打印几个 batch 的 label 看是不是 0/1确认训练和验证用了同一套Normalize参数。5.3 加载权重报 key 不匹配现象load_state_dict抛Missing key(s)或Unexpected key(s)。原因保存权重时用了nn.DataParallel键名多了module.前缀或者模型结构改过。解决用state_dict {k.replace(module., ): v for k, v in state_dict.items()}去掉前缀再load_state_dict(state_dict, strictFalse)。5.4 数据加载报找不到文件现象FileNotFoundError或ImageFolder报Found 0 files。原因data.txt里的路径是作者机器的绝对路径或者目录层级不符合ImageFolder约定。解决把数据路径改成相对路径确认每个类别文件夹下确实有图片ImageFolder只认「类别名做文件夹名」这一种结构。5.5 训练速度异常慢现象一个 epoch 要跑几十分钟。原因模型和数据没搬到 GPU或者num_workers设成 0 导致数据加载成瓶颈。解决确认model.to(device)和imgs.to(device)都执行了DataLoader的num_workers设成 4 或 8pin_memoryTrue。6. 进阶技巧把这份源码改成你自己的项目跑通只是第一步真正有价值的是把它改成能写进论文或落地的东西。第一个技巧是替换数据集做迁移把dataset/train/cat和dog换成你的两类数据比如缺陷/正常、戴口罩/未戴口罩get_data.py里的类别数不用改因为ImageFolder会自动按文件夹名生成类别索引。第二个技巧是用混淆矩阵替代单一准确率猫狗数据均衡时准确率够用但你的数据一旦不均衡准确率会骗人加一段sklearn.metrics.confusion_matrix输出能看清模型到底偏向哪一类。from sklearn.metrics import confusion_matrix, classification_report all_preds, all_labels [], [] model.eval() with torch.no_grad(): for imgs, labels in val_loader: preds model(imgs.to(device)).argmax(dim1).cpu() all_preds.extend(preds.tolist()) all_labels.extend(labels.tolist()) print(confusion_matrix(all_labels, all_preds)) print(classification_report(all_labels, all_preds, target_names[cat, dog]))逻辑说明把验证集所有预测收集起来一次性算指标比逐个 batch 算准确率更可靠。classification_report会给出每一类的 precision、recall、f1哪一类拖后腿一目了然。参数上target_names要和类别索引顺序对应ImageFolder默认按文件夹名字母序cat在前dog在后写反了报告就全错。第三个技巧是冻结层数的消融实验论文里想证明迁移学习的价值可以对比「只训练分类头」「解冻最后两个 stage」「全量微调」三组用同一份数据跑把准确率列成表。这是审稿人爱看的对照也是你真正理解迁移学习边界的过程。我一般会固定随机种子torch.manual_seed(42)否则三组之间的差异可能被随机性淹没得出错误结论。最后一个习惯每次改完代码先拿 100 张图的小子集跑 2 个 epoch确认流程通了再上全量数据。从那以后我每次动数据管道或换模型都强制走一遍这个小闭环省下的调试时间比什么都值。希望这份拆解帮到你源码包里的说明文档和论文可以对照着看遇到配置问题先查环境版本对应关系八成能自己解决。本文还有配套的精品资源点击获取

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询