You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于深度学习的蘑菇物种识别模型精度优化技术问询

蘑菇物种识别模型优化方案

问题背景

我正在开发一个基于深度学习的蘑菇物种识别程序,已构建包含26类蘑菇的数据集,每类约500张图片,但当前模型的最高识别精度仅约45%,希望通过优化措施或调整模型结构提升性能。

现有代码

import keras.losses
import tensorflow as tf
import matplotlib.pyplot as plt
import pathlib

def create_model():
    num_classes = len(train_ds.class_names)
    model = tf.keras.Sequential([
        tf.keras.layers.Rescaling(1. / 255),
        tf.keras.layers.Conv2D(32, 3, activation='relu'),
        tf.keras.layers.MaxPooling2D(),
        tf.keras.layers.Conv2D(32, 3, activation='relu'),
        tf.keras.layers.MaxPooling2D(),
        tf.keras.layers.Dropout(0.4),
        tf.keras.layers.Conv2D(32, 3, activation='relu'),
        tf.keras.layers.MaxPooling2D(),
        tf.keras.layers.Flatten(),
        tf.keras.layers.Dense(128, activation="relu"),
        tf.keras.layers.Dropout(0.2),
        tf.keras.layers.Dense(num_classes)
    ])
    model.compile(
        optimizer='adam',
        loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
        metrics=['accuracy']
    )
    return model

def load_data():
    data_dir = pathlib.Path("Mushrooms")
    image_count = len(list(data_dir.glob('*/*')))
    train_ds = tf.keras.utils.image_dataset_from_directory(
        data_dir,
        validation_split=0.2,
        subset="training",
        seed=123,
        image_size=(img_height, img_width),
        batch_size=batch_size
    )

    val_ds = tf.keras.utils.image_dataset_from_directory(
        data_dir,
        validation_split=0.2,
        subset="validation",
        seed=123,
        image_size=(img_height, img_width),
        batch_size=batch_size
    )
    return train_ds, val_ds

# hyperparams
batch_size = 32
img_height = 100
img_width = 100
epochs = 12

train_ds, val_ds = load_data()
model = create_model()
#model = tf.keras.models.load_model('model')

history = model.fit(
    train_ds,
    validation_data=val_ds,
    epochs=epochs
)

acc = history.history['accuracy']
val_acc = history.history['val_accuracy']

loss = history.history['loss']
val_loss = history.history['val_loss']

epochs_range = range(epochs)

plt.figure(figsize=(8, 8))
plt.subplot(1, 2, 1)
plt.plot(epochs_range, acc, label='Training Accuracy')
plt.plot(epochs_range, val_acc, label='Validation Accuracy')
plt.legend(loc='lower right')
plt.title('Training and Validation Accuracy')

plt.subplot(1, 2, 2)
plt.plot(epochs_range, loss, label='Training Loss')
plt.plot(epochs_range, val_loss, label='Validation Loss')
plt.legend(loc='upper right')
plt.title('Training and Validation Loss')
plt.show()

model.save("model", overwrite=True, save_format='tf')

优化措施

1. 改用预训练模型做迁移学习

当前自定义CNN结构过于简单,3层Conv2D仅用32个滤波器,无法提取蘑菇的精细特征。建议基于ImageNet预训练模型(如EfficientNetB0、ResNet50)做迁移学习,大幅提升特征提取能力:

def create_model():
    num_classes = len(train_ds.class_names)
    # 加载预训练模型,冻结顶层以外的参数
    base_model = tf.keras.applications.EfficientNetB0(
        input_shape=(img_height, img_width, 3),
        include_top=False,
        weights='imagenet'
    )
    base_model.trainable = False
    
    model = tf.keras.Sequential([
        tf.keras.layers.Rescaling(1./255),
        base_model,
        tf.keras.layers.GlobalAveragePooling2D(),  # 替代Flatten,减少参数
        tf.keras.layers.Dense(256, activation='relu'),
        tf.keras.layers.Dropout(0.5),
        tf.keras.layers.Dense(num_classes)
    ])
    model.compile(
        optimizer='adam',
        loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
        metrics=['accuracy']
    )
    return model

训练完顶层后,可解冻预训练模型的最后10-20层,用1e-5的小学习率继续微调,进一步挖掘特征潜力。

2. 添加数据增强扩充训练样本

蘑菇存在角度、光照、大小差异,数据增强能有效提升模型泛化能力:

# 定义增强策略
data_augmentation = tf.keras.Sequential([
    tf.keras.layers.RandomFlip("horizontal_and_vertical"),
    tf.keras.layers.RandomRotation(0.2),
    tf.keras.layers.RandomZoom(0.2),
    tf.keras.layers.RandomContrast(0.2)
])

# 仅对训练集应用增强
train_ds = train_ds.map(lambda x, y: (data_augmentation(x, training=True), y))

3. 调整核心超参数

  • 图片尺寸:当前100x100太小,建议改为224x224(适配多数预训练模型输入),保留更多细节特征。
  • 训练轮次:12轮不足以充分训练,建议先训练30轮,搭配早停机制避免过拟合:
early_stopping = tf.keras.callbacks.EarlyStopping(
    monitor='val_accuracy',
    patience=5,
    restore_best_weights=True
)
# 在model.fit中添加callbacks=[early_stopping]
  • 学习率调度:默认Adam学习率0.001,微调预训练层时需降到1e-5~1e-4;也可添加学习率衰减:
lr_scheduler = tf.keras.callbacks.ReduceLROnPlateau(
    monitor='val_loss',
    factor=0.5,
    patience=3,
    min_lr=1e-6
)

4. 优化数据加载效率

开启缓存与预取,减少训练等待时间:

train_ds = train_ds.cache().prefetch(buffer_size=tf.data.AUTOTUNE)
val_ds = val_ds.cache().prefetch(buffer_size=tf.data.AUTOTUNE)

同时检查数据集是否存在类别不平衡,若有差异可使用class_weight参数加权训练。

5. 增强模型正则化

除现有Dropout外,可添加:

  • L2正则化:在Conv2D和Dense层添加kernel_regularizer=tf.keras.regularizers.L2(0.001)
  • BatchNormalization:在Conv2D层后添加tf.keras.layers.BatchNormalization(),稳定训练过程

6. 错误分析针对性优化

训练完成后,生成混淆矩阵分析易混淆类别,针对性补充样本或调整模型:

import numpy as np
from sklearn.metrics import confusion_matrix
import seaborn as sns

# 提取验证集数据与标签
val_images = []
val_labels = []
for x, y in val_ds:
    val_images.append(x)
    val_labels.append(y)
val_images = np.concatenate(val_images)
val_labels = np.concatenate(val_labels)

# 预测并生成混淆矩阵
predictions = model.predict(val_images)
pred_labels = np.argmax(predictions, axis=1)
cm = confusion_matrix(val_labels, pred_labels)

# 可视化混淆矩阵
plt.figure(figsize=(12,12))
sns.heatmap(cm, annot=True, fmt='d', xticklabels=train_ds.class_names, yticklabels=train_ds.class_names)
plt.xlabel('Predicted')
plt.ylabel('True')
plt.show()

内容的提问来源于stack exchange,提问作者Amadeusz

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 02:50:21