不调用yield()实现Keras fit_generator方法的生成器是否合理?
嘿,这个问题问得特别实在,我来帮你把这事儿掰扯明白~
Keras fit_generator 与生成器的疑问解答
为啥所有示例都用 yield?
首先得明确,fit_generator()(在TensorFlow 2.x里其实已经被整合进model.fit()了,但核心逻辑没变)的核心需求是:一个能按需返回批次数据的可迭代对象。
用yield写生成器函数,是Python里实现这种“按需生成”最顺手的方式——它不用一次性把所有数据塞进内存,每次迭代时才生成一批数据,用完就释放空间,完美贴合大数据集的场景。示例都用它,主要因为:
- 语法极简,不用手动写
__iter__和__next__这些迭代器方法 - 天然支持“暂停-恢复”,每
yield一次就返回一批数据,下次迭代接着上次的进度来 - 内存效率拉满,完全匹配
fit_generator的设计初衷
不用 yield 的“生成器”合理吗?
这里得先澄清:如果你的“生成器”没用到yield,那它其实不是Python标准的生成器函数,而是自定义迭代器(比如实现了__iter__和__next__方法的类),或者返回可迭代对象的普通函数。
只要你的自定义迭代器满足以下要求,那用来适配fit_generator()完全合理:
- 每次迭代返回
(x_batch, y_batch)的元组(如果需要样本权重,也可以返回三元组) - 能循环生成数据,支持多轮epoch训练(比如迭代完所有数据后重置索引)
- 不会一次性加载全量数据,而是按需加载当前批次
给你举个自定义迭代器的例子:
class BatchDataIterator: def __init__(self, data_file_paths, batch_size): self.file_paths = data_file_paths self.batch_size = batch_size self.current_idx = 0 def __iter__(self): return self def __next__(self): # 迭代完一轮后重置索引,支持多epoch if self.current_idx >= len(self.file_paths): self.current_idx = 0 raise StopIteration # 只加载当前批次的数据 batch_files = self.file_paths[self.current_idx:self.current_idx+self.batch_size] x_batch = self._load_inputs(batch_files) # 自定义加载输入的函数 y_batch = self._load_labels(batch_files) # 自定义加载标签的函数 self.current_idx += self.batch_size return x_batch, y_batch def _load_inputs(self, paths): # 这里写你的数据加载逻辑,比如读取图片、预处理等 pass def _load_labels(self, paths): # 加载对应标签的逻辑 pass
然后直接把这个迭代器传给fit_generator就行:
train_iter = BatchDataIterator(train_file_list, batch_size=32) model.fit_generator( generator=train_iter, steps_per_epoch=len(train_file_list)//32, epochs=10 )
这种方式完全合规,yield只是实现迭代器的便捷途径,绝非唯一方式。只要你的可迭代对象符合Keras的要求,就能正常跑起来。
小提醒
如果用的是TensorFlow 2.x的新版本,官方已经标记fit_generator()为废弃了,推荐直接用model.fit(),传入tf.data.Dataset对象或者自定义迭代器就行——逻辑和之前一致,但API更统一。要是还在维护旧代码,那用fit_generator也没问题。
内容的提问来源于stack exchange,提问作者Jeremy Friesner
相关产品推荐
相关产品推荐

