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

如何校验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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 18:09:01