如何从TensorFlow的EMNIST PrefetchDataset拆分x_train、y_train等数据?
解决PrefetchDataset无法下标提取图像和标签的问题
嘿,这个问题我遇到过!你之所以碰到TypeError: 'PrefetchDataset' object is not subscriptable,是因为TensorFlow Datasets(TFDS)返回的train_ds和test_ds是tf.data.Dataset的子类(这里是PrefetchDataset),它和Keras内置MNIST直接返回的numpy数组不一样,不能用['image']这种下标方式访问特征。
下面给你两种常用的解决方案,按需选择:
方案1:保持Dataset格式(推荐用于TF流水线)
如果你之后还要用TensorFlow的Dataset API做数据增强、批量处理等操作,推荐用map方法拆分图像和标签,或者加载时直接指定as_supervised=True:
import tensorflow as tf import tensorflow_datasets as tfds # 加载时设置as_supervised=True,直接得到(image, label)元组形式的Dataset train_ds, test_ds = tfds.load( 'emnist', split=['train', 'test'], shuffle_files=True, as_supervised=True # 关键参数!开启后每个元素是(image, label) ) # 如果已经加载了没加as_supervised=True,用map转换: # train_ds = train_ds.map(lambda sample: (sample['image'], sample['label'])) # test_ds = test_ds.map(lambda sample: (sample['image'], sample['label'])) # 之后就可以像这样批量处理了 train_ds = train_ds.batch(32).prefetch(tf.data.AUTOTUNE) test_ds = test_ds.batch(32).prefetch(tf.data.AUTOTUNE)
方案2:转换成numpy数组(和Keras MNIST用法一致)
如果你确实需要像Keras MNIST那样得到独立的numpy数组(比如和Scikit-learn等库配合使用),可以用tfds.as_numpy把Dataset转换成numpy迭代器,再收集成数组:
import tensorflow as tf import tensorflow_datasets as tfds import numpy as np # 先加载数据集,不管有没有as_supervised=True都可以 train_ds, test_ds = tfds.load('emnist', split=['train', 'test'], shuffle_files=True) # 如果没开as_supervised=True,先转成(image, label)元组 train_ds = train_ds.map(lambda x: (x['image'], x['label'])) test_ds = test_ds.map(lambda x: (x['image'], x['label'])) # 转换成numpy数组 x_train, y_train = [], [] for img, lbl in tfds.as_numpy(train_ds): x_train.append(img) y_train.append(lbl) x_train = np.array(x_train) y_train = np.array(y_train) # 同样处理测试集 x_test, y_test = [], [] for img, lbl in tfds.as_numpy(test_ds): x_test.append(img) y_test.append(lbl) x_test = np.array(x_test) y_test = np.array(y_test)
更高效的numpy转换方式
如果数据集规模不大,也可以用tf.concat一次性合并所有元素,再转numpy:
# 合并训练集图像和标签 x_train = tf.concat([img for img, _ in train_ds], axis=0).numpy() y_train = tf.concat([lbl for _, lbl in train_ds], axis=0).numpy() # 合并测试集 x_test = tf.concat([img for img, _ in test_ds], axis=0).numpy() y_test = tf.concat([lbl for _, lbl in test_ds], axis=0).numpy()
内容的提问来源于stack exchange,提问作者vel
相关产品推荐
相关产品推荐

