多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

BLAS矩阵乘法优化:从分块到选型,性能提升30倍

BLAS矩阵乘法优化:从分块到选型,性能提升30倍 1. 从一次矩阵乘法卡顿说起BLAS到底在算什么很多人第一次接触BLAS是在某个程序跑得特别慢的时候。比如你写了一段Python做矩阵乘法数据量一上来循环套循环跑个几分钟都出不来结果。换成NumPy的np.dot同样的数据量零点几秒就完事了。这中间的差距很大程度上就是BLAS在起作用。BLAS的全称是Basic Linear Algebra Subprograms翻译过来叫“基础线性代数子程序”。名字听起来很学术但你可以把它理解成一套专门做向量和矩阵运算的“标准工具箱”。它定义了一组接口规范把线性代数里最常用的操作——向量加法、点积、矩阵乘向量、矩阵乘矩阵——都封装成了标准函数。任何语言、任何平台只要按照这套规范去调用就能获得经过高度优化的计算性能。那为什么需要这么一个标准因为线性代数运算在科学计算、机器学习、图形渲染、信号处理这些领域里出现得太频繁了。如果每个项目都自己手写矩阵乘法不仅开发效率低性能也很难做到极致。BLAS的出现相当于把“矩阵运算”这件事从应用层剥离出来交给专门的底层库去处理。应用层只管调用接口底层库负责把CPU的缓存、指令集、多核并行这些硬件特性榨干。BLAS本身只是一个接口标准真正干活的是它的各种实现。常见的有开源的OpenBLAS、Intel的MKL、AMD的BLIS等等。不同实现针对不同的CPU架构做了不同的优化性能差异可能达到几倍甚至十几倍。这也是为什么同一个NumPy程序在不同机器上跑出来的速度完全不一样——它背后链接的BLAS实现不同。BLAS把运算分成了三个层级这个分层逻辑非常关键理解了它你就理解了BLAS的设计哲学。Level 1是向量与向量之间的运算比如两个向量相加、做点积、求范数。这类操作的特点是数据访问量小计算量也小瓶颈通常在内存带宽上。你读一堆数据进来做一次简单的加减乘除然后就结束了。内存读写的速度决定了整体性能。Level 2是矩阵与向量之间的运算典型的就是矩阵乘向量。这类操作的计算量比Level 1大了一个量级但数据复用率仍然不够高。矩阵的每一行都要和向量做一次运算矩阵被完整读取一遍但向量可能被反复读取。性能瓶颈介于内存和计算之间。Level 3是矩阵与矩阵之间的运算代表就是矩阵乘法。这是BLAS里最核心、优化最狠的部分。矩阵乘法有一个天然优势数据复用率极高。两个矩阵相乘时每个元素都会被多次使用这就给缓存优化留下了巨大的空间。BLAS实现里最复杂的代码、最精巧的分块策略几乎都集中在Level 3上。这个分层不是随便分的它直接对应了不同的优化策略。Level 1优化内存访问模式Level 2想办法提高数据复用Level 3则要综合考虑缓存分块、寄存器分配、指令级并行。你在调用BLAS的时候知道自己用的是哪个层级的函数就能大致判断性能瓶颈在哪里。2. 矩阵乘法的性能密码为什么分块能快十倍矩阵乘法是BLAS里最值得深挖的部分。表面上看C A × B就是三重循环写起来不到十行代码。但为什么手写的三重循环慢得离谱而BLAS实现能快几十倍答案藏在计算机的存储层次结构里。现代CPU的存储层次大致是这样的寄存器最快但容量极小L1缓存次之然后是L2、L3最后是主内存。主内存的访问延迟可能是L1缓存的几十倍甚至上百倍。矩阵乘法如果按照最朴素的方式写每次计算C的一个元素都要从主内存里把A的一整行和B的一整列读进来。当矩阵规模变大时这些数据根本装不进缓存CPU大部分时间都在等内存计算单元反而闲着。BLAS的解决方案是分块。把大矩阵切成小块每块的大小精心设计确保参与运算的子矩阵能完整装进L1或L2缓存。然后在缓存内部完成这些小矩阵的乘法再把结果写回主内存。这样主内存的访问次数大幅减少CPU的计算单元能持续保持忙碌。分块的大小不是随便定的。它取决于目标CPU的缓存容量、缓存行大小、寄存器数量。比如一个典型的L1缓存是32KB那么分块后的子矩阵总大小就不能超过这个数。BLAS实现里会有一组经过实测调优的参数针对不同的矩阵尺寸选择不同的分块策略。这也是为什么BLAS在不同CPU上性能表现不同——分块参数需要针对具体硬件调优。除了分块BLAS还大量使用了SIMD指令。SIMD是“单指令多数据”的缩写简单说就是一条指令同时处理多个数据。比如你要做四个浮点数的加法普通指令要执行四次SIMD指令一次就能搞定。现代CPU普遍支持AVX、AVX-512这类SIMD指令集BLAS实现会针对这些指令集写专门的汇编代码把计算吞吐量拉到硬件极限。还有一个容易被忽略的点是内存对齐。SIMD指令加载数据时如果数据在内存里是对齐的加载效率最高。BLAS在分配矩阵内存时会特意按照缓存行大小对齐确保每次加载都能命中最佳路径。这个细节在应用层完全感知不到但对性能的影响是实打实的。我做过一个简单的对比测试用C语言手写三重循环做1024×1024的矩阵乘法和调用OpenBLAS的cblas_dgemm做同样的运算。手写版本跑了大约8秒OpenBLAS版本只用了0.3秒左右。差距接近30倍。这个差距不是算法层面的而是工程优化层面的——分块、SIMD、多线程、内存对齐每一项都在贡献性能。如果你在自己的程序里发现矩阵运算成了瓶颈第一件事不是去改算法而是确认你用的底层库是不是一个经过优化的BLAS实现。很多时候换一个BLAS后端就能带来数倍的性能提升。3. 选OpenBLAS还是MKL一次真实的选型对比BLAS的实现有好几个选哪个往往让人纠结。我拿两个最常见的实现——OpenBLAS和MKL——做过一轮比较这里把实测数据和选型逻辑分享一下。OpenBLAS是开源的社区维护支持几乎所有主流CPU架构。它的优势在于免费、可定制、跨平台。你可以自己编译针对特定CPU做优化。MKL是Intel的商业库对Intel自家CPU的优化非常到位尤其是最新的AVX-512指令集MKL往往能比OpenBLAS多榨出10%到20%的性能。但MKL在AMD CPU上的表现就不一定了有时候反而不如OpenBLAS。我测试的环境是一台搭载Intel处理器的机器矩阵规模从256到4096不等分别测试单线程和多线程下的双精度矩阵乘法性能。结果大致是这样的矩阵规模OpenBLAS单线程MKL单线程OpenBLAS多线程MKL多线程256×2560.8ms0.7ms0.3ms0.3ms1024×102452ms45ms8ms7ms2048×2048420ms360ms55ms48ms4096×40963400ms2900ms420ms370ms从数据看MKL在Intel平台上确实有优势但差距没有想象中那么大大约在10%到15%之间。多线程的加速比在两种实现上都比较理想4096规模下接近8倍加速。选型的时候除了性能还要考虑几个实际因素。许可证是一个关键点。MKL虽然可以免费使用但它是闭源的某些场景下可能有许可限制。OpenBLAS是BSD许可证用起来更自由。部署便利性也要考虑MKL的安装包比较大OpenBLAS相对轻量。CPU兼容性方面如果你的程序要在多种CPU上跑OpenBLAS的通用性更好。还有一个容易被忽略的点是线程管理。BLAS内部会自己开多线程如果你的应用层也开了多线程两者可能会打架导致性能不升反降。我遇到过这种情况一个多线程程序调用MKL做矩阵乘法结果因为线程数设置不当性能比单线程还差。解决办法是设置环境变量控制BLAS的线程数比如OMP_NUM_THREADS让BLAS的线程数和应用层的线程数协调好。如果你不确定选哪个我的建议是先在目标机器上把两个都装一遍跑一下你的实际业务数据用真实性能说话。不要只看跑分因为不同矩阵形状、不同数据类型的表现可能完全不同。4. 从C到PythonBLAS在不同语言里的调用方式BLAS的接口是C和Fortran的但实际使用中大多数人不会直接写C去调BLAS。更多时候我们是通过上层语言或框架间接使用BLAS。理解这个调用链路能帮你在遇到性能问题时快速定位。在**C/C**里你可以直接链接BLAS库调用cblas_dgemm这类函数。需要包含头文件链接时指定库路径。编译命令大概长这样gcc -o myapp myapp.c -lopenblas -lpthread调用的时候要注意参数顺序。以cblas_dgemm为例它有一堆参数矩阵布局、转置标志、矩阵维度、缩放系数、矩阵指针、前导维度等等。前导维度这个参数特别容易搞错它表示矩阵在内存中每一行的实际长度如果矩阵是更大矩阵的子块前导维度就不等于列数。搞错了会导致计算结果完全错误。在Python里NumPy和SciPy底层都链接了BLAS。你调用np.dot或scipy.linalg.blas.dgemm时实际执行的就是BLAS。NumPy默认会链接一个BLAS实现通常是OpenBLAS。你可以用numpy.show_config()查看当前链接的是哪个。如果想换可以设置环境变量或者在编译NumPy时指定。Python层面还有一个选择是直接调用scipy.linalg.blas模块它提供了对BLAS函数的直接封装。这样做的好处是可以绕过NumPy的一些额外开销直接控制BLAS的调用参数。比如你可以指定是否转置、是否覆盖输入矩阵等。对于性能敏感的场景这个层级的控制很有价值。在Java里情况稍微复杂一些。标准库没有直接提供BLAS接口但可以通过JNI或者第三方库来调用。常见的有netlib-java、jblas等。这些库封装了BLAS的本地调用让你在Java里也能享受到BLAS的性能。不过JNI调用本身有开销对于小规模矩阵运算可能还不如纯Java实现快。在Rust里有blas和blas-src这两个crate。blas-src负责链接具体的BLAS实现blas提供安全的Rust接口。Rust的生态还在发展中BLAS相关的库不如Python和C那么成熟但基本功能已经可用。不管用什么语言有一个原则是通用的尽量让BLAS处理大块运算而不是频繁调用小规模函数。BLAS的函数调用有固定开销如果矩阵很小开销占比就很高。比如你循环调用一千次2×2矩阵乘法可能还不如自己手写一个内联函数快。BLAS的优势在大矩阵上才能充分发挥。5. 那些年我踩过的BLAS坑从链接错误到性能反转用BLAS的过程中我踩过不少坑有些是配置问题有些是理解偏差。这里挑几个典型的分享一下希望能帮你省点时间。第一个坑是链接顺序。在Linux下用gcc链接BLAS时库的顺序很重要。如果写成-lopenblas -lm可能没问题但如果写成-lm -lopenblas就可能报未定义符号。这是因为链接器处理库的顺序是从左到右后面的库要能解析前面库的未定义符号。BLAS依赖数学库所以-lm要放在-lopenblas后面。这个规则在静态链接时尤其严格。第二个坑是线程数设置。OpenBLAS默认会根据CPU核心数自动开线程。但在容器环境里它可能读到的是宿主机的核心数而不是容器的限制。结果就是开了太多线程上下文切换开销巨大性能反而下降。解决办法是显式设置OPENBLAS_NUM_THREADS环境变量把它限制在合理范围内。第三个坑是数据类型不匹配。BLAS的函数名里包含了数据类型信息。s开头的是单精度浮点d开头的是双精度c是单精度复数z是双精度复数。如果你用单精度数据去调双精度函数结果会完全错误而且编译器不一定报错。我见过有人因为这个原因调试了一整天都没找到问题。第四个坑是矩阵存储顺序。BLAS默认使用列主序而C语言的多维数组是行主序。如果你直接把C数组传给BLAS而不做转置或指定正确的布局参数计算结果就会出错。NumPy在这方面做了封装用户感知不到但如果你直接调C接口就一定要注意。第五个坑是性能反转。有时候你费尽心思优化了BLAS调用结果发现整体程序反而变慢了。原因可能是BLAS的多线程和你应用层的多线程产生了资源竞争。或者BLAS的线程创建开销超过了计算本身的收益。对于小规模运算关掉BLAS的多线程反而更快。这个需要根据实际场景做权衡。排查BLAS相关问题时一个有用的技巧是用LD_DEBUGlibs环境变量运行程序看看实际加载的是哪个BLAS库。有时候系统里装了多个BLAS链接器选中的可能不是你期望的那个。6. 自己动手验证一个最小化的BLAS性能测试光看理论不够直观我写了一个最小化的测试程序你可以自己跑一下感受BLAS和手写循环的差距。这个程序用C语言写分别用朴素三重循环和BLAS做矩阵乘法对比耗时。#include stdio.h #include stdlib.h #include time.h #include cblas.h #define N 1024 void naive_gemm(double *A, double *B, double *C) { for (int i 0; i N; i) { for (int j 0; j N; j) { double sum 0.0; for (int k 0; k N; k) { sum A[i * N k] * B[k * N j]; } C[i * N j] sum; } } } int main() { double *A malloc(N * N * sizeof(double)); double *B malloc(N * N * sizeof(double)); double *C malloc(N * N * sizeof(double)); for (int i 0; i N * N; i) { A[i] (double)(i % 100) / 100.0; B[i] (double)(i % 100) / 100.0; } clock_t start clock(); naive_gemm(A, B, C); double naive_time (double)(clock() - start) / CLOCKS_PER_SEC; start clock(); cblas_dgemm(CblasRowMajor, CblasNoTrans, CblasNoTrans, N, N, N, 1.0, A, N, B, N, 0.0, C, N); double blas_time (double)(clock() - start) / CLOCKS_PER_SEC; printf(Naive: %.3f sec\n, naive_time); printf(BLAS: %.3f sec\n, blas_time); printf(Speedup: %.1fx\n, naive_time / blas_time); free(A); free(B); free(C); return 0; }编译命令gcc -O2 -o bench bench.c -lopenblas -lpthread -lm在我的机器上朴素版本跑了大约7.8秒BLAS版本0.28秒加速比约28倍。这个差距在更大的矩阵上还会拉大。你可以把N改成2048或4096试试但注意内存占用会成平方增长。这个测试还揭示了一个细节BLAS的cblas_dgemm函数有一个beta参数我传的是0.0表示C矩阵的初始值不参与计算直接覆盖。如果传1.0就是累加。这个参数在迭代算法里很有用可以避免额外的矩阵加法操作。另外注意CblasRowMajor这个参数。它告诉BLAS矩阵是按行主序存储的。如果你的数据是列主序就改成CblasColMajor。这个参数搞错了结果就全错了。NumPy在底层会自动处理这些但直接调C接口时一定要小心。7. 当BLAS不够用时什么场景需要更专业的方案BLAS虽然强大但它不是万能的。有些场景下你需要考虑更专业的方案。稀疏矩阵是一个典型场景。BLAS处理的是稠密矩阵矩阵里大部分元素都是非零的。但实际应用中很多矩阵是稀疏的比如社交网络的关系矩阵、有限元分析的刚度矩阵。对稀疏矩阵用稠密BLAS会浪费大量内存和计算。这时候需要用稀疏BLAS比如Sparse BLAS标准或者专门的稀疏线性代数库。大规模并行是另一个场景。BLAS的多线程是在单机内利用多核。如果你的矩阵大到单机内存装不下或者需要跨多台机器计算就需要分布式线性代数库比如ScaLAPACK。它把矩阵分块分布到多个节点上通过网络通信协调计算。这类库的编程模型比BLAS复杂得多但能处理BLAS无法企及的规模。GPU加速也值得考虑。BLAS有对应的GPU版本比如cuBLAS。GPU的并行度远高于CPU对于大规模矩阵乘法GPU能提供数倍甚至数十倍的性能。但GPU编程有额外的复杂性数据传输、显存管理、核函数编写。如果你的应用已经在GPU上跑用cuBLAS是很自然的选择如果只是偶尔做一次矩阵乘法把数据搬到GPU再搬回来可能得不偿失。自动微分场景下BLAS的接口就不太够用了。深度学习框架需要的是能自动求导的线性代数操作。这些框架通常会在BLAS之上再封装一层比如PyTorch的torch.matmul它内部调用BLAS但对外提供自动微分能力。如果你在做深度学习相关的开发直接用框架的接口就好不需要直接调BLAS。混合精度计算是近年来的一个趋势。有些场景下用半精度浮点做矩阵乘法速度能快好几倍而精度损失可以接受。BLAS标准本身没有定义半精度接口但一些实现扩展了半精度支持。如果你的应用对精度要求不那么苛刻可以关注一下这方面的进展。选型的时候核心问题是你的矩阵是什么形态规模多大精度要求如何硬件环境是什么把这些想清楚答案自然就出来了。BLAS是基础但不是终点。
返回列表