R版Keras用flow_images_from_directory仍报TensorFlow输入数据耗尽错误
问题根因
你当前使用的R版Keras版本已默认将flow_images_from_directory()输出的生成器包装为TensorFlow Dataset对象,原生生成器的无限循环特性被Dataset的默认单次迭代逻辑覆盖,因此跑完一轮原始训练集(约63个batch)就会触发数据耗尽报错,和官方文档描述的旧版行为不一致。
可行解决方案
方案1:调整步数参数适配数据集大小(最简单兼容方案)
直接把steps_per_epoch和validation_steps改为和数据集匹配的数值,不需要修改生成器逻辑:
- 训练集2000样本,batch size 32:
steps_per_epoch = ceiling(2000/32) = 63 - 验证集1000样本,batch size32:
validation_steps = ceiling(1000/32) = 32
修改后的fit代码:
history <- model %>% fit( train_generator, steps_per_epoch = 63, epochs = 100, validation_data = validation_generator, validation_steps = 32 )
该方案下每个epoch会跑完所有原始样本一次,数据增强依然有效,每个样本每次加载都会应用不同的增强变换。
方案2:给数据集添加repeat()实现无限迭代(匹配原书设置的100步/epoch)
如果要严格复用原书的steps_per_epoch=100、validation_steps=50的设置,只需要将生成器转为TF Dataset对象后调用repeat()即可,R版Keras原生支持该操作:
# 给训练和验证生成器添加无限重复逻辑 train_ds <- train_generator %>% repeat() val_ds <- validation_generator %>% repeat() # 用修改后的数据集训练 history <- model %>% fit( train_ds, steps_per_epoch = 100, epochs = 100, validation_data = val_ds, validation_steps = 50 )
该方案完全匹配原书的代码逻辑,每个epoch会跑满100个batch,生成器会持续输出增强后的样本不会中断。
方案3:降级Keras版本到2.2.5(和原书写作时的版本一致)
如果不需要使用新版Keras的特性,可以直接安装和《Deep Learning with R》第一版匹配的Keras版本,原生支持flow_images_from_directory()的无限循环逻辑,不需要修改任何代码即可运行。
内容的提问来源于stack exchange,提问作者Nickopotamus
相关产品推荐
相关产品推荐

