ResNet50迁移学习二分类模型predict结果异常排查求助
基于ResNet50迁移学习二分类模型的预测异常排查
问题现象
- 测试集
x_test预测时,准确率和混淆矩阵始终固定不变 - 单图预测输出始终为1,图片显示异常(部分全黑、部分有白线条、部分正常),曾触发输入形状不匹配报错:
ValueError: Input 0 of layer "sequential" is incompatible with the layer: expected shape=(None, 224, 224, 3), found shape=(32, 224, 3)
数据处理相关代码
测试集生成器
test_datagen = ImageDataGenerator(rescale=1./255) test_generator = test_datagen.flow_from_directory( test_dir, target_size=(224, 224), batch_size=32, class_mode='binary', shuffle=True,seed=42 )
额外打乱操作
indices = list(range(len(test_generator))) random.shuffle(indices) shuffled_images = [test_generator[index][0] for index in indices] shuffled_labels = [test_generator[index][1] for index in indices] y_true=np.concatenate(shuffled_labels) x_test=np.concatenate(shuffled_images)
单图处理代码
img = cv2.imread(path) img.resize([1,224,224,3]) prediction = model.predict(img)
图片显示代码
img.resize([224,224,3]) plt.imshow(img)
问题排查与修复
1. 测试集准确率/混淆矩阵固定的问题
你额外的打乱操作存在逻辑错误:test_generator[index]会重复生成数据(generator是迭代器,每次访问索引会从头循环生成,并非提取已有批次),导致x_test和y_true中的样本大量重复,甚至完全重复,最终预测结果固定。
修复方式二选一:
# 方式1:直接用generator预测 y_pred = model.predict(test_generator, verbose=1) y_true = test_generator.classes # 方式2:一次性提取完整测试集(适合数据量不大的场景) x_test, y_true = [], [] for _ in range(len(test_generator)): imgs, labels = next(test_generator) x_test.append(imgs) y_true.append(labels) x_test = np.concatenate(x_test) y_true = np.concatenate(y_true) # 提取后再打乱 indices = np.random.permutation(len(x_test)) x_test = x_test[indices] y_true = y_true[indices]
2. 单图预测的问题
- 形状不匹配:
cv2.imread读取的图片形状为(h, w, 3),直接用resize([1,224,224,3])会破坏数据维度顺序(resize按元素总数调整,而非按轴调整),导致形状错误。正确做法:
img = cv2.imread(path) # 先缩放到目标尺寸 img = cv2.resize(img, (224, 224)) # 增加批次维度(对应模型输入的None维度) img = np.expand_dims(img, axis=0) # 同步训练时的归一化操作 img = img / 255.0 prediction = model.predict(img)
- 输出始终为1:大概率是单图未做归一化(训练时用了
rescale=1./255,预测时必须同步处理),导致模型输入分布与训练时不一致;也可能是模型训练不足,或测试集与训练集分布差异过大。 - 图片显示异常:
resize方法直接修改数组形状,会破坏像素空间结构导致混乱。正确显示步骤:
img = cv2.imread(path) img = cv2.resize(img, (224, 224)) # cv2读取为BGR格式,转成RGB后再用plt显示 img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) plt.imshow(img_rgb) plt.show()
3. 额外注意点
- 训练与预测的图像预处理逻辑必须完全一致(包括归一化、颜色空间转换)
- 若使用ResNet50预训练权重,建议用官方预处理方法
tf.keras.applications.resnet50.preprocess_input替代简单的1./255,避免输入分布不匹配导致模型输出异常
内容的提问来源于stack exchange,提问作者Heba
相关产品推荐
相关产品推荐

