news 2026/9/29 2:04:31

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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会按字母序给类别编号,daisy=0、dandelion=1、rose=2、sunflower=3、tulip=4。如果你后面要输出预测结果,这个映射关系必须记牢,否则会出现“模型说 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_split=0.2, subset="training", seed=SEED, image_size=IMG_SIZE, batch_size=BATCH_SIZE, label_mode="int" ) val_ds = tf.keras.utils.image_dataset_from_directory( data_dir, validation_split=0.2, subset="validation", seed=SEED, image_size=IMG_SIZE, batch_size=BATCH_SIZE, label_mode="int" )

参数逐个说清楚。validation_split=0.2表示 20% 做验证,这个比例在几千张量级的数据集上比较稳;seed必须两边一致,否则训练集和验证集会重叠,这是血泪经验里最常见的一种“精度虚高”;image_size设成 224×224 是因为后面要接的 MobileNetV2 或 ResNet50 默认输入就是 224;label_mode="int"输出整数标签,配合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_size=AUTOTUNE) val_ds = val_ds.cache().prefetch(buffer_size=AUTOTUNE)

cache()把解码后的图片存内存或本地缓存文件,第二次 epoch 不用重新读盘;shuffle(1000)维护一个 1000 张的缓冲区做乱序,比全量 shuffle 省内存;prefetch让数据准备和模型计算重叠。这三件套加上去,同样硬件下每个 epoch 能快 20% 到 40%。数据集不大的话,cache()直接放内存就行;如果内存吃紧,可以改成cache(filename="cache.tf-data")落盘。

3. 迁移学习建模:MobileNetV2 微调与训练参数怎么定

3.1 为什么选 MobileNetV2 而不是从零搭 CNN

5 类花卉、几千张图这个量级,从零搭卷积网络也能跑,但精度和收敛速度都不如迁移学习。MobileNetV2 在 ImageNet 上预训练过,底层卷积已经学会了边缘、纹理、颜色块这些通用特征,花卉分类恰好吃这一套。它的参数量约 340 万,比 ResNet50 的 2500 万小一个量级,CPU 上也能推理,适合快速验证。

base_model = tf.keras.applications.MobileNetV2( input_shape=(224, 224, 3), include_top=False, weights="imagenet" ) 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, activation="softmax") ])

include_top=False去掉原模型最后的 1000 类分类头,只保留特征提取部分;trainable=False先冻结预训练权重,只训练新加的分类层。Rescaling(1./127.5, offset=-1)把像素从 [0,255] 映射到 [-1,1],这是 MobileNetV2 要求的输入范围,漏掉这一步精度会明显掉。GlobalAveragePooling2D把特征图压成向量,比Flatten参数少得多。最后Dense(5, activation="softmax")输出 5 类概率。

3.2 编译参数:优化器、学习率与损失函数

model.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=1e-3), loss="sparse_categorical_crossentropy", metrics=["accuracy"] )

冻结阶段学习率用 1e-3 比较合适,因为只训练分类头,梯度不会太大。损失函数用sparse_categorical_crossentropy对应整数标签。如果你在 2.2 里用了label_mode="categorical",这里必须换成categorical_crossentropy,否则会报形状不匹配。评估指标先看 accuracy,但花卉类别如果分布不均,后面要补上混淆矩阵。

3.3 两阶段训练:先冻结,再解冻微调

直接解冻全部层一起训,预训练权重容易被大梯度破坏,这是新手常踩的坑。稳妥做法是分两阶段:

# 阶段一:只训练分类头 history1 = model.fit( train_ds, validation_data=val_ds, epochs=10 ) # 阶段二:解冻顶部若干层做微调 base_model.trainable = True fine_tune_at = 100 for layer in base_model.layers[:fine_tune_at]: layer.trainable = False model.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5), loss="sparse_categorical_crossentropy", metrics=["accuracy"] ) history2 = model.fit( train_ds, validation_data=val_ds, epochs=10 )

阶段一跑 10 个 epoch,让分类头先收敛。阶段二把 MobileNetV2 的前 100 层继续冻结,只解冻后面的层,学习率降到 1e-5——比阶段一低两个数量级,目的是微调而不是重训。fine_tune_at=100这个值不是固定的,MobileNetV2 一共 154 层,从 100 往后解冻是常见做法;如果你的数据集和 ImageNet 差异大,可以往前调到 80。

注意:阶段二重新compile是必须的,否则学习率不会生效。很多人忘了这一步,结果微调阶段用的还是 1e-3,loss 直接飞掉。

3.4 训练过程监控与早停

callbacks = [ tf.keras.callbacks.EarlyStopping( monitor="val_accuracy", patience=5, restore_best_weights=True ), tf.keras.callbacks.ReduceLROnPlateau( monitor="val_loss", factor=0.5, patience=3 ) ]

EarlyStopping在验证精度 5 个 epoch 不提升时停掉并恢复最优权重,避免过拟合;ReduceLROnPlateau在验证损失停滞时把学习率砍半。这两个回调加上去,能省掉不少手动调参的时间。把callbacks=callbacks传进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_mode="int"(输出形状(batch,)),但损失函数用了categorical_crossentropy(期望(batch, 5));或者反过来。

解决:label_mode="int"配sparse_categorical_crossentropy,label_mode="categorical"配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(filename="cache.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, verbose=0) 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负责读图并 resize,img_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, verbose=0) y_true.extend(labels.numpy()) y_pred.extend(np.argmax(preds, axis=1)) cm = confusion_matrix(y_true, y_pred) sns.heatmap(cm, annot=True, fmt="d", xticklabels=class_names, yticklabels=class_names) plt.xlabel("预测") plt.ylabel("真实") plt.show() print(classification_report(y_true, y_pred, target_names=class_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(f"TFLite 模型大小: {len(tflite_model) / 1024:.1f} KB")

Optimize.DEFAULT会做动态范围量化,把权重从 float32 压到 int8,模型体积通常能降到原来的四分之一左右,精度损失一般在 1% 以内。转换完拿几张测试图对比一下 TFLite 和原模型的输出,确认没有明显偏差再上线。

从那以后我每次拿到新数据集,第一件事不是写模型,而是先跑目录统计和图片完整性检查,再确认标签映射关系。这三步花不到十分钟,但能挡掉后面八成的玄学问题。希望帮到你。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/29 2:04:01

BUCK电路从原理到实测:电感选型、同步整流与环路补偿

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/29 2:03:11

Jessibuca 播放器底部控制栏完全指南:5个开关自定义按钮与自动隐藏

Jessibuca 播放器底部控制栏完全指南:5个开关自定义按钮与自动隐藏 【免费下载链接】jessibuca Jessibuca 是一款开源的纯H5直播流播放器,通过Emscripten将音视频解码库编译成Js(wasm)运行于浏览器之中。兼容几乎所有浏览器,可以运…

作者头像 李华
网站建设 2026/9/29 2:03:10

TMS AI Studio:Delphi桌面应用接入大语言模型对话能力的组件之道

简介:这是一套专为 Delphi 13 环境打造的 AI 与机器学习控件套件,面向希望在桌面、移动或 Web 应用中快速引入智能功能的 Delphi 开发者,也适合有一定基础、希望深入 AI 集成的中高级程序员。套件整合了模型训练、数据预处理、自然语言处理、…

作者头像 李华
网站建设 2026/9/29 2:02:48

Ubuntu服务器自建Git仓库实战:从SSH配置到Hook自动部署

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/29 2:02:15

AGX Orin载板设计避坑指南:电源时序、存储选型与高速接口布局

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华