如何高效实现N×N×N数组的b_ij=∑ₖa_{i+k,j+k,k}计算(NumPy/JAX)
三维数组特定求和的高效NumPy实现问题
给定一个$N\times N \times N$的数组$a$,需要高效计算公式:
$$ b_{ij} = \sum_k a_{i+k,j+k,k} $$
当前采用的Python实现如下:
import numpy as np b = np.zeros((N, N)) for i in range(N): for j in range(N): m = np.maximum(i, j) b[i,j] = np.einsum('iii', a[i:N-m+i,j:N-m+j,:N-m])
该实现效率较低,请问是否可以不使用Cython,仅通过NumPy(或jax.numpy等兼容接口)完成高效实现?
编辑1:补充了缺失的显式边界。
内容的提问来源于stack exchange,提问作者Uroc327
相关产品推荐
相关产品推荐

