如何避免嵌套定义函数时出现pickling错误(多进程场景)
多进程并行化时Pickle本地函数报错的解决方法
问题描述
在机器学习项目中创建了一组带专属参数的控制器函数,需多次运行评估性能,为提速采用多进程并行处理,却遭遇Pickle错误。简化后的代码如下:
import multiprocessing as mp class task(): def parallelizationWrap(self): poolSize = 5 with mp.Pool(poolSize) as pool: for _ in pool.imap(self.parallelizationFunc, range(poolSize)): pass def serialWrap(self): for _ in range(5): self.parallelizationFunc() def setup(self, unusedVar=None): vallist = [1,2,3,4,5] self.funclist = [] for i in range(5): def tempfunc(argument, parameter=vallist[i]): print(parameter*argument) self.funclist.append(tempfunc) def parallelizationFunc(self, unuserVar=None): for step in range(25): for j in range(5): result = self.funclist[j](step) simulation.sendSignalToCorrectAgent(result) if __name__ == "__main__": mp.freeze_support() c1 = task() c1.setup() c1.parallelizationWrap() # c1.serialWrap()
运行后报错:
AttributeError: Can't pickle local object 'task.setup.<locals>.tempfunc'
尝试仅保存参数,但无法满足函数动态变化的灵活性;改为全局函数后仍报错:
_pickle.PicklingError: Can't pickle <function tempfunc at 0x000001861D0A3E20>: it's not the same object as __main__.tempfunc
解决方案
方法1:用类封装函数逻辑(推荐,兼顾灵活性与可Pickle性)
将带参数的函数逻辑封装为可Pickle的类实例,通过实现__call__方法让类实例像函数一样被调用:
import multiprocessing as mp class ControllerFunc: def __init__(self, parameter): self.parameter = parameter def __call__(self, argument): print(self.parameter * argument) class task(): def parallelizationWrap(self): poolSize = 5 with mp.Pool(poolSize) as pool: for _ in pool.imap(self.parallelizationFunc, range(poolSize)): pass def serialWrap(self): for _ in range(5): self.parallelizationFunc() def setup(self, unusedVar=None): vallist = [1,2,3,4,5] self.funclist = [] for i in range(5): # 用类实例替代闭包函数 self.funclist.append(ControllerFunc(vallist[i])) def parallelizationFunc(self, unuserVar=None): for step in range(25): for j in range(5): result = self.funclist[j](step) # simulation.sendSignalToCorrectAgent(result) if __name__ == "__main__": mp.freeze_support() c1 = task() c1.setup() c1.parallelizationWrap() # c1.serialWrap()
方法2:使用cloudpickle替代默认Pickle
若坚持使用闭包函数,可借助cloudpickle库,它支持序列化更多Python对象(包括本地函数)。先安装依赖:
pip install cloudpickle
修改多进程池初始化逻辑,指定用cloudpickle完成序列化:
import multiprocessing as mp import cloudpickle def cloudpickle_register(): import pickle pickle.Pickler = cloudpickle.Pickler class task(): def parallelizationWrap(self): poolSize = 5 # 初始化池时注册cloudpickle with mp.Pool(poolSize, initializer=cloudpickle_register) as pool: for _ in pool.imap(self.parallelizationFunc, range(poolSize)): pass def serialWrap(self): for _ in range(5): self.parallelizationFunc() def setup(self, unusedVar=None): vallist = [1,2,3,4,5] self.funclist = [] for i in range(5): def tempfunc(argument, parameter=vallist[i]): print(parameter*argument) self.funclist.append(tempfunc) def parallelizationFunc(self, unuserVar=None): for step in range(25): for j in range(5): result = self.funclist[j](step) # simulation.sendSignalToCorrectAgent(result) if __name__ == "__main__": mp.freeze_support() c1 = task() c1.setup() c1.parallelizationWrap() # c1.serialWrap()
方法3:调整全局函数定义方式(避免对象引用不匹配)
若使用全局函数,不能在循环中动态定义,需通过参数绑定传递专属值:
import multiprocessing as mp from functools import partial # 定义全局函数 def tempfunc(argument, parameter): print(parameter * argument) class task(): def parallelizationWrap(self): poolSize = 5 with mp.Pool(poolSize) as pool: for _ in pool.imap(self.parallelizationFunc, range(poolSize)): pass def serialWrap(self): for _ in range(5): self.parallelizationFunc() def setup(self, unusedVar=None): vallist = [1,2,3,4,5] self.funclist = [] for i in range(5): # 用partial绑定参数 self.funclist.append(partial(tempfunc, parameter=vallist[i])) def parallelizationFunc(self, unuserVar=None): for step in range(25): for j in range(5): result = self.funclist[j](step) # simulation.sendSignalToCorrectAgent(result) if __name__ == "__main__": mp.freeze_support() c1 = task() c1.setup() c1.parallelizationWrap() # c1.serialWrap()
原理说明
Python默认Pickle模块无法序列化本地函数(类方法内部的闭包),这类函数没有全局唯一的引用标识。改用类实例或cloudpickle,本质是为可调用对象提供了Pickle能识别的序列化路径。
内容的提问来源于stack exchange,提问作者The odor
相关产品推荐
相关产品推荐

