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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 13:54:04