Python使用mpi4py拼接并行收集数据的错误修复与优化
问题背景
- 使用
mpi4py开展并行计算时,通过列表append存储各进程计算结果,尝试在根节点(root==0)按频率顺序拼接汇总数据后保存。 - 按参考建议修改代码后脚本可正常运行,但数据拼接结果不符合预期,输出文件内容异常。
- 除修复拼接错误外,当前Python脚本实现效率较低,需要更高效的同类问题解决方案。
现有实现代码
依赖导入与基础配置
import numpy as np from scipy import signal from mpi4py import MPI import random import cmath, math import matplotlib.pyplot as plt import time # 结果存储路径 save_results_to = 'File storing path'
参数定义与串行函数实现
count_day = 1 count_hour = 1 arr_x = [0, 8.49, 0.0, -8.49, -12.0, -8.49, -0.0, 8.49, 12.0] arr_y = [0, 8.49, 12.0, 8.49, 0.0, -8.49, -12.0, -8.49, -0.0] M = len(arr_x) N = len(arr_y) np.random.seed(12345) total_rows = 50000 raw_data=np.reshape(np.random.rand(total_rows*N),(total_rows,N)) # 互功率谱计算(循环实现) fs = 500 # 采样频率 def csdMat(data): dat, cols = data.shape total_csd = [] for i in range(cols): col_csd =[] for j in range(cols): freq, Pxy = signal.csd(data[:,i], data[:, j], fs=fs, window='hann', nperseg=100, noverlap=70, nfft=5000) col_csd.append(Pxy) total_csd.append(col_csd) pxy = np.array(total_csd) return freq, pxy # 计算CSD t0 = time.time() freq, csd = csdMat(raw_data) print('CSD数据维度:', csd.shape) print('CSD循环计算耗时:{} 秒'.format(time.time()-t0)) kf=1*2*np.pi/10 resolution = 50 # 分辨率参数,值越高计算耗时越长 grid_size = N * resolution kx = np.linspace(-kf, kf, ) # 波数x向量(原代码此处缺参数) ky = np.linspace(-kf, kf, grid_size) # 波数y向量 # 二维DFT(四层循环实现) def DFT2D(data): P=len(kx) Q=len(ky) dft2d = np.zeros((P,Q), dtype=complex) for k in range(P): for l in range(Q): sum_matrix = 0.0 for m in range(M): for n in range(N): e = cmath.exp(-1j*((((dx[m]-dx[n])*kx[l])/1) + (((dy[m]-dy[n])*ky[k])/1))) sum_matrix += data[m, n] * e dft2d[k,l] = sum_matrix return dft2d dx = arr_x[:]; dy = arr_y[:] # MPI初始化 comm = MPI.COMM_WORLD size = comm.Get_size() rank = comm.Get_rank()
并行计算与结果保存逻辑
data = [] start_freq = 100 end_freq = 109 freq_range = np.arange(start_freq,end_freq) no_of_freq = len(freq_range) for fr_count in range(start_freq, end_freq): if fr_count % size == rank: spec_csd = csd[:,:, fr_count] dft = DFT2D(spec_csd) spec = np.array(np.real(dft)) print('单频结果维度:', spec.shape) data.append(spec) np.seterr(invalid='ignore') data = comm.gather(data, root =0) print("进程Rank:", rank, ",结果维度:\n", spec.shape) if rank == 0: output_data = np.concatenate(data, axis = 0) dft_tot = np.array((output_data), dtype='object') res = np.zeros((grid_size, grid_size)) for k in range(size): for i in range(no_of_freq): jj = np.around(freq[freq_range[i]], decimals = 2) res[i * size + k] = dft_tot[k][i] data = np.array(res) np.savetxt(save_results_to + f'Day_{count_day}_hour_{count_hour}_f_{jj}_hz.txt', data.view(float))
Slurm作业提交脚本
通过sbatch my_file.sh提交作业,脚本内容如下:
#! /bin/bash -l #SBATCH -J testmvapich2 #SBATCH -N 1 #SBATCH --ntasks=10 #SBATCH --cpus-per-task=1 #SBATCH --mem-per-cpu=3000MB #SBATCH --time=00:20:00 #SBATCH -p para #SBATCH --output="stdout.txt" #SBATCH --error="stderr.txt" #SBATCH -A camk eval "$(conda shell.bash hook)" conda activate myenv cd $SLURM_SUBMIT_DIR mpirun python3 mpi_test.py
修复与优化方案
1. 现有bug修复
(1)波数向量参数错误
原代码中kx = np.linspace(-kf, kf, )未指定采样点数,默认生成50个点,与ky的grid_size长度不匹配,会导致DFT结果维度异常,修改为:
kx = np.linspace(-kf, kf, grid_size)
(2)根节点数据拼接逻辑错误
原拼接逻辑存在两个问题:一是comm.gather返回的是按进程分组的嵌套列表,直接沿axis=0拼接会打乱结果顺序、出现维度不匹配;二是结果存储数组res的维度与实际需要保存的多频结果不匹配,索引赋值会出现越界、覆盖问题。
替换根节点处理逻辑如下:
if rank == 0: # 按频率点数量初始化结果容器,保证顺序一一对应 all_spec = [None]*no_of_freq for proc_id in range(size): for local_idx, spec in enumerate(data[proc_id]): # 计算当前分片对应的全局频率位置 global_pos = (proc_id - start_freq % size) + local_idx * size if 0 <= global_pos < no_of_freq: all_spec[global_pos] = spec # 按频率顺序逐文件保存 for idx, spec in enumerate(all_spec): jj = np.around(freq[freq_range[idx]], decimals=2) np.savetxt(save_results_to + f'Day_{count_day}_hour_{count_hour}_f_{jj}_hz.txt', spec)
2. 性能优化方案
(1)向量化替换四层循环DFT
纯Python循环的DFT2D是核心性能瓶颈,用numpy广播+einsum实现向量化运算,速度可提升100倍以上:
def DFT2D_vec(data): dx_diff = dx[:, None] - dx[None, :] dy_diff = dy[:, None] - dy[None, :] phase = -1j * (dx_diff[..., None, None] * kx[None, None, None, :] + dy_diff[..., None, None] * ky[None, None, :, None]) exp_mat = np.exp(phase) dft2d = np.einsum('mn,mnlk->lk', data, exp_mat) return dft2d
(2)MPI通信优化
替换列表append+gather的通信方式,提前分配连续内存的numpy数组,使用Gatherv直接传输二进制数据,可减少30%以上的序列化开销;单节点运行时也可使用共享内存,直接让各进程将结果写入共享内存对应位置,完全省略根节点拼接步骤。
(3)CSD计算优化
原双重循环的CSD计算可替换为scipy的向量化接口,或提前对信号做分窗、预计算FFT,减少重复运算。
内容的提问来源于stack exchange,提问作者CEB
相关产品推荐
相关产品推荐

