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

兼容Numba njit的Numpy无for循环点积累积和计算方案

Numba兼容的Numpy向量化优化方案

首先我们可以先对原式做代数化简,降低计算量:
原式每一项为 xi[i] @ R[i] @ R[i].T @ xi[i],根据矩阵乘法结合律,等价于 (xi[i] @ R[i]) @ (xi[i] @ R[i]).T,也就是长度为3的向量自身的内积,等于该向量所有元素的平方和,无需计算开销更大的N×N中间矩阵R[i] @ R[i].T。

以下两种实现均在Numba支持的Numpy函数范围内,可直接在@njit装饰的函数中使用:

方案1:基于einsum实现

import numpy as np

h = np.sum(np.einsum('ij,ijk->ik', xi, R) ** 2)
  • np.einsum('ij,ijk->ik', xi, R) 会批量计算每个i对应的 xi[i] @ R[i],输出形状为(N, 3)
  • 对输出逐元素平方后求和,即可得到最终结果h

方案2:基于批量矩阵乘法实现

如果更习惯matmul写法,可使用等价实现:

h = np.sum((xi[:, np.newaxis, :] @ R).squeeze() ** 2)
  • 给xi增加维度得到形状为(N, 1, N)的数组,和形状为(N, N, 3)的R做批量矩阵乘法,得到形状为(N, 1, 3)的结果
  • squeeze去掉多余维度后平方求和,和方案1结果完全一致

两种实现均为纯向量化操作,无Python层面循环,计算复杂度从原实现的O(N³)降低到O(N²),在N=100的场景下性能明显优于原for循环实现,且无需开启并行不会产生额外开销。

内容的提问来源于stack exchange,提问作者Francesco Musso

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 01:45:04