矩阵乘法(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²) | 三层循环 |
| Strassen | O(n^2.807) | O(n²) | 分治优化 |
| Coppersmith-Winograd | O(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加速
🎨 应用场景
- 科学计算:大规模矩阵运算
- 图形学:3D变换矩阵
- 机器学习:神经网络权重计算
- 信号处理:卷积运算