Python multiprocessing Pool.map大数组计算卡顿无报错无输出问题
脚本卡死、无输出的核心原因
- 缺少多进程入口保护:multiprocessing在spawn启动模式下(macOS/Windows默认,Linux下mpirun启动场景也会触发),如果主流程代码没有放在
if __name__ == '__main__':块内,子进程启动时会递归执行主模块代码,反复创建进程池,直接触发死锁。 - 内存占用爆炸:
pool.map会一次性将所有任务参数序列化后塞入进程通信队列,大数组场景下每个切片的序列化/反序列化会生成多份内存副本,10进程场景下内存占用是原数组的10倍以上,触发系统OOM Killer直接杀死子进程,主进程永远等不到子进程返回结果,就表现为无报错卡住。- 用
mpirun python3 sample_prob_func.py启动属于多进程嵌套:mpirun会先启动多个Python主进程,每个主进程内部又创建10个工作进程,进程数直接翻倍,内存和CPU瞬间占满。 - 原代码
my_func3切分的数组切片如果是原数组的视图,序列化时会连带整个原始大数组一起打包传输,进一步放大内存开销。
- 冗余API调用触发死锁:已经用
with contextlib.closing(mp.Pool(...))上下文管理池生命周期,上下文退出时会自动执行pool.close()和pool.join(),手动重复调用这两个方法在部分Python版本下会触发锁等待。 - 低级笔误:
mydata_list = [my_data1,my_data3,my_data3]第二个元素错误引用了my_data3,会导致计算结果不符合预期。 - 路径问题:没有提前创建输出目录,如果
save_results_to路径不存在,np.savetxt会直接报错,但多进程场景下子进程报错可能没有被捕获输出,表现为无文件生成。
优化后可直接运行的代码
优化点包括:增加主入口保护、用共享内存传递大数组避免多份副本、用imap替代map流式提交任务、移除冗余池操作、自动创建输出目录、修正笔误、进程数适配CPU核心数。
注意:直接用
python3 sample_prob_func.py运行,不要加mpirun,mpirun是MPI并行框架的启动命令,和Python内置multiprocessing冲突。
import numpy as np import multiprocessing as mp from scipy import signal import contextlib import os import time # 配置项 save_results_to = './result/' # 替换为实际存储路径 os.makedirs(save_results_to, exist_ok=True) # 自动创建输出目录,避免路径不存在写失败 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] N = len(arr_x) def my_func1(data): # 替换为你实际的CSD计算逻辑,返回维度需为[行数, N, 总频率点数] # 示例返回随机3D数组适配流程 total_freq = 110 return np.random.rand(data.shape[0], N, total_freq) def my_func2(args): csd_shared, arr_shape, fr_count = args # 从共享内存还原数组,无内存复制 csd = np.frombuffer(csd_shared, dtype=np.float64).reshape(arr_shape) csd_single = csd[:, :, fr_count] # 替换为你实际的单频点计算逻辑 return csd_single * 2 if __name__ == '__main__': np.random.seed(12345) total_rows = 5000 arr = np.reshape(np.random.rand(total_rows*N),(total_rows, N)) arr1 = np.reshape(np.random.rand(total_rows*N),(total_rows, N)) arr2 = np.reshape(np.random.rand(total_rows*N),(total_rows, N)) t0 = time.time() my_data1 = my_func1(arr) my_data2 = my_func1(arr1) my_data3 = my_func1(arr2) print(f'单进程CSD计算耗时: {time.time()-t0:.2f}s') # 修正原笔误 mydata_list = [my_data1, my_data2, my_data3] start_freq = 100 stop_freq = 110 freq_range = np.around(np.linspace(start_freq, stop_freq, 11)/10, decimals=2) no_of_freq = len(freq_range) count_day = 1 count_hour = 0 # 进程数设为不超过CPU物理核心数,避免过度调度 process_num = min(mp.cpu_count(), 8) for count in range(3): count_hour += 1 current_csd = mydata_list[count] print(f'开始处理第{count_hour}组数据,数组维度: {current_csd.shape}') # 将大数组放入共享内存,所有子进程共享同一块内存,无副本 csd_shared = mp.RawArray('d', current_csd.size) csd_np = np.frombuffer(csd_shared, dtype=np.float64).reshape(current_csd.shape) np.copyto(csd_np, current_csd) # 任务仅传共享内存指针、数组维度、频率索引,不传实际数组 task_list = [(csd_shared, current_csd.shape, fr_count) for fr_count in range(start_freq, stop_freq)] t0 = time.time() with contextlib.closing(mp.Pool(processes=process_num)) as pool: # imap流式提交任务,chunksize平衡调度开销和内存占用 dft = list(pool.imap(my_func2, task_list, chunksize=2)) data_dft = np.array(dft) print(f'第{count_hour}组并行计算耗时: {time.time()-t0:.2f}s,结果维度: {data_dft.shape}') # 写入结果文件 for i in range(no_of_freq-1): jj = freq_range[i] dft_1hr_complx = data_dft[i, :, :] save_path = os.path.join(save_results_to, f'csd_Day_{count_day}_Hour_{count_hour}_f_{jj}_hz.txt') np.savetxt(save_path, dft_1hr_complx.view(float)) print(f'第{count_hour}组共{no_of_freq-1}个结果文件写入完成')
调优建议
- 如果数组规模进一步增大,可以适当调大
chunksize参数(比如设为4),减少进程间调度开销,不要设为1,否则调度成本会抵消多进程收益。 - 如果计算逻辑中存在大量GIL释放操作(比如scipy/numpy的底层C函数),进程数和CPU核心数1:1即可,不要开过多进程,否则会增加内存开销和上下文切换成本。
- 如果需要捕获子进程报错,可以在
my_func2中加try-except块打印异常信息,方便定位问题。
内容的提问来源于stack exchange,提问作者CEB
相关产品推荐
相关产品推荐

