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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 15:01:25