【发布时间】:2016-07-17 07:44:25
【问题描述】:
这个问题与我不久前发布的一个问题有关:
Python, numpy, einsum multiply a stack of matrices
我试图理解为什么在将一堆矩阵相乘时以特定方式使用 Numba 时会得到加速。和以前一样,我放入一个 (500,201,2,2) 数组,沿第一个轴将最后的 (2x2) 矩阵相乘(所以 500 次乘法),得到一个 (201,2,2) 数组作为结果.
这是 Python 代码:
from numba import jit # numba 0.24, numpy 1.9.3, python 2.7.11
Arr = rand(500,201,2,2)
def loopMult(Arr):
ArrMult = Arr[0]
for i in range(1,len(Arr)):
ArrMult = np.einsum('fij,fjk->fik', ArrMult, Arr[i])
return ArrMult
@jit(nopython=True)
def loopMultJit(Arr):
ArrMult = np.empty(shape=Arr.shape[1:], dtype=Arr.dtype)
for i in range(0, Arr.shape[1]):
ArrMult[i] = Arr[0, i]
for j in range(1, Arr.shape[0]):
ArrMult[i] = np.dot(ArrMult[i], Arr[j, i])
return ArrMult
@jit(nopython=True)
def loopMultJit_2X2(Arr):
ArrMult = np.empty(shape=Arr.shape[1:], dtype=Arr.dtype)
for i in range(0, Arr.shape[1]):
ArrMult[i] = Arr[0, i]
for j in range(1, Arr.shape[0]):
x1 = ArrMult[i,0,0] * Arr[j,i,0,0] + ArrMult[i,0,1] * Arr[j,i,1,0]
y1 = ArrMult[i,0,0] * Arr[j,i,0,1] + ArrMult[i,0,1] * Arr[j,i,1,1]
x2 = ArrMult[i,1,0] * Arr[j,i,0,0] + ArrMult[i,1,1] * Arr[j,i,1,0]
y2 = ArrMult[i,1,0] * Arr[j,i,0,1] + ArrMult[i,1,1] * Arr[j,i,1,1]
ArrMult[i,0,0] = x1
ArrMult[i,0,1] = y1
ArrMult[i,1,0] = x2
ArrMult[i,1,1] = y2
return ArrMult
A1 = loopMult(Arr)
A2 = loopMultJit(Arr)
A3 = loopMultJit_2X2(Arr)
print np.allclose(A1, A2)
print np.allclose(A1, A3)
%timeit loopMult(Arr)
%timeit loopMultJit(Arr)
%timeit loopMultJit_2X2(Arr)
这是输出:
True
True
10 loops, best of 3: 40.5 ms per loop
10 loops, best of 3: 36 ms per loop
1000 loops, best of 3: 808 µs per loop
在上一个问题中,接受的答案表明,使用 f2py 可以将速度提高 8 倍,而无需进行详细优化。在这里,使用 Numba,我在 einsum 循环上使用 numba 获得了大约 10% 的加速,但如果不是在循环中使用 np.dot,我只需手动执行 2x2 矩阵乘法,我将获得 45 倍的加速。为什么是这样?我应该提到我已经实现了这两个带有适当类型签名的 jit 函数作为 guvectorize 版本,它基本上提供了相同的加速因子,所以我把它们排除在外。迭代 201,500,2,2 矩阵的加速也很小。
【问题讨论】:
-
我认为这只是调用
np.dot(检查类型、分配numpy数组等)的合理Python开销。尽管np.dot已经非常优化,但 2x2 数组很小,不值得开销。因此,您可以通过自己使用 numba 轻松完成(跳过所有 Python 开销)。您可能会发现np.dot在(例如)100x100 矩阵乘法上很难被击败。 -
我认为来自 BLAS 和样板的开销。检查 Numba 源代码中的this file。还有一般性。通常矩阵乘法至少包含 3 个循环,但您将它们全部展开在
loopMultiJit_2x2。