深度学习实战-基于EfficientNetB0的观赏鱼图像分类识别模型

🤵♂️ 个人主页:@艾派森的个人主页
✍🏻作者简介:Python学习者
🐋 希望大家多多支持,我们一起进步!😄
如果文章对你有帮助的话,
欢迎评论 💬点赞👍🏻 收藏 📂加关注+
目录

1.项目背景
在现代水族养殖与水生生物学的精准化管理中,鱼类物种的快速且准确识别是确保水务生态监测、智能化精准投喂以及观赏鱼市场规范化交易的核心环节。然而,观赏鱼通常具有体型多变、色彩斑斓且游动迅速的特点,加之水族箱内水体折射、水草摆动以及光影错综复杂,传统的肉眼辨识不仅耗时耗力,且在面对体态或色系相近的品种时极易产生视觉误判。随着数字化和智能农业的深入发展,如何利用计算机视觉技术替代繁琐的人工鉴定,并在复杂的水体环境下实现稳健、高精度的品种分类,已成为构建智慧水族生态系统进程中亟待攻克的实战课题。
本项目针对观赏鱼表面斑斓错综的纹理与空间形态的多样性,展开了多架构卷积神经网络的深度应用与横向对比研究。实验从构建基础的 Custom CNN(自定义卷积神经网络) 架构出发,建立起分类任务的性能基准线;随后全面引入迁移学习策略,深度部署了轻量化倒残差网络 MobileNetV2、基于复合缩放机制的 EfficientNetB0 核心网络,以及具备强大深度残差拟合能力的 ResNet50 架构。通过在这四个具有代表性的网络模型间展开全方位的拟合对比与跑分评测,本实战不仅展示了深度特征提取器在剥离背景噪声、锁定鱼体核心视觉特征方面的卓越表现,更通过混淆矩阵等多维指标深度透视了各模型在相似品种间的判别边界,为开发便携式水生生物智能鉴定设备或嵌入式巡检系统提供了可落地的算法参考与技术闭环。
2.数据集介绍
本实验数据集来源于Kaggle,原始数据集为水族箱鱼类分类数据集,是一个精心整理的图像数据集,专为鱼类物种识别和图像分类任务而设计。它包含1016张水族箱鱼类的RGB图像,这些鱼类被分为六个物种类别。该数据集旨在支持涉及水生生物物种识别的机器学习、深度学习和计算机视觉研究。它既适用于学习图像分类的初学者,也适用于评估迁移学习和卷积神经网络架构的研究人员。
Classes:
- Bete (194 images)
- Cray (80 images)
- Discuss (201 images)
- Gold (207 images)
- Guppy (189 images)
- Oscar (145 images)
数据集统计信息
图片总数:1016张
3.技术工具
Python版本:3.9
代码编辑器:jupyter notebook
4.实验过程
4.1导入数据
在搭建深度学习模型的初期,环境的配置与物理数据的结构化解析是决定流水线能否高效运转的关键。我们首先集成了数值计算、文件系统操作以及深度学习的核心库,并引入了包括 EfficientNetB0 在内的多种主流经典架构的预处理接口,为后续的迁移学习对比实验做好铺垫。针对观赏鱼数据集的存储特性,本阶段通过 Python 接口对目标路径执行自动化检索,过滤掉潜在的系统杂质文件,精准提取出各观赏鱼品种的文件夹名称,并对整个数据集的样本总量及各类别的分布基数进行统计。这种数据基准的建立,不仅能让我们在训练前洞察是否存在数据不平衡问题,也为后续通过张量流流式加载影像奠定了坚实的逻辑基础。
import os
import random
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import tensorflow as tf
from tensorflow.keras import layers, models
from tensorflow.keras.preprocessing import image_dataset_from_directory
from tensorflow.keras.applications import MobileNetV2, EfficientNetB0, ResNet50
from tensorflow.keras.applications.mobilenet_v2 import preprocess_input as mobilenet_preprocess
from tensorflow.keras.applications.efficientnet import preprocess_input as efficientnet_preprocess
from tensorflow.keras.applications.resnet50 import preprocess_input as resnet_preprocess
from tensorflow.keras.callbacks import EarlyStopping, ReduceLROnPlateau, ModelCheckpoint
from sklearn.metrics import classification_report, confusion_matrix
# --- 1. 全局超参数与数据路径配置 ---
DATASET_PATH = "/kaggle/input/datasets/jannatulferdaues/aquarium-fish-classification/Final (°_°)"
IMG_SIZE = (224, 224) # 适配 EfficientNetB0 的标准输入分辨率
BATCH_SIZE = 32 # 批大小
SEED = 42 # 随机种子,确保实验可重复性
# 验证根路径是否存在
print("Dataset path exists:", os.path.exists(DATASET_PATH))
# --- 2. 自动化类别名提取 ---
# 扫描目录并排序,过滤出所有的子文件夹作为观赏鱼的品种标签
class_names = sorted([
item for item in os.listdir(DATASET_PATH)
if os.path.isdir(os.path.join(DATASET_PATH, item))
])
print("Classes:", class_names)
# --- 3. 数据集规模与样本分布统计 ---
total_images = 0
class_counts = {}
valid_exts = (".jpg", ".jpeg", ".png", ".bmp", ".webp") # 定义合法的图像扩展名
# 遍历每个品种文件夹,统计有效观赏鱼图片数量
for cls in class_names:
cls_path = os.path.join(DATASET_PATH, cls)
images = [
img for img in os.listdir(cls_path)
if img.lower().endswith(valid_exts)
]
count = len(images)
class_counts[cls] = count
total_images += count
# 输出统计结果
print("Total Images:", total_images)
print("Class Counts:", class_counts)

4.2数据可视化

观赏鱼品种数据分布量化分析
在训练模型之前,掌握数据集中每个观赏鱼品种的样本占比至关重要。如果某些珍稀鱼类的图像数量远少于常见品种,模型在优化过程中很容易产生类别偏好。为此,我们首先将统计得到的品种字典转化为 Pandas 的 DataFrame 结构,并联合调用 Matplotlib 构建了柱状图与饼图。柱状图清晰地反映了各分类的绝对图像绝对数量,而饼图则以百分比的形式展现了其相对权重,这种双轨制的图表互补能让我们从统计学层面确认数据集的健康程度。
# --- 1. 将样本统计字典转化为结构化 DataFrame ---
df_counts = pd.DataFrame({
"Class": list(class_counts.keys()),
"Images": list(class_counts.values())
})
# --- 2. 绘制条形图:直观审视各品种绝对数量 ---
plt.figure(figsize=(9, 5))
plt.bar(df_counts["Class"], df_counts["Images"])
plt.title("Class Distribution")
plt.xlabel("Fish Class")
plt.ylabel("Number of Images")
plt.xticks(rotation=45) # 标签倾斜 45 度,防止品种名称相互重叠
plt.show()
# --- 3. 绘制饼图:客观评估类别相对占比 ---
plt.figure(figsize=(7, 7))
plt.pie(df_counts["Images"], labels=df_counts["Class"], autopct="%1.1f%%")
plt.title("Class Percentage")
plt.show()


随机样本抽检与生物形态审视
定量的统计只能告诉我们数据的多寡,而定性的视觉抽检才能揭示数据质量的优劣。利用 random.choice 机制,我们从每个观赏鱼品种文件夹中随机抽取了一张具有代表性的物理图片,并以 2 x 3的网格矩阵进行集中渲染。通过移除多余的物理坐标轴,我们可以排除背景杂质的视觉干扰,将焦点锁定在金鱼的尾鳍、神仙鱼的扁平侧躯以及各色热带鱼独特的斑纹上。这种直观的样本回显有助于我们核验数据集是否存在严重的噪点、模糊或错误的标签划分。
# --- 4. 动态网格组装:随机抽检各品种图像特征 ---
plt.figure(figsize=(15, 10))
for i, cls in enumerate(class_names):
cls_path = os.path.join(DATASET_PATH, cls)
# 动态获取当前品种目录下所有符合规范的图像列表
img_files = [
img for img in os.listdir(cls_path)
if img.lower().endswith(valid_exts)
]
# 随机挑选当前分类下的一张鱼类照片
img_name = random.choice(img_files)
img_path = os.path.join(cls_path, img_name)
# 读取图像矩阵并分配到对应的子图网格
img = plt.imread(img_path)
plt.subplot(2, 3, i + 1)
plt.imshow(img)
plt.title(cls) # 将子图标题设定为对应的鱼类品种
plt.axis("off") # 移除坐标轴,净化图像版面
plt.tight_layout()
plt.show()

4.3特征工程
流式数据划分与异步预取机制
本环节我们通过调用 image_dataset_from_directory 接口,将底层的物理图像流式切分为 80% 的训练集与 20% 的验证集,并将尺寸统一重置为 224 x 224 的标准分辨率。为了消除硬件 I/O 带来的性能瓶颈,我们构建了一个基于数据层面的特征增强序列 data_augmentation,包含随机水平翻转、旋转、缩放与对比度调整。随后,利用 shuffle(1000) 打乱训练序列,并全面引入 prefetch(tf.data.AUTOTUNE) 动态预取技术。这一优化使得 CPU 能够在 GPU 进行当前批次反向传播时,异步在内存中完成下一批次图像的矩阵解码与增量增强,从而实现流水线级别的高效并行计算。
# --- 1. 构建实时数据增强序列 ---
# 在 Keras 序贯模型中嵌入随机算子,使数据在训练时动态产生空间与色彩扰动
data_augmentation = tf.keras.Sequential([
layers.RandomFlip("horizontal"), # 模拟鱼群由右向左或由左向右游动
layers.RandomRotation(0.15), # 模拟水流扰动引起的拍摄倾斜
layers.RandomZoom(0.15), # 模拟拍摄距离的远近变化
layers.RandomContrast(0.15), # 模拟水族箱内部不同强度的光源折射
], name="data_augmentation")
# --- 2. 流式加载与训练/验证集自动切分 ---
# 从指定物理路径流式加载数据,并划分出 80% 的训练集
train_ds = image_dataset_from_directory(
DATASET_PATH,
validation_split=0.2,
subset="training",
seed=SEED,
image_size=IMG_SIZE,
batch_size=BATCH_SIZE
)
# 从同一物理路径下流式划分出 20% 的验证集
val_ds = image_dataset_from_directory(
DATASET_PATH,
validation_split=0.2,
subset="validation",
seed=SEED,
image_size=IMG_SIZE,
batch_size=BATCH_SIZE
)
# 提取并解析分类标签信息
class_names = train_ds.class_names
NUM_CLASSES = len(class_names)
print("Class Names:", class_names)
print("Number of Classes:", NUM_CLASSES)
# --- 3. 数据流水线性能深度优化 ---
AUTOTUNE = tf.data.AUTOTUNE
# shuffle: 大内存打乱防止连续批次过拟合;prefetch: 异步预取避免 GPU 空转
train_ds = train_ds.shuffle(1000, seed=SEED).prefetch(AUTOTUNE)
val_ds = val_ds.prefetch(AUTOTUNE)
训练监控机制与性能可视化函数封装
为了让模型的收敛过程更加稳健,我们封装了 get_callbacks 函数,它集成了三大核心回调机制:EarlyStopping 监控验证集准确率,在连续 5 轮未见改善时自动触发保护性熔断并回滚至历史最优权重;ReduceLROnPlateau 监控验证集损失,在连续 3 轮陷入平台期时将学习率缩减为原先的 20%,赋予模型在狭窄地形中的微调能力;ModelCheckpoint 则作为安全策略,实时将表现最好的模型以 .keras 格式固化到磁盘中。最后,通过 plot_history 函数对模型生成的训练日志执行动态图表渲染,以便在后续步骤中多维度复盘准确率与损失函数的演进轨迹。
# --- 4. 封装多维训练监控回调策略 ---
def get_callbacks(model_name):
"""
配置智能化训练监控器组件
"""
return [
# 早期停止:防止无效训练,避免模型过度拟合特定样本
EarlyStopping(
monitor="val_accuracy",
patience=5, # 容忍轮次
restore_best_weights=True # 触发时自动还原历史最高权重
),
# 学习率动态衰减:应对梯度平原陷阱,确保平滑收敛
ReduceLROnPlateau(
monitor="val_loss",
factor=0.2, # 缩减因子
patience=3, # 容忍轮次
min_lr=1e-7 # 学习率下限
),
# 模型权重实时持久化检查点
ModelCheckpoint(
f"{model_name}.keras",
monitor="val_accuracy",
save_best_only=True # 仅保存验证集准确率最高的那一轮模型
)
]
# --- 5. 封装性能曲线可视化算子 ---
def plot_history(history, title):
"""
绘制并对比训练集与验证集的准确率与损失值轨迹
"""
# 渲染准确率趋势子图
plt.figure(figsize=(8, 5))
plt.plot(history.history["accuracy"], label="Train Accuracy")
plt.plot(history.history["val_accuracy"], label="Validation Accuracy")
plt.title(title + " Accuracy")
plt.xlabel("Epoch")
plt.ylabel("Accuracy")
plt.legend()
plt.show()
# 渲染损失函数收敛子图
plt.figure(figsize=(8, 5))
plt.plot(history.history["loss"], label="Train Loss")
plt.plot(history.history["val_loss"], label="Validation Loss")
plt.title(title + " Loss")
plt.xlabel("Epoch")
plt.ylabel("Loss")
plt.legend()
plt.show()
4.4构建并训练CNN模型
本阶段我们通过 Keras 的 Sequential 序贯模型,从零组装了一套 4 层的标准二维卷积网络。输入端的图像矩阵在历经特征增强后,首先通过 Rescaling(1./255) 算子执行标准化的归一化映射。随后,网络通过 3 x 3 的特征提取卷积核,将通道数从 32、64 逐步推进至 128 和 256,利用浅层捕捉金鱼鱼鳍的硬棘边缘,深层锁定其独特的斑纹。网络末端舍弃了臃肿的全连接展平,改用 GlobalAveragePooling2D(全局平均池化)以大幅砍掉参数冗余,最终配合 Dropout(0.4) 抵御过拟合,并依托 Adam 优化器在 20 轮的最大预期周期内展开拟合探索。
# --- 1. 组装自定义 CNN 拓扑网络 ---
custom_cnn = models.Sequential([
layers.Input(shape=(224, 224, 3)), # 定义标准的输入三维张量
data_augmentation, # 嵌入在特征工程中配置的实时增强序列
layers.Rescaling(1./255), # 将 [0, 255] 的像素值缩放到 [0, 1] 空间
# 卷积特征提取层组 1:捕捉图像低频浅层纹理
layers.Conv2D(32, 3, activation="relu"),
layers.MaxPooling2D(),
# 卷积特征提取层组 2:提取边缘与局部几何特征
layers.Conv2D(64, 3, activation="relu"),
layers.MaxPooling2D(),
# 卷积特征提取层组 3:抽象出多维空间结构
layers.Conv2D(128, 3, activation="relu"),
layers.MaxPooling2D(),
# 卷积特征提取层组 4:聚合深层高维语义
layers.Conv2D(256, 3, activation="relu"),
layers.MaxPooling2D(),
# 顶层分类决策头:使用全局平均池化压缩特征图尺寸
layers.GlobalAveragePooling2D(),
layers.Dense(256, activation="relu"),
layers.Dropout(0.4), # 随机失活 40% 的神经元,防止死记硬背样本
layers.Dense(NUM_CLASSES, activation="softmax") # 输出层,计算各观赏鱼类别的归一化概率
])
# --- 2. 编译模型:配置底层的损失度量与优化算法 ---
custom_cnn.compile(
optimizer=tf.keras.optimizers.Adam(learning_rate=1e-3),
# 注意:由于使用了 image_dataset_from_directory 默认返回的整数标签,故采用稀疏多分类交叉熵
loss="sparse_categorical_crossentropy",
metrics=["accuracy"]
)
# --- 3. 启动端到端训练流水线 ---
history_cnn = custom_cnn.fit(
train_ds,
validation_data=val_ds,
epochs=20,
callbacks=get_callbacks("custom_cnn") # 挂载动态学习率调节与安全检查点组件
)
# --- 4. 渲染性能趋势图表 ---
plot_history(history_cnn, "Custom CNN")



4.5构建并训练MobileNetV2模型
本环节我们通过编写 build_transfer_model 函数,实现了一套高内聚、低耦合的迁移学习骨架生成器。该函数首先将传入的 MobileNetV2 骨干网络(Base Model)完全冻结,锁死其在海量分类任务中训练好的卷积滤波器权重。在数据流向设计上,输入的观赏鱼张量在历经 data_augmentation 的实时空间扰动后,会直接喂入 MobileNetV2 专有的数据归一化算子 mobilenet_preprocess。通过设置 training=False 确保批归一化(Batch Normalization)层在特征提取时保持稳定,最后将高阶语义流灌入由 Dense 和 Dropout(0.4) 组成的轻量分类头中,以 0.0001 的低精细学习率开启参数拟合。
# --- 1. 封装通用的迁移学习流水线构建器 ---
def build_transfer_model(base_model, preprocess_func, model_name):
# 冻结骨干网络的所有卷积权重,防止预训练的通用特征提取器遭到破坏
base_model.trainable = False
inputs = layers.Input(shape=(224, 224, 3))
# 挂载数据增强层,模拟复杂的物理采样噪声
x = data_augmentation(inputs)
# 挂载模型专属的预处理函数(如 MobileNetV2 的像素值缩放)
x = preprocess_func(x)
# 提取预训练骨干网络的特征图,注意设定 training=False 以稳定 BN 层的统计量
x = base_model(x, training=False)
# 压缩特征空间并挂载自定义密集分类层
x = layers.GlobalAveragePooling2D()(x)
x = layers.Dense(256, activation="relu")(x)
x = layers.Dropout(0.4)(x) # 引入 40% 的随机失活率防止特征依赖
outputs = layers.Dense(NUM_CLASSES, activation="softmax")(x)
# 封装最终的端到端计算图模型
model = models.Model(inputs, outputs, name=model_name)
# 编译模型:针对微调任务,选用更加保守细致的微小学习率
model.compile(
optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4),
loss="sparse_categorical_crossentropy",
metrics=["accuracy"]
)
return model
# --- 2. 实例化预训练 MobileNetV2 骨干网络 ---
# 去掉顶层千分类全连接头,锁定标准的 224x224x3 输入拓扑
mobilenet_base = MobileNetV2(
weights="imagenet",
include_top=False,
input_shape=(224, 224, 3)
)
# --- 3. 组装并训练 MobileNetV2 鱼类诊断模型 ---
mobilenet_model = build_transfer_model(
mobilenet_base,
mobilenet_preprocess,
"MobileNetV2"
)
# 打印模型宏观结构与可训练参数分布
mobilenet_model.summary()
# 启动拟合流程并持久化最佳检查点权重
history_mobilenet = mobilenet_model.fit(
train_ds,
validation_data=val_ds,
epochs=20,
callbacks=get_callbacks("mobilenetv2")
)
# 渲染训练过程的收敛状态图表
plot_history(history_mobilenet, "MobileNetV2")



4.6构建并训练EfficientNetB0模型
本环节我们正式引入预训练的 EfficientNetB0 进行特征空间的适配整合。我们复用了先前封装的高效迁移学习流构建器 build_transfer_model,将去掉千分类顶层的 ImageNet 权重实例与 EfficientNet 专用的预处理算子 efficientnet_preprocess 深度绑定。这种流式配置能自动将输入的观赏鱼矩阵调整至网络最舒适的数值分布空间。通过冻结主干卷积特征提取器,我们强迫自定义的顶层分类头在 0.0001 的精细梯度探索下,将 EfficientNetB0 输出的高阶语义张量转化为对水族品种的精确分类概率,并在 20 轮的最大生命周期内展开拟合。
# --- 1. 实例化预训练 EfficientNetB0 骨干网络 ---
# 剔除原有千分类顶层,注入标准的 224x224x3 生物影像张量
efficientnet_base = EfficientNetB0(
weights="imagenet",
include_top=False,
input_shape=(224, 224, 3)
)
# --- 2. 依托通用流水线组装 EfficientNetB0 核心模型 ---
efficientnet_model = build_transfer_model(
efficientnet_base,
efficientnet_preprocess,
"EfficientNetB0"
)
# 打印并核验该核心模型的参数拓扑与层级连接
efficientnet_model.summary()
# --- 3. 启动高增益参数拟合流程 ---
# 挂载包含自动熔断(EarlyStopping)与动态降速(ReduceLROnPlateau)的专用回调组件
history_efficientnet = efficientnet_model.fit(
train_ds,
validation_data=val_ds,
epochs=20,
callbacks=get_callbacks("efficientnetb0")
)
# --- 4. 实时渲染训练收敛趋势图表 ---
plot_history(history_efficientnet, "EfficientNetB0")



4.7构建并训练ResNet50模型
本阶段我们正式加载预训练的 ResNet50 架构。为了保持对照实验的严谨性,我们继续沿用先前封装的经典迁移学习流水线构建器 build_transfer_model,将剔除了原始千分类全连接头的 ResNet50 骨干网络与专用的特征标准化算子 resnet_preprocess 进行深度绑定。通过将主干网络的参数全部冻结,强迫分类头在 0.0001 的精细梯度探索下,专职优化针对观赏鱼 5 类品种的分类映射。模型同样挂载了自动熔断与动态学习率调节组件,在 20 轮的最大预期生命周期内展开拟合。
# --- 1. 实例化预训练 ResNet50 骨干网络 ---
# 剔除原有千分类全连接顶层,锁定标准的 224x224x3 输入拓扑
resnet_base = ResNet50(
weights="imagenet",
include_top=False,
input_shape=(224, 224, 3)
)
# --- 2. 依托通用流水线组装 ResNet50 鱼类分类模型 ---
resnet_model = build_transfer_model(
resnet_base,
resnet_preprocess,
"ResNet50"
)
# 打印并核验 ResNet50 模型的参数分布与结构拓扑
resnet_model.summary()
# --- 3. 启动端到端参数拟合流程 ---
# 挂载动态调速与安全检查点组件,确保模型平滑收敛
history_resnet = resnet_model.fit(
train_ds,
validation_data=val_ds,
epochs=20,
callbacks=get_callbacks("resnet50")
)
# --- 4. 实时渲染训练收敛趋势图表 ---
plot_history(history_resnet, "ResNet50")



4.8模型评估
多架构横向跑分与最优模型自动筛选
本环节我们首先构建了一个包含四大模型的字典 models_dict,并利用循环结构流式调用 model.evaluate 对验证集 val_ds 进行集中跑分测试。为了让对比一目了然,我们利用 Pandas 将提取出的验证集准确率(Validation Accuracy)封装为结构化的 DataFrame,并联合 Matplotlib 绘制了直观的柱状图。系统随后通过对准确率执行降序排列(sort_values),以全自动的逻辑锁定了本次实验的表现最优模型(best_model),从而消除了人工挑选的主观偏差。
# --- 1. 初始化评估容器与对照字典 ---
results = {}
models_dict = {
"Custom CNN": custom_cnn,
"MobileNetV2": mobilenet_model,
"EfficientNetB0": efficientnet_model,
"ResNet50": resnet_model
}
# --- 2. 自动化循环跑分测试 ---
for name, model in models_dict.items():
# 静默评估验证集指标
loss, acc = model.evaluate(val_ds, verbose=0)
results[name] = acc
# --- 3. 结构化表格转换与可视化图表渲染 ---
results_df = pd.DataFrame({
"Model": list(results.keys()),
"Validation Accuracy": list(results.values())
})
# 在控制台或 Notebook 中以表格形式直观展现跑分
display(results_df)
# 绘制多模型性能对比柱状图
plt.figure(figsize=(9, 5))
plt.bar(results_df["Model"], results_df["Validation Accuracy"])
plt.title("Model Comparison")
plt.xlabel("Model")
plt.ylabel("Validation Accuracy")
plt.ylim(0, 1) # 将纵坐标固定在 [0, 1] 区间,确保对比公平性
plt.xticks(rotation=30)
plt.show()
# --- 4. 自动检索性能最优的骨干网络模型 ---
best_model_name = results_df.sort_values(
by="Validation Accuracy",
ascending=False
).iloc[0]["Model"]
best_model = models_dict[best_model_name]
print("Best Model:", best_model_name)


最优模型深度透视:混淆矩阵与精细化分类报告
为了对筛选出的最优模型进行深入解构,我们首先遍历验证集,通过 np.argmax 锁定模型预测置信度最高的类别,并将真实标签与预测标签分别归集入 y_true 与 y_pred 阵列。随后,我们利用 confusion_matrix 算子绘制了双向混淆矩阵。矩阵的对角线代表预测正确的样本数,而其余方格则量化了各品种间产生误判的绝对数量。最后,通过调用 classification_report 打印出针对各个观赏鱼品种的精确率(Precision)、召回率(Recall) 以及 F1-Score,为模型的工程落地提供坚实的数理背书。
# --- 5. 遍历验证集提取真实标签与预测预测值 ---
y_true = []
y_pred = []
for images, labels in val_ds:
# 预测当前批次的概率分布
preds = best_model.predict(images, verbose=0)
pred_labels = np.argmax(preds, axis=1)
# 将张量数据平铺扩展到标准的 Python 列表中
y_true.extend(labels.numpy())
y_pred.extend(pred_labels)
# 转换为 NumPy 数组以适配后续评估接口
y_true = np.array(y_true)
y_pred = np.array(y_pred)
# =========================================================
# 14. Confusion Matrix (混淆矩阵可视化)
# =========================================================
# 计算多分类混淆矩阵矩阵
cm = confusion_matrix(y_true, y_pred)
plt.figure(figsize=(8, 6))
plt.imshow(cm) # 使用内置色彩热图渲染矩阵
plt.title(f"Confusion Matrix - {best_model_name}")
plt.colorbar()
plt.xticks(np.arange(NUM_CLASSES), class_names, rotation=45)
plt.yticks(np.arange(NUM_CLASSES), class_names)
# 嵌套循环:在矩阵的各个格子正中心实时打印样本频数文本
for i in range(NUM_CLASSES):
for j in range(NUM_CLASSES):
plt.text(j, i, cm[i, j], ha="center", va="center")
plt.xlabel("Predicted")
plt.ylabel("Actual")
plt.tight_layout()
plt.show()
# =========================================================
# 15. Classification Report (多指标详细分类报告)
# =========================================================
# 自动计算并输出针对各品种的 Precision, Recall 和 F1 指标
report = classification_report(
y_true,
y_pred,
target_names=class_names
)
print(report)


最后可以保存模型
best_model.save("best_aquarium_fish_classifier.keras")
5.总结
本实验围绕水族箱鱼类物种识别任务,利用包含 1016 张 RGB 高清影像的观赏鱼数据集,成功搭建并系统评估了四种不同的图像分类网络。该数据集涵盖了 Bete、Cray、Discuss、Gold、Guppy 和 Oscar 六类典型的水生生物物种,在经历了严谨的数据增强与流式分发流水线后,分别喂入自定义 CNN 基准网络以及三大主流迁移学习架构中。横向跑分与多维评估结果表明,相较于从头训练且验证集准确率仅为 63.05% 的自定义 CNN,引入 ImageNet 预训练先验的迁移学习网络展现出了压倒性的特征提取优势。其中,基于复合缩放机制的 EfficientNetB0 骨干网络凭借其特有的通道注意力机制,成功克服了水体折射、水草背景等非目标噪声的干扰,以 97.04% 的最高验证集准确率夺得本次实验的最优模型桂型。在对最优模型的细粒度解构中,分类报告显示其加权平均 F1-score 达到了 0.97,特别是对 Discuss(神仙鱼)品种实现了 100% 的精准识别与完全召回,而对其他品种也保持了极高且均衡的判别边界。本实战不仅有力论证了轻量化复合缩放网络在复杂生物特征识别中的卓越泛化性能,也为未来将算法移植到移动端设备或智能水族巡检硬件上提供了极具工程落地价值的参考范例。
资料获取,更多粉丝福利,关注下方公众号获取

更多推荐

所有评论(0)