Keras传入yield生成器到model.fit报cannot pickle 'generator'错如何解决
问题根因
你遇到的报错核心原因是:Python 原生生成器实例无法被 pickle 序列化,而你在 model.fit 中开启了 use_multiprocessing=True 配置,Windows 系统下 Python 多进程默认采用 spawn 模式,需要把主进程中的生成器对象序列化后传递给子进程,触发了序列化失败的报错。
两种解决方案
方案1:改用 Keras 官方支持多进程的 Sequence 结构(推荐)
tf.keras.utils.Sequence是 Keras 专门为数据加载设计的基类,天然支持序列化和多进程加载,还支持 epoch 结束自动回调打乱数据等扩展能力,修改后代码如下:import numpy as np import tensorflow as tf class DataGenerator(tf.keras.utils.Sequence): def __init__(self, x_df, y_df, batch_size): self.x_df = x_df self.y_df = y_df self.batch_size = batch_size def __len__(self): # 返回每轮epoch的总步数 return int(np.ceil(len(self.x_df) / self.batch_size)) def __getitem__(self, idx): # 按批次返回数据 batch_x = self.x_df[idx * self.batch_size : (idx+1) * self.batch_size] batch_y = self.y_df[idx * self.batch_size : (idx+1) * self.batch_size] return np.array(batch_x), np.array(batch_y) # 实例化后调用fit,原多进程配置可以保留 train_generator = DataGenerator(x_train_df, y_train_df, batch_size) history = model.fit(train_generator, epochs=epochs_no, steps_per_epoch=number_of_rows_input/batch_size, verbose=1, max_queue_size=100, validation_data=None, workers=8, use_multiprocessing=True )方案2:关闭多进程加载
如果你不想修改原有生成器代码,直接关闭多进程开关即可,线程模式下不需要对数据对象做序列化,不会触发报错:
history = model.fit(train_generator, epochs=epochs_no, steps_per_epoch=number_of_rows_input/batch_size, verbose=1, max_queue_size=100, validation_data=None, workers=8, use_multiprocessing=False # 仅修改这一行配置 )
补充说明
如果使用 Sequence 方案,还可以实现 on_epoch_end 方法,在每轮训练结束后自动打乱数据集,提升模型泛化能力:
def on_epoch_end(self): # 打乱数据集索引,示例代码需要你提前把索引存到类属性中使用 np.random.shuffle(self.indexes)
内容的提问来源于stack exchange,提问作者Liviu Barbu
相关产品推荐
相关产品推荐

