如何从TensorFlow BatchDataSet中获取固定的前N个样本而非可变迭代器
解决TensorFlow BatchDataSet获取固定样本的问题
你遇到的核心问题是:ds.take(9)返回的是惰性计算的数据集对象,而非实际存储的样本数据。每次遍历这个对象时,TensorFlow会重新执行数据读取流程,导致每次得到的样本不同。要获取固定的N个样本,需要把数据从数据集中提取出来,转为内存中存储的numpy数组。
解决方案代码
import tensorflow as tf import numpy as np import matplotlib.pyplot as plt # 加载数据集 ds = tf.keras.utils.image_dataset_from_directory( "Images", validation_split=0.2, seed=123, subset="training") class_names = ds.class_names # 提取前9个固定样本 test_ds = ds.take(9) images_list = [] labels_list = [] # 遍历take到的数据集,提取所有样本和标签 for imgs, lbls in test_ds: images_list.append(imgs.numpy()) labels_list.append(lbls.numpy()) # 合并为单个numpy数组,并截取前9个样本 fixed_images = np.concatenate(images_list, axis=0)[:9] fixed_labels = np.concatenate(labels_list, axis=0)[:9] # 第一次绘制样本 plt.figure(figsize=(4, 4)) for i in range(9): ax = plt.subplot(3, 3, i + 1) plt.imshow(fixed_images[i].astype("uint8")) plt.title(class_names[fixed_labels[i]]) plt.axis("off") # 第二次绘制,会和第一次完全一致 plt.figure(figsize=(4, 4)) for i in range(9): ax = plt.subplot(3, 3, i + 1) plt.imshow(fixed_images[i].astype("uint8")) plt.title(class_names[fixed_labels[i]]) plt.axis("off") plt.show()
关键说明
- 惰性计算特性:TensorFlow的数据集是按需生成数据的,
take()方法只是定义了一个数据读取的操作逻辑,并没有把数据实际存储下来。每次遍历都会重新执行这个逻辑,导致样本变化。 - 转为numpy数组:通过
numpy()方法将张量转为numpy数组,数据会被存储到内存中,后续使用时直接读取内存中的数据,不会再从原数据集重新读取,因此样本固定不变。 - 处理批量数据:如果你的数据集设置了
batch_size(默认值为32),take(9)会获取9个批量的数据,而非9个单样本。因此需要用np.concatenate合并所有批量数据后,再截取你需要的前N个样本。
另一种提取方式(逐样本获取)
如果你想精确控制获取的样本数量,也可以用迭代器逐样本提取:
# 使用numpy迭代器逐样本获取 iterator = ds.as_numpy_iterator() fixed_images = [] fixed_labels = [] while len(fixed_images) < 9: imgs_batch, lbls_batch = iterator.next() # 遍历当前批量的样本 for img, lbl in zip(imgs_batch, lbls_batch): if len(fixed_images) >= 9: break fixed_images.append(img) fixed_labels.append(lbl) # 转为numpy数组 fixed_images = np.array(fixed_images) fixed_labels = np.array(fixed_labels)
内容的提问来源于stack exchange,提问作者Tumbleweed53
相关产品推荐
相关产品推荐

