You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.01 01:15:04