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

如何遍历TensorFlow测试数据集 展示图像并输出对应预测结果

实现方法

首先纠正一个常见误用:model.fit()是模型训练接口,会根据输入数据更新模型权重,不能用来做预测,预测请用model.predict()或者直接调用模型实例。
你的测试集是通过tf.keras.utils.image_dataset_from_directory生成的,本身按(图像批次张量, 真实标签批次张量)的格式返回数据,直接遍历即可实现需求,完整流程如下:


前置准备

首先确保测试集构建逻辑和训练时完全对齐,避免预处理不一致导致预测结果错误:

  • 测试集的image_size参数必须和训练时输入模型的图像尺寸完全一致
  • 如果训练时做了像素归一化(比如将0-255像素值缩放到0-1区间)、通道转换等预处理,测试集也要做完全相同的处理
  • 建议测试集构建时设置shuffle=False,固定样本遍历顺序,方便核对图像和预测结果的对应关系
  • 可以直接从数据集对象中拿到类别名称列表,不需要手动编写:class_names = test_dataset.class_names,这个列表的顺序和模型输出的类别索引一一对应。

完整可运行代码

import tensorflow as tf
import matplotlib.pyplot as plt

# 1. 构建测试集(和你之前的逻辑对齐,此处仅作示例)
test_dataset = tf.keras.utils.image_dataset_from_directory(
    directory="./car_test", # 替换成你的测试集根目录
    image_size=(224, 224), # 和训练时输入尺寸保持一致
    batch_size=32,
    shuffle=False, # 关闭样本乱序
)
class_names = test_dataset.class_names # 得到类别列表,如['sedan', 'suv', 'truck'...]

# 2. 遍历数据集、展示图像、输出预测结果
plt.figure(figsize=(14, 8)) # 设置展示画布大小

# 遍历数据集,.take(1)表示只取1个批次展示,要遍历全量测试集删掉.take(1)即可
for batch_idx, (batch_images, batch_real_labels) in enumerate(test_dataset.take(1)):
    # 对当前批次做预测,得到每个样本对应各类别的置信度
    batch_pred_probs = model.predict(batch_images, verbose=0)
    # 取置信度最高的索引作为预测类别
    batch_pred_labels = tf.argmax(batch_pred_probs, axis=1).numpy()

    # 逐个展示批次内的图像
    for img_idx in range(len(batch_images)):
        # 按4行8列排布子图(总共32个位置,和batch_size=32对应,可自行调整行列数)
        ax = plt.subplot(4, 8, img_idx + 1)
        # 显示图像,若你的图像已经归一化到0-1区间,可去掉.astype("uint8")
        plt.imshow(batch_images[img_idx].numpy().astype("uint8"))
        # 子图标题显示预测结果,要同时显示真实标签可以改成下面的写法
        # plt.title(f"预测:{class_names[batch_pred_labels[img_idx]]}\n真实:{class_names[batch_real_labels[img_idx]]}")
        plt.title(f"预测:{class_names[batch_pred_labels[img_idx]]}", fontsize=9)
        plt.axis("off") # 隐藏坐标轴

plt.tight_layout() # 自动调整子图间距
plt.show()

# 如果不需要可视化图像,只想在控制台打印映射关系,用下面的逻辑替换画图部分即可
# for img_idx, (pred_label_idx, real_label_idx) in enumerate(zip(batch_pred_labels, batch_real_labels)):
#     print(f"批次{batch_idx} 样本{img_idx} | 预测结果:{class_names[pred_label_idx]} | 真实类别:{class_names[real_label_idx]}")

注意事项

  • 不要在预测阶段调用model.fit(),否则会用测试集数据更新模型权重,破坏你已经训练好的模型
  • 如果测试集批次大小设置过大,一次性展示上百张图像会导致画布拥挤、文字重叠,建议单批次大小设为16-32,逐批次查看即可
  • 若预测结果和预期偏差极大,优先检查预处理逻辑是否和训练阶段对齐,90%以上的这类问题都是输入数据格式不匹配导致的。

内容的提问来源于stack exchange,提问作者Aman Vishnoi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 21:42:15