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

资讯详情

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

TensorFlow花卉识别实战:5类数据集迁移学习与MobileNetV2微调

TensorFlow花卉识别实战:5类数据集迁移学习与MobileNetV2微调 简介这份资源面向深度学习入门者与计算机视觉方向的在校学生提供一套可直接上手的花卉图像五分类训练素材帮助解决从零搭建物体分类模型时缺少规范数据与配套指导的问题。压缩包共约2000个文件以1972张jpg花卉图片为主体另含14个Python脚本、7个txt说明、5个xml标注及md、pdf文档整体约596.75MB图片按类别组织脚本覆盖数据读取、模型构建与训练流程文档则补充环境配置与参数说明。目前已有14841人学习下载热度较高。读者可借助TensorFlow代码与作者录制的B站视频快速跑通数据预处理、卷积网络搭建、训练评估到预测的完整链路并对照教程理解每步实现细节与常见报错处理适合作为课程作业、入门练手或分类任务迁移的参考方案。1. 花卉识别数据集5类从零训练一个能用的分类器拿到一个标注好的花卉数据集最怕的不是模型跑不起来而是跑起来之后发现类别对不上、图片损坏、训练集和验证集混在一起。这个资源包解决的就是这类问题5 类花卉图像配好了 TensorFlow 训练代码和一份能照着走的教程解压就能开始跑。适合两类人——刚接触物体分类、想找一个干净数据集练手的新手以及手头有业务场景比如植物科普 App、园艺电商的品类初筛需要快速验证分类可行性的从业者。5 类这个量级不算大但恰好卡在“能跑通完整流程”和“不至于在数据清洗上耗三天”之间。下面按实际拆包和训练的顺序把这份资源讲透。2. 数据集结构与 TensorFlow 读取管线先搞清楚目录长什么样2.1 五类花卉的目录组织与文件格式这类分类数据集最常见的组织方式是每个类别一个子目录目录名就是标签名。解压后大概率是下面这种结构flower_dataset/ ├── daisy/ │ ├── 001.jpg │ ├── 002.jpg │ └── ... ├── dandelion/ ├── rose/ ├── sunflower/ └── tulip/五个类别分别是雏菊、蒲公英、玫瑰、向日葵、郁金香这是花卉分类任务里最经典的一组。图片格式以 JPG 为主尺寸不统一常见的是 320×240 到 500×500 之间。这里有个容易被忽略的点目录名直接决定标签顺序。tf.keras.utils.image_dataset_from_directory会按字母序给类别编号daisy0、dandelion1、rose2、sunflower3、tulip4。如果你后面要输出预测结果这个映射关系必须记牢否则会出现“模型说 2你以为是雏菊”的翻车。先跑一段代码确认目录结构和类别数import os import pathlib data_dir pathlib.Path(flower_dataset) class_names sorted([d.name for d in data_dir.iterdir() if d.is_dir()]) print(类别列表:, class_names) print(类别数:, len(class_names)) for name in class_names: count len(list((data_dir / name).glob(*.jpg))) print(f{name}: {count} 张)这段代码做三件事列出所有子目录作为类别、统计类别总数、逐类统计 JPG 数量。如果某个类别数量明显偏少比如不到其他类的一半训练时就会出现类别不平衡后面评估指标会失真。常见做法是先跑这一步心里有数再决定要不要做重采样。2.2 用 image_dataset_from_directory 构建训练/验证管线TensorFlow 读取这种目录结构的数据集最省事的方式是image_dataset_from_directory。它自动完成三件事按目录分配标签、划分训练集和验证集、把图片统一 resize 到指定尺寸。import tensorflow as tf IMG_SIZE (224, 224) BATCH_SIZE 32 SEED 42 train_ds tf.keras.utils.image_dataset_from_directory( data_dir, validation_split0.2, subsettraining, seedSEED, image_sizeIMG_SIZE, batch_sizeBATCH_SIZE, label_modeint ) val_ds tf.keras.utils.image_dataset_from_directory( data_dir, validation_split0.2, subsetvalidation, seedSEED, image_sizeIMG_SIZE, batch_sizeBATCH_SIZE, label_modeint )参数逐个说清楚。validation_split0.2表示 20% 做验证这个比例在几千张量级的数据集上比较稳seed必须两边一致否则训练集和验证集会重叠这是血泪经验里最常见的一种“精度虚高”image_size设成 224×224 是因为后面要接的 MobileNetV2 或 ResNet50 默认输入就是 224label_modeint输出整数标签配合sparse_categorical_crossentropy损失函数用。如果你改成categorical标签会变成 one-hot损失函数也得换成categorical_crossentropy两者必须匹配。提示image_dataset_from_directory默认不打乱文件顺序再划分而是先打乱索引再取子集。如果你手动用os.listdir划分务必先 shuffle否则同一类图片可能全进训练集。2.3 缓存与预取让 GPU 不等数据数据管线搭好之后如果不做优化训练时 GPU 会频繁等 CPU 读图。标准做法是加cache()和prefetch()AUTOTUNE tf.data.AUTOTUNE train_ds train_ds.cache().shuffle(1000).prefetch(buffer_sizeAUTOTUNE) val_ds val_ds.cache().prefetch(buffer_sizeAUTOTUNE)cache()把解码后的图片存内存或本地缓存文件第二次 epoch 不用重新读盘shuffle(1000)维护一个 1000 张的缓冲区做乱序比全量 shuffle 省内存prefetch让数据准备和模型计算重叠。这三件套加上去同样硬件下每个 epoch 能快 20% 到 40%。数据集不大的话cache()直接放内存就行如果内存吃紧可以改成cache(filenamecache.tf-data)落盘。3. 迁移学习建模MobileNetV2 微调与训练参数怎么定3.1 为什么选 MobileNetV2 而不是从零搭 CNN5 类花卉、几千张图这个量级从零搭卷积网络也能跑但精度和收敛速度都不如迁移学习。MobileNetV2 在 ImageNet 上预训练过底层卷积已经学会了边缘、纹理、颜色块这些通用特征花卉分类恰好吃这一套。它的参数量约 340 万比 ResNet50 的 2500 万小一个量级CPU 上也能推理适合快速验证。base_model tf.keras.applications.MobileNetV2( input_shape(224, 224, 3), include_topFalse, weightsimagenet ) base_model.trainable False model tf.keras.Sequential([ tf.keras.layers.Rescaling(1./127.5, offset-1), base_model, tf.keras.layers.GlobalAveragePooling2D(), tf.keras.layers.Dropout(0.2), tf.keras.layers.Dense(5, activationsoftmax) ])include_topFalse去掉原模型最后的 1000 类分类头只保留特征提取部分trainableFalse先冻结预训练权重只训练新加的分类层。Rescaling(1./127.5, offset-1)把像素从 [0,255] 映射到 [-1,1]这是 MobileNetV2 要求的输入范围漏掉这一步精度会明显掉。GlobalAveragePooling2D把特征图压成向量比Flatten参数少得多。最后Dense(5, activationsoftmax)输出 5 类概率。3.2 编译参数优化器、学习率与损失函数model.compile( optimizertf.keras.optimizers.Adam(learning_rate1e-3), losssparse_categorical_crossentropy, metrics[accuracy] )冻结阶段学习率用 1e-3 比较合适因为只训练分类头梯度不会太大。损失函数用sparse_categorical_crossentropy对应整数标签。如果你在 2.2 里用了label_modecategorical这里必须换成categorical_crossentropy否则会报形状不匹配。评估指标先看 accuracy但花卉类别如果分布不均后面要补上混淆矩阵。3.3 两阶段训练先冻结再解冻微调直接解冻全部层一起训预训练权重容易被大梯度破坏这是新手常踩的坑。稳妥做法是分两阶段# 阶段一只训练分类头 history1 model.fit( train_ds, validation_dataval_ds, epochs10 ) # 阶段二解冻顶部若干层做微调 base_model.trainable True fine_tune_at 100 for layer in base_model.layers[:fine_tune_at]: layer.trainable False model.compile( optimizertf.keras.optimizers.Adam(learning_rate1e-5), losssparse_categorical_crossentropy, metrics[accuracy] ) history2 model.fit( train_ds, validation_dataval_ds, epochs10 )阶段一跑 10 个 epoch让分类头先收敛。阶段二把 MobileNetV2 的前 100 层继续冻结只解冻后面的层学习率降到 1e-5——比阶段一低两个数量级目的是微调而不是重训。fine_tune_at100这个值不是固定的MobileNetV2 一共 154 层从 100 往后解冻是常见做法如果你的数据集和 ImageNet 差异大可以往前调到 80。注意阶段二重新compile是必须的否则学习率不会生效。很多人忘了这一步结果微调阶段用的还是 1e-3loss 直接飞掉。3.4 训练过程监控与早停callbacks [ tf.keras.callbacks.EarlyStopping( monitorval_accuracy, patience5, restore_best_weightsTrue ), tf.keras.callbacks.ReduceLROnPlateau( monitorval_loss, factor0.5, patience3 ) ]EarlyStopping在验证精度 5 个 epoch 不提升时停掉并恢复最优权重避免过拟合ReduceLROnPlateau在验证损失停滞时把学习率砍半。这两个回调加上去能省掉不少手动调参的时间。把callbackscallbacks传进model.fit即可。4. 避坑与排查训练花卉分类器时最容易翻车的五件事4.1 验证集精度 99% 但预测全错现象训练日志里 val_accuracy 很快到 0.99但拿新图片预测结果离谱。原因训练集和验证集划分时 seed 不一致或者手动划分时没 shuffle导致验证集图片和训练集高度重叠甚至完全相同。解决确认两次image_dataset_from_directory的seed参数完全一致如果是手动划分先random.shuffle文件列表再切分。另外检查数据集本身有没有重复图片用 MD5 去重一遍。4.2 损失函数与标签模式不匹配报错现象model.fit一启动就报形状错误类似ValueError: Shapes (None, 5) and (None, 1) are incompatible。原因数据管线用了label_modeint输出形状(batch,)但损失函数用了categorical_crossentropy期望(batch, 5)或者反过来。解决label_modeint配sparse_categorical_crossentropylabel_modecategorical配categorical_crossentropy。两者必须成对出现改一个就得改另一个。4.3 图片损坏导致训练中途崩溃现象训练到某个 batch 突然报InvalidArgumentError: Unknown image file format或truncated JPEG。原因数据集里混入了下载不完整或格式损坏的图片image_dataset_from_directory在解码时才报错。解决训练前先跑一遍完整性检查from PIL import Image import pathlib bad_files [] for img_path in pathlib.Path(flower_dataset).rglob(*.jpg): try: img Image.open(img_path) img.verify() except Exception as e: bad_files.append((str(img_path), str(e))) print(f损坏文件数: {len(bad_files)}) for f, e in bad_files: print(f, e)把损坏文件删掉或替换后再训练。这个检查花不了几分钟但能省掉训练到一半崩溃的后悔药。4.4 显存不够OOM但 batch size 已经很小现象ResourceExhaustedError: OOM when allocating tensor即使 batch size 降到 8 还是报。原因cache()把整个数据集解码后存内存如果图片分辨率高、数量多内存先爆或者prefetch缓冲区设得太大。解决把cache()改成cache(filenamecache.tf-data)落盘prefetch的 buffer_size 从AUTOTUNE改成固定值 2同时确认image_size没有设得过大224 够用就别上 512。4.5 微调后精度反而下降现象阶段一 val_accuracy 到 0.85阶段二解冻微调后掉到 0.78。原因解冻层数太多或者学习率没降下来预训练权重被破坏。解决减少解冻层数把fine_tune_at从 100 调到 120确认微调学习率是 1e-5 而不是 1e-3微调 epoch 控制在 10 以内。如果还降说明数据集太小不适合微调直接用阶段一的结果就行。5. 从训练到落地导出 SavedModel、混淆矩阵与单图预测5.1 导出模型并在新图片上推理训练完的模型要能脱离训练脚本独立使用标准做法是存成 SavedModel 格式model.save(flower_classifier_v1)加载和预测import numpy as np import tensorflow as tf loaded_model tf.keras.models.load_model(flower_classifier_v1) class_names [daisy, dandelion, rose, sunflower, tulip] def predict_image(img_path): img tf.keras.utils.load_img(img_path, target_size(224, 224)) img_array tf.keras.utils.img_to_array(img) img_array tf.expand_dims(img_array, 0) predictions loaded_model.predict(img_array, verbose0) score tf.nn.softmax(predictions[0]) idx np.argmax(score) return class_names[idx], float(score[idx]) label, conf predict_image(test_rose.jpg) print(f预测: {label}, 置信度: {conf:.4f})load_img负责读图并 resizeimg_to_array转成 numpy 数组expand_dims加一个 batch 维度。softmax把输出转成概率分布argmax取最大概率的索引再映射回类别名。这套流程可以直接嵌到 Flask 或 FastAPI 里做推理接口。5.2 用混淆矩阵看模型到底错在哪accuracy 只告诉你整体对多少混淆矩阵才能看出哪两类容易混。花卉里玫瑰和郁金香在低分辨率下确实容易混from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns import matplotlib.pyplot as plt y_true [] y_pred [] for images, labels in val_ds: preds loaded_model.predict(images, verbose0) y_true.extend(labels.numpy()) y_pred.extend(np.argmax(preds, axis1)) cm confusion_matrix(y_true, y_pred) sns.heatmap(cm, annotTrue, fmtd, xticklabelsclass_names, yticklabelsclass_names) plt.xlabel(预测) plt.ylabel(真实) plt.show() print(classification_report(y_true, y_pred, target_namesclass_names))classification_report会给出每一类的 precision、recall、f1-score。如果某一类 recall 明显低说明这类被大量误判成别的类常见原因是这类样本太少或者图片风格和其他类差异大。针对性地补这类样本比盲目加 epoch 有效得多。5.3 一个具体技巧用 TFLite 把模型压到手机能跑如果最终目标是端侧部署SavedModel 还不够小。用 TFLite 转换并量化converter tf.lite.TFLiteConverter.from_saved_model(flower_classifier_v1) converter.optimizations [tf.lite.Optimize.DEFAULT] tflite_model converter.convert() with open(flower_classifier.tflite, wb) as f: f.write(tflite_model) print(fTFLite 模型大小: {len(tflite_model) / 1024:.1f} KB)Optimize.DEFAULT会做动态范围量化把权重从 float32 压到 int8模型体积通常能降到原来的四分之一左右精度损失一般在 1% 以内。转换完拿几张测试图对比一下 TFLite 和原模型的输出确认没有明显偏差再上线。从那以后我每次拿到新数据集第一件事不是写模型而是先跑目录统计和图片完整性检查再确认标签映射关系。这三步花不到十分钟但能挡掉后面八成的玄学问题。希望帮到你。本文还有配套的精品资源点击获取
返回列表