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

如何最大化利用NumPy实现带可变维度的张量高效乘法?

问题描述

我需要计算张量集合 $R = {R₁, R₂, ..., R_M}$ 与 $X = {X₁, X₂, ..., X_M}$ 的乘积,其中每个 $R_i$ 是3×3矩阵,每个 $X_i$ 是3×$N_i$ 矩阵。目标是最大化利用NumPy的功能完成 $R_i × X_i$ 的计算,最终将所有结果按列拼接。

我的初始实现代码如下:

import numpy as np

np.random.seed(0)

M = 5
R = [np.random.rand(3, 3) for _ in range(M)]
X = []
for i in range(M):
    N_i = np.random.randint(1, 6)
    X_i = np.random.rand(3, N_i)
    X.append(X_i)
    
result = np.zeros((3, 0))
for i in range(M):
    R_i = R[i]
    X_i = X[i]
    result = np.hstack((result, np.dot(R_i, X_i)))

print(result)

补充场景1:N_i取值极小,且X为整体矩阵+索引向量

实际场景中,$N_i$ 大多是1或2,很少超过5。另外,我实际存储的是一个3×N的整体矩阵X,而非分块的 $X_i$,X的列通过索引向量p匹配对应的R矩阵(即X的第i列对应R[p[i]])。

我尝试用einsum实现,对比直接循环的版本:

import numpy as np

M = 30
N = 100

np.random.seed(0)
p = np.random.randint(M, size=N)
R = np.random.rand(M, 3, 3)
X = np.random.rand(3, N)

result_einsum = np.einsum('ijk,ki->ji', R[p], X)

result_direct = np.zeros((3, N))
for i in range(N):
    result_direct[:, i] = np.dot(R[p[i]], X[:, i])

print(np.allclose(result_einsum, result_direct))

补充场景2:Numba加速效果测试

测试发现用Numba编译循环代码能获得大幅加速,代码及测试结果如下:

import numpy as np
import numba
from timeit import Timer

M = 30
N = 100

np.random.seed(0)
p = np.random.randint(M, size=N)
R = np.random.rand(M, 3, 3)
X = np.random.rand(3, N)

@numba.njit
def numba_direct(R, p, X, result_direct, N):
    for i in range(N):
        p_i = p[i]
        for j in range(3):
            res = 0.0
            for k in range(3):
                res += R[p_i, j, k] * X[k, i]
            result_direct[j, i] = res

result_direct = np.zeros((3, N))
numba_direct(R, p, X, result_direct, N)
result_einsum = np.einsum('ijk,ki->ji', R[p], X)
print(np.allclose(result_einsum, result_direct))

ntimes = 10000

einsum_timer = Timer(lambda: np.einsum('ijk,ki->ji', R[p], X))
einsum_time = einsum_timer.timeit(number=ntimes)

numba_direct_timer = Timer(lambda: numba_direct(R, p, X, result_direct, N))
numba_direct_time = numba_direct_timer.timeit(number=ntimes)

print(f'Einsum runtime: {einsum_time:.4f} seconds')
print(f'Numba direct runtime: {numba_direct_time:.4f} seconds')

执行结果:

Einsum runtime: 0.0979 seconds
Numba direct runtime: 0.0129 seconds

优化方案总结

  1. 初始分块场景:如果必须保留分块的X列表,可通过列表推导式结合np.hstack简化代码,减少循环中反复拼接的内存拷贝开销:

    result = np.hstack([np.dot(R[i], X[i]) for i in range(M)])
    

    逻辑和原循环一致,但代码更简洁,且一次性完成拼接,效率更高。

  2. 整体X+索引p场景:

    • 当N较大时,einsum是纯NumPy的高效方案,避免了Python显式循环的开销;
    • 但当单批次处理的列数(N_i)极小(如1、2),Python循环的解释器开销占比高,此时Numba编译的手动循环能获得数倍甚至一个数量级的加速——Numba将Python循环编译为机器码,消除了解释器开销,同时针对小矩阵乘法做了底层优化。
  3. 纯NumPy的中间性能方案:如果需要纯NumPy实现且追求比einsum更快的速度,可以预先将R按p索引展开,再用np.matmul批量计算:

    R_expanded = R[p]  # shape (N,3,3)
    X_reshaped = X.T[:, :, np.newaxis]  # shape (N,3,1)
    result_matmul = np.matmul(R_expanded, X_reshaped).squeeze().T  # shape (3,N)
    

    该方案性能介于einsum和Numba版本之间,适合不能引入Numba依赖的场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 18:46:04