Strassen矩阵乘法Python实现报错:矩阵尺寸不匹配与递归超限
Strassen矩阵乘法错误排查与修复
问题根源分析
你遇到的两个错误是连锁问题:
- 递归深度超限:核心原因是
strassen函数缺失递归终止条件,导致矩阵被无限分块,直到触发递归栈上限。即使调大递归限制,也会因为无限递归无法终止。 - 矩阵尺寸不匹配:分块函数
divide实现错误,导致生成的子矩阵尺寸不一致,或者递归到极小矩阵时(如1x1)继续分块得到空矩阵,触发辅助函数的尺寸检查异常。
具体修复步骤
1. 补充递归终止条件
Strassen算法的终止条件是当矩阵为1x1时,直接返回两元素的乘积,无需继续分块。在strassen函数最开头添加:
def strassen(a, b): # 递归终止条件:1x1矩阵直接相乘 if len(a) == 1: return [[a[0][0] * b[0][0]]] # 后续分块、递归逻辑...
2. 修复分块函数divide
确保分块后生成的四个子矩阵都是n/2 x n/2的尺寸,避免空矩阵或尺寸不一致:
def divide(matrix): n = len(matrix) mid = n // 2 # 按行和列的中点分割矩阵 quad1 = [row[:mid] for row in matrix[:mid]] # 左上 quad2 = [row[mid:] for row in matrix[:mid]] # 右上 quad3 = [row[:mid] for row in matrix[mid:]] # 左下 quad4 = [row[mid:] for row in matrix[mid:]] # 右下 return quad1, quad2, quad3, quad4
3. 确保辅助函数的尺寸检查严谨
matrix_addition和matrix_subtraction需要同时检查行数和列数是否一致:
def matrix_addition(a, b): if len(a) != len(b) or any(len(row_a) != len(row_b) for row_a, row_b in zip(a, b)): raise ValueError("Matrices must have the same size") return [[a[i][j] + b[i][j] for j in range(len(a[0]))] for i in range(len(a))] def matrix_subtraction(a, b): if len(a) != len(b) or any(len(row_a) != len(row_b) for row_a, row_b in zip(a, b)): raise ValueError("Matrices must have the same size") return [[a[i][j] - b[i][j] for j in range(len(a[0]))] for i in range(len(a))]
完整修复代码示例
def divide(matrix): n = len(matrix) mid = n // 2 quad1 = [row[:mid] for row in matrix[:mid]] quad2 = [row[mid:] for row in matrix[:mid]] quad3 = [row[:mid] for row in matrix[mid:]] quad4 = [row[mid:] for row in matrix[mid:]] return quad1, quad2, quad3, quad4 def matrix_addition(a, b): if len(a) != len(b) or any(len(row_a) != len(row_b) for row_a, row_b in zip(a, b)): raise ValueError("Matrices must have the same size") return [[a[i][j] + b[i][j] for j in range(len(a[0]))] for i in range(len(a))] def matrix_subtraction(a, b): if len(a) != len(b) or any(len(row_a) != len(row_b) for row_a, row_b in zip(a, b)): raise ValueError("Matrices must have the same size") return [[a[i][j] - b[i][j] for j in range(len(a[0]))] for i in range(len(a))] def strassen(a, b): # 递归终止条件 if len(a) == 1: return [[a[0][0] * b[0][0]]] # 分块 quad1_a, quad2_a, quad3_a, quad4_a = divide(a) quad1_b, quad2_b, quad3_b, quad4_b = divide(b) # 计算7个中间矩阵 p1 = strassen(matrix_addition(quad1_a, quad4_a), matrix_addition(quad1_b, quad4_b)) p2 = strassen(matrix_addition(quad3_a, quad4_a), quad1_b) p3 = strassen(quad1_a, matrix_subtraction(quad2_b, quad4_b)) p4 = strassen(quad4_a, matrix_subtraction(quad3_b, quad1_b)) p5 = strassen(matrix_addition(quad1_a, quad2_a), quad4_b) p6 = strassen(matrix_subtraction(quad3_a, quad1_a), matrix_addition(quad1_b, quad2_b)) p7 = strassen(matrix_subtraction(quad2_a, quad4_a), matrix_addition(quad3_b, quad4_b)) # 计算结果矩阵的四个分块 quad1_c = matrix_addition(matrix_subtraction(matrix_addition(p1, p4), p5), p7) quad2_c = matrix_addition(p3, p5) quad3_c = matrix_addition(p2, p4) quad4_c = matrix_addition(matrix_subtraction(matrix_addition(p1, p3), p2), p6) # 合并分块 top = [row1 + row2 for row1, row2 in zip(quad1_c, quad2_c)] bottom = [row1 + row2 for row1, row2 in zip(quad3_c, quad4_c)] return top + bottom # 测试案例 A = [[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12], [13, 14, 15, 16]] B = [[17, 18, 19, 20], [21, 22, 23, 24], [25, 26, 27, 28], [29, 30, 31, 32]] result = strassen(A, B) for row in result: print(row)
运行结果
执行后会输出预期的结果:
[250, 260, 270, 280] [618, 644, 670, 696] [986, 1028, 1070, 1112] [1354, 1412, 1470, 1528]
内容的提问来源于stack exchange,提问作者death_templar
相关产品推荐
相关产品推荐

