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

如何正确使用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默认吞噬异常,看起来像函数未触发。
  • 无异常捕获机制:子进程出错时没有任何提示,无法排查问题。
  • 回调函数冗余:临时矩阵的设计没必要,增加了数据传递的开销。

修复方案
  1. 传递data_dict给子进程:用Pool的initializer和initargs参数,在子进程启动时初始化全局变量,确保子进程能访问数据。
  2. 添加异常回调:设置error_callback捕获子进程异常,方便排查问题。
  3. 简化逻辑:去掉冗余的临时矩阵,直接返回需要的相关系数值,减少数据传递成本。

修复后的代码
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 03:35:45