如何用mpi4py向子进程Scatter大型Pandas DataFrame?
解决mpi4py分散大型Pandas DataFrame时的挂起问题
当DataFrame规模较大时,直接用comm.scatter()传递Pandas对象会因pickle序列化的低效导致进程阻塞挂起。以下是几种可行的解决方案:
方案一:转为NumPy数组传输后重构DataFrame
mpi4py对NumPy数组的传输支持更高效,可避免Pandas对象序列化的额外开销。
修改master.py
from mpi4py import MPI import numpy as np import sys import pandas as pd comm = MPI.COMM_SELF.Spawn(sys.executable, args=['child.py'], maxprocs=5) class Model: def __init__(self, num_draw): self.num_draw = num_draw def test(self): id = np.arange(14) id_repeat = np.repeat(id, self.num_draw) draws = np.arange(self.num_draw*14) data = pd.DataFrame({"id": id_repeat, "draws": draws}) # 转换为NumPy数组并记录列名 data_np = data.to_numpy() columns = data.columns.tolist() # 拆分数组块 chunks = np.array_split(data_np, 5) # 先广播列名到所有子进程 comm.bcast(columns, root=MPI.ROOT) # 分散数组块 comm.scatter(chunks, root=MPI.ROOT) sum_gather = None sum_gather = comm.gather(sum_gather, root=MPI.ROOT) print('parent', sum_gather) model=Model(500) model.test() comm.Disconnect()
修改child.py
from mpi4py import MPI import pandas as pd comm = MPI.Comm.Get_parent() rank = comm.Get_rank() # 接收列名 columns = comm.bcast(None, root=0) # 接收数组块 data_np = comm.scatter(None, root=0) # 重构DataFrame data = pd.DataFrame(data_np, columns=columns) def li(data): data_sum = data.sum() return data_sum data_sum=li(data) print("child rank", rank, "after scatter", data_sum) comm.gather(data_sum, root=0) comm.Disconnect()
方案二:手动Send/Recv控制传输
对于超大规模DataFrame,手动控制数据传输流程,避免scatter的自动序列化瓶颈。
修改master.py的test方法
def test(self): id = np.arange(14) id_repeat = np.repeat(id, self.num_draw) draws = np.arange(self.num_draw*14) data = pd.DataFrame({"id": id_repeat, "draws": draws}) chunks = np.array_split(data, 5) # 手动发送每个块到对应子进程 for i in range(5): # 先发送列名和数据形状 comm.send((chunks[i].columns.tolist(), chunks[i].shape), dest=i, tag=0) # 发送数据数组 comm.Send(chunks[i].to_numpy(), dest=i, tag=1) sum_gather = None sum_gather = comm.gather(sum_gather, root=MPI.ROOT) print('parent', sum_gather)
修改child.py
from mpi4py import MPI import pandas as pd import numpy as np comm = MPI.Comm.Get_parent() rank = comm.Get_rank() # 接收元信息(列名、形状) columns, shape = comm.recv(source=0, tag=0) # 初始化数组接收数据 data_np = np.empty(shape, dtype=np.int64) comm.Recv(data_np, source=0, tag=1) # 重构DataFrame data = pd.DataFrame(data_np, columns=columns) def li(data): data_sum = data.sum() return data_sum data_sum=li(data) print("child rank", rank, "after scatter", data_sum) comm.gather(data_sum, root=0) comm.Disconnect()
方案三:提升pickle序列化效率
通过指定更高版本的pickle协议,优化默认序列化的性能:
在master.py开头添加:
import pickle # 使用最高版本的pickle协议 MPI.pickle.__init__(protocol=pickle.HIGHEST_PROTOCOL)
此方法效果弱于前两种,仅作为辅助优化手段。
内容的提问来源于stack exchange,提问作者jasmine
相关产品推荐
相关产品推荐

