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

如何将自定义图像数据生成器传入Keras的model.fit()完成全量训练

问题解决方法

1. 报错直接原因及修复

你遇到的image_gen object is not callable错误,是因为你错误对image_gen类的实例gen加了()调用:gen是类实例而非可调用函数,本身没有实现__call__方法,不能直接加括号执行。
你需要将get_images()方法返回的迭代器传入model.fit,同时配置steps_per_epoch参数(指定每个epoch跑多少个batch,生成器模式下必填),不需要传batch_size参数(你生成器内部已经按固定batch生成数据了)。
正确训练代码如下:

batch_size=5
idir=r'D:\\train'
odir=r'D:\\Train\\train'
shuffle=True
gen=image_gen(idir,odir,batch_size,shuffle=True)
# 计算每个epoch需要的step数:总样本数//batch_size
steps_per_epoch = len(gen.i_list) // batch_size
# 传入生成器迭代器训练
model.fit(
    gen.get_images(),
    steps_per_epoch=steps_per_epoch,
    validation_data = (valin_images,valout_images),
    epochs = 100
)

2. 现有生成器的问题修复建议

你的get_images方法存在逻辑缺陷,会导致训练只能跑1个epoch就停止,且采样逻辑可能出现样本重复/遗漏:

  • 你同时写了while True和遍历完所有样本就break的逻辑,跑完1轮所有样本后迭代器就终止了,多epoch训练会报错
  • 随机采样逻辑是每次在batch内随机抽索引,可能出现同一个epoch内同一样本被多次抽到、部分样本从未被抽到的问题
  • os.listdir返回的文件名顺序不固定,可能出现输入输出图像配对错误
    建议修改后的生成器逻辑如下:
import os
import numpy as np
import cv2
class image_gen():
    def __init__(self, idir,odir,batch_size, shuffle = True):          
        self.idir=idir
        self.odir=odir
        self.batch_size=batch_size   
        self.shuffle=shuffle
        # 先排序输入输出文件名,保证配对正确
        self.i_list = sorted(os.listdir(self.idir))
        self.o_list = sorted(os.listdir(self.odir))
        self.sample_count = len(self.i_list)
        
    def get_images(self): 
        while True: # 无限循环适配多epoch训练
            # 每个epoch开始前打乱样本顺序
            if self.shuffle:
                indices = np.random.permutation(self.sample_count)
            else:
                indices = np.arange(self.sample_count)
            # 按batch遍历所有样本
            for start in range(0, self.sample_count, self.batch_size):
                end = min(start + self.batch_size, self.sample_count)
                batch_indices = indices[start:end]
                input_image_batch=[]
                output_image_batch=[]
                for idx in batch_indices:
                    path_to_in_img=os.path.join(self.idir,self.i_list[idx])
                    path_to_out_img=os.path.join(self.odir,self.o_list[idx])
                    input_image=cv2.imread(path_to_in_img)
                    input_image=cv2.resize(input_image,(3200,3200))/255.0
                    output_image=cv2.imread(path_to_out_img)
                    output_image=cv2.resize(output_image,(3200,3200))/255.0
                    input_image_batch.append(input_image)
                    output_image_batch.append(output_image)
                yield np.array(input_image_batch), np.array(output_image_batch)

内容的提问来源于stack exchange,提问作者user123

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 23:15:03