如何使用tf.keras.image_dataset_from_directory输出训练自编码器?
解决自编码器训练时的数据集结构问题
问题核心在于你当前的数据集仅包含输入图像,而Keras的fit方法在使用Dataset作为输入时,要求数据集本身具备输入-目标的二元组结构(对应自编码器的输入和标签为同一图像)。直接传入y=train_ds会触发错误,因为当x是Dataset类型时,fit不支持单独的y参数;去掉y参数后,数据集又缺少模型所需的目标数据。
修改方案
通过tf.data.Dataset.map方法,将单元素的图像数据集转换为(输入图像, 目标图像)的二元组结构,其中目标图像与输入图像完全一致:
# 转换训练集结构:每个元素变为(input, input)二元组 train_ds = train_ds.map(lambda x: (x, x)) # 转换验证集结构 validate_ds = validate_ds.map(lambda x: (x, x))
修改后的完整训练代码
train_ds, validate_ds = tf.keras.utils.image_dataset_from_directory( data_dir, labels=None, validation_split=0.1, subset="both", shuffle=True, seed=123, image_size=(img_height, img_width), batch_size=batch_size, crop_to_aspect_ratio=True) # 调整数据集结构以适配自编码器训练 train_ds = train_ds.map(lambda x: (x, x)) validate_ds = validate_ds.map(lambda x: (x, x)) # 训练模型,直接传入处理后的数据集即可 history = autoencoder.fit( train_ds, validation_data=validate_ds, epochs=epochs )
补充细节
- 无需在
fit中重复指定batch_size,image_dataset_from_directory已经为数据集设置了批次大小。 - 如果需要对图像做预处理(比如归一化到[0,1]范围),可以把逻辑整合到
map函数中:def preprocess(x): x = tf.cast(x, tf.float32) / 255.0 return (x, x) train_ds = train_ds.map(preprocess) validate_ds = validate_ds.map(preprocess)
内容的提问来源于stack exchange,提问作者user3170530
相关产品推荐
相关产品推荐

