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

预加载Keras模型的predict方法无法并行运行,有哪些可行方案?

首先修正你代码中MyClass的语法问题,load和inference方法都缺少第一个self参数,修正后定义如下:

class MyClass():
    def __init__(self):
        self.model = None

    def load(self, path):
        self.model = tf.keras.models.load_model(path)

    def inference(self, data):
        #...
        pred = self.model.predict(data)
        #...
        return pred
报错核心原因

你遇到的cannot pickle 'weakref' object错误,本质是Keras模型内部的弱引用、计算图等对象无法被pickle序列化,而joblib的loky多进程后端需要把主进程的myobj实例序列化后传递给子进程,因此触发报错。threading后端因为不需要跨进程传递序列化对象,所以能运行,但受Python GIL全局解释器锁限制,CPU密集的预测任务无法真正并行,所以速度和串行差不多甚至更慢。

可行解决方案
  • 方案1:子进程内单独加载模型
    不要在主进程预加载模型再传给子进程,而是把模型加载逻辑放到每个子进程的执行函数里,每个进程单独加载一份模型,避开序列化问题,能真正实现多进程并行。如果显存/内存充足,这是改造成本最低的方案:

    from joblib import Parallel, delayed
    
    def single_inference(data, model_path):
        myobj = MyClass()
        myobj.load(model_path)
        return myobj.inference(data)
    
    n_jobs = 8
    results = Parallel(n_jobs=n_jobs)(delayed(single_inference)(d, <Path_to_model>) for d in mydata)
    

    如果显存不足,可以适当降低n_jobs数量,避免显存溢出。

  • 方案2:用Keras原生批预测代替手动多进程
    Keras的predict方法本身已经做了底层的并行优化,手动开多进程调用predict反而会因为TensorFlow内部线程池抢占变慢。你可以直接把所有输入数据打包成批次调用predict,效率远高于手动多进程实现:

    import numpy as np
    # 拼接所有输入数据
    all_data = np.concatenate(mydata, axis=0)
    # 调整batch_size适配你的显存大小
    all_pred = myobj.model.predict(all_data, batch_size=32)
    # 按原始输入的长度拆分预测结果
    split_idx = np.cumsum([len(d) for d in mydata[:-1]])
    results = np.split(all_pred, split_idx)
    
  • 方案3:本地部署TensorFlow Serving做批量预测
    把模型部署为本地TF Serving服务,主进程通过gRPC/HTTP请求调用预测接口,TF Serving内部会自动做请求批处理、GPU并行计算,适合大批量数据预测场景,性能远高于手动多进程调用。

  • 方案4:用Python原生multiprocessing的spawn模式
    如果你不想用joblib,也可以用Python原生多进程模块,设置启动方式为spawn,同样在子进程内加载模型,也可以避开序列化问题。

其他并行库说明

所有Python多进程并行库本质都绕不开跨进程对象序列化的限制,只要需要把主进程预加载的Keras模型传递给子进程,都会遇到相同的pickle错误,不需要额外寻找其他并行库,优先调整模型加载逻辑即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 16:09:03