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

Keras EfficientNet模型全样本预测为同一类别问题排查

图像分类模型全预测为sad类问题排查求助
  • 基于Udemy教程搭建图像分类模型,使用PyCharm作为开发环境(未用Jupyter)
  • 测试过多个版本的EfficientNet模型,无论训练1、10还是20个epoch,所有测试图像的预测结果均为sad类别
  • 验证集表现:
    • angry类别:0/100正确(准确率0.00%)
    • happy类别:0/100正确(准确率0.00%)
    • sad类别:100/100正确(准确率100.00%)
    • 整体验证准确率:33.33%
  • 已尝试调整模型结构、添加数据增强操作,但问题未得到解决
  • 以下是模型构建、训练及验证的完整代码,求帮忙排查问题根源:
# 模型构建、训练及验证完整代码
import tensorflow as tf
from tensorflow.keras.applications import EfficientNetB0
from tensorflow.keras.layers import Dense, GlobalAveragePooling2D
from tensorflow.keras.models import Model
from tensorflow.keras.preprocessing.image import ImageDataGenerator
from sklearn.metrics import classification_report

# 1. 模型定义
base_model = EfficientNetB0(weights='imagenet', include_top=False, input_shape=(224, 224, 3))
x = base_model.output
x = GlobalAveragePooling2D()(x)
predictions = Dense(3, activation='softmax')(x)
model = Model(inputs=base_model.input, outputs=predictions)

# 2. 数据加载与预处理
train_datagen = ImageDataGenerator(rescale=1./255, rotation_range=20, width_shift_range=0.2, height_shift_range=0.2, horizontal_flip=True)
val_datagen = ImageDataGenerator(rescale=1./255)

train_generator = train_datagen.flow_from_directory(
    './train',
    target_size=(224, 224),
    batch_size=32,
    class_mode='categorical'
)

val_generator = val_datagen.flow_from_directory(
    './val',
    target_size=(224, 224),
    batch_size=32,
    class_mode='categorical'
)

test_generator = val_datagen.flow_from_directory(
    './test',
    target_size=(224, 224),
    batch_size=1,
    class_mode='categorical',
    shuffle=False
)

# 3. 模型编译与训练
model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
model.fit(train_generator, epochs=10, validation_data=val_generator)

# 4. 验证与结果输出
predictions = model.predict(test_generator)
predicted_classes = tf.argmax(predictions, axis=1)
true_classes = test_generator.classes
class_labels = list(test_generator.class_indices.keys())

print(classification_report(true_classes, predicted_classes, target_names=class_labels))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 14:25:01