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

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()

请求排查拼接报错的原因并提供解决方案。


问题原因分析

报错核心是递归分割时出现维度不匹配的子矩阵:

  1. 将5x5矩阵填充为6x6后,第一次分割得到3x3的子矩阵(6//2=3)
  2. 递归处理3x3矩阵时,split函数会将其拆分为1x1和2x2的混合维度子矩阵(3//2=1),比如a是1x1,b是1x2
  3. 后续计算如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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 05:48:09