【发布时间】:2022-01-17 10:58:27
【问题描述】:
首先让我们考虑一个黑盒子,其中 m 和 n 是两个变量(其中 m 是 n) 它输出一个形状为 (m*n, m*n) 的二维矩阵。现在,需要将这个 2D 矩阵转换为形状为 (m, m, n, n)。我不确定以书面形式描述这一点的最佳方式,但数据的结构方式是在 2D (m*n)x(m *n) 矩阵中存在 m 很多 (nxn) 个“瓦片”每个方向。考虑一个示例数组 a,在这种情况下,我们有 m = 3 和 n = 2 所以传入的 2D 矩阵是 6x6:
print(a)
[[0 1 4 5 8 9 ]
[2 3 6 7 10 11]
[12 13 16 17 20 21]
[14 15 18 19 22 23]
[24 25 28 29 32 33]
[26 27 30 31 34 35]]
然后将 this 传递给某个函数:
b = some_func(a)
所需的输出将是 4D 数组:
print(b)
[[[[ 0 1]
[ 2 3]]
[[ 4 5]
[ 6 7]]
[[ 8 9]
[10 11]]]
[[[12 13]
[14 15]]
[[16 17]
[18 19]]
[[20 21]
[22 23]]]
[[[24 25]
[26 27]]
[[28 29]
[30 31]]
[[32 33]
[34 35]]]]
简而言之,我们需要在较大的二维数组中分离出“nxn”块。这种情况的实际含义是我们有mxm矩阵,其中每个条目实际上是一个n xn 矩阵,创建一个 4D 矩阵,然后我们可以做以下工作。这是一个高度简化的示例,用于演示一个更复杂的系统,其中还有很多事情要做。在我的例子中,还有一个额外的轴,m = 256,矩阵中的条目很复杂(64 位),我们非常关注性能,但是这些细节与问题无关。如果它有帮助n = 2的情况是唯一我们关心的情况,但是我希望有更通用的解决方案。
我可以合理地构想出一个使用 for 循环、索引、模运算等的解决方案,但是这在 Python 中会非常低效。
可能的解决方案?
- 头脑会立即跳到类似 np.reshape() 的东西上,但是我们不能简单地使用 a.reshape(m, m, n, n),因为 np.reshape() 的方式不会保留正确的顺序首先解开数组,如similar issue 中所述,但我坚信在经过大量 审议之后,该问题的解决方案在这种情况下将不起作用。我认为一个 np.reshape()、一个 np.swapaxes()、另一个 np.reshape() 和一个 np.swapaxes() 返回可能会起作用,但是唉,即使这样的方法会起作用,它似乎效率也很低。完全可以想象 np.reshapes()、np.swapaxes() 的一些混合物会提供解决方案,但我没有成功。
- 一位同事建议 np.einsum() 是一种非常强大且可通用的方法来执行矩阵运算(?),但是我没有成功。
- 最有可能的解决方案:我缺少一种特定的“Pythonic”做事方式 - 我不知道的某个 numpy 函数可以做到这一点!
我希望我对问题的描述已经足够了。围绕这个问题的背景非常复杂(射电天文图像处理),并且提供完整的细节会非常麻烦,请在提供任何解决方案时从表面上看待问题和初始假设。
这是重现测试问题的代码行。
a = np.array([[0, 1, 4, 5, 8, 9], [2, 3, 6, 7, 10, 11], [12, 13, 16, 17, 20, 21], [14, 15, 18, 19, 22, 23], [24, 25, 28, 29, 32, 33], [26, 27, 30, 31, 34, 35]])
【问题讨论】:
-
使用 reshape
a.reshape(3, 2, -1, 2).swapaxes(1,2).reshape(3, 3, 2, 2),但我认为np.einsum可能更简单。据我了解,reshape和swapaxes是 O(1) 操作且高效。如果我错了,请纠正我。 -
操作 reshape 和 swapaxes,如果可能,不要执行复制操作。在 Michael 描述的操作中,创建了数组
a的视图,以防万一具有恒定的复杂性。可以检查例如通过修改结果中的值,并查看值是否也在 a 中被修改。np.einsum确实是一个非常强大的操作,但它并不能解决手头的问题。这对于对整形结果执行的操作可能很有用,我需要更多信息来评论。 -
@LiamRyan 实际上,如果输入数组是连续的,您不会通过
a.reshape(3, 2, -1, 2).swapaxes(1,2).reshape(3, 3, 2, 2)创建数组(只是视图)。因此,性能与数组大小无关。对于这个相当小的示例,您不会看到复制和视图之间有太大区别。如果您将示例放大 10 000 倍,您肯定会看到差异。 -
einsum 解决方案看起来像
A=A.reshape(3, 2, -1, 2) #view array as (m,n,m,n) ,B=np.einsum('ijkl->ikjl',A) #the same as the swap axis.最简单的方法是查看数组的标志方法。如果它不拥有数据,它只是一个视图。 -
@max9111,
swapaxes之后的重塑将(通常)创建一个副本。reshape说把它想象成首先做一个解谜。在交换之后执行ravel,您会看到不同的元素顺序(与原始顺序相比)。swapaxes/transpose通过更改shape和strides工作,不需要复制。但由于值不再是“c”连续的,因此进一步处理通常会生成副本。
标签: python arrays numpy performance matrix