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

如何避免嵌套定义函数时出现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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 21:43:13