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

自建数据集训练的图像识别模型对未知图像预测结果一致问题求助

图像识别模型预测结果固定问题排查与解决

项目概况

  • 自建图像数据集:共563个文件,分为2个类别
  • 数据集划分:训练集395张,验证集168张
  • 训练采用的CNN模型代码:
# creating datagen object
base_dir = '/content/dataset/emotional_people'
train_data_gen = ImageDataGenerator(rescale = 1/255, validation_split=0.3)
val_data_gen = ImageDataGenerator(rescale = 1/255, validation_split = 0.3)

train_data = train_data_gen.flow_from_directory(base_dir,
                                                target_size = (256,256),
                                                class_mode = 'binary',
                                                subset = 'training',
                                                batch_size = 32)

val_data = val_data_gen.flow_from_directory(base_dir,
                                            target_size = (256,256),
                                            class_mode = 'binary',
                                            subset = 'validation',
                                            batch_size = 32)


# Data augmentation
data_augmentation = keras.Sequential([
    layers.RandomFlip("horizontal", input_shape=(256,256,3)),
    layers.RandomFlip("vertical", input_shape=(256,256,3)),
    layers.RandomRotation(0.1),
    layers.RandomZoom(0.1)])

model = Sequential()
model.add(Conv2D(16, (3,3), 1, activation='relu', input_shape=(256,256,3)))
model.add(MaxPooling2D())
model.add(Conv2D(32, (3,3), 1, activation='relu'))
model.add(MaxPooling2D())
model.add(Conv2D(16, (3,3), 1, activation='relu'))
model.add(MaxPooling2D())
model.add(Flatten())
model.add(Dense(128, activation='relu'))
model.add(Dense(1, activation='sigmoid'))

model.compile('adam', loss=tf.losses.BinaryCrossentropy(), metrics=['accuracy'])

hist = model.fit(train_data, epochs=50, validation_data=val_data)
  • 训练结果:训练精度0.9796,验证精度0.7143,但对未见过的图像预测结果始终固定;更换模型架构、调整验证集比例、添加数据增强后问题仍存在

可能的原因分析

  1. 预测预处理不匹配
    训练时对图像做了rescale=1/255归一化,若预测时未执行相同操作,或图像尺寸、通道顺序(RGB/BGR)与训练时不一致,会导致模型输入数据分布偏差极大,输出固定值。
  2. 数据集类别不平衡
    若两个类别样本占比差异过大(如某类占比90%以上),模型会倾向于输出占比高的类别,尤其是泛化能力不足时,对新数据的预测会固化到多数类。
  3. 模型严重过拟合
    训练精度远高于验证精度,说明模型学到的是训练集的噪声而非通用特征,面对新数据无法有效提取特征,只能输出固定结果。另外代码中定义的data_augmentation序列未加入模型,训练时实际未启用数据增强,加剧过拟合。
  4. 预测代码逻辑错误
    比如重复使用同一预处理后的数据、模型加载错误、输出结果处理逻辑异常(如始终取第一个结果)。

解决方法

1. 统一预测与训练的预处理逻辑

确保预测时的图像处理步骤和训练完全一致:

import cv2
import numpy as np

def preprocess_image(img_path):
    # 读取图像(OpenCV默认BGR,需转为RGB匹配ImageDataGenerator的读取格式)
    img = cv2.imread(img_path)
    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
    # 调整尺寸至训练时的target_size
    img = cv2.resize(img, (256, 256))
    # 归一化
    img = img / 255.0
    # 增加batch维度(模型默认接受批量输入)
    img = np.expand_dims(img, axis=0)
    return img

# 执行预测
processed_img = preprocess_image("test_image.jpg")
prediction = model.predict(processed_img)

2. 解决类别不平衡问题

  • 先统计两个类别的样本数量,若差异显著:
    • 采用过采样(复制少数类样本)或欠采样(减少多数类样本)平衡数据集
    • 训练时给少数类设置权重,通过class_weight参数实现:
      # 示例:假设类别0有100样本,类别1有463样本,权重按样本占比倒数设置
      class_weight = {0: 4.63, 1: 1.0}
      hist = model.fit(train_data, epochs=50, validation_data=val_data, class_weight=class_weight)
      

3. 缓解模型过拟合

  • 启用数据增强:将定义好的data_augmentation加入模型开头:
    model = Sequential()
    model.add(data_augmentation)  # 加入数据增强层
    model.add(Conv2D(16, (3,3), 1, activation='relu', input_shape=(256,256,3)))
    # 后续层保持不变
    
  • 添加Dropout层:在全连接层前加入Dropout抑制过拟合:
    model.add(Flatten())
    model.add(Dropout(0.5))  # 随机丢弃50%神经元
    model.add(Dense(128, activation='relu'))
    model.add(Dense(1, activation='sigmoid'))
    
  • 降低模型复杂度:减少卷积层滤波器数量或全连接层神经元数量,避免模型容量过剩
  • 加入正则化:在卷积层或全连接层添加L2正则化:
    from tensorflow.keras import regularizers
    model.add(Conv2D(16, (3,3), 1, activation='relu', input_shape=(256,256,3), kernel_regularizer=regularizers.l2(0.01)))
    

4. 检查预测代码逻辑

  • 确认每次预测都重新读取并预处理图像,而非复用同一变量
  • 验证模型加载正确性:训练完成后用model.save()保存,预测时重新加载
  • 打印预测输入的图像数据,确认维度(需为(1,256,256,3))和数值范围(0-1之间)是否正确

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 21:35:03