Strassen矩阵乘法拼接报错:奇数维度填充后6x6矩阵运行异常
参考Strassen矩阵乘法实现,当矩阵维度为奇数时填充一行一列0(比如5x5转6x6),运行时出现拼接报错:
Traceback (most recent call last):
File "/home/surfacepro/Downloads/strassen.py", line 82, in
main()
File "/home/surfacepro/Downloads/strassen.py", line 77, in main
print(strassen(matrixA, matrixB))
^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/surfacepro/Downloads/strassen.py", line 33, in strassen
p1 = strassen(a, f - h)
^^^^^^^^^^^^^^^^^^
File "/home/surfacepro/Downloads/strassen.py", line 35, in strassen
p3 = strassen(c + d, e)
^^^^^^^^^^^^^^^^^^
File "/home/surfacepro/Downloads/strassen.py", line 49, in strassen
c = np.vstack((np.hstack((c11, c12)), np.hstack((c21, c22))))
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/usr/lib64/python3.11/site-packages/numpy/core/shape_base.py", line 289, in vstack
return _nx.concatenate(arrs, 0, dtype=dtype, casting=casting)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
ValueError: all the input array dimensions except for the concatenation axis must match exactly, but along dimension 1, the array at index 0 has size 1 and the array at index 1 has size 0
完整实现代码如下:
# Version 3.6 import numpy as np import re, math def split(matrix): """ Splits a given matrix into quarters. Input: n1 xn matrix Output: tuple containing 4 n/2 x n/2 matrices corresponding to a, b, c, d """ row, col = matrix.shape row2, col2 = row//2, col//2 return matrix[:row2, :col2], matrix[:row2, col2:], matrix[row2:, :col2], matrix[row2:, col2:] def strassen(x, y): """ Computes matrix product by divide and conquer approach, recursively. Input: nxn matrices x and y Output: nxn matrix, product of x and y """ # Base case when size of matrices is 1x1 if len(x) == 1: return x * y # Splitting the matrices into quadrants. This will be done recursively # until the base case is reached. a, b, c, d = split(x) e, f, g, h = split(y) # Computing the 7 products, recursively (p1, p2...p7) p1 = strassen(a, f - h) p2 = strassen(a + b, h) p3 = strassen(c + d, e) p4 = strassen(d, g - e) p5 = strassen(a + d, e + h) p6 = strassen(b - d, g + h) p7 = strassen(a - c, e + f) # Computing the values of the 4 quadrants of the final matrix c c11 = p5 + p4 - p2 + p6 c12 = p1 + p2 c21 = p3 + p4 c22 = p1 + p5 - p3 - p7 # Combining the 4 quadrants into a single matrix by stacking horizontally and vertically. c = np.vstack((np.hstack((c11, c12)), np.hstack((c21, c22)))) return c def main(): with open('matrix_A.txt', 'r') as f: flatMatrix = f.read() with open('matrix_B.txt', 'r') as f: flatMatrix2 = f.read() numbers = re.compile("-?\\d+") result = list(map(int, numbers.findall(flatMatrix))) result2 = list(map(int, numbers.findall(flatMatrix2))) matrix_dimension = int(math.sqrt(len(result))) matrixA = np.array(result).reshape(matrix_dimension, matrix_dimension) matrixB = np.array(result2).reshape(matrix_dimension, matrix_dimension) if matrix_dimension % 2 != 0: matrix_dimension = matrix_dimension + 1 matrixA = np.pad(matrixA, [(0, 1), (0, 1)], mode='constant', constant_values=0) matrixB = np.pad(matrixB, [(0, 1), (0, 1)], mode='constant', constant_values=0) print(matrixA) print(matrixB) print(strassen(matrixA, matrixB)) print("The sum is: ", strassen(matrixA, matrixB).sum()) if __name__ == '__main__': main()
请求排查拼接报错的原因并提供解决方案。
报错核心是递归分割时出现维度不匹配的子矩阵:
- 将5x5矩阵填充为6x6后,第一次分割得到3x3的子矩阵(6//2=3)
- 递归处理3x3矩阵时,
split函数会将其拆分为1x1和2x2的混合维度子矩阵(3//2=1),比如a是1x1,b是1x2 - 后续计算如
a + b得到1x2矩阵,与h的2x2矩阵传入strassen后,递归分割会导致子矩阵维度混乱,最终触发vstack/hstack的维度匹配错误
本质是Strassen算法要求每一步分割后的子矩阵均为偶数维度,但填充为6x6后,递归到3x3时无法满足该要求,进而引发连锁维度问题。
1. 修改填充逻辑:填充到最小2的幂次维度
将矩阵填充到大于等于原维度的最小2的幂次(比如5x5填充到8x8,而非6x6),确保每一次递归分割都能得到偶数维度子矩阵:
# 替换main函数中的填充代码 if matrix_dimension % 2 != 0 or (matrix_dimension & (matrix_dimension - 1)) != 0: # 计算大于等于原维度的最小2的幂次 new_dim = 1 while new_dim < matrix_dimension: new_dim *= 2 pad_rows = new_dim - matrix_dimension pad_cols = new_dim - matrix_dimension matrixA = np.pad(matrixA, [(0, pad_rows), (0, pad_cols)], mode='constant', constant_values=0) matrixB = np.pad(matrixB, [(0, pad_rows), (0, pad_cols)], mode='constant', constant_values=0)
2. 修正base case判断逻辑
原代码用len(x) == 1判断1x1矩阵,但len(x)取的是numpy数组第一维度长度,1xN矩阵也会误触发。改为直接判断矩阵形状:
# 替换strassen函数中的base case if x.shape == (1, 1): return x * y
3. 裁剪结果回原维度
计算完成后,裁剪掉填充的部分,得到原维度的正确结果:
# 修改main函数中的输出代码 result = strassen(matrixA, matrixB) # 裁剪回原矩阵维度 original_dim = int(math.sqrt(len(result))) result = result[:original_dim, :original_dim] print(result) print("The sum is: ", result.sum())
内容的提问来源于stack exchange,提问作者ryan714

