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

多输出蘑菇图像分类模型训练报错求助(MobileNetV2+TFLite)

蘑菇多输出分类模型维度不匹配问题排查与修正

核心错误原因

报错logits_size=[20,20] labels_size=[20,2]本质是标签维度与模型输出不匹配,根源在以下三点:

  1. 数据集目录结构与flow_from_directory逻辑冲突:你当前的目录是train/Edible和train/Non-Edible两级,flow_from_directory只会识别最顶层的2个类别,因此生成的标签y是形状为(20,2)的one-hot向量(对应可食用/不可食用),而非你预期的包含20个种类的标签。
  2. 自定义生成器标签拆分逻辑错误:你错误地认为y包含22列(2+20),但实际只有2列,导致mushroom_labels = y[:,0:20]取到的仍是(20,2)的向量,与模型mushroom_class输出的20类维度不匹配。
  3. 损失函数与输出层不匹配:edibility_output用softmax输出2类,却搭配binary_crossentropy,虽能运行但不符合二分类常规配置。

解决方案与修正代码

第一步:调整数据集目录结构

将数据集改为按蘑菇种类划分一级目录,同时维护一个可食用性映射表:

train/
├── Agaricus_bisporus/  # 可食用
├── Amanita_phalloides/ # 不可食用
├── ...  # 剩余18个种类文件夹
test/
├── Agaricus_bisporus/
├── Amanita_phalloides/
├── ...

第二步:修正模型与编译代码

建议二分类输出用sigmoid更高效,同时保持种类分类的softmax:

import tensorflow as tf
from tensorflow.keras.applications import MobileNetV2
from tensorflow.keras.layers import Dense, GlobalAveragePooling2D
from tensorflow.keras.models import Model
from tensorflow.keras.optimizers import Adam

# 模型搭建
IMG_SHAPE = (224, 224, 3)
base_model = MobileNetV2(weights='imagenet', include_top=False, input_shape=IMG_SHAPE)
base_model.trainable = False

x = base_model.output
x = GlobalAveragePooling2D()(x)
x = Dense(512, activation='relu')(x)

# 20分类输出(蘑菇种类)
mushroom_class = Dense(20, activation='softmax', name='mushroom_class')(x)
# 二分类输出(可食用性,用sigmoid更合理)
edibility_output = Dense(1, activation='sigmoid', name='edibility_output')(x)

model = Model(inputs=base_model.input, outputs=[edibility_output, mushroom_class])
model.summary()

# 模型编译:二分类用binary_crossentropy,多分类用categorical_crossentropy
model.compile(optimizer=Adam(),
              loss={'edibility_output': 'binary_crossentropy',
                    'mushroom_class': 'categorical_crossentropy'},
              metrics={'mushroom_class': 'accuracy',
                       'edibility_output': 'accuracy'})

第三步:修正数据生成器逻辑

添加可食用性映射,从20类标签中推导可食用性标签:

from tensorflow.keras.preprocessing.image import ImageDataGenerator
import os

# 定义每个种类的可食用性:0=不可食用,1=可食用
edibility_map = {
    'Agaricus_bisporus': 1,
    'Amanita_phalloides': 0,
    # 补充剩余18个种类的映射关系
}

# 数据生成器配置
train_dir = 'mushroom3/MO_95/mushroom_dataset/train'
validation_dir = 'mushroom3/MO_95/mushroom_dataset/test'

train_datagen = ImageDataGenerator(
    rescale=1./255,
    rotation_range=40,
    width_shift_range=0.2,
    height_shift_range=0.2,
    shear_range=0.2,
    zoom_range=0.2,
    horizontal_flip=True,
    fill_mode='nearest'
)

validation_datagen = ImageDataGenerator(rescale=1./255)

# 生成20类的one-hot标签
train_generator = train_datagen.flow_from_directory(
    train_dir,
    target_size=(224, 224),
    batch_size=20,
    class_mode='categorical',
    shuffle=True
)

validation_generator = validation_datagen.flow_from_directory(
    validation_dir,
    target_size=(224, 224),
    batch_size=20,
    class_mode='categorical',
    shuffle=True
)

# 获取类别名称列表
class_names = list(train_generator.class_indices.keys())

# 自定义生成器:从20类标签生成可食用性标签
def custom_generator(generator):
    while True:
        x, y_20class = generator.next()
        # 从one-hot标签中获取类别索引
        class_idx_batch = tf.argmax(y_20class, axis=1).numpy()
        # 生成可食用性标签(数值型,对应sigmoid输出)
        edible_labels = tf.convert_to_tensor(
            [edibility_map[class_names[idx]] for idx in class_idx_batch],
            dtype=tf.float32
        )
        yield x, {'edibility_output': edible_labels, 'mushroom_class': y_20class}

train_generator_custom = custom_generator(train_generator)
validation_generator_custom = custom_generator(validation_generator)

第四步:模型训练(无需修改)

history = model.fit(
    train_generator_custom,
    steps_per_epoch=len(train_generator),
    epochs=200,
    validation_data=validation_generator_custom,
    validation_steps=len(validation_generator)
)

第五步:导出TFLite模型(适配移动端)

# 保存Keras模型
model.save('mushroom_classifier.h5')

# 转换为TFLite并优化
converter = tf.lite.TFLiteConverter.from_keras_model(model)
# 开启默认优化
converter.optimizations = [tf.lite.Optimize.DEFAULT]
# 启用FP16量化(减小模型体积,适配移动端)
converter.target_spec.supported_types = [tf.float16]

tflite_model = converter.convert()

# 保存TFLite模型
with open('mushroom_classifier.tflite', 'wb') as f:
    f.write(tflite_model)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 17:44:55