如何从图片文件夹创建Prefetch数据集适配TensorFlow CycleGAN训练
问题原因
你当前用列表存储图片路径/图像数据的方式和TensorFlow官方要求的tf.data.Dataset类型不匹配,所以调用cache()、map()这类Dataset专属方法时会报错。同时你写的grab_path函数存在逻辑错误:每次循环都取文件夹下的第一个文件路径,最终返回的列表全是同一张图片的路径。
最简解决方案
直接用TensorFlow内置API构建符合要求的Prefetch Dataset,不需要手动用OpenCV读图片,完全适配官方CycleGAN代码的输入要求:
步骤1:正确收集两类图片的路径
import os import tensorflow as tf # 数据集路径 ART_DIR = "/content/abstract-art-gallery/Abstract_gallery/Abstract_gallery/" LAND_DIR = "/content/landscape-pictures/" def get_all_img_paths(folder): res = [] for filename in os.listdir(folder): if filename.lower().endswith(('.jpg', '.png', '.jpeg')): res.append(os.path.join(folder, filename)) return res art_paths = get_all_img_paths(ART_DIR) land_paths = get_all_img_paths(LAND_DIR)
步骤2:构建tf.data.Dataset格式数据集
直接复用官方CycleGAN教程里的preprocess_image_train和preprocess_image_test预处理函数,不用改原有逻辑:
# 定义加载图片的工具函数 def load_img(path): img = tf.io.read_file(path) img = tf.image.decode_jpeg(img, channels=3) return tf.cast(img, tf.float32) # 构建艺术图训练集 train_art = tf.data.Dataset.from_tensor_slices(art_paths) train_art = train_art.map(load_img, num_parallel_calls=tf.data.AUTOTUNE) # 构建风景画训练集 train_land = tf.data.Dataset.from_tensor_slices(land_paths) train_land = train_land.map(load_img, num_parallel_calls=tf.data.AUTOTUNE) # 对接官方原有预处理逻辑,最终生成的就是符合要求的Prefetch Dataset BUFFER_SIZE = 1000 BATCH_SIZE = 1 AUTOTUNE = tf.data.AUTOTUNE train_art = train_art.cache().map( preprocess_image_train, num_parallel_calls=AUTOTUNE).shuffle( BUFFER_SIZE).batch(BATCH_SIZE).prefetch(AUTOTUNE) train_land = train_land.cache().map( preprocess_image_train, num_parallel_calls=AUTOTUNE).shuffle( BUFFER_SIZE).batch(BATCH_SIZE).prefetch(AUTOTUNE)
之后把官方代码里的train_horses、train_zebras替换成上面的train_art、train_land即可直接运行。
内容的提问来源于stack exchange,提问作者Rithwik Babu
相关产品推荐
相关产品推荐

