如何为Keras CycleGAN示例将自定义数据集加载到tfds中
自定义数据集适配Keras CycleGAN示例方案
你当前的目录结构完全符合tfds.folder_dataset.ImageFolder的加载规则,无需调整目录结构,可按以下步骤获取对应域的数据集并适配官方示例的输入要求:
1. 拆分获取trainA、trainB、testA、testB数据集
你调用as_dataset()后返回的是嵌套字典结构,第一层key是split名(train、test),第二层是对应split下的样本,每个样本包含image和label两个字段:
label=0对应每个split下第一个子文件夹的样本(train下为trainA、test下为testA)label=1对应每个split下第二个子文件夹的样本(train下为trainB、test下为testB)
直接用以下代码提取四个子集:
import tensorflow_datasets as tfds import tensorflow as tf # 加载数据集,shape参数会自动把所有图片resize到指定尺寸 data = tfds.folder_dataset.ImageFolder('Images', shape=(256, 256, 3)) ds = data.as_dataset() # 提取训练集两个域的数据 train_ds_a = ds['train'].filter(lambda x: x['label'] == 0).map(lambda x: x['image']) train_ds_b = ds['train'].filter(lambda x: x['label'] == 1).map(lambda x: x['image']) # 提取测试集两个域的数据 test_ds_a = ds['test'].filter(lambda x: x['label'] == 0).map(lambda x: x['image']) test_ds_b = ds['test'].filter(lambda x: x['label'] == 1).map(lambda x: x['image'])
2. 适配CycleGAN示例的输入流水线
官方示例后续的预处理、批处理逻辑可以直接复用,只需要把示例里的预置数据集替换为上面提取的四个数据集即可,对应代码如下:
# 像素值归一化逻辑,和官方示例保持一致 def normalize_img(img): img = tf.cast(img, tf.float32) return (img / 127.5) - 1.0 # 训练集预处理:随机水平翻转+归一化 def preprocess_train_img(img): img = tf.image.random_flip_left_right(img) img = normalize_img(img) return img # 测试集预处理:仅归一化 def preprocess_test_img(img): return normalize_img(img) # 构建训练、测试流水线 BATCH_SIZE = 1 # 可根据显存大小调整 AUTOTUNE = tf.data.AUTOTUNE train_ds_a = train_ds_a.map(preprocess_train_img, num_parallel_calls=AUTOTUNE).cache().shuffle(1000).batch(BATCH_SIZE) train_ds_b = train_ds_b.map(preprocess_train_img, num_parallel_calls=AUTOTUNE).cache().shuffle(1000).batch(BATCH_SIZE) test_ds_a = test_ds_a.map(preprocess_test_img, num_parallel_calls=AUTOTUNE).cache().batch(BATCH_SIZE) test_ds_b = test_ds_b.map(preprocess_test_img, num_parallel_calls=AUTOTUNE).cache().batch(BATCH_SIZE)
处理完成后,四个数据集的结构和官方示例加载的数据集完全一致,后续模型训练代码不需要做任何修改即可直接运行。
内容的提问来源于stack exchange,提问作者Saksham
相关产品推荐
相关产品推荐

