尝试复制您的 nxn 块细分。我使用了blk 辅助函数来复制二维数组切片:
import numpy as np
def matrix_mulitplication(A, B):
def blk(x,i1,i2):
return [row[i2] for row in x[i1]]
n = len(A) # A.shape[0]
if n == 1:
#print(A)
return A[0][0] * B[0][0]
else:
i = int(n / 2)
i1, i2 = slice(None,i), slice(i,None)
#C = np.zeros((n, n), dtype=np.int)
C1 = matrix_mulitplication(blk(A, i1, i1), blk(B,i1,i1)) +\
matrix_mulitplication(blk(A, i1, i2), blk(B,i2,i1))
C2 = matrix_mulitplication(blk(A, i1, i1), blk(B,i1,i2)) +\
matrix_mulitplication(blk(A, i1, i2), blk(B,i2,i2))
C3 = matrix_mulitplication(blk(A, i2, i1), blk(B,i1,i1)) +\
matrix_mulitplication(blk(A, i2, i2), blk(B,i2,i1))
C4 = matrix_mulitplication(blk(A, i2, i1), blk(B,i1,i2)) +\
matrix_mulitplication(blk(A, i2, i2), blk(B,i2,i2))
C = [[C1,C2],[C3,C4]]
return C
x = np.array([[1, 2], [3, 4]])
y = np.arange(16).reshape(4,4)
z = matrix_mulitplication([[1]],[[2]])
print(z)
z = matrix_mulitplication(x.tolist(), x.tolist())
print(z)
print(x@x)
z = matrix_mulitplication(y.tolist(), y.tolist())
print(z)
print(y@y)
这适用于一级递归,但不适用于两级:
1253:~/mypy$ python3 stack50552791.py
2
[[7, 10], [15, 22]]
[[ 7 10]
[15 22]]
[[[[4, 5], [20, 29], [52, 57], [132, 145]], [[6, 7], [38, 47], [62, 67], [158, 171]]], [[[36, 53], [52, 77], [212, 233], [292, 321]], [[70, 87], [102, 127], [254, 275], [350, 379]]]]
[[ 56 62 68 74]
[152 174 196 218]
[248 286 324 362]
[344 398 452 506]]
第2层的问题是matrix_mulitplication返回一个嵌套列表,而列表+被定义为concatenate,而不是元素添加。所以我必须定义另一个辅助函数(或 2 个)来正确解决这个问题。
好一点
def matrix_mulitplication(A, B):
def blk(x,i1,i2):
return [row[i2] for row in x[i1]]
def add(x,y):
if isinstance(x, list):
return [add(i,j) for i,j in zip(x,y)]
else:
return x+y
n = len(A) # A.shape[0]
if n == 1:
return A[0][0] * B[0][0]
else:
i = int(n / 2)
i1, i2 = slice(None,i), slice(i,None)
#C = np.zeros((n, n), dtype=np.int)
C1 = add(matrix_mulitplication(blk(A, i1, i1), blk(B,i1,i1)) ,\
matrix_mulitplication(blk(A, i1, i2), blk(B,i2,i1)))
C2 = add(matrix_mulitplication(blk(A, i1, i1), blk(B,i1,i2)) ,\
matrix_mulitplication(blk(A, i1, i2), blk(B,i2,i2)))
C3 = add(matrix_mulitplication(blk(A, i2, i1), blk(B,i1,i1)) ,\
matrix_mulitplication(blk(A, i2, i2), blk(B,i2,i1)))
C4 = add(matrix_mulitplication(blk(A, i2, i1), blk(B,i1,i2)) ,\
matrix_mulitplication(blk(A, i2, i2), blk(B,i2,i2)))
C = [[C1,C2],[C3,C4]]
return C
对于 4x4 的情况
[[[[56, 62], [152, 174]], [[68, 74], [196, 218]]], [[[248, 286], [344, 398]], [[324, 362], [452, 506]]]]
[[ 56 62 68 74]
[152 174 196 218]
[248 286 324 362]
[344 398 452 506]]
数字都在那里,但在一个 4 深度的嵌套列表中。
因此,不仅将嵌套列表分解为二维块比使用数组更尴尬,重新组装它们也很尴尬。
https://en.wikipedia.org/wiki/Strassen_algorithm
在没有详细阅读 Strassen 算法的情况下,很明显,任何关于其效率的声明都假定A_ij 索引(对于单个元素或块)在获取值和设置 (C_ij) 时都是有效的。对于列表,A[i] 或 A[i:j] 相当有效,但 [row[i2] for row in x[i1] 则不然。
用[[a,b],[c,d]] 组装块是可以的,但是任何与[[C_11,C_12],[C_21,C_22]] 可比的东西,其中C 元素是块而不是标量,都是复杂的。
似乎 Strassen 的目标是减少所需的乘法次数。这假设(标量)乘法是矩阵乘法中最昂贵的部分。对于 Python 列表,情况显然并非如此。访问元素的成本更高。
麻木作弊
我可以将 4x4 机箱中的 z 改造成
arr = np.array(z)
arr = arr.transpose(0,2,1,3).reshape(4,4)
这在numpy 中要容易得多。