TensorFlow使用自定义数据集训练GAN时格式报错解决方法
问题根因
你能跑通fashion_mnist版本,是因为这套代码最终输出的数据集完全匹配GAN训练的输入要求:
- 单张图片形状为
(28, 28, 1),单通道灰度 - 像素值归一化到
[0,1]区间,数据类型为float32 - 最终封装为批次大小64的
tf.data.Dataset对象,单个批次形状为(64, 28, 28, 1)
自有数据集处理到第16行dataset = np.reshape(dataset, (-1, 28, 28, 1))报错,基本就两个原因:
- 读入的自有图片总元素数和
样本数*28*28*1不匹配:比如用了RGB三通道图没转灰度、原图尺寸不是28*28、读入时误把图片拉成一维没做尺寸校验 - 数据集里混了形状不一致的图片:比如部分是3232、部分是2828,reshape时维度对不上
自有数据集处理正确写法
别直接硬套reshape逻辑,按下面步骤写,逐环节保证格式和fashion_mnist对齐:
- 先写单图预处理函数,强制对齐尺寸、通道、数值范围
def process_img(img_path): img = tf.io.read_file(img_path) # 是png格式就把decode_jpeg换成decode_png,channels=1强制转单通道灰度 img = tf.image.decode_jpeg(img, channels=1) # 统一缩放到28*28,不管原图是什么尺寸 img = tf.image.resize(img, [28, 28]) # 转float32+归一化到0-1,和fashion_mnist处理逻辑一致 img = tf.cast(img, tf.float32) / 255.0 return img
- 批量加载构造数据集,后续逻辑和你原来的代码完全兼容
# 替换成你自己的数据集目录,默认目录下全是训练图片 img_folder = "./your_own_dataset" img_paths = [os.path.join(img_folder, f) for f in os.listdir(img_folder) if f.lower().endswith(('.jpg','.jpeg','.png'))] BATCH_SIZE = 64 dataset = tf.data.Dataset.from_tensor_slices(img_paths) # 多线程并行预处理,速度比用numpy循环读快很多 dataset = dataset.map(process_img, num_parallel_calls=tf.data.AUTOTUNE) # 下面的shuffle、batch逻辑和你原来写的完全一样,不用改 dataset = dataset.shuffle(buffer_size=1024).batch(BATCH_SIZE)
格式校验方法
处理完跑下面的代码抽一个批次检查,输出符合要求就不会再报格式错:
for test_batch in dataset.take(1): print(test_batch.shape) # 正确输出:(64, 28, 28, 1) print(test_batch.dtype) # 正确输出:<dtype: 'float32'> print(tf.reduce_max(test_batch).numpy(), tf.reduce_min(test_batch).numpy()) # 正确值在0~1之间
如果你自己改了GAN的输入结构,不是用28*28单通道输入,只要把上面代码里的resize尺寸、channels参数改成和模型输入层匹配的值就行,核心是保证数据集输出的张量形状和模型输入要求完全对应。
内容的提问来源于stack exchange,提问作者somethingidk
相关产品推荐
相关产品推荐

