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

如何高效实现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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 20:22:04