尧图精选

CNN与VGG实现珊瑚识别:三个PyTorch脚本打通图像分类全流程

🕒 发布时间:2026/10/2 2:35:41 📁 来源:尧图网络
简介PyTorch环境下的VGG卷积神经网络珊瑚种类识别代码包面向深度学习入门者与海洋生物图像分类应用场景适合希望掌握CNN训练全流程的读者。压缩包共8个文件包含3个Python脚本01生成txt.py负责生成训练列表02CNN训练数据集.py完成模型训练03pyqt界面.py提供可视化交互、3张数据集类别提示图、1个环境依赖requirements.txt及1个说明文档整体仅213KB结构精简易读。已有70人学习下载可用于教学实践或小型图像分类任务。代码逐行配有中文注释并附说明文档指导环境安装与数据准备用户只需自行搜集珊瑚图片放入对应类别文件夹即可直接训练且类别文件夹可自由增删便于扩展分类任务。整体不包含数据集图片轻量实用是快速上手PyTorch图像分类的理想参考。1. VGG与CNN珊瑚识别一份能直接跑通的三脚本资源做珊瑚种类识别很多人不是缺模型而是缺一个能改得动的VGG模型工程。这套资源把CNN卷积神经网络的图像分类链路拆成三个py文件——数据列表生成、模型训练、PyQt界面推理——每一行都带中文注释环境装好就能照着跑。它解决的是“代码读得懂、跑得通、换数据也能用”的问题适合第一次用PyTorch做分类的新手也适合课程设计和毕设里需要快速出结果的人。注意它不含数据集图片你需要按类别文件夹自己放图反过来讲类别数量完全由你决定不止脑珊瑚、软珊瑚、扇形珊瑚这三种。2. 网络结构与训练流程先看懂三个Py文件的职责边界这套代码虽然只有三个文件但它们不是并列关系而是流水线。01生成txt.py负责把文件夹里的图片整理成训练列表02CNN训练数据集.py读这份列表驱动VGG网络在GPU或CPU上跑几十个epoch03pyqt界面.py把训练好的权重包装成小窗口程序。跑的时候必须按01→02→03的顺序来03依赖02产出的pth文件02依赖01产出的txt任何一个跳步都会报“文件不存在”或者加载时维度对不上。2.1 01生成txt.py把文件夹变成长度对齐的训练列表很多人不理解为什么训练前要先“生成txt”直接用torchvision的ImageFolder不行吗ImageFolder能自动按文件夹结构取标签但划分训练集和验证集要靠random_split比例、随机种子都得额外控制中途想看某个类有多少张图还得自己数。01把路径和标签先落盘成txt训练脚本每次读同一份文件结果可复现。实际代码跟下面这个结构基本一致import os from sklearn.model_selection import train_test_split data_root 数据集 # 选出所有子文件夹作为类别 classes [d for d in sorted(os.listdir(data_root)) if os.path.isdir(os.path.join(data_root, d))] class_to_id {c: i for i, c in enumerate(classes)} print(类别映射, class_to_id) lines [] for cls in classes: cls_dir os.path.join(data_root, cls) count 0 # 逐个统计并写入每张图片的路径和标签 for img_name in sorted(os.listdir(cls_dir)): if img_name.lower().endswith((.jpg, .jpeg, .png)): count 1 path os.path.join(cls_dir, img_name) lines.append(f{path} {class_to_id[cls]}) print(f类别 {cls}: {count} 张) # 按 8:2 切分训练集和验证集 train_lines, val_lines train_test_split( lines, test_size0.2, random_state42 ) with open(train.txt, w, encodingutf-8) as f: f.write(\n.join(train_lines)) with open(val.txt, w, encodingutf-8) as f: f.write(\n.join(val_lines)) print(f总图片数: {len(lines)}训练: {len(train_lines)}验证: {len(val_lines)})逻辑说明sorted()对类别排序是必要的因为dict的枚举顺序取决于文件夹名字名字一旦变动标签编号就全变了后面训练脚本里的对照关系也会跟着错。后缀过滤把jpg、jpeg、png之外的格式挡在门外避免把系统生成的临时文件收进训练集也避免非图片文件在PIL打开时直接中断训练。train_test_split返回两条列表test_size表示验证集占比对小样本的珊瑚识别来说0.2够用。这里要特别留意数据集文件夹里自带的提示图脑珊瑚1.jpg这类如果不删也会被写进txt导致该类别样本量虚高。所以跑之前先手动删掉或替换成真实训练样本。random_state42固定随机种子保证每次运行01得到完全相同的划分这是后面调参对照实验结果的前提。test_size这个参数几万张的大数据集可以放到0.3珊瑚识别这种小数据集0.2已经够验证用了。2.2 02CNN训练数据集.pyVGG骨架与迁移学习的落点02是整个资源的重头戏。它做的事情可以概括为把在ImageNet上预训练好的VGG16搬过来替换最后那个输出1000类的全连接层换成输出类别数的新分类头然后用珊瑚图片做微调。这是迁移学习在深度学习CNN图像分类里的标准套路也是小数据集场景下最稳的做法——让你从头训练一个VGG几千张图远远不够但微调预训练权重每类几十张就能看效果。核心训练代码结构类似这样import torch import torch.nn as nn from torchvision import models, transforms from torch.utils.data import Dataset, DataLoader from PIL import Image class CoralDataset(Dataset): def __init__(self, txt_path, transformNone): self.samples [] # 逐行读取 图片路径 标签编号 with open(txt_path, r, encodingutf-8) as f: for line in f: parts line.strip().split() if len(parts) 2: self.samples.append((parts[0], int(parts[1]))) self.transform transform def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] img Image.open(path).convert(RGB) if self.transform: img self.transform(img) return img, label # VGG 的标准预处理缩放、归一化 transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 加载预训练VGG16替换最后一层分类器 model models.vgg16(pretrainedTrue) num_classes 3 model.classifier[6] nn.Linear(4096, num_classes) train_loader DataLoader(CoralDataset(train.txt, transform), batch_size16, shuffleTrue) val_loader DataLoader(CoralDataset(val.txt, transform), batch_size16, shuffleFalse) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr0.001, momentum0.9) for epoch in range(30): model.train() running_loss 0.0 for imgs, labels in train_loader: optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() print(fepoch {epoch1:02d}, loss {running_loss/len(train_loader):.4f})参数说明classifier[6]是VGG16分类分支的最后一个全连接层原来输出1000这里改成num_classes。这个索引是VGG系列最容易写错的地方VGG19的classifier结构跟VGG16一样索引同样是6。pretrainedTrue时torchvision会自动下载在ImageNet上训好的权重这一步值得等别跳过下载完成后存在本地缓存目录后面不再重复拉取。CoralDataset从01生成的train.txt逐行读数据__getitem__里用PIL打开图片convert(RGB)保证三通道统一即使源图是灰度图也不会报错。transform里的Normalize用的是ImageNet数据集的均值和标准差这是VGG预训练模型的标配不能用错否则特征分布对不上。SGD的momentum0.9是多年经验值lr0.001对小数据集是个偏保守但稳定的起点。epoch设30图片少时几十个epoch就能看到趋势不需要一上来就跑100轮。提示val_loader在这里先保留训练循环里暂时没用它。后面做调参时会在验证集上统计准确率这一行就是那个位置的地基。2.3 03pyqt界面.py推理部署的调用链03是三个文件里最容易被忽略但最出活的一个。训练脚本产出的pth权重只是冷冰冰的字典03用PyQt5做了一个小窗口点按钮选图片界面上直接显示预测的类别和置信度。这对课程设计的演示环节几乎是刚需一张截图就能让评委看到从图片输入到分类输出的完整闭环。推理部分的逻辑比训练简单因为不需要反向传播也不需要Dataset核心代码是这些import torch from torchvision import models, transforms from PIL import Image # 先从pth里读出配置信息再按配置重建模型 checkpoint torch.load(coral_vgg.pth, map_locationcpu) model models.vgg16(pretrainedFalse) model.classifier[6] torch.nn.Linear(4096, checkpoint[num_classes]) model.load_state_dict(checkpoint[state_dict]) model.eval() transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) def predict(img_path): img Image.open(img_path).convert(RGB) img_tensor transform(img).unsqueeze(0) with torch.no_grad(): logits model(img_tensor) prob torch.softmax(logits, dim1) cls_id torch.argmax(prob, dim1).item() return cls_id, prob[0][cls_id].item()这里的关键是load_state_dict之前模型结构必须和保存时完全一致。pretrainedFalse不是为了省下载时间而是避免又拉一次ImageNet权重——你要的是本地那个训练好的pth。map_locationcpu保证在没有CUDA的机器上也能加载GPU训练出来的权重这是演示环境最常见的坑。model.eval()会把dropout和batch norm切到推理模式不写这行预测结果每次都不一样是个典型的玄学错误。界面部分就是PyQt5的标准三件套QFileDialog选图、QPushButton触发predict、QLabel显示结果。拿到cls_id后通过pth里的class_to_id映射反查类别名显示成“脑珊瑚 87.2%”这种格式。映射从训练时保存的配置里读不要写死在界面代码里否则以后新增类别时界面也要跟着改。3. 数据集整理与环境安装类别文件夹和requirement.txt的约定很多人在装环境上花的时间比训练还多。这份资源不含数据集图片只给了每个类别的提示图占位所以你要先解决两件事图片从哪来、环境怎么搭。这两件事都有固定套路按顺序来基本不会卡住。3.1 类别文件夹不是写死的脑珊瑚、软珊瑚之外怎么加类数据集文件夹的约定很简单每一层子文件夹就是一个类别文件夹名就是类别名。原始结构是脑珊瑚、软珊瑚、扇形珊瑚各一个文件夹每个文件夹里有一张提示图告诉你图片该放哪。要做的就是把搜集到的图片直接拖进对应文件夹脚本不限制数量一张也好一百张也好。数据集/ ├── 脑珊瑚/ │ ├── 脑珊瑚1.jpg提示图换成自己的训练图 │ └── img_001.jpg ├── 软珊瑚/ │ ├── 软珊瑚1.jpg提示图 │ └── img_002.jpg └── 扇形珊瑚/ ├── 扇形珊瑚1.jpg提示图 └── img_003.jpg新增类别的操作更直接在数据集文件夹下新建一个目录把图放进去。01生成txt.py用os.listdir遍历新目录会自动进入类别列表class_to_id映射自动多一号。这样设计的好处是代码里不需要维护写死的类别数组坏处是任何人都能通过改文件夹名来改标签所以训练前务必把文件夹名字敲定中途别改名。注意一个容易忽略的点类别顺序由文件名排序决定不是文件放入顺序。脑珊瑚排在第0软珊瑚第1扇形珊瑚第2。如果你加了“石珊瑚”文件夹它会按字符顺序插进去编号可能不是你以为的那个。我一般会先把类别名加上数字前缀比如“0_脑珊瑚、1_软珊瑚”这样排序结果好预期训练脚本里category_id的连续性也不会乱。每类图片数量的经验值每类少于20张VGG的迁移学习优势发挥不出来过拟合几乎是必然的每类50张以上结果基本能看想达到演示级别的稳定每类最好100张上下。图片质量比数量重要模糊的、严重偏色的样本宁可不放。3.2 图片数量与尺寸VGG输入224x224背后的预处理逻辑VGG的卷积层对输入尺寸其实没有硬性要求但模型尾部是全连接层它的权重矩阵在预训练时固定成224x224对应的2048维和4096维所以输入必须统一缩放到224。更深一层的原因是预训练权重在224x224上提取的特征统计已经固化你送一张512x512的图进去理论上也能跑通但效果和稳定性都不如224。手机原图直接送进网络不现实。常见做法分两种一是transforms.Resize((224,224))直接拉伸速度快但改变宽高比珊瑚是近似团块状的物体拉伸导致的畸变影响不算大二是Resize到256再CenterCrop得到224保留更多原图比例信息等于顺便做了一次数据增强。两种写法在2.2和2.3里都用过我建议新手上手先选第一种后面想提升准确率再换第二种。这里有一个细节CenterCrop裁剪的是图片正中心如果珊瑚主体在图片边缘裁掉就亏了。所以用CenterCrop之前最好先人工检查一下数据集的构图习惯。如果图片里珊瑚不在中心老老实实用Resize拉伸不要为了数据增强反而丢掉主体。3.3 requirement.txt环境安装Python和PyTorch版本怎么配环境这一块摘要里已经点明推荐路径Anaconda装Python 3.7或3.8PyTorch装1.7.1或1.8.1。这套版本组合的好处是torchvision的API稳定pretrained参数还没被废弃网上教程对得上号踩坑概率最低。用Anaconda管理还有一个好处环境坏了删掉重建就行不祸害系统Python。conda create -n coral python3.8 conda activate coral pip install -r requirement.txtrequirement.txt的内容大致是这些包版本号按PyTorch 1.8.1的兼容矩阵来配torch1.8.1 torchvision0.9.1 numpy Pillow scikit-learn PyQt5torch和torchvision这两个版本必须配对1.8.1对应0.9.1装错会报底层符号找不到一类的错误排查起来很费劲。用CPU训练的话把PyTorch的CPU版本装好就行GPU机器想用显卡加速按自己的CUDA版本重新装对应编号的torch。没有NVIDIA显卡就安心用cpu版珊瑚这种小数据量CPU跑几十个epoch也就一顿饭的工夫不影响出结果。4. 调参思路与训练日志epoch、学习率、loss震荡怎么判对新手来说能跑通只是第一步。跑完之后怎么判断模型好坏怎么把准确率从70%拉到90%靠的就是那几个参数和读懂日志的能力。这一章把数据划分、学习率、日志判读三个最常见的操作点拆开说。4.1 数据划分比例小样本下的train-val-test拆分珊瑚识别这种场景每类图片通常不会很多数据划分直接决定模型可信度。01生成txt.py默认按8:2切分这在图片总量超过100张时够用但图片少时会遇到验证集太小的问题——验证集只有10张图正确率80%和90%的差别就是一张图结果很不可靠。我的划分建议按下表来每类图片量建议划分操作要点200张以上训练8 : 验证1 : 测试1另留一个测试文件夹训练全部结束后只碰一次50到200张训练8 : 验证2每类手动挑20张当测试50张以下K折交叉验证sklearn的KFold切5折每折训练一次取平均小数据集尤其要注意测试集必须在调参全部结束后再使用否则拿测试集反复试探测试成绩就是过拟合的产物。我一般会在改test_size的同时换一个random_state做多组对比这样能看出哪些指标是运气、哪些是模型真实能力。改划分比例时动的是01生成txt.py里的test_size改完必须重新生成txt再重新训练这一点很容易忘。另一个细节如果训练脚本里写死了class_to_id {脑珊瑚:0, 软珊瑚:1, 扇形珊瑚:2}而01按文件夹名排序生成的顺序恰好不同标签就对不上了。最稳妥的办法是让训练脚本读取01输出的映射关系或者干脆在01里把映射打印出来训练前用肉眼看一眼有没有错位。4.2 学习率与batch size小数据集上的稳妥起点迁移学习的调参有个不成文的规矩预训练权重已经在海量图片上学到了通用的边缘、纹理、形状特征微调时不能一把梭用同一个学习率。常见做法是把模型拆成两组参数——骨架部分用小学习率微调新加的分类头用稍大学习率从头学。params [ {params: model.features.parameters(), lr: 0.0001}, {params: model.classifier.parameters(), lr: 0.001}, ] optimizer torch.optim.SGD(params, momentum0.9)逻辑说明model.features是VGG前13个卷积层加池化的集合它的特征对识别珊瑚的纹理和轮廓仍然有效用0.0001这样的小步长去微调避免把预训练知识冲掉model.classifier里的老权重被换成了新分类头随机初始化起点差需要0.001更高一点的学习率才能快速收敛。这两个数值是图像分类迁移学习里出现频率最高的起点不是唯一正确值但作为基准非常可靠。batch size在小数据集上的选择逻辑越大越稳定但越容易过拟合越小的batch梯度噪声大反而起到正则作用。16是个入门友好值显存不够降到8分类头效果不稳定可以升到32观察。学习率调度lr_scheduler在epoch较多时会有帮助比如第20轮把学习率除以10让loss在后期更精细地下降。PyTorch里用scheduler.step()一行就能实现不需要手动改学习率。4.3 训练日志怎么读loss不降、acc震荡的几种可能训练日志是判断深度学习CNN代码跑没跑对的唯一依据。02训练脚本目前只打印训练loss这远远不够至少要把验证集准确率也打出来否则根本不知道模型是不是在过拟合。给训练循环加一小段验证逻辑model.eval() val_correct 0 val_total 0 with torch.no_grad(): for imgs, labels in val_loader: out model(imgs) pred torch.argmax(out, dim1) val_correct (pred labels).sum().item() val_total labels.size(0) val_acc val_correct / val_total print(fepoch {epoch1:02d} | train loss {running_loss/len(train_loader):.4f} | val acc {val_acc:.4f})这段代码是判断模型状态的仪表盘。train loss持续下降但val acc上不去说明过拟合先怀疑图片数量太少或数据增强不够train loss从一开始就不降先怀疑学习率太高导致梯度爆炸再看是不是pretrained权重没加载成功——pretrainedFalse加上load_state_dict路径写错模型就会用随机权重从头学VGG在几百张图上从零训练几乎不可能收敛val acc每一轮剧烈震荡基本是batch太小或学习率太大模型在最优解附近反复弹跳。日志里还有一个常被忽略的细节torch.argmax在类别数大于10时输出可能和标签编号错位先打印pred和labels的前20个值肉眼核对一次。出现NaN的loss也很常见原因多半是图片像素值没有归一化直接送进网络检查transform里有没有ToTensor这一行。5. CNN珊瑚识别的避坑记录五条血泪经验从哪里来这些坑我在做图像分类课程设计时几乎全踩过一遍有些是数据问题有些是环境问题共同特点是报错信息不直观新手容易在错误方向浪费时间。每条都按现象、原因、解决的顺序写照着排查就行。5.1 图片搜集与格式后缀、尺寸、损坏文件的三个坑坑一提示图被当成训练样本模型学了个寂寞。 现象训练时每类样本数量虚高训练结束后拿一张新珊瑚图片测试预测结果机械地偏向某个固定类别。 原因数据集文件夹里自带的脑珊瑚1.jpg、软珊瑚1.jpg、扇形珊瑚1.jpg提示图没删被01生成txt.py一起收进训练列表。占位图只有一张导致这一类样本比例失调分类器学会了“凡是模糊的就猜这一类”。 解决放训练数据前把提示图删掉或直接覆盖顺手对照01脚本打印出来的“类别 XXX: N 张”发现某类只有1张就说明占位图还在。坑二图片格式不统一训练中途PIL报错。 现象训练跑了几个batch突然弹出PIL.UnidentifiedImageError或者OSError: image file is truncated。 原因网上搜集的图片里混了webp、gif动图或者从聊天软件里另存的文件后缀是.jpg但实际编码不对PIL打开失败。 解决写个小脚本提前扫一遍数据集把坏文件隔离出来。代码很短from PIL import Image import os for root, _, files in os.walk(数据集): for name in files: path os.path.join(root, name) try: img Image.open(path) img.verify() except Exception as e: print(f损坏文件: {path} - {e})img.verify()不会真正解码像素速度快检查几百张图几秒钟就完事。凡是打印出来的路径要么删除要么重新下载。另一个更隐蔽的情况是文件扩展名正确但通道是灰度图convert(RGB)可以兜底所以Dataset里这行别省。坑三训练和推理的预处理不一致测试准确率崩盘。 现象训练日志显示val acc有90%但把图片丢进03界面里预测结果错得离谱。 原因训练脚本transform用的ResizeToTensorNormalize推理代码里漏了Normalize或者Normalize的参数少打了个小数点。输入分布不同模型接收的特征和训练时不在一个量级上。 解决把transform定义抽出来放到两个脚本都能引用的公共位置或者直接复制粘贴后逐行对比。我后来的习惯是训练和推理共用同一个transform常量绝不各写一份。5.2 环境与训练过程版本不一致和权重复用的问题坑四PyTorch版本API变化老代码跑出新报错。 现象运行02时torchvision弹出一行UserWarning内容大致是pretrained参数已废弃后续加载权重也可能失败。 原因PyTorch 1.13之后torchvision改变了预训练权重的加载方式pretrained参数被替换成weights枚举。老代码的pretrainedTrue虽然还能触发下载但行为已经不一致。 解决最省事的方法是按摘要推荐装1.7.1或1.8.1老API完全兼容。如果一定要用新版把models.vgg16(pretrainedTrue)改成models.vgg16(weightstorchvision.models.VGG16_Weights.IMAGENET1K_V1)两者含义相同。坑五保存和加载时的模型结构对不上。 现象03界面启动时加载pth文件报size mismatch for classifier.6后面跟着一行shape对不上的信息。 原因训练时num_classes设成3pth里最后一层是3x4096加载时如果分类头初始化成了默认1000类结构不匹配load_state_dict自然拒绝。 解决保存权重时把配置信息一起存进去加载时先读配置再重建模型。推荐的保存写法torch.save({ state_dict: model.state_dict(), class_to_id: class_to_id, num_classes: num_classes, }, coral_vgg.pth)加载时先从pth里读num_classes再建模型结构永远对得上。顺手把class_to_id也存下来03界面显示类别名就不用写死了。这个习惯是我最值钱的一条经验因为每次换数据集只要改文件夹和num_classes界面代码一行不用动。6. 用十张新图片验证模型取证比训练更花心思训练脚本打印的val acc再好看也说明不了最终效果——验证集来自同一个数据集文件夹分布和训练集太接近。我习惯在全部训练结束后单独拿十张网络上找来的、没参与训练的珊瑚图片做一次盲测这个动作花不了几分钟但能暴露一大半隐藏问题。批量验证的脚本和02的推理部分几乎一样区别是按类别文件夹遍历后逐张预测并打印import torch import os from torchvision import models, transforms from PIL import Image transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) checkpoint torch.load(coral_vgg.pth, map_locationcpu) model models.vgg16(pretrainedFalse) model.classifier[6] torch.nn.Linear(4096, checkpoint[num_classes]) model.load_state_dict(checkpoint[state_dict]) model.eval() test_root 验证集 for cls in os.listdir(test_root): cls_dir os.path.join(test_root, cls) for img_name in os.listdir(cls_dir): path os.path.join(cls_dir, img_name) img Image.open(path).convert(RGB) with torch.no_grad(): prob torch.softmax(model(transform(img).unsqueeze(0)), dim1) pred_id torch.argmax(prob, dim1).item() conf prob[0][pred_id].item() print(f{cls} - 预测{list(checkpoint[class_to_id].keys())[pred_id]} 置信度{conf:.2f})看结果时我最关注两个点一是置信度低于0.6的预测我直接视为模型没认清——它可能靠运气蒙对了类别但特征提取是失败的二是搞混的组合软珊瑚和扇形珊瑚在纹理上很接近如果反复在这两类之间出错说明VGG的浅层特征对纹理区分不够此时加数据增强随机旋转、水平翻转、色彩抖动比加epoch更有效。等十张图全部预测正确且置信度都在0.7以上我才会认为模型的训练闭环真正成立。我最早做珊瑚识别时训练acc冲上95%就高兴得直接拿去演示结果现场被一张水下偏蓝的照片狠狠打脸——模型把脑珊瑚认成了软珊瑚。从那以后我每次训练完都强制走一遍十图抽检先看置信度再看预测类别低于0.6的全部视为没认清这个习惯已经救了我好几次。这套资源三个py文件按01、02、03的顺序跑通配合说明文档把环境装好按顺序走一遍就能看到完整效果。希望帮到你。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联 返回资讯列表 →