如何最大化利用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
优化方案总结
初始分块场景:如果必须保留分块的X列表,可通过列表推导式结合
np.hstack简化代码,减少循环中反复拼接的内存拷贝开销:result = np.hstack([np.dot(R[i], X[i]) for i in range(M)])逻辑和原循环一致,但代码更简洁,且一次性完成拼接,效率更高。
整体X+索引p场景:
- 当N较大时,
einsum是纯NumPy的高效方案,避免了Python显式循环的开销; - 但当单批次处理的列数(N_i)极小(如1、2),Python循环的解释器开销占比高,此时Numba编译的手动循环能获得数倍甚至一个数量级的加速——Numba将Python循环编译为机器码,消除了解释器开销,同时针对小矩阵乘法做了底层优化。
- 当N较大时,
纯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
相关产品推荐
相关产品推荐

