You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Strassen矩阵乘法Python实现报错:矩阵尺寸不匹配与递归超限

Strassen矩阵乘法错误排查与修复

问题根源分析

你遇到的两个错误是连锁问题:

  1. 递归深度超限:核心原因是strassen函数缺失递归终止条件,导致矩阵被无限分块,直到触发递归栈上限。即使调大递归限制,也会因为无限递归无法终止。
  2. 矩阵尺寸不匹配:分块函数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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 03:38:10