使用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
相关产品推荐
相关产品推荐

