矩阵乘法(Strassen算法)

📌 定义

矩阵乘法是线性代数中的基本运算。对于两个n×n矩阵A和B,传统方法计算A×B需要O(n³)时间,而Strassen算法通过分治思想将复杂度优化到O(n

示例: 2×2矩阵乘法
A = [1 2]    B = [5 6]
    [3 4]        [7 8]

A × B = [1×5+2×7  1×6+2×8] = [19 22]
        [3×5+4×7  3×6+4×8]   [43 50]

传统方法: 8次乘法
Strassen: 7次乘法

核心思路

将矩阵分成4个子矩阵,使用分治法减少乘法次数:

对于两个n×n矩阵 A 和 B (n为2的幂):

分块:
A = [A11 A12]    B = [B11 B12]
    [A21 A22]        [B21 B22]

传统方法:
C = A × B = [C11 C12]
            [C21 C22]

其中:
C11 = A11×B11 + A12×B21  (2次乘法)
C12 = A11×B12 + A12×B22  (2次乘法)
C21 = A21×B11 + A22×B21  (2次乘法)
C22 = A21×B12 + A22×B22  (2次乘法)
总共: 8次乘法

Strassen方法:
计算7个辅助矩阵:
M1 = (A11 + A22) × (B11 + B22)
M2 = (A21 + A22) × B11
M3 = A11 × (B12 - B22)
M4 = A22 × (B21 - B11)
M5 = (A11 + A12) × B22
M6 = (A21 - A11) × (B11 + B12)
M7 = (A12 - A22) × (B21 + B22)

组合:
C11 = M1 + M4 - M5 + M7
C12 = M3 + M5
C21 = M2 + M4
C22 = M1 - M2 + M3 + M6

只需7次乘法!

复杂度分析

方法时间复杂度空间复杂度说明
朴素算法O(n³)O(n²)三层循环
StrassenO(n^2.807)O(n²)分治优化
Coppersmith-WinogradO(n^2.376)O(n²)理论最优
当前最优O(n^2.3728596)O(n²)2020年

其中 2.807 ≈ log₂(7)

Go 代码

Go 实现

package main
 
import "fmt"
 
type Matrix [][]int
 
func strassen(A, B Matrix) Matrix {
    n := len(A)
 
    // 基准情况
    if n <= 64 {
        return naiveMultiply(A, B)
    }
 
    mid := n / 2
 
    // 分块
    A11, A12, A21, A22 := split(A)
    B11, B12, B21, B22 := split(B)
 
    // 7个辅助矩阵
    M1 := strassen(add(A11, A22), add(B11, B22))
    M2 := strassen(add(A21, A22), B11)
    M3 := strassen(A11, sub(B12, B22))
    M4 := strassen(A22, sub(B21, B11))
    M5 := strassen(add(A11, A12), B22)
    M6 := strassen(sub(A21, A11), add(B11, B12))
    M7 := strassen(sub(A12, A22), add(B21, B22))
 
    // 组合
    C11 := add(sub(add(M1, M4), M5), M7)
    C12 := add(M3, M5)
    C21 := add(M2, M4)
    C22 := add(sub(add(M1, M2), M3), M6)
 
    return merge(C11, C12, C21, C22)
}
 
func naiveMultiply(A, B Matrix) Matrix {
    n := len(A)
    C := make(Matrix, n)
    for i := range C {
        C[i] = make([]int, n)
    }
 
    for i := 0; i < n; i++ {
        for j := 0; j < n; j++ {
            for k := 0; k < n; k++ {
                C[i][j] += A[i][k] * B[k][j]
            }
        }
    }
    return C
}
 
func add(A, B Matrix) Matrix {
    n := len(A)
    C := make(Matrix, n)
    for i := 0; i < n; i++ {
        C[i] = make([]int, n)
        for j := 0; j < n; j++ {
            C[i][j] = A[i][j] + B[i][j]
        }
    }
    return C
}
 
func sub(A, B Matrix) Matrix {
    n := len(A)
    C := make(Matrix, n)
    for i := 0; i < n; i++ {
        C[i] = make([]int, n)
        for j := 0; j < n; j++ {
            C[i][j] = A[i][j] - B[i][j]
        }
    }
    return C
}
 
func split(M Matrix) (Matrix, Matrix, Matrix, Matrix) {
    n := len(M)
    mid := n / 2
 
    M11 := make(Matrix, mid)
    M12 := make(Matrix, mid)
    M21 := make(Matrix, mid)
    M22 := make(Matrix, mid)
 
    for i := 0; i < mid; i++ {
        M11[i] = M[i][:mid]
        M12[i] = M[i][mid:]
        M21[i] = M[mid+i][:mid]
        M22[i] = M[mid+i][mid:]
    }
 
    return M11, M12, M21, M22
}
 
func merge(C11, C12, C21, C22 Matrix) Matrix {
    mid := len(C11)
    n := 2 * mid
    C := make(Matrix, n)
 
    for i := 0; i < n; i++ {
        C[i] = make([]int, n)
    }
 
    for i := 0; i < mid; i++ {
        for j := 0; j < mid; j++ {
            C[i][j] = C11[i][j]
            C[i][j+mid] = C12[i][j]
            C[i+mid][j] = C21[i][j]
            C[i+mid][j+mid] = C22[i][j]
        }
    }
 
    return C
}
 
func main() {
    A := Matrix{{1, 2}, {3, 4}}
    B := Matrix{{5, 6}, {7, 8}}
 
    C := strassen(A, B)
 
    fmt.Println("A × B (Strassen):")
    for _, row := range C {
        fmt.Println(row)
    }
}

思路展开

复杂度推导

T(n) = 递归求解n×n矩阵乘法的时间

朴素算法:
T(n) = O(n³)

分块朴素算法:
T(n) = 8T(n/2) + O(n²)  (8次乘法,O(n²)加法)
     = O(n³)  (无改进)

Strassen算法:
T(n) = 7T(n/2) + O(n²)  (7次乘法,O(n²)加法)

根据主定理:
a = 7, b = 2, f(n) = O(n²)
log_b(a) = log_2(7) ≈ 2.807

因为 f(n) = O(n²) < O(n^2.807)
所以 T(n) = O(n^log_2(7)) = O(n^2.807)

经典题目

应用场景

  • 科学计算
  • 图形学变换
  • 深度学习(矩阵运算)
  • 数值分析

相关算法

  • Coppersmith-Winograd算法
  • 矩阵链乘法(动态规划)
  • 稀疏矩阵乘法

⚖️ 优缺点

优点

  • ✅ 渐近最优:O(n^2.807) vs O(n³)
  • ✅ 分治经典:展示分治思想的威力
  • ✅ 可并行:子问题可以并行计算

缺点

  • ❌ 常数因子大:小矩阵时不如朴素算法
  • ❌ 数值稳定性:减法操作可能导致舍入误差
  • ❌ 实际应用少:现代库使用优化的朴素算法或GPU加速

🎨 应用场景

  1. 科学计算:大规模矩阵运算
  2. 图形学:3D变换矩阵
  3. 机器学习:神经网络权重计算
  4. 信号处理:卷积运算

相关主题


返回:分治算法 | 算法学习导航