【发布时间】:2019-09-10 01:53:32
【问题描述】:
我正在尝试使用 Numpy 在 Python 中高效地实现 n 模式张量矩阵乘积(由 Kolda 和 Bader 定义:https://www.sandia.gov/~tgkolda/pubs/pubfiles/SAND2007-6702.pdf)。该操作有效地归结为(对于矩阵 U、张量 X 和轴/模式 k):
通过折叠所有其他轴,从 X 中提取沿轴 k 的所有向量。
使用标准矩阵乘法将左侧的这些向量乘以 U。
使用相同的形状再次将向量插入输出张量,除了 X.shape[k],它现在等于 U.shape[0](最初,X.shape[k] 必须是等于 U.shape[1],作为矩阵乘法的结果)。
一段时间以来,我一直在使用显式实现,它分别执行所有这些步骤:
转置张量以将轴 k 带到前面(在我的完整代码中,我添加了一个例外情况,以防 k == X.ndim - 1,在这种情况下,将其保留在那里并转置所有未来的操作会更快,或者至少在我的应用程序中,但这与此处无关)。
重塑张量以折叠所有其他轴。
计算矩阵乘法。
重塑张量以重建所有其他轴。
将张量转回原来的顺序。
我认为这个实现会创建很多不必要的(大)数组,所以一旦我发现 np.einsum,我认为这会大大加快速度。但是使用下面的代码我得到了更糟糕的结果:
import numpy as np
from time import time
def mode_k_product(U, X, mode):
transposition_order = list(range(X.ndim))
transposition_order[mode] = 0
transposition_order[0] = mode
Y = np.transpose(X, transposition_order)
transposed_ranks = list(Y.shape)
Y = np.reshape(Y, (Y.shape[0], -1))
Y = U @ Y
transposed_ranks[0] = Y.shape[0]
Y = np.reshape(Y, transposed_ranks)
Y = np.transpose(Y, transposition_order)
return Y
def einsum_product(U, X, mode):
axes1 = list(range(X.ndim))
axes1[mode] = X.ndim + 1
axes2 = list(range(X.ndim))
axes2[mode] = X.ndim
return np.einsum(U, [X.ndim, X.ndim + 1], X, axes1, axes2, optimize=True)
def test_correctness():
A = np.random.rand(3, 4, 5)
for i in range(3):
B = np.random.rand(6, A.shape[i])
X = mode_k_product(B, A, i)
Y = einsum_product(B, A, i)
print(np.allclose(X, Y))
def test_time(method, amount):
U = np.random.rand(256, 512)
X = np.random.rand(512, 512, 256)
start = time()
for i in range(amount):
method(U, X, 1)
return (time() - start)/amount
def test_times():
print("Explicit:", test_time(mode_k_product, 10))
print("Einsum:", test_time(einsum_product, 10))
test_correctness()
test_times()
适合我的时间:
显式:3.9450525522232054
Einsum:15.873924326896667
这是正常的还是我做错了什么?我知道在某些情况下存储中间结果可以降低复杂性(例如链式矩阵乘法),但是在这种情况下,我想不出任何重复的计算。矩阵乘法是否如此优化以至于它消除了不转置的好处(从技术上讲,它的复杂性较低)?
【问题讨论】:
-
以下是我的系统(3.6.5,NumPy 1.14.3)上的时间:
Explicit: 1.1669329166412354和Einsum: 1.553536319732666 -
我的时间是:
Explicit: 0.781634783744812 Einsum: 0.8676517248153687在 Python 3.5 和 NumPy 1.16.1 上。你用的是什么 NumPy 版本? -
在
k_product中,重担在@产品中。转置和重塑几乎没有成本。根据einsum的形状,优化可能最终也会使用matmul。因此,时间安排相似也就不足为奇了。 -
谢谢你的cmets,看来我的软件可能有点过时了。我将 Python 3.6.7 与 Numpy 1.13.3 一起使用。我会看看我是否可以更新我的 Numpy 安装。不过我的电脑还是很慢,这可能与硬件有关吗?我正在使用 Intel(R) Core(TM) i7-4700HQ CPU @ 2.40GHz,两种方法似乎只使用一个内核。
-
einsum 中
optimize参数的使用在最近的numpy版本中发生了很大变化。所以是的,最好获得一个最新版本,并测试几个替代方案。