TensorFlow使用自定义生成器训练出现数据耗尽报错如何修复
问题根源
- 自定义生成器逻辑冲突:外层虽然写了
while True想要无限循环输出数据,但末尾加了if batch_index * batch_size > sample_count: break判断,遍历完一轮样本后就直接终止生成器,导致第二个epoch没有新的batch输入 model.fit中传入了无效的batch_size参数:自定义生成器返回的已经是按batch组装好的数据,fit的batch_size参数不会生效,反而可能造成计数逻辑冲突steps_per_epoch计算值和实际生成器能提供的batch数不匹配
解决方案
方案1:直接修复自定义生成器逻辑
修改生成器代码,移除终止逻辑,保证可以无限生成batch:
import os import numpy as np import cv2 def generator(idir,odir,batch_size,shuffle ): i_list=os.listdir(idir) o_list=os.listdir(odir) sample_count=len(i_list) while True: input_image_batch=[] output_image_batch=[] # 每次生成一个batch直接取对应数量的样本即可 for _ in range(batch_size): if shuffle: # 修正原逻辑少取最后一个样本的问题,randint的high为开区间 m=np.random.randint(low=0, high=sample_count, dtype=int) else: # 如果不需要打乱,维护索引指针遍历完自动重置 if not hasattr(generator, 'idx'): generator.idx = 0 m = generator.idx generator.idx = (generator.idx + 1) % sample_count path_to_in_img=os.path.join(idir,i_list[m]) path_to_out_img=os.path.join(odir,o_list[m]) input_image=cv2.imread(path_to_in_img) input_image=cv2.resize(input_image,(3200,3200)) output_image=cv2.imread(path_to_out_img) output_image=cv2.resize(output_image,(3200,3200)) input_image_batch.append(input_image) output_image_batch.append(output_image) input_val1image_array=np.array(input_image_batch) / 255.0 output_val2image_array=np.array(output_image_batch) / 255.0 yield (input_val1image_array, output_val2image_array)
修改model.fit调用,删除无效的batch_size参数:
idir = r"D:\\image\\" odir=r"D:\\image1\\" batch_size = 4 train = generator(idir,odir,batch_size,True) model.compile(optimizer="adam", loss='mean_squared_error', metrics=['mean_squared_error']) # steps_per_epoch按实际样本数计算,如有560个样本就填560//batch_size model.fit(train, validation_data = (valin_images,valout_images), epochs = 20, steps_per_epoch = 560//batch_size)
方案2:转为TF Dataset使用内置repeat方法
如果需要使用TensorFlow官方的repeat()、prefetch()等性能优化方法,可将自定义生成器包装为标准数据集:
import tensorflow as tf # 定义生成器输出的张量形状和类型 output_signature = ( tf.TensorSpec(shape=(None, 3200, 3200, 3), dtype=tf.float32), tf.TensorSpec(shape=(None, 3200, 3200, 3), dtype=tf.float32) ) # 包装自定义生成器 train_ds = tf.data.Dataset.from_generator( lambda: generator(idir,odir,4,True), output_signature=output_signature ) # 调用repeat()实现无限重复,加prefetch优化训练性能 train_ds = train_ds.repeat().prefetch(tf.data.AUTOTUNE) # 训练时传入包装好的数据集即可 model.fit(train_ds, validation_data = (valin_images,valout_images), epochs = 20, steps_per_epoch = 560//4)
内容的提问来源于stack exchange,提问作者user123
相关产品推荐
相关产品推荐

