如何校验Keras模型在image_dataset_from_directory生成测试集上的预测结果
Keras 测试集预测结果与真实标签匹配校验方法
前置修正
首先需要修改测试集生成配置,关闭shuffle,保证预测结果和真实标签的顺序严格对应:
ds_test = tf.keras.preprocessing.image_dataset_from_directory( 'test/', labels = 'inferred', label_mode = 'categorical', color_mode = 'rgb', batch_size = batch_size, image_size = (img_height, img_width), shuffle = False, # 核心修改,避免打乱样本顺序 seed = 123, )
方法1:遍历批次收集标签(推荐,内存占用低)
该方法不需要一次性加载所有测试数据到内存,适合大数据集场景:
import numpy as np # 获取类别与索引的映射关系 class_names = ds_test.class_names # 收集所有真实标签的索引 true_labels = [] for images, labels in ds_test: # categorical格式标签转成类别索引,和预测结果格式对齐 true_labels.extend(np.argmax(labels.numpy(), axis=1)) # 生成预测结果,顺序与真实标签完全对应 predictions = model.predict(ds_test) pred_labels = np.argmax(predictions, axis=1)
校验用示例代码:
# 计算测试集准确率 test_acc = np.sum(np.array(true_labels) == pred_labels) / len(true_labels) print(f"测试集准确率:{test_acc:.4f}") # 转换为实际类别名称查看 true_class = [class_names[i] for i in true_labels] pred_class = [class_names[i] for i in pred_labels]
方法2:一次性获取全量测试数据
如果数据集较小,可以直接拆分出全量的样本和标签:
# 拆分批次,逐个读取所有样本 all_data = list(ds_test.unbatch().as_numpy_iterator()) x_test = np.array([item[0] for item in all_data]) y_test = np.array([np.argmax(item[1]) for item in all_data]) # 直接预测全量数据 pred_labels = np.argmax(model.predict(x_test), axis=1)
常见错误说明
ds_test.class_names是存储类别映射的列表,不是可调用方法,不需要传入参数image_dataset_from_directory返回的是tf.data.Dataset批次迭代对象,不是包含全量数据的元组,无法直接用(x_test, y_test)赋值- 预测操作不要放在遍历数据集的循环内,每次
model.predict调用都会遍历整个测试集,不仅冗余还会导致标签和预测结果顺序不匹配 - 测试集开启
shuffle=True时,样本顺序会随机打乱,标签和预测结果无法一一对应
内容的提问来源于stack exchange,提问作者Nardz59
相关产品推荐
相关产品推荐

