请求为scipy.signal.correlate添加axis参数以支持二维数组指定轴的互相关计算
scipy.signal.correlate添加axis参数以支持二维数组指定轴的互相关计算
你提的这个需求真的很戳痛点啊!确实现在scipy.signal.correlate在处理二维数组时,没法直接指定轴来做逐行/逐列的互相关,只能靠Python循环硬扛,大数据集下慢得离谱,而且临时 workaround 也不怎么直观,完全能理解你想要加个axis参数的诉求。
关于官方支持这个参数的可行性
从技术实现角度看,这个需求完全是可以落地的:
- Scipy里不少信号处理函数(比如
scipy.signal.convolve)已经支持axis参数了,实现思路可以直接参考——把指定轴之外的所有维度当成“批量维度”,对每个批量分组单独计算互相关,最后再把结果拼接起来,和你示例里预期的correlate(A, B, axis=1, mode='full')行为完全一致。 - 底层不管是用直接卷积还是FFT实现,只要对指定轴做向量化的批量处理,就能避免Python循环的开销,效率和原生1D计算差不多。
官方支持前的高效临时替代方案
如果等不及官方更新,这里有个基于FFT的高效实现,能避免Python循环的性能损耗,完全模拟你想要的axis参数行为:
import numpy as np from scipy.fft import rfft, irfft def axis_wise_correlate(A, B, axis=1, mode='full'): # 确保输入数组形状一致 if A.shape != B.shape: raise ValueError("A and B must have the same shape") # 计算full模式下的输出长度 n = A.shape[axis] + B.shape[axis] - 1 # 对指定轴执行FFT,利用FFT的批量处理能力 fft_A = rfft(A, n=n, axis=axis) # 互相关等价于:A和翻转后的B做卷积,所以先翻转B的指定轴 fft_B = rfft(np.flip(B, axis=axis), n=n, axis=axis) # 频域相乘后逆FFT得到时域结果 result = irfft(fft_A * np.conj(fft_B), n=n, axis=axis) # 根据mode参数裁剪结果 if mode == 'same': start_idx = (n - A.shape[axis]) // 2 end_idx = start_idx + A.shape[axis] result = np.take(result, range(start_idx, end_idx), axis=axis) elif mode == 'valid': start_idx = B.shape[axis] - 1 end_idx = A.shape[axis] result = np.take(result, range(start_idx, end_idx), axis=axis) # mode='full'直接返回原结果 return result # 测试验证 A = np.random.rand(1000, 100) B = np.random.rand(1000, 100) # 自定义函数计算轴-wise互相关 res_custom = axis_wise_correlate(A, B, axis=1, mode='full') # 循环计算作为对照 res_loop = np.array([np.correlate(a_row, b_row, mode='full') for a_row, b_row in zip(A, B)]) # 验证结果一致性 print(np.allclose(res_custom, res_loop)) # 应输出True
这个函数的效率比Python循环高得多,尤其是当你有大量行/列需要处理时,完全能hold住大数据集。
推动官方支持的建议
如果你想让这个功能尽快加入Scipy,可以试试这两个方式:
- 去Scipy的官方代码仓库提交Feature Request,把你的使用场景、示例代码(就是你提问里的那段)都附上,开发团队会评估需求的优先级。
- 先搜一下仓库里的已有issues,看看有没有人提过类似需求,如果有,你可以补充自己的使用场景,给需求加权重。
备注:内容来源于stack exchange,提问作者Habtie27
相关产品推荐
相关产品推荐

