如何不本地下载TensorFlow数据集并流式传输至模型?遇AssertionError求助
解决TensorFlow Datasets加载Open Images V4时的AssertionError问题
问题场景
尝试不本地下载,直接从TensorFlow Datasets加载Open Images V4数据集并构建批量数据管道,但调用tfds.load()时触发AssertionError。
原代码
import tensorflow as tf import tensorflow_datasets as tfds # tf.__version__ # '2.18.0' (train_data, test_data), info= tfds.load('open_images_v4', split=['train','test'], download=False, try_gcs=True, as_supervised=True, shuffle_files=True, with_info = True, )
报错信息
--------------------------------------------------------------------------- AssertionError Traceback (most recent call last) <ipython-input-9-5483ac7d6ce0> in <cell line: 0>() ----> 1 (train_data, test_data), info= tfds.load('open_images_v4', 2 split=['train','test'], 3 download=False, 4 try_gcs=True, 5 as_supervised=True, 3 frames /usr/local/lib/python3.11/dist-packages/tensorflow_datasets/core/dataset_builder.py in as_dataset(self, split, batch_size, shuffle_files, decoders, read_config, as_supervised) 1002 # pylint: enable=line-too-long 1003 if not self.data_path.exists(): -> 1004 raise AssertionError( 1005 "Dataset %s: could not find data in %s. Please make sure to call " 1006 "dataset_builder.download_and_prepare(), or pass download=True to " AssertionError: Dataset open_images_v4: could not find data in /root/tensorflow_datasets. Please make sure to call dataset_builder.download_and_prepare(), or pass download=True to tfds.load() before trying to access the tf.data.Dataset object.
问题原因
Open Images V4并不支持直接从GCS(Google Cloud Storage)流式加载。try_gcs=True仅对部分预先在GCS存储了预处理后TFRecord文件的数据集有效,而Open Images V4不在此列。设置download=False时,TFDS会在本地路径查找预处理好的数据集文件,找不到就抛出该错误。
解决方案
虽然无法完全跳过本地下载,但可以通过download=True让TFDS自动从GCS下载源数据并完成预处理,之后即可正常加载为tf.data.Dataset对象。如果想减少本地存储占用,可选择加载数据集的子集切片,或配合TF数据API进行流式处理。
修改后的代码:
import tensorflow as tf import tensorflow_datasets as tfds # TF版本2.18.0 (train_data, test_data), info = tfds.load( 'open_images_v4', split=['train', 'test'], download=True, # 允许TFDS自动下载并预处理数据集 try_gcs=True, # 优先从GCS获取源数据 as_supervised=True, shuffle_files=True, with_info=True, )
批量数据管道构建示例
下载完成后,可对数据集进行预处理并构建批量管道:
def preprocess_image(image, label): # 调整图片尺寸至目标大小 image = tf.image.resize(image, (224, 224)) # 归一化像素值到[0,1]区间 image = tf.cast(image, tf.float32) / 255.0 return image, label batch_size = 32 # 训练集管道:预处理+打乱+批量+预取 train_dataset = train_data.map(preprocess_image) train_dataset = train_dataset.shuffle(1000).batch(batch_size).prefetch(tf.data.AUTOTUNE) # 测试集管道:预处理+批量+预取 test_dataset = test_data.map(preprocess_image) test_dataset = test_dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)
内容的提问来源于stack exchange,提问作者Sid
相关产品推荐
相关产品推荐

