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

如何从给定的tf.keras图像数据集生成代码中获取x_test和y_test

从tf.data.Dataset提取测试集数据与标签

tf.keras.preprocessing.image_dataset_from_directory返回的是tf.data.Dataset对象,要提取出x_test(图像数据)和y_test(标签数据),可以通过以下方法将其转换为numpy数组:

方法1:循环遍历批次拼接

import numpy as np

# 初始化空数组,通道数3对应RGB图像,灰度图改为1
x_test = np.empty((0, img_height, img_width, 3))
# 列数对应类别数量,由test_generator.class_names长度决定
y_test = np.empty((0, len(test_generator.class_names)))

# 遍历数据集的每个批次,拼接数据
for images, labels in test_generator:
    x_test = np.concatenate((x_test, images.numpy()), axis=0)
    y_test = np.concatenate((y_test, labels.numpy()), axis=0)

方法2:简洁式转换

利用列表推导式直接提取所有批次的数据,再拼接:

import numpy as np

# 提取所有图像和标签批次
image_batches = list(test_generator.map(lambda x, y: x))
label_batches = list(test_generator.map(lambda x, y: y))

# 拼接成完整数组
x_test = np.concatenate([batch.numpy() for batch in image_batches], axis=0)
y_test = np.concatenate([batch.numpy() for batch in label_batches], axis=0)

补充说明

  • 由于你设置了label_mode='categorical',y_test会是独热编码格式的二维数组,行数等于测试集样本总数,列数等于类别数量。
  • img_height和img_width需与创建test_generator时设置的尺寸保持一致。

内容的提问来源于stack exchange,提问作者Monika Roopak

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 16:25:13