十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

VGG16迁移学习+Pytorch实现珊瑚图像分类完整项目

VGG16迁移学习+Pytorch实现珊瑚图像分类完整项目 简介面向深度学习初学者与珊瑚研究相关的开发者一套基于PyTorch的CNN珊瑚种类识别代码包目的是让用户避开复杂环境配置直接体验从图片整理到模型训练的完整流程。整个压缩包共8个文件体积仅213KB包含3个Python脚本对应生成训练列表、CNN训练、PyQt可视化界面、requirements.txt环境依赖、说明文档以及3张用于指示数据存放位置的示例图片结构紧凑且分工清晰。已有70人学习下载。代码中每一行均附有中文注释并专门讲解了Anaconda、Python与PyTorch的版本搭配建议如Python3.7/3.8配合PyTorch1.7.1/1.8.1新手也能按说明自行完成环境搭建和数据准备。下载后配合自备图片按脚本提示即可完成一次完整的CNN图像分类训练非常适合作为课程设计、毕业设计或深度学习入门练习。1. 为什么珊瑚识别要用 VGG16 迁移学习而不是从零训练 CNN珊瑚种类识别看起来比猫狗分类简单但实际做起来坑很多。水下照片的光照偏色、水流带来的模糊、拍摄角度差异都会让同类珊瑚的纹理和颜色产生较大变化。如果直接拿随机初始化的 CNN 去训练几百张图片根本撑不住几百万参数很容易在训练集上过拟合到 99%换一张真实环境图就完全失灵。VGG16 的预训练权重本身来自 ImageNet前几层已经学会了边缘、纹理、色彩过渡等通用特征在珊瑚这种细粒度识别任务上只需要把最后分类层换成自己的类别数再微调部分卷积层就能用小规模数据集稳定收敛。这个项目正好选了 PyTorch 实现三个脚本把“生成标签 - 训练模型 - 图形界面识别”串在一起代码带逐行注释适合想完整跑通一条 CNN 工程链路的人。你不需要懂花哨的新架构把 VGG16 的迁移学习吃透就足够处理这种十到二十类的图像分类需求。2. 数据集目录设计与 01生成txt.py 的标签管线2.1 先按文件夹区分类别再生成路径清单这个项目刻意不打包数据集图片而是要求你自己按类别建文件夹。最常见的组织方式是在项目根目录下放一个data文件夹里面每个子文件夹代表一个珊瑚类别比如brain_coral、soft_coral、fan_coral。子文件夹内部直接放该类别对应的.jpg图片。这么做的好处是文件夹名就是标签你不用额外维护一张 CSV 表新增一个类别只需要新建一个文件夹重新跑一次脚本。01生成txt.py 干的事就是扫描这些子文件夹把每张图片的完整路径和它的类别索引写到文本文件里。为什么必须先生成 txt 而不是直接在训练时扫文件夹因为 PyTorch 的ImageFolder虽然也能按文件夹读数据但它在排序、缓存和跨设备复现上不如自定义Dataset txt 列表可控。尤其当你要做分层采样或过滤坏图时txt 路径清单更加直观中途出现问题也容易定位到具体是哪一行数据。2.2 01生成txt.py 的逐行拆解下面是一个符合项目描述的典型实现代码里已经写了中文注释方便对照你的源文件理解import os # 数据集根目录下面每个子文件夹是一个类别 data_root data # 输出的训练列表文件 output_file train.txt # 需要过滤的图片后缀 valid_ext (.jpg, .jpeg, .png) with open(output_file, w, encodingutf-8) as f: # 按类名排序保证类别索引稳定 class_names sorted(os.listdir(data_root)) for class_idx, class_name in enumerate(class_names): class_dir os.path.join(data_root, class_name) if not os.path.isdir(class_dir): continue for img_name in os.listdir(class_dir): if not img_name.lower().endswith(valid_ext): continue img_path os.path.join(class_dir, img_name) # 每行格式图片绝对路径 空格 类别索引 f.write(f{img_path} {class_idx}\n)这段代码的核心逻辑是os.listdir遍历文件夹用enumerate给每个类别分配一个从 0 开始的数字索引。路径没有写成绝对路径而是用相对路径data/brain_coral/1.jpg这样项目挪动位置后不需要改代码。训练脚本读取这个文件时会把图片路径和标签分别解析出来。如果你希望路径是绝对的可以把os.path.join换成os.path.abspath但要注意换机器后路径失效的问题。01生成txt.py里的一个关键点是类别索引和文件夹名的映射。排序用sorted很重要否则每次运行生成的索引顺序会不同比如本来brain_coral是 0下次可能变成 1导致训练时的标签和界面显示完全对不上。项目里每个文件夹内有一张“提示图”它的作用只是告诉你把图片放哪里不是用来训练的所以一定要过滤掉文件名中带有_提示或.png但不属于正式图片的文件。上面代码通过判断文件后缀的方式可以顺带忽略那些非图片格式的临时文件。2.3 自定义类别时只需要改文件夹名如果你不想用三种珊瑚想改成六种不需要动训练脚本。只需要在data下新建三个文件夹把图片丢进去重新跑01生成txt.py。新的类别索引会自动分配但是类别顺序是按文件夹名字母排序来的所以建议文件夹名用你能一眼识别出的英文或拼音不要用 “类别1”、“类别2” 这种无意义命名。训练脚本里一般还会有一个类别名称列表用来在预测时把数字索引翻译成中文显示。这个列表需要和文件夹的排序保持一致例如类别索引文件夹名界面显示名称0brain_coral脑珊瑚1soft_coral软珊瑚2fan_coral扇形珊瑚当你新增类别时记得同步更新训练脚本或界面脚本里的这个列表。如果文件夹名直接用中文也可以让代码从文件夹名读取显示名称但 Windows 和 Linux 对中文编码的处理不同项目里的requirement.txt和说明文档应该也标注过这一点。我自己的习惯是文件夹用英文界面显示用中文这样既避免due to source code encoding的问题又方便后续打包成 exe。3. 02CNN训练数据集.py数据增强、VGG 微调和关键超参3.1 数据加载与图像预处理训练脚本第一步是从train.txt读取图片路径和标签然后交给 PyTorch 的DataLoader。由于 VGG16 的输入尺寸是 224x224但珊瑚图片原始尺寸可能从几百像素到几千像素都有所以需要先做 Resize 再 CenterCrop。常见做法是先缩放到 256x256再随机裁剪到 224x224相当于给模型提供了轻微的位置扰动比直接拉伸更稳。数据增强部分不能只靠翻转。海洋照片通常有偏色和亮度不均我建议增加ColorJitter让模型对色温变化不敏感。下面是典型的预处理代码from torchvision import transforms # 训练集增强 train_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.3, contrast0.3, saturation0.3), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 验证集和测试集只做缩放和裁剪不做随机增强 val_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])Normalize里的均值和标准差是 ImageNet 预训练模型的标准值迁移学习时不能随意改。如果使用自己的均值相当于把模型输入分布改变了预训练权重会失效。这里RandomCrop和CenterCrop的尺寸都是 224对应 VGG16 的fc6层要求。训练时用随机裁剪验证时用中心裁剪这是图像分类任务的标准做法避免了验证时随机性带来的指标波动。3.2 替换 VGG16 分类头冻结卷积基VGG16 的原始输出是 1000 类要改成我们的珊瑚类别数。在 PyTorch 中通常加载torchvision.models.vgg16(pretrainedTrue)然后把classifier[6]替换成nn.Linear(4096, num_classes)。是否冻结卷积基取决于数据量。如果每类图片只有几十张冻结前面所有卷积层只训练分类头如果每类有几百张以上可以解冻最后两个卷积块做微调。import torch.nn as nn from torchvision import models # 加载 ImageNet 预训练权重 model models.vgg16(pretrainedTrue) # 冻结所有卷积层参数 for param in model.features.parameters(): param.requires_grad False # 替换分类器的最后一层 num_classes 3 in_features model.classifier[6].in_features model.classifier[6] nn.Linear(in_features, num_classes)设置requires_grad False后卷积层在反向传播时不再计算梯度显存占用会显著下降。但要注意即使冻结了参数前向传播仍然会流过这些层所以显存不会减少太多只是优化器只更新分类头的参数。如果你想微调部分层可以循环model.features找到你想解冻的层之后设置requires_grad True。项目里如果只给了三个类别我建议全程冻结卷积层把训练重心放在classifier上这样训练速度快而且不容易过拟合。3.3 训练循环与模型保存训练循环本身不算复杂但有几个容易忽略的点。优化器建议使用SGD加上动量学习率从 0.001 开始每若干个 epoch 乘以 0.1 衰减。CrossEntropyLoss会自动把类别索引变成 one-hot 计算所以train.txt里的标签直接就是 0、1、2 即可。import torch.optim as optim from torch.utils.data import DataLoader, Dataset from PIL import Image class CoralDataset(Dataset): def __init__(self, txt_path, transformNone): self.samples [] self.transform transform with open(txt_path, r) as f: for line in f.readlines(): img_path, label line.strip().split() self.samples.append((img_path, int(label))) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label self.samples[idx] image Image.open(img_path).convert(RGB) if self.transform: image self.transform(image) return image, label # 训练参数 batch_size 8 learning_rate 0.001 epochs 30 train_dataset CoralDataset(train.txt, train_transform) train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue, num_workers0, drop_lastTrue) criterion nn.CrossEntropyLoss() optimizer optim.SGD(model.classifier.parameters(), lrlearning_rate, momentum0.9, weight_decay1e-4) for epoch in range(epochs): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) epoch_loss running_loss / len(train_dataset) print(fEpoch {epoch1}/{epochs}, Loss: {epoch_loss:.4f}) if (epoch 1) % 10 0: torch.save(model.state_dict(), fcoral_vgg16_epoch{epoch1}.pth) torch.save(model.state_dict(), coral_vgg16_final.pth)batch_size设成 8 是因为 VGG16 很耗显存如果你用的是 6GB 显存的显卡8 已经是极限如果显存不够可以降到 4同时把num_workers设为 0 避免 Windows 上出现多进程报错。drop_lastTrue会在最后一批样本数量不足时丢弃防止 BatchNorm 层计算不稳定的情况。VGG16 没有 BatchNorm 原生实现但如果微调版本里有 BN 层这个参数就很重要。3.4 关键超参表和显存控制下面这个表格是我复现这个项目时常用的参数组合不同数据量对应不同策略每类图片数学习率epoch优化器是否冻结卷积层预期准确率1030 张0.000520Adam全部冻结80% 左右50100 张0.00130SGD momentum冻结前 13 层90% 左右200 张以上0.00150SGD momentum解冻最后 2 个 block95% 以上值得注意的是torchvision里 VGG16 的pretrainedTrue参数在 PyTorch 2.x 中被标记为废弃建议用weightsmodels.VGG16_Weights.IMAGENET1K_V1这种写法。如果你用的 PyTorch 版本是 1.7.1 或 1.8.1pretrainedTrue还能用不会报错。项目里requirement.txt应该锁定了 torch 和 torchvision 的版本强烈建议不要用最新版 torch 2.3 去跑这个项目因为vgg16的接口有变化同时预训练权重下载地址也换了老代码可能直接连接超时。4. 03pyqt界面.py把训练好的模型变成可点选的分类器4.1 界面布局和信号槽第三个脚本是 PyQt5 写的图形界面功能是选择一张珊瑚图片加载训练好的模型输出类别名称和置信度。界面布局通常包含一个图片预览区、一个按钮、一个文本框。按钮点击信号clicked关联到select_and_predict方法。PyQt 的主线程是界面线程如果直接在槽函数里跑模型推理图片较大时会出现界面卡顿最简单的做法是先把图片压缩到 224x224 再预测这样单次推理时间在 CPU 上也只有几百毫秒。4.2 加载模型与预测函数模型加载时要保持和训练时一样的结构。你不能只加载state_dict必须先实例化一个 VGG16 模型替换分类层再把权重load_state_dict进去。这里有一个常见错误训练保存的是完整模型还是状态字典。项目里保存的是model.state_dict()所以加载代码必须重新定义模型结构。from PyQt5.QtWidgets import QApplication, QLabel, QPushButton, QVBoxLayout, QWidget, QFileDialog from PyQt5.QtGui import QPixmap import torch from torchvision import transforms, models from PIL import Image import torch.nn as nn # 定义模型结构和类别名 class_names [脑珊瑚, 软珊瑚, 扇形珊瑚] def load_model(model_path, num_classes3): model models.vgg16(pretrainedFalse) in_features model.classifier[6].in_features model.classifier[6] nn.Linear(in_features, num_classes) model.load_state_dict(torch.load(model_path, map_locationcpu)) model.eval() return model # 预处理与预测 def predict_image(model, img_path): transform transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) image Image.open(img_path).convert(RGB) tensor transform(image).unsqueeze(0) with torch.no_grad(): output model(tensor) prob torch.softmax(output, dim1) top_prob, top_idx torch.max(prob, dim1) return class_names[top_idx.item()], top_prob.item()上面代码中map_locationcpu很关键当你利用 GPU 训练后把模型拿到没有 NVIDIA 显卡的电脑上去加载不加这个参数就会报KeyError: cuda之类的错误。model.eval()必须调用因为它会关闭Dropout层和 BatchNorm 的统计更新。VGG16 原始分类器里有两个Dropout层如果忘记切到 eval 模式预测结果会被随机扰动每次点击图片得到不同置信度。4.3 单线程预测的注意事项PyQt 界面中如果图片分辨率很大QPixmap加载后显示会非常消耗内存建议在显示时调用.scaled()缩放。另外QFileDialog打开文件时要设置图片过滤器避免用户选择到 txt 或模型文件。预测函数里最好加一个异常捕获比如except Exception as e弹出QMessageBox提示用户图片格式不支持或模型加载失败。这个界面脚本的功能虽然简单但如果你想把它做得更专业可以把模型加载放在初始化时完成而不是每次点击按钮都加载一次那样会慢很多。5. 复现这个珊瑚分类项目时的四个实用验证技巧5.1 用两张图快速验证数据管道是否通不要一上来就训练 30 个 epoch。先找两个类别各两张图把batch_size设为 2跑一个 epoch看 loss 是否能下降。如果 loss 始终不变大概率是标签和图片没对上。这时可以在CoralDataset.__getitem__里临时打印img_path和label确认train.txt里的路径在当前环境下存在。很多报错FileNotFoundError或者因图片损坏导致的PIL.UnidentifiedImageError都在这个小规模测试中暴露。5.2 用 torchsummary 打印参数量确认冻结是否生效微调 VGG16 时如果冻结失败参数量会包含所有卷积层的大量可学习参数。执行torchsummary.summary(model, (3, 224, 224))观察Trainable params和Non-trainable params的比例。如果冻结成功可训练参数应该在 1 亿以下而不可训练参数约 1.38 亿。如果你看到所有参数都可训练说明requires_gradFalse没有生效或模型结构重新定义后覆盖了原来的冻结设置。5.3 用混淆矩阵看不同珊瑚种类的混淆倾向训练结束后不要只看整体准确率。珊瑚种类之间可能存在相似纹理的误判比如某些软珊瑚和扇形珊瑚在颜色上接近。用sklearn.metrics.confusion_matrix对验证集全部预测一遍能直观看到哪两类互相混淆。如果脑珊瑚被误判为软珊瑚的比例很高可以针对性收集这两类的边界样本或者增加RandomErasing数据增强让模型学会忽略局部遮挡。from sklearn.metrics import confusion_matrix import numpy as np all_preds [] all_labels [] model.eval() with torch.no_grad(): for images, labels in val_loader: outputs model(images) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) cm confusion_matrix(all_labels, all_preds) print(cm)5.4 用 ONNX 导出加快 CPU 推理如果 PyQt 界面在 CPU 上用 PyTorch 推理太慢可以把模型导出成 ONNX 格式再用onnxruntime跑。导出前要把模型切到 eval 模式并固定输入尺寸。这样在普通办公电脑上单张图片推理时间可以从 800ms 降到 100ms 以内。导出命令也很简单pip install onnxruntime python -c import torch; from torchvision import models; mmodels.vgg16(pretrainedFalse); m.load_state_dict(torch.load(coral_vgg16_final.pth)); m.eval(); torch.onnx.export(m, torch.randn(1,3,224,224), coral.onnx)5. 复现这个珊瑚分类项目时的四个实用验证技巧5.1 用两张图快速验证数据管道是否通不要一上来就训练 30 个 epoch。先找两个类别各两张图把batch_size设为 2跑一个 epoch看 loss 是否能下降。如果 loss 始终不变大概率是标签和图片没对上。这时可以在CoralDataset.__getitem__里临时打印img_path和label确认train.txt里的路径在当前环境下存在。很多报错FileNotFoundError或者因图片损坏导致的PIL.UnidentifiedImageError都在这个小规模测试中暴露。5.2 用 torchsummary 打印参数量确认冻结是否生效微调 VGG16 时如果冻结失败参数量会包含所有卷积层的大量可学习参数。执行torchsummary.summary(model, (3, 224, 224))观察Trainable params和Non-trainable params的比例。如果冻结成功可训练参数应该在 1 亿以下而不可训练参数约 1.38 亿。如果你看到所有参数都可训练说明requires_gradFalse没有生效或模型结构重新定义后覆盖了原来的冻结设置。5.3 用混淆矩阵看不同珊瑚种类的混淆倾向训练结束后不要只看整体准确率。珊瑚种类之间可能存在相似纹理的误判比如某些软珊瑚和扇形珊瑚在颜色上接近。用sklearn.metrics.confusion_matrix对验证集全部预测一遍能直观看到哪两类互相混淆。如果脑珊瑚被误判为软珊瑚的比例很高可以针对性收集这两类的边界样本或者增加RandomErasing数据增强让模型学会忽略局部遮挡。from sklearn.metrics import confusion_matrix import numpy as np all_preds [] all_labels [] model.eval() with torch.no_grad(): for images, labels in val_loader: outputs model(images) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) cm confusion_matrix(all_labels, all_preds) print(cm)5.4 用 ONNX 导出加快 CPU 推理如果 PyQt 界面在 CPU 上用 PyTorch 推理太慢可以把模型导出成 ONNX 格式再用onnxruntime跑。导出前要把模型切到 eval 模式并固定输入尺寸。这样在普通办公电脑上单张图片推理时间可以从 800ms 降到 100ms 以内。导出命令也很简单pip install onnxruntime python -c import torch; from torchvision import models; mmodels.vgg16(pretrainedFalse); m.load_state_dict(torch.load(coral_vgg16_final.pth)); m.eval(); torch.onnx.export(m, torch.randn(1,3,224,224), coral.onnx)本文还有配套的精品资源点击获取
返回列表