Strassen算法与普通矩阵乘法:C++实现与性能对比分析

发布时间:2026/7/26 4:37:07
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::vectorstd::vectordouble data; // 构造函数 Matrix(int r, int c, double initVal 0.0) : rows(r), cols(c), data(r, std::vectordouble(c, initVal)) {} // 拷贝构造函数 Matrix(const Matrix other) : rows(other.rows), cols(other.cols), data(other.data) {} // 从向量构造方便测试 Matrix(const std::vectorstd::vectordouble 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_caststd::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_caststd::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::vectorint 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~4500.27xStrassen 慢递归开销主导64~900~11000.82x接近阈值Strassen 仍稍慢128~7000~55001.27xStrassen 开始显现优势256~56000~380001.47x优势扩大512~450000~2650001.70x优势明显结果分析小矩阵n 64普通乘法完胜。Strassen 算法的递归调用、大量的矩阵加法和内存分配/拷贝开销完全抵消了减少一次乘法带来的理论收益。这就是设置递归基的重要性。中等矩阵n ≈ 128Strassen 算法开始反超。当矩阵规模足够大使得减少的乘法递归成本高于额外的加法和管理开销时理论上的复杂度优势转化为实际的性能优势。大矩阵n 256Strassen 算法的优势变得显著且稳定。随着 n 增大O(n^2.81) 和 O(n³) 的差距在绝对计算时间上体现得越来越明显。重要提示这个对比是基于我们实现的、未深度优化的版本。工业级的高性能线性代数库如 OpenBLAS中的普通矩阵乘法通过使用汇编级别优化、循环分块、多线程和 SIMD其性能可以达到我们朴素实现的数十倍甚至上百倍。因此我们的 Strassen 实现要超越高度优化的普通乘法需要的矩阵规模阈值会大得多可能要到 n1000 甚至更大。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更多的是理解其分治思想和复杂度分析的方法这是算法工程师内功的重要组成部分。在面试中能够清晰阐述其原理、实现以及优缺点远比死记硬背代码更有价值。