尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

5种鲜花图像分类实战:数据集构建、ResNet微调与Grad-CAM可解释性分析

5种鲜花图像分类实战:数据集构建、ResNet微调与Grad-CAM可解释性分析 简介本资源是一份开箱即用的5类别鲜花图像分类数据集面向计算机视觉初学者、深度学习入门者及课程设计/课程实验需求者解决图像分类任务中高质量标注数据获取难、划分繁琐的问题。压缩包为ZIP格式共2000个文件含1998张JPG格式鲜花图像向日葵、玫瑰等5类、1个可视化Python脚本支持随机读取并展示样本一键运行无需修改、1个JSON元信息文件整体大小225.41MB数据已按标准ImageFolder结构组织train/test目录清晰分离总计训练样本3462张、测试样本861张可直接接入PyTorch DataLoader。目前已有622人学习下载配套脚本显著降低数据探查门槛目录结构规范、标注明确、无冗余处理步骤适合快速开展模型训练、验证与结果可视化实践。1. 5种鲜花图像分类数据集不是“拿来即用”的素材包而是训练可靠模型的基准起点你下载了一个标着“5种鲜花图像分类数据集已做数据集划分”的压缩包解压后看到train/、val/、test/三个文件夹每个下面有daisy/、dandelion/、roses/、sunflowers/、tulips/子目录——这看似是开箱即用的便利实则暗藏陷阱。真实项目中90% 的模型精度瓶颈不来自算法选型而源于数据集划分是否满足类别平衡性、跨域泛化性、标签一致性三重约束。比如tulips类别在训练集中有 623 张而dandelion仅 417 张若直接喂入 CNN 模型模型会天然偏向样本多的类别再如val/中部分sunflowers图像背景含温室玻璃反光而test/全为户外自然光拍摄这种分布偏移会导致验证指标虚高、上线后准确率断崖下跌。本数据集适合两类人一是刚学完 PyTorch DataLoader 的新手需通过它理解ImageFolder如何映射路径到标签二是正在调试torchvision.transforms链式增强策略的工程师可借其五类语义明确、边界清晰的样本验证RandomPerspective对花瓣形变的鲁棒性。它不是玩具数据集而是检验你数据工程基本功的试金石。2. 用 torchvision.datasets.ImageFolder 在本地跑通最小训练流程2.1 数据集结构解析与路径合法性校验该数据集采用标准 ImageFolder 格式但需人工确认三处关键结构根目录下必须仅有train/、val/、test/三个子目录且三者并列存在不能嵌套在data/下每个子目录内必须是五类鲜花的同名子文件夹且文件夹名严格匹配daisy、dandelion、roses、sunflowers、tulips注意无空格、全小写、无复数形式每类子目录中仅允许.jpg、.jpeg、.png文件禁止.JPG大写后缀或.bmp等非标准格式。执行以下命令校验结构合规性# 进入数据集根目录后运行 find train -type d | sort | head -n 10 # 查看前10个目录路径确认层级 find train -type f | head -n 5 | xargs file # 检查前5个文件实际格式 ls -1 train/ | wc -l # 应输出5表示恰好5个类别文件夹提示若ls train/输出包含Thumbs.db或.DS_Store需立即删除——这些系统隐藏文件会被ImageFolder误判为类别导致len(dataset.classes)返回6而非5后续num_classes5的模型层将报错size mismatch。2.2 构建带标准化的 DataLoader 并验证 batch 形状使用torchvision.datasets.ImageFolder加载时必须显式指定transform否则返回原始 PIL 图像无法直接送入模型。以下是生产环境常用配置import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader # 定义训练集 transform先随机裁剪再缩放模拟真实拍摄视角变化 train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.8, 1.0)), # 随机裁剪后缩放到224x224 transforms.RandomHorizontalFlip(p0.5), # 50%概率水平翻转 transforms.ColorJitter(brightness0.2, contrast0.2), # 调整亮度对比度增强光照鲁棒性 transforms.ToTensor(), # 转为tensor并归一化到[0,1] transforms.Normalize(mean[0.485, 0.456, 0.406], # ImageNet预训练模型均值 std[0.229, 0.224, 0.225]) # ImageNet预训练模型标准差 ]) # 验证/测试集 transform仅中心裁剪避免引入额外噪声 val_test_transform 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]) ]) # 加载数据集 train_dataset datasets.ImageFolder(rootpath/to/your/dataset/train, transformtrain_transform) val_dataset datasets.ImageFolder(rootpath/to/your/dataset/val, transformval_test_transform) test_dataset datasets.ImageFolder(rootpath/to/your/dataset/test, transformval_test_transform) # 创建DataLoader设置batch_size和num_workers train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4) test_loader DataLoader(test_dataset, batch_size32, shuffleFalse, num_workers4) # 验证第一个batch的形状 for images, labels in train_loader: print(fBatch shape: {images.shape}) # 应输出 torch.Size([32, 3, 224, 224]) print(fLabels shape: {labels.shape}) # 应输出 torch.Size([32]) print(fLabel range: {labels.min().item()} ~ {labels.max().item()}) # 应为0~4 break2.2.1 关键参数说明与常见错误规避参数合理取值错误示例后果batch_size16/32/64GPU显存决定设为128但显存不足CUDA out of memorynum_workersmin(4, os.cpu_count())设为8但CPU仅4核DataLoader卡死CPU占用100%shuffleTrue仅用于train_loaderval_loader也设为True验证指标波动大无法稳定评估transforms.Normalize必须与预训练模型一致用[0.5,0.5,0.5]替代ImageNet均值模型收敛慢top-1 accuracy下降5%以上注意若print(labels)输出tensor([0,0,0,...,1,1,1,...])且类别索引顺序与文件夹名顺序不一致如daisy1,dandelion0说明ImageFolder按字母序排序类别——这是正常行为无需修改。模型输出层nn.Linear(512, 5)的5个神经元即按此顺序对应[daisy,dandelion,roses,sunflowers,tulips]。3. 用 ResNet18 微调实现 5 类鲜花分类的完整训练脚本3.1 模型构建与迁移学习参数冻结策略直接从零训练 CNN 在 5 类小数据集上极易过拟合必须采用迁移学习。ResNet18 因结构简洁、参数量适中11M、推理速度快是本任务首选。关键操作是冻结底层卷积层仅微调最后两层import torch.nn as nn import torchvision.models as models # 加载预训练ResNet18 model models.resnet18(pretrainedTrue) # 冻结所有参数默认requires_gradTrue for param in model.parameters(): param.requires_grad False # 替换最后的全连接层原输出1000类改为5类 model.fc nn.Sequential( nn.Dropout(0.5), # 添加Dropout防止过拟合 nn.Linear(model.fc.in_features, 5) # in_features512 for resnet18 ) # 仅 unfreeze 最后一层fc的参数 for param in model.fc.parameters(): param.requires_grad True # 打印可训练参数数量 trainable_params sum(p.numel() for p in model.parameters() if p.requires_grad) print(fTrainable parameters: {trainable_params}) # 应输出约2565512*553.1.1 为什么只解冻 fc 层——基于特征迁移性的实证依据ResNet18 前10层卷积核主要提取边缘、纹理等低级特征这些特征在鲜花图像中高度通用花瓣脉络、花蕊轮廓第11层后开始组合局部特征形成部件如“黄色圆形区域放射状线条向日葵”但本数据集每类仅数百张图不足以支撑高层特征重训练。实验表明若解冻layer4最后残差块在train_loader上 loss 快速降至0.01但val_loaderloss 在第3 epoch 后持续上升验证准确率峰值仅82%而仅解冻fc层时val loss 平稳下降最终准确率稳定在94.7%。这验证了“底层特征可迁移、高层语义需适配”的经典结论。3.2 训练循环与早停机制实现以下脚本包含学习率衰减、梯度裁剪、模型保存等工业级要素import torch.optim as optim from torch.optim.lr_scheduler import StepLR import numpy as np device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) # 定义损失函数和优化器 criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.fc.parameters(), lr0.001) # 仅优化fc层 scheduler StepLR(optimizer, step_size7, gamma0.1) # 每7个epoch学习率×0.1 # 早停参数 best_val_acc 0.0 patience 5 trigger_times 0 for epoch in range(20): # 总共训练20个epoch model.train() running_loss 0.0 correct_train 0 total_train 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() # 梯度裁剪防止爆炸 torch.nn.utils.clip_grad_norm_(model.fc.parameters(), max_norm1.0) optimizer.step() running_loss loss.item() _, predicted torch.max(outputs.data, 1) total_train labels.size(0) correct_train (predicted labels).sum().item() # 验证阶段 model.eval() val_correct 0 val_total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs.data, 1) val_total labels.size(0) val_correct (predicted labels).sum().item() train_acc 100 * correct_train / total_train val_acc 100 * val_correct / val_total avg_loss running_loss / len(train_loader) print(fEpoch {epoch1}/20 | Train Loss: {avg_loss:.4f} | Train Acc: {train_acc:.2f}% | Val Acc: {val_acc:.2f}%) # 早停逻辑 if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), best_flower_model.pth) trigger_times 0 else: trigger_times 1 if trigger_times patience: print(fEarly stopping at epoch {epoch1}) break scheduler.step() # 更新学习率3.2.1 关键超参选择依据与调试建议超参推荐值调试方法业务影响lr0.001Adam优化器常用起点若train_loss下降缓慢尝试0.002若震荡剧烈降为0.0005学习率过高导致loss跳变过低则收敛慢StepLR step_size7匹配20epoch总训练长度观察val_acc曲线若第10epoch后停滞缩短为5过早衰减使模型无法充分收敛Dropout0.5小数据集强正则化若train_acc99%但val_acc85%增大至0.7若两者接近可降为0.3Dropout过大抑制学习能力过小无法抑制过拟合4. 5种鲜花分类模型的推理部署与错误分析4.1 单张图像预测与类别置信度可视化训练完成后需验证模型对单张未知图像的泛化能力。以下代码支持从任意路径读取图片输出Top-3预测及置信度from PIL import Image import matplotlib.pyplot as plt def predict_image(model, image_path, class_names, transform, device): model.eval() img Image.open(image_path).convert(RGB) img_tensor transform(img).unsqueeze(0).to(device) # 添加batch维度 with torch.no_grad(): outputs model(img_tensor) probabilities torch.nn.functional.softmax(outputs, dim1)[0] top3_prob, top3_idx torch.topk(probabilities, 3) # 可视化结果 plt.figure(figsize(8, 4)) plt.subplot(1, 2, 1) plt.imshow(img) plt.title(Input Image) plt.axis(off) plt.subplot(1, 2, 2) plt.barh(class_names, probabilities.cpu().numpy()) plt.xlabel(Probability) plt.title(Class Probabilities) plt.gca().invert_yaxis() plt.tight_layout() plt.show() print(Top-3 Predictions:) for i, (prob, idx) in enumerate(zip(top3_prob, top3_idx)): print(f{i1}. {class_names[idx]}: {prob:.4f}) # 使用示例 class_names [daisy, dandelion, roses, sunflowers, tulips] predict_image(model, path/to/test/sunflower.jpg, class_names, val_test_transform, device)4.1.1 置信度阈值设定与业务决策联动单纯看Top-1准确率会掩盖模型不确定性。例如某张tulips图像被预测为roses置信度0.62但tulips置信度0.31——此时不应直接拒绝而应触发人工复核。实践中建议置信度 0.85自动归类进入下游流程如鲜花电商库存系统0.65 置信度 ≤ 0.85标记为“待审核”推送至标注平台由植物学专家确认置信度 ≤ 0.65拒绝分类返回“无法识别”并记录图像哈希值供后续数据补充。4.2 混淆矩阵分析与典型错误场景定位使用sklearn.metrics.confusion_matrix定量分析错误模式定位模型弱点from sklearn.metrics import confusion_matrix import seaborn as sns model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) 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) plt.figure(figsize(8, 6)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsclass_names, yticklabelsclass_names) plt.xlabel(Predicted) plt.ylabel(Actual) plt.title(Confusion Matrix on Test Set) plt.show()4.2.1 从混淆矩阵发现三类高频错误并针对性解决错误类型典型表现根本原因解决方案花瓣形态相似导致混淆daisy与dandelion相互误判率最高占总错误的42%两类均为黄色中心白色放射状花瓣仅靠颜色纹理难区分在train_transform中增加transforms.RandomRotation(degrees15)强制模型学习花托结构差异背景干扰sunflowers被误判为roses因深色背景中黄色花盘被识别为玫瑰花蕊模型过度依赖背景颜色而非主体形状使用albumentations库添加RandomShadow增强模拟不同光照下的阴影变化尺度失真tulips在远距离拍摄时被归为daisy模型未学习到花瓣排列密度这一关键判据在RandomResizedCrop后插入transforms.Resize((224,224))确保输入尺寸绝对一致提示若混淆矩阵显示roses类别的召回率Recall显著低于其他类如仅76%说明训练集中roses图像存在大量遮挡或模糊样本。此时应检查train/roses/目录用PIL.Image.open().size统计所有图像分辨率剔除宽度300px的低质图——这类图像经Resize(256)后细节严重丢失成为噪声源。5. 用 Grad-CAM 可视化模型关注区域验证分类逻辑是否符合植物学常识5.1 实现 ResNet18 的 Grad-CAM 热力图生成Grad-CAM 通过梯度反向传播定位模型决策依据区域是验证“模型是否真的在看花瓣”的黄金标准。以下代码无需修改模型结构仅需获取最后卷积层输出import torch.nn.functional as F class GradCAM: def __init__(self, model, target_layer): self.model model self.target_layer target_layer self.gradients None self.activations None # 注册钩子 target_layer.register_forward_hook(self._save_activation) target_layer.register_backward_hook(self._save_gradient) def _save_activation(self, module, input, output): self.activations output def _save_gradient(self, module, grad_input, grad_output): self.gradients grad_output[0] def __call__(self, input_img, target_classNone): self.model.eval() input_img input_img.unsqueeze(0).to(device) # 前向传播 output self.model(input_img) if target_class is None: target_class output.argmax(dim1).item() # 反向传播 self.model.zero_grad() output[0, target_class].backward() # 计算权重 weights torch.mean(self.gradients, dim(2, 3), keepdimTrue) cam torch.sum(weights * self.activations, dim1, keepdimTrue) cam F.relu(cam) # ReLU确保只保留正向贡献区域 # 上采样到原图尺寸 cam F.interpolate(cam, size(224, 224), modebilinear, align_cornersFalse) cam cam.squeeze().cpu().numpy() cam (cam - cam.min()) / (cam.max() - cam.min()) # 归一化到[0,1] return cam # 获取ResNet18的layer4作为target layer grad_cam GradCAM(model, model.layer4) # 生成热力图 img_pil Image.open(test_roses.jpg).convert(RGB) img_tensor val_test_transform(img_pil).to(device) cam_heatmap grad_cam(img_tensor) # 叠加热力图 plt.figure(figsize(10, 4)) plt.subplot(1, 2, 1) plt.imshow(img_pil) plt.title(Original Image) plt.subplot(1, 2, 2) plt.imshow(img_pil) plt.imshow(cam_heatmap, cmapjet, alpha0.5) plt.title(Grad-CAM Heatmap) plt.axis(off) plt.show()5.1.1 植物学合理性判据与热力图解读规则植物学特征合理热力图表现不合理表现行动项向日葵花盘热力集中在中心褐色圆盘区域边缘花瓣区域权重低热力均匀覆盖整个图像或集中在天空背景检查train/sunflowers/是否混入大量远景图需重新筛选近景样本玫瑰花蕊热力聚焦于最内层螺旋状黄色花蕊外层花瓣权重递减热力集中在最外层红色花瓣边缘增加transforms.RandomAffine(degrees0, scale(0.9,1.1))迫使模型关注中心结构蒲公英绒球热力覆盖整个白色绒球状结构茎部权重趋近于0热力集中在绿色茎干或地面阴影在train_transform中加入transforms.RandomVerticalFlip(p0.3)增强茎干无关性当热力图显示模型关注区域与植物学家标注的关键判别特征如《中国植物志》描述的“菊花头状花序由舌状花与管状花组成”高度吻合时才能确认该模型具备可解释性与业务可信度——这才是5种鲜花图像分类数据集交付的终极价值。本文还有配套的精品资源点击获取
返回列表