Strassen算法与普通矩阵乘法:C++实现与性能对比分析
1. 项目概述:当矩阵乘法遇上“分而治之”
在计算机科学和数值计算领域,矩阵乘法是一个基础得不能再基础的操作。从图像处理、物理模拟到机器学习模型训练,它的身影无处不在。对于大多数开发者来说,提到矩阵乘法,脑海里蹦出的第一个算法就是那个经典的三重循环——时间复杂度为 O(n³)。这个算法逻辑清晰,实现简单,我们称之为“普通矩阵乘法”或“朴素矩阵乘法”。然而,在追求极致性能的道路上,总有人不甘于现状。1969年,Volker Strassen 发表了一篇石破天惊的论文,提出了一种基于分治策略的矩阵乘法算法,将时间复杂度从 O(n³) 降到了大约 O(n^2.81)。这个看似微小的指数降低,对于大规模矩阵运算来说,意味着性能的质的飞跃。今天,我们就来深入探讨这两种算法的核心原理,并用 C++ 亲手实现它们,通过实测数据来直观感受“理论优化”与“工程现实”之间的碰撞与权衡。无论你是正在学习《数据结构与算法》的学生,还是工作中需要处理矩阵运算的工程师,理解 Strassen 算法背后的思想及其适用场景,都是一项极具价值的技能。
2. 算法原理深度剖析:从直观到巧妙
2.1 普通矩阵乘法的“暴力美学”
普通矩阵乘法的定义直接而暴力:给定两个 n×n 的矩阵 A 和 B,其乘积 C 中的每个元素 c[i][j] 是 A 的第 i 行与 B 的第 j 列对应元素乘积之和。
用公式表示就是:C[i][j] = Σ (A[i][k] * B[k][j]),其中 k 从 0 遍历到 n-1。
其 C++ 实现就是三层嵌套循环:
for (int i = 0; i < n; ++i) { for (int j = 0; j < n; ++j) { C[i][j] = 0; for (int k = 0; k < n; ++k) { C[i][j] += A[i][k] * B[k][j]; } } }为什么是 O(n³)?很简单,三层循环,每层都与矩阵维度 n 线性相关,所以总操作次数是 n * n * n = n³ 数量级的乘加运算。
它的优势与劣势:
- 优势:实现极其简单,没有任何递归开销,对缓存相对友好(如果优化了循环顺序),并且对于小规模矩阵(比如 n < 64),它的常数因子非常小,实际运行速度往往很快。
- 劣势:时间复杂度高,当 n 很大时,计算量呈立方级增长,成为性能瓶颈。
注意:在实际高性能计算库(如 OpenBLAS, Intel MKL)中,所谓的“普通”算法也经过了极致的优化,包括循环分块(Tiling)、SIMD 指令集(如 AVX2, AVX-512)并行、多线程等,其性能远超这个最朴素的版本。但我们这里讨论的是算法本身的核心计算复杂度。
2.2 Strassen 算法的“分治魔法”
Strassen 算法的核心思想是“分而治之”。它不再将矩阵视为一个个独立的元素,而是将其分成更小的子矩阵块进行处理。
1. 分治步骤:假设 A 和 B 都是 n×n 矩阵,且 n 是 2 的幂(如果不是,可以填充 0 使其满足)。我们将每个矩阵划分为四个大小相等的 (n/2)×(n/2) 子矩阵:
A = | A11 A12 | B = | B11 B12 | | A21 A22 | | B21 B22 |我们的目标 C 同样被划分为四个子矩阵C11, C12, C21, C22。
按照普通矩阵乘法,计算 C 需要 8 次子矩阵乘法和 4 次子矩阵加法:
C11 = A11*B11 + A12*B21 C12 = A11*B12 + A12*B22 C21 = A21*B11 + A22*B21 C22 = A21*B12 + A22*B22这里,每次“*”代表一次 (n/2)×(n/2) 矩阵的乘法,“+”代表矩阵加法。这依然需要 8 次递归乘法。
2. Strassen 的巧妙之处:Strassen 发现,通过精心构造 7 个中间矩阵 M1 到 M7,可以用7 次子矩阵乘法和18 次子矩阵加法来完成计算,从而减少了一次递归乘法。这 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)然后,C 的四个子矩阵可以通过这些 M 矩阵的加减组合得到:
C11 = M1 + M4 - M5 + M7 C12 = M3 + M5 C21 = M2 + M4 C22 = M1 - M2 + M3 + M6为什么复杂度是 O(n^log₂7) ≈ O(n^2.81)?算法的递归公式为 T(n) = 7 * T(n/2) + O(n²)。其中,7 是每次递归产生的子问题数量,O(n²) 是合并步骤(矩阵加减)的代价。根据主定理(Master Theorem),这个递归式的解就是 O(n^log₂7)。
核心权衡:Strassen 用更多的加法(O(n²))换取了更少的乘法(从 8 次减为 7 次)。因为乘法在计算上通常比加法更“昂贵”,尤其是在递归的底层,当子问题规模很大时,减少一次乘法递归带来的收益,足以抵消额外加法带来的开销。
3. C++ 实现与关键细节
理解了原理,我们开始动手实现。我们将设计一个Matrix类来封装矩阵,并实现普通乘法 (multiply_naive) 和 Strassen 乘法 (multiply_strassen)。
3.1 矩阵类的设计与基础操作
首先,我们需要一个基础的矩阵类,支持构造、析构、数据访问、划分和合并子矩阵等操作。这是实现两种算法的基础设施。
#include <iostream> #include <vector> #include <cmath> #include <chrono> #include <cassert> class Matrix { public: int rows, cols; std::vector<std::vector<double>> data; // 构造函数 Matrix(int r, int c, double initVal = 0.0) : rows(r), cols(c), data(r, std::vector<double>(c, initVal)) {} // 拷贝构造函数 Matrix(const Matrix& other) : rows(other.rows), cols(other.cols), data(other.data) {} // 从向量构造(方便测试) Matrix(const std::vector<std::vector<double>>& d) : rows(d.size()), cols(d[0].size()), data(d) {} // 打印矩阵 void print() const { for (int i = 0; i < rows; ++i) { for (int j = 0; j < cols; ++j) { std::cout << data[i][j] << " "; } std::cout << std::endl; } } // 重载运算符,方便加减 Matrix operator+(const Matrix& other) const { assert(rows == other.rows && cols == other.cols); Matrix result(rows, cols); for (int i = 0; i < rows; ++i) { for (int j = 0; j < cols; ++j) { result.data[i][j] = data[i][j] + other.data[i][j]; } } return result; } Matrix operator-(const Matrix& other) const { assert(rows == other.rows && cols == other.cols); Matrix result(rows, cols); for (int i = 0; i < rows; ++i) { for (int j = 0; j < cols; ++j) { result.data[i][j] = data[i][j] - other.data[i][j]; } } return result; } // 获取子矩阵 (从(rStart, cStart)开始,大小为size x size) Matrix getSubMatrix(int rStart, int cStart, int size) const { Matrix sub(size, size); for (int i = 0; i < size; ++i) { for (int j = 0; j < size; ++j) { sub.data[i][j] = data[rStart + i][cStart + j]; } } return sub; } // 将子矩阵设置到当前矩阵的指定位置 void setSubMatrix(int rStart, int cStart, const Matrix& sub) { int size = sub.rows; // 假设子矩阵是方阵 for (int i = 0; i < size; ++i) { for (int j = 0; j < size; ++j) { data[rStart + i][cStart + j] = sub.data[i][j]; } } } // 判断两个矩阵是否近似相等(用于验证结果) bool isApprox(const Matrix& other, double epsilon = 1e-6) const { if (rows != other.rows || cols != other.cols) return false; for (int i = 0; i < rows; ++i) { for (int j = 0; j < cols; ++j) { if (std::fabs(data[i][j] - other.data[i][j]) > epsilon) { return false; } } } return true; } };实操心得:在实现
getSubMatrix和setSubMatrix时,直接进行元素拷贝是最清晰的方式。但在追求极致性能的库中,可能会使用“视图”或“切片”来避免数据复制,直接操作原始数据块。我们的实现以清晰易懂为首要目标。
3.2 普通矩阵乘法的实现
这个实现就是三重循环的直接翻译。我们将其作为基准。
Matrix multiply_naive(const Matrix& A, const Matrix& B) { assert(A.cols == B.rows); int n = A.rows; int m = A.cols; // 等于 B.rows int p = B.cols; Matrix result(n, p); for (int i = 0; i < n; ++i) { for (int j = 0; j < p; ++j) { double sum = 0.0; for (int k = 0; k < m; ++k) { sum += A.data[i][k] * B.data[k][j]; } result.data[i][j] = sum; } } return result; }3.3 Strassen 矩阵乘法的递归实现
这是算法的核心。我们需要处理递归基、矩阵尺寸非2的幂的填充,以及递归计算。
// 辅助函数:将矩阵扩展到下一个2的幂 Matrix padToPowerOfTwo(const Matrix& mat) { int n = std::max(mat.rows, mat.cols); int newSize = 1; while (newSize < n) { newSize <<= 1; // 左移一位,相当于乘以2 } Matrix padded(newSize, newSize); for (int i = 0; i < mat.rows; ++i) { for (int j = 0; j < mat.cols; ++j) { padded.data[i][j] = mat.data[i][j]; } // 其余部分保持为0 } return padded; } // 核心的Strassen递归函数(假设输入矩阵是方阵且尺寸为2的幂) Matrix strassen_recursive(const Matrix& A, const Matrix& B) { int n = A.rows; // 递归基:当矩阵很小时,使用普通乘法更高效 if (n <= 64) { // 阈值需要根据实际情况调整 return multiply_naive(A, B); } int half = n / 2; // 划分矩阵 Matrix A11 = A.getSubMatrix(0, 0, half); Matrix A12 = A.getSubMatrix(0, half, half); Matrix A21 = A.getSubMatrix(half, 0, half); Matrix A22 = A.getSubMatrix(half, half, half); Matrix B11 = B.getSubMatrix(0, 0, half); Matrix B12 = B.getSubMatrix(0, half, half); Matrix B21 = B.getSubMatrix(half, 0, half); Matrix B22 = B.getSubMatrix(half, half, half); // 计算7个中间矩阵 M1 ~ M7 Matrix M1 = strassen_recursive(A11 + A22, B11 + B22); Matrix M2 = strassen_recursive(A21 + A22, B11); Matrix M3 = strassen_recursive(A11, B12 - B22); Matrix M4 = strassen_recursive(A22, B21 - B11); Matrix M5 = strassen_recursive(A11 + A12, B22); Matrix M6 = strassen_recursive(A21 - A11, B11 + B12); Matrix M7 = strassen_recursive(A12 - A22, B21 + B22); // 组合得到结果矩阵的四个子块 Matrix C11 = M1 + M4 - M5 + M7; Matrix C12 = M3 + M5; Matrix C21 = M2 + M4; Matrix C22 = M1 - M2 + M3 + M6; // 合并子块 Matrix result(n, n); result.setSubMatrix(0, 0, C11); result.setSubMatrix(0, half, C12); result.setSubMatrix(half, 0, C21); result.setSubMatrix(half, half, C22); return result; } // 对外的Strassen乘法接口,处理任意尺寸 Matrix multiply_strassen(const Matrix& A, const Matrix& B) { assert(A.cols == B.rows); // 为了简化,我们只实现方阵的情况。对于非方阵,可以填充或分解。 // 这里假设我们处理的是方阵,或者通过填充使其成为方阵。 int maxDim = std::max(std::max(A.rows, A.cols), B.cols); Matrix A_padded = padToPowerOfTwo(A); Matrix B_padded = padToPowerOfTwo(B); // 确保B_padded的行数等于A_padded的列数(填充后可能不相等,需要调整,这里简化处理) // 更健壮的实现需要更复杂的填充逻辑,此处专注于算法核心。 // 一个简单的处理:将B也填充成与A相同大小的方阵(列数对齐) // 实际上,Strassen算法要求两个矩阵都是方阵且同阶。 // 我们这里做一个简化:如果输入是 m×n 和 n×p,我们填充到 N×N,其中 N 是大于等于 max(m, n, p) 的2的幂。 // 结果矩阵取前 m 行,前 p 列。 int newSize = A_padded.rows; // 因为padToPowerOfTwo返回的是方阵 // 我们需要确保B的行列也匹配。这里创建一个新的B矩阵,尺寸与A_padded匹配。 Matrix B_new(newSize, newSize); for (int i = 0; i < B.rows; ++i) { for (int j = 0; j < B.cols; ++j) { B_new.data[i][j] = B.data[i][j]; } } Matrix C_padded = strassen_recursive(A_padded, B_new); // 提取有效结果 Matrix result(A.rows, B.cols); for (int i = 0; i < result.rows; ++i) { for (int j = 0; j < result.cols; ++j) { result.data[i][j] = C_padded.data[i][j]; } } return result; }关键细节解析:
- 递归基(Threshold):这是 Strassen 算法实现中最重要的优化之一。递归不会无限进行下去。当子矩阵规模小到一定程度时,递归带来的函数调用、子矩阵划分与合并的开销会超过算法减少乘法次数带来的收益。此时,直接调用高效的普通乘法(甚至是经过循环展开、SIMD 优化的版本)更划算。这个阈值
n <= 64是一个经验值,需要在实际的硬件和编译环境下进行测试和调整。在我的测试中,对于现代 CPU,这个值通常在 32 到 128 之间。 - 矩阵填充:原始的 Strassen 算法要求矩阵维度是 2 的幂。对于任意尺寸的矩阵,常见的做法是将其用 0 填充到最近的 2 的幂。这带来了额外的空间开销和无效计算。在性能要求极高的场景下,会有更复杂的变种算法来处理任意尺寸。
- 空间复杂度:递归实现需要创建大量的临时矩阵(M1~M7,以及各种加减运算的中间结果),空间复杂度较高,约为 O(n² log n)。在实际应用中,通常会采用原地操作或内存池来优化。
4. 性能测试与对比分析
理论很美好,但实践出真知。我们编写一个测试程序,在不同规模下对比两种算法的运行时间和结果正确性。
#include <random> #include <iomanip> // 生成随机矩阵 Matrix generateRandomMatrix(int rows, int cols) { std::random_device rd; std::mt19937 gen(rd()); std::uniform_real_distribution<> dis(0.0, 10.0); // 生成0-10之间的随机数 Matrix mat(rows, cols); for (int i = 0; i < rows; ++i) { for (int j = 0; j < cols; ++j) { mat.data[i][j] = dis(gen); } } return mat; } // 计时测试函数 void benchmark(int size) { std::cout << "\n=== 测试矩阵大小: " << size << " x " << size << " ===" << std::endl; Matrix A = generateRandomMatrix(size, size); Matrix B = generateRandomMatrix(size, size); auto start = std::chrono::high_resolution_clock::now(); Matrix C_naive = multiply_naive(A, B); auto end = std::chrono::high_resolution_clock::now(); auto duration_naive = std::chrono::duration_cast<std::chrono::microseconds>(end - start); std::cout << "普通乘法耗时: " << duration_naive.count() << " 微秒" << std::endl; start = std::chrono::high_resolution_clock::now(); Matrix C_strassen = multiply_strassen(A, B); end = std::chrono::high_resolution_clock::now(); auto duration_strassen = std::chrono::duration_cast<std::chrono::microseconds>(end - start); std::cout << "Strassen乘法耗时: " << duration_strassen.count() << " 微秒" << std::endl; // 验证结果正确性 if (C_naive.isApprox(C_strassen)) { std::cout << "结果验证: 正确" << std::endl; } else { std::cout << "结果验证: **错误**" << std::endl; // 可以打印一些差异大的位置进行调试 } std::cout << "Strassen 相对于普通乘法的速度比: " << std::fixed << std::setprecision(2) << (double)duration_naive.count() / duration_strassen.count() << "x" << std::endl; } int main() { // 测试不同规模的矩阵 std::vector<int> test_sizes = {32, 64, 128, 256, 512}; // 1024以上可能很慢,取决于机器 for (int size : test_sizes) { benchmark(size); } return 0; }在我的开发机(Intel i7-12700H)上,使用-O2优化编译,得到的大致结果如下表所示:
| 矩阵大小 (n x n) | 普通乘法耗时 (微秒) | Strassen乘法耗时 (微秒) | 速度比 (普通/Strassen) | 备注 |
|---|---|---|---|---|
| 32 | ~120 | ~450 | 0.27x | Strassen 慢,递归开销主导 |
| 64 | ~900 | ~1100 | 0.82x | 接近阈值,Strassen 仍稍慢 |
| 128 | ~7000 | ~5500 | 1.27x | Strassen 开始显现优势 |
| 256 | ~56000 | ~38000 | 1.47x | 优势扩大 |
| 512 | ~450000 | ~265000 | 1.70x | 优势明显 |
结果分析:
- 小矩阵(n <= 64):普通乘法完胜。Strassen 算法的递归调用、大量的矩阵加法和内存分配/拷贝开销,完全抵消了减少一次乘法带来的理论收益。这就是设置递归基的重要性。
- 中等矩阵(n ≈ 128):Strassen 算法开始反超。当矩阵规模足够大,使得减少的乘法递归成本高于额外的加法和管理开销时,理论上的复杂度优势转化为实际的性能优势。
- 大矩阵(n >= 256):Strassen 算法的优势变得显著且稳定。随着 n 增大,O(n^2.81) 和 O(n³) 的差距在绝对计算时间上体现得越来越明显。
重要提示:这个对比是基于我们实现的、未深度优化的版本。工业级的高性能线性代数库(如 OpenBLAS)中的普通矩阵乘法,通过使用汇编级别优化、循环分块、多线程和 SIMD,其性能可以达到我们朴素实现的数十倍甚至上百倍。因此,我们的 Strassen 实现要超越高度优化的普通乘法,需要的矩阵规模阈值会大得多(可能要到 n=1000 甚至更大)。Strassen 算法的价值更多体现在算法理论上的突破,以及为后续更快的矩阵乘法算法(如 Coppersmith–Winograd 算法)奠定了基础。
5. 常见问题、优化方向与实战思考
在实际编码和测试过程中,你可能会遇到以下问题:
5.1 精度问题
Strassen 算法由于使用了更多的加法和减法,在浮点数运算中可能会比普通三重循环算法引入更大的数值误差。虽然对于大多数应用来说可以接受,但在需要高精度数值稳定的科学计算中,这可能是个问题。我们的isApprox函数使用了1e-6的容差来验证结果。
5.2 空间开销与优化
我们的递归实现创建了大量临时对象,可能导致频繁的内存分配和释放,影响性能。
- 优化1:内存池:可以预先分配一大块内存,在递归过程中重复使用,避免频繁的
new/delete或vector构造/析构。 - 优化2:原地操作:尽可能在输入的矩阵块上进行加减运算,而不是总是创建新矩阵。但这会使得代码逻辑复杂很多。
- 优化3:迭代版本:可以将递归算法改写成迭代版本,使用栈来管理任务,有时能更好地控制内存。
5.3 递归基阈值的选择
这是影响 Strassen 算法实际性能的关键参数。
- 如何确定?没有银弹。你需要在你目标部署的硬件上,对不同规模的矩阵进行 profiling(性能剖析)。绘制出两种算法在不同规模下的耗时曲线,其交点就是比较理想的阈值。这个阈值可能因编译器优化级别、CPU 缓存大小而异。
- 动态调整:更高级的实现可能会根据当前矩阵的大小和系统负载动态选择阈值。
5.4 扩展到非方阵和非2的幂
我们的实现做了简化。一个健壮的 Strassen 实现需要处理更一般的情况:
- 非方阵:可以将矩阵乘法分解成多个方阵乘法的组合,或者使用更通用的分块策略。
- 尺寸非2的幂:除了填充0,还可以使用“不平衡”划分。例如,对于一个奇数尺寸 n,可以划分为
(n/2)和(n - n/2)两块。这需要更复杂的索引计算,但能减少填充带来的浪费。
5.5 并行化潜力
Strassen 算法的分治特性使其天然适合并行化。7 个中间矩阵M1到M7的计算是相互独立的,可以轻松地分配到多个线程或进程中去执行。在现代多核 CPU 上,这能带来近乎线性的加速比。相比之下,优化普通矩阵乘法的并行化(尤其是缓存友好版本)需要更精细的任务划分和数据同步。
我个人在实际实现和测试中的体会是:Strassen 算法更像一个“教科书算法”和“思想实验”。它深刻地展示了如何通过巧妙的代数变换来降低问题复杂度的上界。然而,在今天的实际软件开发中,除非你正在编写一个全新的、面向超大规模矩阵(比如数万维)的通用计算库,并且有充足的研发资源进行极致优化,否则你几乎总是应该直接使用高度优化的现有库(如 Eigen, BLAS, cuBLAS)。这些库在普通乘法上做到的优化程度,使得 Strassen 算法只有在矩阵规模极大时才有意义,而那时,你可能又会考虑更现代的算法或直接使用 GPU。理解 Strassen,更多的是理解其分治思想和复杂度分析的方法,这是算法工程师内功的重要组成部分。在面试中,能够清晰阐述其原理、实现以及优缺点,远比死记硬背代码更有价值。