预加载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

