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

使用Numba加速Python代码时MPI4py报错的排查与解决

问题:Numba+mpi4py并行代码出现MemoryError

复现代码

import numpy as np
from numba import njit, prange
from mpi4py import MPI

#####MPI setting#####
comm=MPI.COMM_WORLD
rank=comm.Get_rank()
size=comm.Get_size()
N_theta_to_scan=1000
n_per_proc=int(N_theta_to_scan/size)
n_more=int(N_theta_to_scan%size)
if rank<n_more:
    start=rank*(n_per_proc+1)
    number_to_cal=n_per_proc+1
else:
    start=n_more*(n_per_proc+1)+(rank-n_more)*n_per_proc
    number_to_cal=n_per_proc

#####function#####
@njit(parallel=True)
def func1(na,nb,nc):
    to_sum=np.zeros(na*nb*nc)
    for a in prange(0,na):
        for b in prange(0,nb):
            for c in prange(0,nc):
                to_sum[a*nb*nc+b*nc+c]=a*b*c
    out=np.sum(to_sum)
    return out

@njit(parallel=True)
def func2(start,number_to_cal):
    to_sum=np.zeros(number_to_cal)
    for i in prange(start,start+number_to_cal):
        to_sum[i-start]=func1(i,i*i,i*i*i)
    out2=np.sum(to_sum)
    return out2

#####main section#####
to_be_gather=np.array([func2(start,number_to_cal)])
gatheres=np.zeros(0)
comm.Reduce(to_be_gather,gatheres,op=MPI.SUM)

报错信息

MemoryError: Allocation failed (probably too large).

The above exception was the direct cause of the following exception:

Traceback (most recent call last):
  File "/share/workspace/wuliang/hanlin/test/test.py", line 38, in <module>
    to_be_gather=func2(start,number_to_cal)
                 ^^^^^^^^^^^^^^^^^^^^^^^^^^
SystemError: CPUDispatcher(<function func2 at 0x2aed19d7b420>) returned a result with an exception set

错误原因

核心问题是func1中创建的数组to_sum=np.zeros(na*nb*nc)规模远超系统内存:

  • 当i取值接近1000时,na=i=1000,nb=i²=1e6,nc=i³=1e9,数组元素总数达到1e18,完全超出任何硬件的内存承载能力,直接触发MemoryError。
  • 你提到单rank0运行时正常,是因为该进程分配到的i范围较小(进程数较多时,rank0仅处理前几百个i值),但随着i增大,最终仍会触发内存错误。

修改方案

无需创建巨量数组存储所有元素,利用数学公式直接计算总和:
sum(a*b*c for a in 0..na-1, b in 0..nb-1, c in 0..nc-1) = sum(a)*sum(b)*sum(c)

同时修正MPI Reduce的目标数组错误(原代码gatheres=np.zeros(0)无法接收数据):

import numpy as np
from numba import njit, prange
from mpi4py import MPI

#####MPI setting#####
comm=MPI.COMM_WORLD
rank=comm.Get_rank()
size=comm.Get_size()
N_theta_to_scan=1000
n_per_proc=int(N_theta_to_scan/size)
n_more=int(N_theta_to_scan%size)
if rank<n_more:
    start=rank*(n_per_proc+1)
    number_to_cal=n_per_proc+1
else:
    start=n_more*(n_per_proc+1)+(rank-n_more)*n_per_proc
    number_to_cal=n_per_proc

#####function#####
@njit
def sum_range(n):
    # 计算0到n-1的整数和:n*(n-1)/2
    return n * (n - 1) // 2

@njit(parallel=True)
def func1(na,nb,nc):
    # 用数学公式直接计算总和,避免创建巨量数组
    sum_a = sum_range(na)
    sum_b = sum_range(nb)
    sum_c = sum_range(nc)
    return sum_a * sum_b * sum_c

@njit(parallel=True)
def func2(start,number_to_cal):
    to_sum=np.zeros(number_to_cal)
    for i in prange(start,start+number_to_cal):
        to_sum[i-start]=func1(i,i*i,i*i*i)
    out2=np.sum(to_sum)
    return out2

#####main section#####
to_be_gather=np.array([func2(start,number_to_cal)])
# 修正Reduce的目标数组:root进程创建大小为1的数组,其他进程设为None
gatheres=np.zeros(1) if rank == 0 else None
comm.Reduce(to_be_gather,gatheres,op=MPI.SUM, root=0)

if rank ==0:
    print("Total sum:", gatheres[0])

内容的提问来源于stack exchange,提问作者Lin Han

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 13:12:27