多输出蘑菇图像分类模型训练报错求助(MobileNetV2+TFLite)
蘑菇多输出分类模型维度不匹配问题排查与修正
核心错误原因
报错logits_size=[20,20] labels_size=[20,2]本质是标签维度与模型输出不匹配,根源在以下三点:
- 数据集目录结构与
flow_from_directory逻辑冲突:你当前的目录是train/Edible和train/Non-Edible两级,flow_from_directory只会识别最顶层的2个类别,因此生成的标签y是形状为(20,2)的one-hot向量(对应可食用/不可食用),而非你预期的包含20个种类的标签。 - 自定义生成器标签拆分逻辑错误:你错误地认为
y包含22列(2+20),但实际只有2列,导致mushroom_labels = y[:,0:20]取到的仍是(20,2)的向量,与模型mushroom_class输出的20类维度不匹配。 - 损失函数与输出层不匹配:
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
相关产品推荐
相关产品推荐

