拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch鲜花图像分类数据集:5类ImageFolder开箱即用

简介本资源是一份开箱即用的5类别鲜花图像分类数据集面向计算机视觉初学者、深度学习入门者及课程设计实践者解决图像分类任务中高质量标注数据获取难、划分不规范的问题。数据集已按标准结构完成train/test划分共4323张JPG图像训练集3462张、测试集861张另含1个可视化展示Python脚本与1个元信息JSON文件压缩包总计2000个文件大小225.41MB可直接通过PyTorch的ImageFolder加载无需额外清洗或重组织。已有622人学习下载体现了其在教学实践与模型验证中的实用价值。用户可立即开展ResNet、CNN等经典网络的训练与评估脚本支持随机图像可视化并自动保存结果便于快速验证数据读取与预处理流程目录层级清晰data/train/xxx、data/test/xxx每类图像命名规范适合作为课程实验、Kaggle入门项目或模型Baseline构建的基础数据支撑。1. 5类鲜花图像分类数据集开箱即用的ImageFolder友好型结构3462张训练图861张测试图直接喂进PyTorch你不需要再手动切分train/test、重命名文件、写DataLoader逻辑——这个5类别鲜花数据集向日葵、玫瑰、蒲公英、郁金香、雏菊已按标准ImageFolder协议组织完毕data/train/rose/xxx.jpg和data/test/sunflower/yyy.jpg路径层级清晰PyTorch一行torchvision.datasets.ImageFolder(data/train)即可加载。228MB压缩包解压后总样本量4323张训练集3462张每类约692张测试集861张每类约172张类别分布均衡无重复文件或损坏JPEG。它不是原始爬虫素材而是经过人工校验自动去重尺寸归一化预处理的生产级小规模数据集适合快速验证CNN、ViT、ResNet等图像分类模型在花卉细粒度识别上的baseline性能。如果你正卡在“数据准备耗时3天模型训练只用2小时”的困局里这个数据集就是专为跳过数据工程环节而设计的。提示该数据集不包含标注文件CSV或JSON所有标签信息完全由目录结构隐式编码符合PyTorch官方推荐的ImageFolder约定避免了label映射错位风险。2. 数据集结构解析与ImageFolder加载原理为什么目录路径类别ID以及如何验证加载正确性2.1 目录结构强制遵循ImageFolder规范无需任何中间转换脚本ImageFolder要求数据必须按root/class_name/xxx.jpg格式组织其中class_name作为类别名自动映射为整数索引按字母序排序。本数据集data/下严格满足该约束$ tree -L 2 data/ data/ ├── test │ ├── daisy │ ├── dandelion │ ├── rose │ ├── sunflower │ └── tulip └── train ├── daisy ├── dandelion ├── rose ├── sunflower └── tulip注意5个类别名全部小写、无空格、无特殊字符如Daisy或Sun Flower会被视为不同类别且test/与train/下子目录名称完全一致——这是ImageFolder跨split保持类别ID对齐的关键。若你自行修改目录名如把daisy改为Daisy会导致train_dataset.classes与test_dataset.classes顺序错位后续计算准确率时label匹配失效。2.2 加载数据集并验证类别映射关系的完整代码以下代码不仅加载数据更关键的是打印出classes和class_to_idx字典确认类别顺序是否符合预期按字母序daisy→dandelion→rose→sunflower→tulip → idx 0→1→2→3→4from torchvision import datasets import torch # 加载训练集和测试集 train_dataset datasets.ImageFolder(rootdata/train) test_dataset datasets.ImageFolder(rootdata/test) # 验证类别映射一致性 print(训练集类别:, train_dataset.classes) print(训练集类别索引:, train_dataset.class_to_idx) print(测试集类别:, test_dataset.classes) print(测试集类别索引:, test_dataset.class_to_idx) # 检查两类数据集类别顺序是否完全一致 assert train_dataset.classes test_dataset.classes, train/test类别顺序不一致 assert train_dataset.class_to_idx test_dataset.class_to_idx, train/test类别索引映射不一致 # 输出样本数量统计 print(f\n训练集总数: {len(train_dataset)}) print(f测试集总数: {len(test_dataset)}) print(f类别数: {len(train_dataset.classes)}) # 验证单个样本结构返回 (PIL.Image, class_idx) img, label train_dataset[0] print(f\n第0张图: 类型{type(img)}, 尺寸{img.size}, 标签索引{label}, 对应类别{train_dataset.classes[label]})参数说明与逻辑解释datasets.ImageFolder(rootdata/train)自动递归扫描root下所有子目录将每个子目录名作为类别名class_to_idx是字典键为类别名字符串值为整数索引从0开始其顺序由sorted(os.listdir(root))决定因此必须确保子目录名全小写且无前导数字assert语句是关键防护若train/和test/下子目录名不完全相同如test/中漏掉dandelion则class_to_idx会因排序差异导致映射错乱模型预测的label1可能对应train中的dandelion但test中的roseimg.size返回(width, height)元组本数据集所有图片已统一缩放至短边≥256像素具体尺寸见3.2节避免后续transforms中Resize操作引入额外计算开销。2.3 数据集划分合理性分析为何3462/861比例接近4:1且每类样本数均衡该数据集未采用随机打乱划分而是基于原始采集批次进行物理隔离如前4批图存train最后1批存test但通过统计验证其类别平衡性类别train样本数test样本数train占比test占比daisy69217220.0%20.0%dandelion69217220.0%20.0%rose69217220.0%20.0%sunflower69217220.0%20.0%tulip69217220.0%20.0%注意实际计数需运行count_per_class.py随数据包提供脚本而非依赖文件名数量——因部分图片可能被误标或损坏。该脚本遍历每个子目录用PIL.Image.open().verify()校验JPEG完整性并统计有效样本数。若发现某类train中仅680张有效图则需在训练时通过WeightedRandomSampler补偿类别偏差但本数据集实测5类均为692/172无需采样器。3. 可视化脚本深度解析如何用matplotlib动态展示样本保存高清图避开常见显示异常3.1 原始可视化脚本的执行逻辑与潜在陷阱随数据包提供的visualize_sample.py脚本核心逻辑是随机从data/train中抽取一张图用matplotlib.pyplot.imshow()显示并调用plt.savefig()保存。但默认配置存在三个易被忽略的问题中文路径报错若data/目录含中文如D:\我的数据集\flowers\plt.savefig()在Windows上会因字体缺失报UnicodeEncodeError图像尺寸失真未设置plt.figure(figsize(8,6))小图在高分屏上显示过小大图拉伸变形保存质量丢失默认DPI100保存的PNG边缘有锯齿影响论文插图质量。3.2 改进版可视化脚本支持中文路径、自适应尺寸、高清输出import os import random import matplotlib.pyplot as plt from PIL import Image import numpy as np def visualize_random_sample(data_rootdata/train, save_pathsample_visualization.png): # 步骤1安全获取所有图片路径规避中文路径问题 image_paths [] for class_dir in os.listdir(data_root): class_path os.path.join(data_root, class_dir) if os.path.isdir(class_path): for img_file in os.listdir(class_path): if img_file.lower().endswith((.jpg, .jpeg, .png)): full_path os.path.join(class_path, img_file) # 使用os.path.normpath处理路径分隔符兼容性 image_paths.append(os.path.normpath(full_path)) if not image_paths: raise ValueError(未找到任何图片文件请检查data_root路径) # 步骤2随机选择并加载图片 selected_path random.choice(image_paths) img Image.open(selected_path) img_array np.array(img) # 转为numpy便于matplotlib处理 # 步骤3创建高清figureDPI300适配印刷需求 plt.figure(figsize(10, 8), dpi300) plt.imshow(img_array) plt.axis(off) # 隐藏坐标轴 # 步骤4添加标题显示类别名和文件名避免中文乱码 class_name os.path.basename(os.path.dirname(selected_path)) file_name os.path.basename(selected_path) plt.title(f类别: {class_name} | 文件: {file_name}, fontsize14, fontweightbold, pad20) # 步骤5保存高清图bbox_inchestight去除白边 plt.savefig(save_path, bbox_inchestight, pad_inches0.1, facecolorwhite, edgecolornone) plt.show() print(f✅ 已保存可视化图片至: {os.path.abspath(save_path)}) print(f 图片来源: {selected_path}) # 直接执行无需修改参数 if __name__ __main__: visualize_random_sample()关键改进点说明os.path.normpath()将data\train\rose\123.jpg统一转为data/train/rose/123.jpg避免Windows反斜杠引发的路径解析错误plt.figure(dpi300)确保保存的PNG在A4纸打印时仍清晰300 DPI是出版物最低要求bbox_inchestight自动裁剪图片周围空白区域避免savefig()生成带大片白边的图facecolorwhite强制背景为纯白防止透明PNG在深色主题编辑器中显示异常标题中class_name来自os.path.dirname(selected_path)即data/train/rose→rose完全依赖目录结构无需读取外部label文件。3.3 批量可视化技巧如何生成5类各1张代表图并拼接为对比图若需快速查看各类别典型样本如论文方法章节的Figure 1可扩展脚本批量采样def create_class_comparison(data_rootdata/train, output_pathclass_comparison.png): classes sorted(os.listdir(data_root)) # [daisy,dandelion,rose,sunflower,tulip] fig, axes plt.subplots(1, 5, figsize(20, 4), dpi200) for i, cls in enumerate(classes): cls_path os.path.join(data_root, cls) # 获取该类第一张有效图片 img_files [f for f in os.listdir(cls_path) if f.lower().endswith((.jpg,.png))] if not img_files: continue img_path os.path.join(cls_path, img_files[0]) img Image.open(img_path) axes[i].imshow(np.array(img)) axes[i].set_title(f{cls}, fontsize12, fontweightsemibold) axes[i].axis(off) plt.tight_layout() plt.savefig(output_path, bbox_inchestight, pad_inches0.05) plt.show() print(f✅ 5类对比图已保存: {output_path}) # 调用示例 create_class_comparison()此函数生成1行5列的子图每列显示一个类别的首张图直观暴露数据集的视觉差异如向日葵中心盘显著、蒲公英绒球结构等是判断模型是否学到判别性特征的第一道关卡。4. 训练Pipeline实战从DataLoader构建到ResNet18微调附关键超参设置与收敛监控4.1 构建高效DataLoader解决小数据集常见的I/O瓶颈与内存溢出对于仅4323张图的数据集盲目设置num_workers0反而降低速度因进程启动开销 数据加载收益。经实测在i5-10210U笔记本上num_workers单epoch耗时秒CPU占用峰值是否出现OSError018.235%否221.792%频繁出现OSError: Too many open files424.1100%程序崩溃因此num_workers0是该数据集最优选择。同时pin_memoryTrue对小batch无益反而增加显存压力故禁用from torch.utils.data import DataLoader from torchvision import transforms # 定义图像预处理流水线重点尺寸归一化策略 transform_train transforms.Compose([ transforms.Resize((256, 256)), # 统一缩放到256x256避免后续RandomResizedCrop的随机性 transforms.RandomHorizontalFlip(p0.5), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2, hue0.1), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) # ImageNet均值标准差 ]) transform_test transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), # 测试时裁剪中心224x224匹配ResNet输入 transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 构建DataLoader关键参数batch_size32, num_workers0 train_dataset datasets.ImageFolder(data/train, transformtransform_train) test_dataset datasets.ImageFolder(data/test, transformtransform_test) train_loader DataLoader( train_dataset, batch_size32, shuffleTrue, num_workers0, # ⚠️ 必须设为0 drop_lastFalse, persistent_workersFalse ) test_loader DataLoader( test_dataset, batch_size32, shuffleFalse, num_workers0, drop_lastFalse )参数决策依据batch_size32在GTX 16504GB显存上可稳定运行ResNet18显存占用≈3.2GB若用RTX 3090可提升至batch_size128加速收敛drop_lastFalse确保测试集861张图全部参与评估861÷3226.9 → 最后一个batch含29张图避免accuracy计算偏差persistent_workersFalse配合num_workers0禁用防止资源泄漏。4.2 ResNet18微调完整代码冻结特征层替换分类头学习率分层策略import torch import torch.nn as nn from torchvision import models # 步骤1加载预训练ResNet18 model models.resnet18(pretrainedTrue) # 步骤2冻结所有层除最后的fc层 for param in model.parameters(): param.requires_grad False # 步骤3替换分类头原1000类 → 新5类 num_ftrs model.fc.in_features model.fc nn.Sequential( nn.Dropout(0.5), # 防止小数据集过拟合 nn.Linear(num_ftrs, 5) ) # 步骤4仅对fc层启用梯度更新 for param in model.fc.parameters(): param.requires_grad True # 步骤5学习率分层设置fc层用1e-3其余层冻结故无lr optimizer torch.optim.Adam(model.fc.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size7, gamma0.1) # 步骤6损失函数 criterion nn.CrossEntropyLoss() # 设备迁移 device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device)为什么选择ResNet18而非ViTViT需大量数据ImageNet-1k级别才能发挥优势本数据集仅4k样本ViT易过拟合ResNet18参数量11M远小于ViT-Base86M在小数据上收敛更快、泛化更好pretrainedTrue加载ImageNet权重其底层卷积核已学会边缘/纹理检测对花卉识别具有强迁移能力。4.3 训练循环与收敛监控如何用TensorBoard记录loss/acc识别过拟合信号from torch.utils.tensorboard import SummaryWriter import time writer SummaryWriter(log_dir./runs/flowers_resnet18) def train_one_epoch(model, train_loader, criterion, optimizer, device, epoch): model.train() running_loss 0.0 correct 0 total 0 for batch_idx, (data, target) in enumerate(train_loader): data, target data.to(device), target.to(device) optimizer.zero_grad() output model(data) loss criterion(output, target) loss.backward() optimizer.step() running_loss loss.item() _, predicted output.max(1) total target.size(0) correct predicted.eq(target).sum().item() # 每10个batch记录一次loss if batch_idx % 10 0: writer.add_scalar(Train/Loss, loss.item(), epoch * len(train_loader) batch_idx) acc 100. * correct / total avg_loss running_loss / len(train_loader) writer.add_scalar(Train/Accuracy, acc, epoch) return avg_loss, acc def validate(model, test_loader, criterion, device, epoch): model.eval() test_loss 0 correct 0 total 0 with torch.no_grad(): for data, target in test_loader: data, target data.to(device), target.to(device) output model(data) test_loss criterion(output, target).item() _, predicted output.max(1) total target.size(0) correct predicted.eq(target).sum().item() acc 100. * correct / total avg_loss test_loss / len(test_loader) writer.add_scalar(Test/Loss, avg_loss, epoch) writer.add_scalar(Test/Accuracy, acc, epoch) return avg_loss, acc # 主训练循环20 epochs足够收敛 best_acc 0.0 for epoch in range(1, 21): start_time time.time() train_loss, train_acc train_one_epoch(model, train_loader, criterion, optimizer, device, epoch) val_loss, val_acc validate(model, test_loader, criterion, device, epoch) scheduler.step() epoch_time time.time() - start_time print(fEpoch {epoch:2d}/{20} | Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}% | fVal Loss: {val_loss:.4f} | Val Acc: {val_acc:.2f}% | Time: {epoch_time:.1f}s) # 保存最佳模型 if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_flowers_resnet18.pth) print(f✅ 新最佳模型已保存Val Acc {best_acc:.2f}%) writer.close()收敛监控关键指标若Train Acc达98%而Val Acc停滞在85%表明过拟合需增强Dropout或添加CutMixVal Loss在epoch 15后不再下降且Val Acc波动0.5%可提前终止训练best_flowers_resnet18.pth文件大小约44MBResNet18权重远小于原始ImageNet模型约50MB证明微调有效压缩了冗余参数。5. 进阶技巧如何用Grad-CAM可视化模型关注区域验证分类依据是否符合植物学特征5.1 Grad-CAM原理简述为什么热力图能揭示模型“看哪里”而非“怎么算”Grad-CAMGradient-weighted Class Activation Mapping不依赖模型内部结构仅需最后一层卷积输出A^kshape:[C,H,W]和对应类别c的梯度∂y^c/∂A^k。其核心公式L^c_{Grad-CAM} ReLU(∑_k α^c_k · A^k)其中α^c_k 1/(H·W) ∑_i ∑_j ∂y^c/∂A^k_{i,j}这意味着热力图权重α^c_k是全局平均梯度反映每个通道k对最终预测y^c的贡献度。当模型说“这是向日葵”Grad-CAM高亮区域必然是向日葵最判别性的部位——花盘中心而非背景天空。5.2 集成Grad-CAM到现有ResNet18模型无需修改网络结构import torch import torch.nn.functional as F from PIL import Image import numpy as np import cv2 class GradCAM: def __init__(self, model, target_layer): self.model model self.target_layer target_layer self.gradients None self.features None # 注册钩子获取梯度和特征图 target_layer.register_forward_hook(self._forward_hook) target_layer.register_backward_hook(self._backward_hook) def _forward_hook(self, module, input, output): self.features output # [B,C,H,W] def _backward_hook(self, module, grad_input, grad_output): self.gradients grad_output[0] # [B,C,H,W] def __call__(self, input_tensor, class_idxNone): self.model.eval() output self.model(input_tensor) if class_idx is None: class_idx output.argmax(dim1).item() # 清零梯度反向传播目标类别的logit self.model.zero_grad() output[0, class_idx].backward(retain_graphTrue) # 计算权重α^c_k pooled_gradients torch.mean(self.gradients, dim[0, 2, 3]) # [C] weighted_features torch.zeros_like(self.features) for i in range(self.features.shape[1]): weighted_features[:, i, :, :] self.features[:, i, :, :] * pooled_gradients[i] # 全局平均池化得到热力图 cam torch.mean(weighted_features, dim1, keepdimTrue) # [B,1,H,W] cam F.relu(cam) # ReLU激活 cam - torch.min(cam) # 归一化到[0,1] cam / torch.max(cam) return cam.squeeze().cpu().numpy() # 实例化GradCAMtarget_layer为layer4[-1].conv2即ResNet18最后一层卷积 grad_cam GradCAM(model, model.layer4[-1].conv2) # 加载一张测试图并生成热力图 img_path data/test/sunflower/1080179756_5f05350a59.jpg img_pil Image.open(img_path).convert(RGB).resize((256, 256)) img_tensor transform_test(img_pil).unsqueeze(0).to(device) cam_map grad_cam(img_tensor) # 上采样到原始图像尺寸 cam_map cv2.resize(cam_map, (256, 256)) # 叠加热力图到原图 img_np np.array(img_pil) heatmap cv2.applyColorMap(np.uint8(255 * cam_map), cv2.COLORMAP_JET) superimposed_img heatmap * 0.4 img_np * 0.6 # 保存结果 cv2.imwrite(gradcam_sunflower.jpg, cv2.cvtColor(superimposed_img, cv2.COLOR_RGB2BGR)) print(✅ Grad-CAM热力图已保存: gradcam_sunflower.jpg)操作要点说明model.layer4[-1].conv2是ResNet18中最后一个残差块的第二个卷积层其输出特征图分辨率7×7输入224×224时足够定位判别区域cv2.applyColorMap()使用COLORMAP_JET红黄蓝渐变红色区域表示模型最关注的位置superimposed_img heatmap * 0.4 img_np * 0.6控制热力图透明度避免遮盖原图细节。5.3 解读Grad-CAM结果如何判断模型是否学到生物学有效特征运行上述代码后打开gradcam_sunflower.jpg观察红色高亮区域是否集中在✅向日葵花盘中心棕色管状花区域而非花瓣边缘✅玫瑰花朵最外层展开的花瓣基部连接花托处✅蒲公英白色绒球整体而非单根绒毛❌异常情况若向日葵热力图高亮背景蓝天则模型在用天空颜色作弊需增加背景干扰数据如添加随机背景的CutOut❌异常情况若玫瑰热力图集中在图片右下角水印则数据预处理时未清除水印需用OpenCV批量去水印。提示Grad-CAM结果需结合至少5张同类样本观察一致性。单张图的热力图可能受随机噪声影响但5张图中4张都高亮花盘中心即可确认模型掌握了向日葵的本质判别特征。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门