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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 14:22:10