You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.03 17:35:30