如何在mpi4py中使用comm.Gather聚合不同形状的数组?
解决mpi4py中不同形状数组的聚合问题
当使用comm.Gather处理不同形状但可拼接的数组时,因它要求所有发送数据的尺寸完全一致,会导致执行失败。这种场景下应该用**comm.Gatherv**,它支持每个进程发送不同长度的数据,无需手动拆分数组再拼接。
实现逻辑
- 每个进程计算自身数组的元素总数(因列数固定,也可仅计算第一维度的行数,后续转换为元素数)
- 根进程收集所有进程的元素数量,以此计算接收缓冲区的偏移量
- 根进程初始化足够大的接收缓冲区
- 调用
comm.Gatherv完成聚合,传入sendcounts(各进程发送的元素数)和displs(接收时的偏移量)参数
修改后的示例代码
from mpi4py import MPI import numpy as np comm = MPI.COMM_WORLD size = comm.Get_size() rank = comm.Get_rank() # 生成不同形状的数组:rank1是(2,3),其他是(5,3) a = np.zeros((2 if rank == 1 else 5, 3), dtype=float) + rank print(f"rank {rank} array shape: {a.shape}") # 计算当前进程数组的元素总数 send_count = a.size # 根进程收集所有进程的元素数量 send_counts = comm.gather(send_count, root=0) if rank == 0: # 计算接收缓冲区总元素数,初始化二维缓冲区 total_elements = sum(send_counts) b = np.zeros((total_elements // 3, 3), dtype=float) - 1 # 计算每个进程数据在接收缓冲区的起始偏移(元素为单位) displs = [sum(send_counts[:i]) for i in range(size)] else: b = None displs = None # 用Gatherv聚合不同大小的数据 comm.Gatherv(sendbuf=a, recvbuf=(b, send_counts, displs, MPI.DOUBLE), root=0) if rank == 0: print("聚合后的数组:") print(b)
关键参数说明
send_counts:根进程收集到的、每个进程要发送的元素数量列表displs:每个进程的数据在接收缓冲区中的起始位置(以元素为单位),确保不同进程的数据按顺序拼接recvbuf元组:格式为(接收数组, 各进程发送元素数, 偏移量, MPI数据类型),这里MPI.DOUBLE对应numpy的float类型
内容的提问来源于stack exchange,提问作者Gabe BAO
相关产品推荐
相关产品推荐

