标签为文件路径时能否使用flow_from_dataframe训练图像重构网络
完全可以实现,用你当前选用的flow_from_dataframe方法即可完成,不需要更换其他生成器方案。
实现步骤
- 构造基础图像生成器
调用flow_from_dataframe时仅指定输入图像路径,关闭默认标签生成:
train_generator = datagen.flow_from_dataframe( dataframe=df, x_col="filepath", class_mode=None, # 禁用默认标签输出 target_size=(img_height, img_width), # 替换为你的图像实际尺寸 batch_size=batch_size, # 替换为你的批次大小 color_mode="rgb" # 灰度图可修改为"grayscale" )
- 包装生成器适配VAE重构需求
写一个简单的迭代器包装函数,将同一份批次数据同时作为输入和标签输出:
def vae_data_generator(generator): for batch_img in generator: yield batch_img, batch_img
- 调用fit方法训练
注意需要手动指定steps_per_epoch参数,告知Keras每轮训练的迭代步数:
vae.fit( vae_data_generator(train_generator), epochs=10, steps_per_epoch=len(train_generator) )
注意事项
如果你在ImageDataGenerator中配置了随机数据增强(随机翻转、裁剪、旋转等),上述方案会更稳妥:它直接将同一份增强后的图像同时作为输入和标签,不会出现输入、标签分别做增强导致的不匹配问题。
如果直接设置x_col和y_col均为filepath来实现,部分Keras版本会对输入、标签分别应用随机增强,导致训练失效。
内容的提问来源于stack exchange,提问作者P M
相关产品推荐
相关产品推荐

