Keras中model.evaluate与.predict结果不一致的故障排查及解决
我此前已查阅过同类问题的相关解答,但尝试所有解决方案后均未解决我的问题。
问题描述
我搭建卷积神经网络(CNN)执行常规图像分类任务,模型编译代码如下:
model.compile(optimizer = keras.optimizers.Adam(learning_rate = exp_learning_rate), loss = tf.keras.losses.SparseCategoricalCrossentropy(), metrics = ['accuracy'])
我在训练数据集上拟合模型,并在验证集上评估,代码如下:
history = model.fit(train_dataset, validation_data = validation_dataset, epochs = 5)
随后我在独立测试集上评估模型:
model.evaluate(test_dataset)
得到的输出为:
4/4 [==============================] - 30s 7s/step - loss: 1.7180 - accuracy: 0.8627
但当我运行如下代码执行预测时:
model.predict(test_dataset)
输出的混淆矩阵显示实际准确率仅为35.39%,和.evaluate返回的86%准确率差距极大。为排除测试集问题,我分别在训练集和验证集上执行预测,得到的准确率也仅为30%左右,但训练过程中记录的训练准确率、验证准确率分别高达96%、87%。
疑问
我不清楚为什么.evaluate和.predict会输出不同结果,看起来调用.predict时没有使用训练得到的权重(3分类场景下该预测效果和随机猜测无异)。我的损失函数配置正确:我按照TensorFlow要求对数据做了标签编码以适配SparseCategoricalCrossentropy,传入的accuracy指标也会自动匹配损失函数计算逻辑,理论上结果应该一致。为什么两者会出现这么大的偏差?我应该信任哪个结果?
已尝试的解决方法
我曾怀疑SparseCategoricalCrossentropy配置有误,因此将目标标签做独热编码后改用CategoricalCrossentropy损失,但问题完全没有得到解决。
顾虑
如果.evaluate的结果不准确,是不是意味着训练过程中输出的训练准确率、验证准确率也不可靠?训练过程的评估逻辑不也是调用.evaluate吗?如果是这样我应该信任什么指标?损失无法直接反映模型效果,众所周知损失最小不代表准确率更高,在准确率指标不可靠的情况下我要如何评估模型的实际效果?我现在没有办法判断模型是否在学习,非常希望有人能帮我理清问题原因。
2021年10月28日 12:26 AM 更新
我补充更多代码以方便排查问题。我最初的数据预处理流程如下:
image_size = (256, 256) batch_size = 16 train_ds = keras.preprocessing.image_dataset_from_directory( directory = image_directory, label_mode = 'categorical', shuffle = True, validation_split = 0.2, subset = 'training', seed = 24, batch_size = batch_size ) val_ds = keras.preprocessing.image_dataset_from_directory( directory = image_directory, label_mode = 'categorical', shuffle = True, validation_split = 0.2, subset = 'validation', seed = 24, batch_size = batch_size )
其中image_directory是存储图片的路径字符串。根据官方文档,image_dataset_from_directory方法会返回tf.data.Dataset对象,包含对应训练、验证数据的多个批次。
我引入VGG16架构做分类,因此调用VGG16对应的预处理函数做数据转换:
preprocess_input = tf.keras.applications.vgg16.preprocess_input train_ds = train_ds.map(lambda x, y: (preprocess_input(x), y)) val_ds = val_ds.map(lambda x, y: (preprocess_input(x), y))
将图片转换为VGG16适配的输入格式后,我对验证集做如下拆分得到验证集和测试集:
val_batches = tf.data.experimental.cardinality(val_ds) test_dataset = val_ds.take(val_batches // 3) validation_dataset = val_ds.skip(val_batches // 3)
随后我对数据做缓存和预取处理:
AUTOTUNE = tf.data.AUTOTUNE train_dataset = train_ds.prefetch(buffer_size = AUTOTUNE) validation_dataset = validation_dataset.prefetch(buffer_size = AUTOTUNE) test_dataset = test_dataset.prefetch(buffer_size = AUTOTUNE)
问题定位
我目前还不能完全确认.evaluate是否能真实反映模型准确率,但我发现当我使用keras.Sequential()模型时,.evaluate和.predict的结果始终一致。我怀疑从keras.applications API导入的VGG16模型不属于keras.Sequential()模型,因此直接传入上述流程处理的数据时,.predict和.evaluate的结果无法对齐(我暂时没有足够的研究验证这个猜想,如果有人了解相关逻辑欢迎补充)。
最终我改用ImageDataGenerator()替换image_dataset_from_directory()解决了该问题,代码如下:
train_datagen = ImageDataGenerator( preprocessing_function = preprocess_input, width_shift_range = 0.2, height_shift_range = 0.2, shear_range = 0.2, zoom_range = 0.2, horizontal_flip = True ) val_datagen = ImageDataGenerator( preprocessing_function = preprocess_input ) train_ds = train_datagen.flow_from_directory( train_image_directory, target_size = (224, 224), batch_size = 16, seed = 24, shuffle = True, classes = ['class1', 'class2', 'class3'], class_mode = 'categorical' ) test_ds = val_datagen.flow_from_directory( test_image_directory, target_size = (224, 224), batch_size = 16, seed = 24, shuffle = False, classes = ['class1', 'class2', 'class3'], class_mode = 'categorical' )
这套流程完成所有预处理后,调用model.evaluate(test_ds)返回的结果和model.predict_generator(test_ds)的结果完全一致。对预测结果做简单处理后,我用如下代码生成混淆矩阵:
Y_pred = model.predict(test_ds) y_pred = np.argmax(Y_pred, axis=1) cf = confusion_matrix(test_ds.classes, y_pred) sns.heatmap(cf, annot= True, xticklabels = class_names, yticklabels = class_names) plt.title('Performance of Model on Testing Set')
此时混淆矩阵计算得到的准确率和model.evaluate(test_ds)的结果完全匹配,偏差问题消失。
总结
如果你在训练图像分类模型时,损失和准确率指标匹配,但预测结果和评估指标存在偏差,可尝试调整数据预处理流程。我通常在keras.sequential()模型上使用image_dataset_from_directory()方法加载数据,但对于非Sequential结构的VGG16模型,使用ImageDataGenerator(...).flow_from_directory(...)加载数据可让预测结果和评估指标保持一致。
简而言之
我没有完全定位到问题的根本原因,但找到了可行的解决方法,希望我的踩坑经验能帮到遇到同类问题的开发者。
内容的提问来源于stack exchange,提问作者AndrewJaeyoung

