如何正确使用multiprocessing填充相关系数矩阵?
问题:多进程计算相关系数时函数未触发,矩阵始终为0
初始化空矩阵用于填充相关系数,使用multiprocessing.Pool提升速度后,calc_corr和pool_initializer函数未触发(无输出),最终矩阵仍全为0,相关代码如下:
# pool initializer function def pool_initializer(result): i, j, new_matr = result corr_matrix[i, j] = new_matr[0,1] corr_matrix[j, i] = new_matr[1,0] def calc_corr(i, j): sig1 = data_dict[i]['data'] sig2 = data_dict[j]['data'] nsig1 = sig1 - np.mean(sig1) nsig2 = sig2 - np.mean(sig2) corr = sp.signal.correlate(nsig1, nsig2, mode='same', method='fft') corr /= (len(sig2) * np.std(sig1) * np.std(sig2)) max_corr = find_highest_peak(corr) temp_matrix = np.zeros((2,2)) temp_matrix[0,1] = max_corr[1] temp_matrix[1,0] = max_corr[1] return (i, j, temp_matrix) def find_highest_peak(data): data_array = np.array(data) highest_peak_index = np.argmax(data_array) highest_peak_value = data_array[highest_peak_index] return highest_peak_index,highest_peak_value if __name__ == '__main__': data_dict = load_np_array_pickle('epoch_dict.pickle')[:300] start = time.perf_counter() corr_len = len(data_dict) corr_matrix = np.zeros((corr_len, corr_len)) status_percentage = int(corr_len / 100) if corr_len > 100 else 1 with Pool(processes=cpu_count()) as p: for i in range(corr_len): for j in range(i + 1, corr_len): p.apply_async(calc_corr, args=(i, j), callback=pool_initializer) p.close() p.join() end = time.perf_counter() print('Aufgebrachte Zeit: {:.2f} s'.format(end - start, 2)) print(corr_matrix)
问题原因分析
- 子进程无法访问全局变量
data_dict:Windows系统下multiprocessing用spawn模式创建子进程,会重新导入模块,主进程的data_dict不会被继承;即使在Linux/macOS的fork模式下,子进程对全局变量的访问也不可靠,导致calc_corr执行时抛出异常,但apply_async默认吞噬异常,看起来像函数未触发。 - 无异常捕获机制:子进程出错时没有任何提示,无法排查问题。
- 回调函数冗余:临时矩阵的设计没必要,增加了数据传递的开销。
修复方案
- 传递
data_dict给子进程:用Pool的initializer和initargs参数,在子进程启动时初始化全局变量,确保子进程能访问数据。 - 添加异常回调:设置
error_callback捕获子进程异常,方便排查问题。 - 简化逻辑:去掉冗余的临时矩阵,直接返回需要的相关系数值,减少数据传递成本。
修复后的代码
import numpy as np import scipy as sp import time from multiprocessing import Pool, cpu_count # 子进程全局变量 data_dict = None def init_worker(worker_data): """子进程初始化函数,传递data_dict""" global data_dict data_dict = worker_data def calc_corr(i, j): print(f"计算对 ({i}, {j})") # 验证函数触发 sig1 = data_dict[i]['data'] sig2 = data_dict[j]['data'] nsig1 = sig1 - np.mean(sig1) nsig2 = sig2 - np.mean(sig2) corr = sp.signal.correlate(nsig1, nsig2, mode='same', method='fft') corr /= (len(sig2) * np.std(sig1) * np.std(sig2)) max_corr = find_highest_peak(corr) return (i, j, max_corr[1]) # 直接返回核心值 def find_highest_peak(data): data_array = np.array(data) highest_peak_index = np.argmax(data_array) highest_peak_value = data_array[highest_peak_index] return highest_peak_index, highest_peak_value def update_corr_matrix(result): """回调函数,更新主进程的相关系数矩阵""" i, j, max_val = result corr_matrix[i, j] = max_val corr_matrix[j, i] = max_val def handle_error(error): """异常回调,打印错误信息""" print(f"发生错误: {error}") if __name__ == '__main__': def load_np_array_pickle(path): # 替换为实际的加载逻辑 import pickle with open(path, 'rb') as f: return pickle.load(f) data_dict_main = load_np_array_pickle('epoch_dict.pickle')[:300] start = time.perf_counter() corr_len = len(data_dict_main) corr_matrix = np.zeros((corr_len, corr_len)) # 创建进程池并初始化子进程数据 with Pool(processes=cpu_count(), initializer=init_worker, initargs=(data_dict_main,)) as p: for i in range(corr_len): for j in range(i + 1, corr_len): p.apply_async( calc_corr, args=(i, j), callback=update_corr_matrix, error_callback=handle_error ) p.close() p.join() end = time.perf_counter() print(f'耗时: {:.2f} s'.format(end - start)) print(corr_matrix)
内容的提问来源于stack exchange,提问作者alpa
相关产品推荐
相关产品推荐

