如何从ImageDataGenerator.flow_from_directory创建tf.data.Dataset?解决output_types参数疑问
从DirectoryIterator创建tf.data.Dataset的方法
刚好我之前处理过类似的需求,给你一步步拆解怎么用tf.data.Dataset.from_generator实现这个转换,重点解决output_types的设置问题:
1. 先准备好DirectoryIterator实例
首先还是按常规方式创建你的ImageDataGenerator和对应的DirectoryIterator:
import tensorflow as tf from tensorflow.keras.preprocessing.image import ImageDataGenerator # 初始化数据生成器(可根据需求添加增强参数) datagen = ImageDataGenerator(rescale=1./255, rotation_range=20) # 生成DirectoryIterator train_iterator = datagen.flow_from_directory( "./train_data", target_size=(224, 224), batch_size=32, class_mode="categorical", # 根据任务选:binary/categorical/sparse等 shuffle=True )
2. 确定output_types和output_shapes
你已经知道Iterator返回(x, y)元组,要确定它们的类型其实很简单:先取出一批样本查看 dtype 即可:
# 获取一批样本 sample_x, sample_y = next(train_iterator) # 查看数据类型 print(f"图像数据类型:{sample_x.dtype}") print(f"标签数据类型:{sample_y.dtype}") # 查看形状 print(f"图像形状:{sample_x.shape}") print(f"标签形状:{sample_y.shape}")
比如做了rescale=1./255后,sample_x.dtype会是float32;如果是categorical类模式,sample_y.dtype也是float32;如果是sparse模式,标签会是int32。
有了这个信息,我们就可以直接设置output_types和output_shapes:
output_types是一个元组,对应(x, y)各自的TensorFlow数据类型,比如(tf.float32, tf.float32)output_shapes对应两者的形状,用tf.TensorShape包裹即可,batch_size维度也可以保留原样本的数值
3. 生成tf.data.Dataset
接下来定义一个无限生成器(因为DirectoryIterator本身是无限迭代的),然后传入from_generator:
def data_generator(): while True: x_batch, y_batch = next(train_iterator) yield x_batch, y_batch # 从样本中提取类型和形状 output_types = (tf.dtypes.as_dtype(sample_x.dtype), tf.dtypes.as_dtype(sample_y.dtype)) output_shapes = (tf.TensorShape(sample_x.shape), tf.TensorShape(sample_y.shape)) # 创建Dataset train_dataset = tf.data.Dataset.from_generator( generator=data_generator, output_types=output_types, output_shapes=output_shapes ) # 可选:添加后续优化处理 train_dataset = train_dataset.repeat() # 匹配Iterator的无限迭代特性 train_dataset = train_dataset.prefetch(tf.data.AUTOTUNE) # 加速数据加载
额外小贴士
- 如果不想手动取样本,也可以直接手动指定类型:比如你明确知道图像是
float32,分类标签是float32,直接写output_types=(tf.float32, tf.float32)就行,但取样本验证更稳妥。 - 其实TensorFlow现在有更便捷的API:
tf.keras.utils.image_dataset_from_directory,可以直接从文件夹生成tf.data.Dataset,不需要先过ImageDataGenerator,如果你没有特别依赖ImageDataGenerator的复杂增强功能,可以试试这个方法。
内容的提问来源于stack exchange,提问作者A T
相关产品推荐
相关产品推荐

