多进程代码报错AttributeError:无法pickle本地对象parallel_operations.<locals>.process原因
错误原因分析
问题代码
def parallel_operations(points, primitives): batch_size,number_of_points,_ = points.shape _,_,number_of_primitives = primitives.shape gradient = torch.zeros(batch_size,number_of_points,number_of_primitives) def process(_lock,i): _lock.acquire() temp_points = points[i,:,:] temp_primitives= primitives[i,:,:].transpose(1,0) #[7,1024] #print("temp_shape{}".format(temp_primitives.shape)) temp = torch.zeros(number_of_points,number_of_primitives) for k in range(number_of_points): for j in range(number_of_primitives): temp[k,j] = torch.norm(temp_points[k,:]*temp_primitives[j,:3]+temp_primitives[j,3:6]) gradient[i,:,:] = temp print("gradient update {} {}" .format(i, gradient)) lock.release() return (i, gradient[i,:,:]) result = [] pool = Pool(multiprocessing.cpu_count()) lock = Manager().Lock() for i in range(10): result.append(pool.apply_async(process,args=(lock,i))) pool.close() pool.join() print(len(result)) for i in result: print(i.get()) if __name__ == "__main__": points = torch.randn(10,3,3) primitives = torch.randn(10,7,3) result1 = parallel_operations(points,primitives)
错误原因
运行代码时抛出AttributeError: Can't pickle local object 'parallel_operations.<locals>.process',核心原因如下:
- Python标准库的
multiprocessing调度子进程时,需要将任务函数序列化(pickle)后传递给子进程,但嵌套在函数内部的局部函数无法被pickle序列化——局部函数依赖外部函数的作用域上下文,pickle无法完整保存这种关联关系,导致序列化失败。 - 另外你在
process里直接引用了外部函数的points、gradient等变量,多进程间内存是隔离的,子进程拿到的是这些变量的副本,即便序列化函数成功,直接修改gradient也无法同步到主进程,当前加锁操作完全无效。
解决思路参考
- 把
process函数移到parallel_operations外部,将所有需要的参数(如points切片、number_of_points等)显式传入,消除对外部作用域的依赖。 - 让子进程计算完局部结果后返回,在主进程中统一合并结果,不要尝试在子进程中直接修改主进程的张量。
- 如果使用PyTorch,优先使用
torch.multiprocessing替代标准库的multiprocessing,它对张量的序列化和跨进程传递支持更友好。
内容的提问来源于stack exchange,提问作者Bill
相关产品推荐
相关产品推荐

