为何ImageDataGenerator生成器会出现数据耗尽的情况?
嘿,这个问题我之前踩过坑!其实问题出在flow_from_directory返回的DirectoryIterator的默认行为上——虽然它理论上可以生成无限的增强数据,但在单个epoch内,它默认只会遍历原始数据集一次,之后就会抛出StopIteration异常,这就是触发报错的核心原因。
具体原因拆解
你单独遍历生成器的时候能无限循环,是因为DirectoryIterator的__next__方法在遍历完一轮原始数据后,会自动重置迭代器(打乱数据,重新生成增强样本)。但model.fit在处理每个epoch时,会严格按照你设置的steps_per_epoch去抓取批次:如果在这个过程中生成器抛出了StopIteration(也就是当前epoch内的原始数据已经遍历完了),TensorFlow就会判定为“数据耗尽”,抛出你看到的报错。
举个例子:你的原始数据是1000张图,batch_size=10,一轮最多能生成100个批次。当你设置steps_per_epoch=101时,在第101步,生成器已经把当前epoch的100个批次都输出完了,就会触发StopIteration,进而导致报错。
解决方案:让生成器无限重复生成
要解决这个问题,只需要让生成器在遍历完一轮后不停止,而是无限重复即可,最简单的方式就是给flow_from_directory的结果加上.repeat():
datagenerator = ImageDataGenerator( rescale=1./255, shear_range=0.1, zoom_range=0.1, # 其他增强参数... ) imageFlow = datagenerator.flow_from_directory( image_dir_with_1000_pcs, target_size=(150, 150), batch_size=10, class_mode='binary' ).repeat() # 加上这一行,让生成器无限重复
之后再运行model.fit,设置任意大的steps_per_epoch都不会再触发数据耗尽的报错了——生成器会在每轮数据遍历完成后自动重置,继续生成新的增强批次。
补充说明
如果你不想让生成器无限重复,也可以手动计算合理的steps_per_epoch(比如等于total_samples // batch_size),这样每个epoch刚好遍历一轮原始数据,也不会报错。但如果你的需求是在单个epoch内多次使用原始数据进行增强训练,.repeat()就是最直接的解决方案。
内容的提问来源于stack exchange,提问作者flipcc

