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

TensorFlow加载CycleGAN自定义图像数据集维度报错解决方案

错误诱因

报错核心是加载返回的数据集结构和CycleGAN教程的预处理逻辑不兼容:

  • 官方教程中tfds.load加载的未分批CycleGAN数据集,单条样本为3维结构(图像高度, 图像宽度, 通道数),和后续random_crop等预处理操作的输入要求完全匹配。
  • 原有加载代码存在两处不兼容问题:
    • 设置batch_size=1时,接口返回带批次维度的4维图像张量(批次大小, 高度, 宽度, 通道数),而random_crop的裁剪参数只适配3维单图输入,维度数不匹配直接触发报错。
    • 未指定label_mode=None,接口默认返回(图像, 类别标签)的二元组,和教程中纯图像张量的输入格式不兼容。
  • 额外注意:image_dataset_from_directory默认要求传入的根目录下必须存在子文件夹作为类别区分,即使只加载单类图像也不能直接把存满图片的文件夹作为根目录传入,否则会触发类别识别错误。
适配CycleGAN训练的本地数据集加载方法

以下两种方案加载后的数据集格式和官方tfds加载结果完全一致,可以直接套入教程的训练逻辑无需额外修改。

方案1:调整后使用image_dataset_from_directory加载

首先调整本地目录结构,每个域的图像文件夹下新建一个子文件夹存放所有对应图片,例如:

dataset/
├── train_horses/
│   └── horse/  # 所有马的训练图存在这个子文件夹里
│       ├── 001.jpg
│       ├── 002.jpg
│       └── ...
├── train_zebras/
│   └── zebra/  # 所有斑马的训练图存在这个子文件夹里
│       ├── 001.jpg
│       ├── 002.jpg
│       └── ...
├── test_horses/
│   └── horse/
└── test_zebras/
    └── zebra/

加载代码如下:

from tensorflow.keras.preprocessing import image_dataset_from_directory
import tensorflow as tf

# 加载训练集,先不分批、不返回标签
train_horses = image_dataset_from_directory(
    './dataset/train_horses',
    color_mode='rgb',
    image_size=(286, 286),
    batch_size=None,  # 保持单样本3维结构,对齐tfds输出
    label_mode=None,  # 不返回标签,只输出图像张量
    shuffle=True,
    seed=123
)
train_zebras = image_dataset_from_directory(
    './dataset/train_zebras',
    color_mode='rgb',
    image_size=(286, 286),
    batch_size=None,
    label_mode=None,
    shuffle=True,
    seed=123
)

# 加载测试集
test_horses = image_dataset_from_directory(
    './dataset/test_horses',
    color_mode='rgb',
    image_size=(286, 286),
    batch_size=None,
    label_mode=None,
    shuffle=False
)
test_zebras = image_dataset_from_directory(
    './dataset/test_zebras',
    color_mode='rgb',
    image_size=(286, 286),
    batch_size=None,
    label_mode=None,
    shuffle=False
)

# 和官方教程一致的预处理逻辑
def normalize_crop_train(image):
    image = tf.cast(image, tf.float32) / 127.5 - 1  # 像素值缩放到[-1,1]
    image = tf.image.random_crop(image, size=[256, 256, 3])
    image = tf.image.random_flip_left_right(image)
    return image

def normalize_test(image):
    image = tf.cast(image, tf.float32) / 127.5 - 1
    image = tf.image.resize(image, (256, 256))
    return image

# 映射预处理
train_horses = train_horses.map(normalize_crop_train, num_parallel_calls=tf.data.AUTOTUNE)
train_zebras = train_zebras.map(normalize_crop_train, num_parallel_calls=tf.data.AUTOTUNE)
test_horses = test_horses.map(normalize_test, num_parallel_calls=tf.data.AUTOTUNE)
test_zebras = test_zebras.map(normalize_test, num_parallel_calls=tf.data.AUTOTUNE)

# 组装训练集、分批、预加载,和教程逻辑完全对齐
train_dataset = tf.data.Dataset.zip((train_horses, train_zebras))
train_dataset = train_dataset.batch(1).prefetch(tf.data.AUTOTUNE)
test_dataset = tf.data.Dataset.zip((test_horses, test_zebras))
test_dataset = test_dataset.batch(1).prefetch(tf.data.AUTOTUNE)

方案2:用tf.data原生接口加载(无需调整目录结构)

如果不想改动现有文件夹结构,可以直接读取文件路径构建数据集,灵活度更高:

import tensorflow as tf
import glob

# 读取对应文件夹下所有图片路径
train_horse_paths = glob.glob('./horses/*.jpg') + glob.glob('./horses/*.png')
train_zebra_paths = glob.glob('./zebras/*.jpg') + glob.glob('./zebras/*.png')
test_horse_paths = glob.glob('./test_horses/*.jpg') + glob.glob('./test_horses/*.png')
test_zebra_paths = glob.glob('./test_zebras/*.jpg') + glob.glob('./test_zebras/*.png')

def read_and_resize(path):
    image = tf.io.read_file(path)
    image = tf.io.decode_image(image, channels=3, expand_animations=False)
    image = tf.image.resize(image, (286, 286))
    return image

# 构建数据集
train_horses = tf.data.Dataset.from_tensor_slices(train_horse_paths)\
    .map(read_and_resize, num_parallel_calls=tf.data.AUTOTUNE)\
    .shuffle(1000, seed=123)
train_zebras = tf.data.Dataset.from_tensor_slices(train_zebra_paths)\
    .map(read_and_resize, num_parallel_calls=tf.data.AUTOTUNE)\
    .shuffle(1000, seed=123)
test_horses = tf.data.Dataset.from_tensor_slices(test_horse_paths)\
    .map(read_and_resize, num_parallel_calls=tf.data.AUTOTUNE)
test_zebras = tf.data.Dataset.from_tensor_slices(test_zebra_paths)\
    .map(read_and_resize, num_parallel_calls=tf.data.AUTOTUNE)

# 后续预处理、组装数据集的逻辑和方案1完全一致

内容的提问来源于stack exchange,提问作者Stephan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 15:24:19