Keras人-马分类CNN模型预测始终输出相同结果的问题排查
解决Keras CNN模型全预测为人类的问题
针对你遇到的模型始终输出1或接近1、将所有图片判定为人类的问题,可按以下步骤排查解决:
1. 确认类别映射与数据集一致性
flow_from_directory会按文件夹字母顺序自动分配标签(比如horses文件夹在humans前,马会被标记为0,人类为1)。先打印类别映射确认对应关系:
print(train_generator.class_indices)
- 如果测试图片中马被错误判定为人类,先检查训练集的文件夹结构是否正确,以及测试图片是否符合训练集的类别定义。
2. 检查训练数据集的平衡性与完整性
从训练日志的Found 256 images belonging to 2 classes可以看出,训练集规模远小于官方人-马数据集(官方训练集共1027张):
- 先确认
training_dir路径下的horses和humans子文件夹是否包含足够数量的图片,且两类数量大致均衡。如果某类图片数量远多于另一类,模型会偏向预测数量多的类别。 - 若存在类别不平衡,在训练时添加类权重:
import os # 假设0是马,1是人类,根据实际class_indices调整 num_horses = len(os.listdir(os.path.join(training_dir, 'horses'))) num_humans = len(os.listdir(os.path.join(training_dir, 'humans'))) class_weight = {0: num_humans/num_horses, 1: 1.0} model.fit(train_generator, epochs=15, validation_data=validation_generator, class_weight=class_weight)
3. 验证预测时的图片预处理是否正确
确保测试图片的预处理与训练集完全一致:
- 检查图片通道数:
image.load_img默认加载RGB图片,打印img.shape确认是(300,300,3),避免灰度图(单通道)导致模型输出异常。 - 确认归一化步骤:训练和预测都做了
img /= 255.,这点你已经正确实现,但要避免额外的预处理操作(比如颜色通道反转)。
4. 解决模型过拟合与训练不稳定问题
从训练日志看,训练准确率达到96%但验证准确率波动极大(最低52%,最高83%),说明模型泛化能力差:
- 添加Dropout层抑制过拟合:
model = tf.keras.models.Sequential([ # 原卷积层和池化层保持不变 tf.keras.layers.Conv2D(16 , (3,3), activation=tf.nn.relu , input_shape = (300,300,3)), tf.keras.layers.MaxPooling2D(2,2), tf.keras.layers.Conv2D(32 , (3,3), activation=tf.nn.relu), tf.keras.layers.MaxPooling2D(2,2), tf.keras.layers.Conv2D(64 , (3,3), activation=tf.nn.relu), tf.keras.layers.MaxPooling2D(2,2), tf.keras.layers.Conv2D(64 , (3,3), activation=tf.nn.relu), tf.keras.layers.MaxPooling2D(2,2), tf.keras.layers.Conv2D(64 , (3,3), activation=tf.nn.relu), tf.keras.layers.MaxPooling2D(2,2), tf.keras.layers.Flatten(), tf.keras.layers.Dropout(0.5), # 添加此层 tf.keras.layers.Dense(512 ,activation=tf.nn.relu ), tf.keras.layers.Dense(1, activation = tf.nn.sigmoid) ]) - 降低学习率:将RMSprop的
learning_rate从0.001改为0.0001,让训练更稳定。 - 简化模型:比如减少卷积层的滤波器数量(如把64改为32),或缩小Dense层的神经元数(如512改为256)。
5. 验证模型加载与基础预测能力
加载模型后,先用验证集的已知图片测试,排除模型加载错误:
val_imgs, val_labels = next(validation_generator) # 取第一张验证图片 pred = model.predict(val_imgs[:1]) print(f"真实标签: {val_labels[0]}, 预测值: {pred[0]}")
- 如果这里预测正确,但你手动选择的图片预测错误,说明测试图片本身存在问题(比如图片内容不符合训练集的类别特征,或格式异常)。
内容的提问来源于stack exchange,提问作者Ahmad Badpey
相关产品推荐
相关产品推荐

