将图像数据集转TFRecord时遇AttributeError:_NumpyIterator无shard属性
解决AttributeError: '_NumpyIterator'对象没有'shard'属性的问题
错误原因
问题出在循环逻辑里:第一次迭代时,你把encode_ds重新赋值为encode_ds.shard(...).as_numpy_iterator(),此时encode_ds已经变成了_NumpyIterator类型的迭代器,而非原来的tf.data.Dataset对象。而shard()是tf.data.Dataset专属的方法,迭代器没有这个属性,所以第二次循环调用.shard()时就会报错。
修复方案
不要在循环中修改原始的encode_ds数据集对象,每次循环都基于原始的encode_ds来执行shard操作,再转换为迭代器。
修复后的完整代码
ds_train = tf.keras.utils.image_dataset_from_directory(some parameters) ds_train = ( ds_train .unbatch() ) def encode_image(image, label): image_converted = tf.image.convert_image_dtype(image, dtype=tf.uint8) image = tf.io.encode_jpeg(image_converted) label = tf.argmax(label) return image, label encode_ds = ( ds_train.map(encode_image) ) NUM_SHARD=10 PATH = "some path" for shard_no in range(NUM_SHARD): # 每次基于原始数据集执行shard,不修改原始对象 shard_ds = ( encode_ds .shard(NUM_SHARD, shard_no) .as_numpy_iterator() ) with tf.io.TFRecordWriter(PATH.format(shard_no)) as file_writer: for image, label in shard_ds: file_writer.write(create_example(image, label))
内容的提问来源于stack exchange,提问作者ZKS
相关产品推荐
相关产品推荐

