不使用Sequence API实现自定义数据生成器是否需要添加while循环?
Keras自定义yield生成器是否需要while循环解答
核心结论
你当前未继承Sequence API的生成器实现,必须添加无限while循环才能正常用于Keras训练流程。
原因说明
- 你现有的
get_images()是单次生成器:调用后仅会执行到yield语句返回1个batch的数据,之后迭代就会终止,Keras训练需要生成器持续输出batch完成多epoch训练,没有循环的话跑完1个batch就会停止,训练直接报错中断。 - 你看到的Sequence实现不需要循环,是因为Sequence类本身实现了Keras要求的可迭代协议,框架会自动按epoch重复调用
__getitem__拉取数据,不需要开发者自己维持循环逻辑。
代码修改方案
仅需要在get_images()方法最外层加无限while循环即可,同时修正原代码导入别名不匹配的问题,其余逻辑保持不变:
import os import numpy as np import cv2 class image_gen(): def __init__(self, sdir, batch_size, shuffle): self.batch_index=0 self.sdir=sdir self.batch_size=batch_size self.shuffle=shuffle self.image_list=[] self.s_list=os.listdir(self.sdir) def get_images(self): while True: # 新增无限循环 input_image_batch=[] sample_count=len(self.s_list) for i in range(self.batch_index * self.batch_size, (self.batch_index + 1) * self.batch_size ): j=i % sample_count if self.shuffle: m=np.random.randint(low=0, high=sample_count-1, size=None, dtype=int) else: m=j path_to_img=os.path.join(self.sdir, self.s_list[m]) input_image=cv2.imread(path_to_img) input_image_batch.append(input_image) input_image_array=np.array(input_image_batch) self.batch_index=self.batch_index + 1 yield input_image_array
使用注意事项
- 调用生成器训练时,需要手动给
fit()方法传入steps_per_epoch参数,值为总样本数除以batch_size向下取整,否则Keras无法判断一个epoch什么时候结束。 - 如果你设置shuffle=True,当前的随机采样逻辑可能导致同个epoch内重复采样同一张图片,建议每遍历完一轮所有样本后统一打乱
s_list列表,比每次随机采样的逻辑更合理。 - 如果用到验证生成器,同样需要手动指定
validation_steps参数。
内容的提问来源于stack exchange,提问作者user123
相关产品推荐
相关产品推荐

