VGG16迁移学习+Pytorch实现珊瑚图像分类完整项目
2026/9/23 1:19:47 网站建设 项目流程

简介:面向深度学习初学者与珊瑚研究相关的开发者,一套基于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_coralsoft_coralfan_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', encoding='utf-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(brightness=0.3, contrast=0.3, saturation=0.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 预训练模型的标准值,迁移学习时不能随意改。如果使用自己的均值,相当于把模型输入分布改变了,预训练权重会失效。这里RandomCropCenterCrop的尺寸都是 224,对应 VGG16 的fc6层要求。训练时用随机裁剪,验证时用中心裁剪,这是图像分类任务的标准做法,避免了验证时随机性带来的指标波动。

3.2 替换 VGG16 分类头,冻结卷积基

VGG16 的原始输出是 1000 类,要改成我们的珊瑚类别数。在 PyTorch 中,通常加载torchvision.models.vgg16(pretrained=True),然后把classifier[6]替换成nn.Linear(4096, num_classes)。是否冻结卷积基取决于数据量。如果每类图片只有几十张,冻结前面所有卷积层,只训练分类头;如果每类有几百张以上,可以解冻最后两个卷积块做微调。

import torch.nn as nn from torchvision import models # 加载 ImageNet 预训练权重 model = models.vgg16(pretrained=True) # 冻结所有卷积层参数 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, transform=None): 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_size=batch_size, shuffle=True, num_workers=0, drop_last=True) criterion = nn.CrossEntropyLoss() optimizer = optim.SGD(model.classifier.parameters(), lr=learning_rate, momentum=0.9, weight_decay=1e-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(f'Epoch {epoch+1}/{epochs}, Loss: {epoch_loss:.4f}') if (epoch + 1) % 10 == 0: torch.save(model.state_dict(), f'coral_vgg16_epoch{epoch+1}.pth') torch.save(model.state_dict(), 'coral_vgg16_final.pth')

batch_size设成 8 是因为 VGG16 很耗显存,如果你用的是 6GB 显存的显卡,8 已经是极限;如果显存不够,可以降到 4,同时把num_workers设为 0 避免 Windows 上出现多进程报错。drop_last=True会在最后一批样本数量不足时丢弃,防止 BatchNorm 层计算不稳定的情况。VGG16 没有 BatchNorm 原生实现,但如果微调版本里有 BN 层,这个参数就很重要。

3.4 关键超参表和显存控制

下面这个表格是我复现这个项目时常用的参数组合,不同数据量对应不同策略:

每类图片数学习率epoch优化器是否冻结卷积层预期准确率
10~30 张0.000520Adam全部冻结80% 左右
50~100 张0.00130SGD + momentum冻结前 13 层90% 左右
200 张以上0.00150SGD + momentum解冻最后 2 个 block95% 以上

值得注意的是,torchvision里 VGG16 的pretrained=True参数在 PyTorch 2.x 中被标记为废弃,建议用weights=models.VGG16_Weights.IMAGENET1K_V1这种写法。如果你用的 PyTorch 版本是 1.7.1 或 1.8.1,pretrained=True还能用,不会报错。项目里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_classes=3): model = models.vgg16(pretrained=False) 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_location='cpu')) 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, dim=1) top_prob, top_idx = torch.max(prob, dim=1) return class_names[top_idx.item()], top_prob.item()

上面代码中map_location='cpu'很关键,当你利用 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_pathlabel,确认train.txt里的路径在当前环境下存在。很多报错FileNotFoundError或者因图片损坏导致的PIL.UnidentifiedImageError,都在这个小规模测试中暴露。

5.2 用 torchsummary 打印参数量,确认冻结是否生效

微调 VGG16 时,如果冻结失败,参数量会包含所有卷积层的大量可学习参数。执行torchsummary.summary(model, (3, 224, 224)),观察Trainable paramsNon-trainable params的比例。如果冻结成功,可训练参数应该在 1 亿以下,而不可训练参数约 1.38 亿。如果你看到所有参数都可训练,说明requires_grad=False没有生效,或模型结构重新定义后覆盖了原来的冻结设置。

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; m=models.vgg16(pretrained=False); 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_pathlabel,确认train.txt里的路径在当前环境下存在。很多报错FileNotFoundError或者因图片损坏导致的PIL.UnidentifiedImageError,都在这个小规模测试中暴露。

5.2 用 torchsummary 打印参数量,确认冻结是否生效

微调 VGG16 时,如果冻结失败,参数量会包含所有卷积层的大量可学习参数。执行torchsummary.summary(model, (3, 224, 224)),观察Trainable paramsNon-trainable params的比例。如果冻结成功,可训练参数应该在 1 亿以下,而不可训练参数约 1.38 亿。如果你看到所有参数都可训练,说明requires_grad=False没有生效,或模型结构重新定义后覆盖了原来的冻结设置。

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; m=models.vgg16(pretrained=False); m.load_state_dict(torch.load('coral_vgg16_final.pth')); m.eval(); torch.onnx.export(m, torch.randn(1,3,224,224), 'coral.onnx')"

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询