
我在做大规模聚类实验时最常遇到的性能瓶颈不是聚类算法本身而是每次迭代都要计算的那张“成对样本平方距离”矩阵。N个样本两两算距离如果老老实实写两层for循环一万个样本基本就卡到没法看很多人换成pdist2数据一旦到十万量级内存和耗时又会一起爆炸。这篇文章我会把一个在MATLAB社区里被反复使用、名字叫sqdistance的函数彻底讲透它利用一个简单的平方和恒等式把成对平方距离的计算从二重循环变成三次矩阵运算代码只有三五行速度却能快上几倍甚至一个数量级。先说明一下sqdistance并不是MATLAB官方主程序里内置的函数官方自带的是pdist2通过squaredeuclidean这个距离标志也能计算平方欧氏距离。但在很多开源工具箱和资深用户的自定义代码里sqdistance这个名字已经被约定俗成指的就是“用矩阵展开法高效求成对平方距离”的这样一个工具函数。它适合所有被距离矩阵折磨过的MATLAB用户你在做K-means、KNN分类、RBF核函数、谱聚类、多维缩放只要涉及成对样本距离都可以直接用这套思路替换。1. 平方距离到底算的是什么又卡在哪些算法里1.1 从欧氏距离到平方距离开根号是不是多余两个向量x和y的欧氏距离定义为sqrt(sum((x - y).^2))而平方距离就是去掉根号的那部分sum((x - y).^2)。很多人会问既然是距离为什么不保留欧氏距离的形式答案很简单在很多算法里根号不仅多余还会引入不必要的计算和数值麻烦。以K-means为例聚类算法本质上是比较“哪个质心离我更近”而不是关心“到底有多远”。sqrt是严格单调递增函数对距离排序没有任何影响而平方距离在数学上等价于欧氏距离的排序。RBF核函数K exp(-||x - y||^2 / (2*sigma^2))更是直接就需要平方距离你如果先开根号再平方等于来回浪费了一次运算。更关键的是平方距离在很多优化目标里是凸函数数学性质比欧氏距离好处理得多。所以平方距离并不是一个“半成品距离”它本身就是有独立价值的计算对象。理解这一点你才知道为什么社区里会专门写一个sqdistance函数去优化它而不是简单调一下pdist2就算了。实际项目中求完平方距离再按需开根号也是完全可控的D sqrt(sqdistance(X, Y))一步就能得到成对欧氏距离计算量并不会比直接算欧氏距离多到哪去。1.2 哪些经典算法的性能瓶颈就在这张距离矩阵上我列一下最常踩到距离矩阵性能坑的场景都是实际项目中反复出现的算法/场景距离矩阵规模计算频率痛点K-means聚类样本数m × 质心数k每轮迭代都算m到十万量级时每次迭代都吃内存KNN最近邻查询集m × 训练集n每次预测都算n大到百万时全矩阵根本无法存RBF核函数n × n一次全部算完核矩阵本身n²无法再额外浪费谱聚类n × n一次n10000时double矩阵已是800MB多维缩放MDSn × n迭代多次每次迭代都要重建距离信息这张表里的每个场景都要面对一个共同问题距离矩阵的输出本身是n²量级运算量是m×n×d量级。pdist2虽然封装得干净但它一次性返回完整矩阵内存占用是所有方法的天花板而循环写法虽然省内存但解释器开销能让你等到怀疑人生。sqdistance的价值就在这里——它通过矩阵展开把最重的计算交给BLAS又可以通过分块把内存峰值压到你想要的范围内。2. sqdistance的加速原理一条恒等式把二重循环干掉2.1 从中学数学到矩阵乘法推导过程拆开看平方距离能矩阵化的核心是一个再普通不过的恒等式||x - y||^2 ||x||^2 ||y||^2 - 2 * dot(x, y)展开推导是这样(x - y)^T (x - y)等于x^T x减2x^T y再加y^T y也就是x自己点自己加上y自己点自己再减两倍的交叉点积。把单个向量的关系推广到成对矩阵问题就变得非常优雅了。设有样本矩阵X大小m×d以及目标矩阵Y大小n×d。我们想要的输出是m×n矩阵D其中D(i,j)是X第i行和Y第j行的平方距离。根据恒等式可以写成D(i,j) rowSum(X(i,:).^2) rowSum(Y(j,:).^2) - 2 * X(i,:) * Y(j,:).把所有“行平方和”先算出来再统一乘交叉项就得到了矩阵形式D X2 * ones(1, n) ones(m, 1) * Y2. - 2 * X * Y.其中X2 sum(X.^2, 2)是一个m×1列向量Y2 sum(Y.^2, 2)是n×1列向量。X * Y.就是标准的矩阵乘法一次性得到所有交叉内积。这个公式就是整个sqdistance的灵魂所有优化版本都是围绕“怎么把这个式子算得更稳、更省内存”展开的。2.2 复杂度没有变低为什么反而变快了从计算机科学的角度sqdistance的时间复杂度仍然是O(m * n * d)和暴力二重循环完全一样。因为这本质上就是把“内积计算”这一步从显式循环挪到了矩阵乘法里。既然复杂度没变快在哪答案是常数项和内存访问方式。暴力循环的每一对i,j都要做一次长度为d的向量减法、逐元素平方、再求和。在MATLAB里这一系列操作要在解释器里逐行解析即使有JIT加速生成的机器码也比不上专门针对CPU指令集优化过的BLAS库。MATLAB的矩阵乘法底层默认链接了Intel MKL或OpenBLAS这类高性能数学库它们对数据缓存、SIMD指令比如AVX、多线程并行都做了深度调优。我把矩阵乘法比作流水线批量生产把循环比作手工单件定制同一条产线批量模式的吞吐量能差出一个数量级。另一个经常被忽略的点是内存操作次数。循环写法需要对X(i,:)、Y(j,:)反复做切片访问中间还会生成大量临时数组矩阵乘法则把所有数据一次性送入BLAS中间结果的复用率和CPU缓存的命中率都高得多。所以结论是sqdistance没有减少算术运算量但它把运算形式变成了硬件最擅长处理的方式。2.3 内存复杂度也要提前心里有数计算X * Y.时输出就是m×n的矩阵这是不可避免的。假设x和y都是double类型每个元素8字节那么输出矩阵占用的内存是8 * m * n字节。mn10000时约800MBmn50000时约20GB这已经超过绝大多数个人电脑的物理内存了。所以光有基础版sqdistance还不够分块版才是处理大数据的关键。分块的核心思想是不一次性生成整个m×n矩阵而是每次只算m×blockSize大小的子块算完存进D的对应列这样内存峰值从8*m*n降到了8*m*blockSize。比如m50000、blockSize1000时每个子块只有400MB多数机器可以轻松扛住。这个思路在后面第3节代码里会完整展开。3. 从基础版到增强版手写一个可复用的sqdistance3.1 基础版三行核心代码先看最简洁的实现这也是很多工具箱里sqdistance的原型function D sqdistance(X, Y) if nargin 2 Y X; end X2 sum(X.^2, 2); Y2 sum(Y.^2, 2); D X2 Y2. - 2 * (X * Y.); end这段代码非常小但每个细节都值得说清楚。X.^2是逐元素平方sum(..., 2)表示沿第二维求和结果是m×1的列向量记录的是每个样本的自身内积||x_i||^2。Y2同理。第三行里X2是m×1Y2.是1×nMATLAB从R2016b开始支持隐式扩展二者相加自动广播成m×n矩阵每个元素(i,j)对应||x_i||^2 ||y_j||^2。后面再减去2 * (X * Y.)就把交叉内积项补齐了。验证方式很简单拿循环写法做参照X rand(100, 10); Y rand(50, 10); D sqdistance(X, Y); ref zeros(100, 50); for i 1:100 for j 1:50 ref(i, j) sum((X(i,:) - Y(j,:)).^2); end end max(abs(D(:) - ref(:)))正常情况下结果应该在1e-12量级如果出现较大的偏差优先检查X、Y里有没有NaN或Inf。3.2 分块版把内存峰值牢牢摁住数据量一大基础版会直接内存不足。分块版代码如下function D sqdistance_block(X, Y, blockSize) if nargin 2 Y X; end if nargin 3 blockSize 1000; end [m, ~] size(X); n size(Y, 1); X2 sum(X.^2, 2); Y2 sum(Y.^2, 2); D zeros(m, n); for j 1:blockSize:n idx j:min(j blockSize - 1, n); D(:, idx) X2 Y2(idx). - 2 * (X * Y(idx, :).); end end这里的思路是按列分块每一轮取出Y的idx对应的blockSize行算出一个m×blockSize的子矩阵填到D的第idx列。内存占用主要由X * Y(idx,:).决定峰值大约是8 * m * blockSize字节加上D本身占用的8 * m * n字节。如果最终D是必须保留的完整矩阵那么D自身的占用没法省但至少不会在计算过程中额外复制一份完整的m×n中间结果。blockSize的选择有讲究。太小了比如只有几十矩阵乘法的尺寸过小BLAS的优势发挥不出来太大了内存又失去控制。我的经验是先从round(1e7 / m)估算保证m * blockSize在千万级别以下再根据实际内存手动调整。比如m50000时blockSize200左右比较稳m5000时blockSize1000则很均衡。3.3 对称矩阵版X与Y相同时可以省一半计算当YX距离矩阵D天然对称理论上只需要算上三角或者下三角。这里给一个对称化实现function D sqdistance_sym(X) n size(X, 1); X2 sum(X.^2, 2); D zeros(n); for j 1:n D(1:j, j) X2(1:j) X2(j) - 2 * (X(1:j, :) * X(j, :).); end D D D. - diag(diag(D)); end这个版本逐列填充上三角最后通过D D.恢复对称结构再减去被重复计算的对角线。它的优点是只算了约一半的内积适合矩阵乘法开销极其昂贵的场景缺点是逐列循环破坏了连续大矩阵乘法的高效性当n比较小比如几千时收益并不明显。我的实际建议是如果想省内存优先用3.2的分块版如果n很大且你只需要上三角比如谱聚类的相似度矩阵可以结合对称分块来写而不是逐列循环。对称优化不是银弹要根据具体算力环境做基准测试。3.4 复数数据与边界输入两个容易写错的地方MATLAB里转置有两种写法。A.是非共轭转置A是共轭转置。对于实数数据二者没有区别对于复数数据平方距离的定义应该是模长平方即|z|^2 z * conj(z)所以矩阵交叉项必须用X * Y而不是X * Y.。同时X2也不能再用X.^2必须写成sum(abs(X).^2, 2)。边界情况也要防一手。如果X或Y为空X * Y.会输出一个空矩阵后续的减法可能产生非预期维度如果只有一行广播和矩阵乘法倒还能正常工作但性能优势不明显。更值得警惕的是大量相同值的样本整体距离矩阵会退化为低秩结构这时候X * Y.照样能算但如果你对数值精度有极端要求建议改用显式循环再验一下关键位置。4. 实测对比sqdistance、pdist2与双循环的差距4.1 一个公平的测试脚本怎么写性能对比最忌讳只测一组数据。我通常固定d50分别测mn1000、mn5000、mn10000三档同时记录时间和峰值内存。测试脚本框架如下m 5000; n 5000; d 50; X randn(m, d); Y randn(n, d); tic; D1 pdist2(X, Y, squaredeuclidean); toc; tic; D2 sqdistance(X, Y); toc; tic; D3 sqdistance_block(X, Y, 500); toc;注意计时前最好先跑一次热身把BLAS线程和内存分配的热状态拉起来不然第一轮计时经常偏大。tic/toc对毫秒级以下的差异不敏感所以小规模测试建议用timeit包一层函数句柄f1 () pdist2(X, Y, squaredeuclidean); f2 () sqdistance(X, Y); timeit(f1) timeit(f2)4.2 结果怎么看量级比绝对值更重要我在自己的台式机上六核CPU、MATLAB R2023a、d50实测的相对关系大概是这样的方法相对耗时峰值内存备注双循环几十倍以上低但慢到不可用解释器开销是主要瓶颈pdist2约基准1倍高有输入检查、参数解析等额外开销sqdistance基础版约为pdist2的1/5到1/3高核心只剩三次矩阵运算sqdistance分块版比基础版略慢但内存可控按blockSize线性下降大矩阵下的必选项必须强调的是这个比例强烈依赖机器、数据规模和BLAS线程数。不要看到“1/5”就以为在任何机器上都成立。唯一稳定可复现的结论是双循环一定是最慢的sqdistance在大中规模下几乎一定比pdist2快分块版能用时间换内存空间。你的数据如果特别小比如只有几百个样本差距可能只有几毫秒用哪个都无所谓。数据一旦到几千以上sqdistance的优势就会明显到肉眼可见。建议你自己跑一遍上面的基准脚本保存一组本机数据后面调参时就心里有底了。4.3 pdist2到底慢在哪为什么自己写反而快pdist2是一个功能极其丰富的通用函数支持几十种距离度量。功能丰富意味着每次调用都要做参数解析、距离名称匹配、输入校验、错误检查这些虽然单次开销很小但在短小的大矩阵计算面前会被放大。更关键的是pdist2内部计算squaredeuclidean时默认返回一个完整的m×n矩阵没有分块选项也不允许你自定义子块计算。自己写sqdistance等于绕开了所有通用性负担只留一个数学核心。你还能根据场景做定制不需要完整距离矩阵时可以只保留每行的最小值可以自己控制分块大小可以选择单精度可以选择只算上三角。这些自由度是pdist2给不了的也正是自定义函数在实战中的价值所在。5. 数值稳定与常见坑位避雷5.1 距离矩阵出现负值先别慌D X2 Y2. - 2 * (X * Y.)的数学形式是精确的但计算机浮点运算做不到完全精确。当两个样本几乎重合时理论上距离应该为0实际计算结果可能出现-1e-12这种微小的负数。这个负值在后续处理中会带来大麻烦尤其是当你执行sqrt(D)时会直接得到NaN。标准解法是加一行修正D max(D, 0);。如果不想额外生成一个新矩阵可以用原地修改D(D 0) 0;对大矩阵更友好。还有一个进阶做法先检查min(D(:))是否在-1e-10量级如果负值过大说明不是单纯浮点误差而是输入数据本身就含NaN或者计算过程有数值溢出得回到数据上排查。5.2 稀疏数据别急着无脑套公式如果X和Y是稀疏矩阵X.^2仍然保持稀疏非零位置不变理论上X * Y.也可以用稀疏矩阵乘法。但有个隐藏陷阱两个稀疏矩阵相乘的结果矩阵密度会迅速上升特别是当每行非零元较多时结果可能不再是稀疏向量能舒服处理的对象存储和计算量同时暴增。我处理高维稀疏特征的经验是先统计每行非零数的分布如果平均非零数超过几十直接算稠密矩阵可能更省心如果非零比例极低比如文本数据的TF-IDF特征则可以用sparse矩阵直接套用公式但要留意输出D的密度。更稳妥的办法是分块读取原始数据把稀疏矩阵若干行先转成稠密子块再调sqdistance避免一次性把所有数据拍扁。5.3 NaN与Inf一小粒老鼠屎毁掉整锅粥距离计算里NaN有传染性只要X或Y任意一行含NaN对应的整行、整列距离都会变成NaN。Inf也一样高维数据中某个元素是1e200量级时平方后会溢出成Inf距离矩阵里会冒出大量Inf导致后续聚类标签全是NaN。我的建议是在进入sqdistance之前先做一次数据体检any(isnan(X(:)))、any(isinf(X(:)))扫一遍有异常就先用fillmissing、rmmissing或按业务规则过滤。不要指望函数内部帮你兜底MATLAB里加了冗余检查反而会拖慢性能把这个职责留在调用层更合理。5.4 单精度内存少一半速度还可能更快数据量巨大时一个非常实用的技巧是全部转成single类型。单精度浮点数占4字节距离矩阵内存直接减半BLAS对单精度矩阵乘法往往也有专门优化路径耗时通常比double更短。代价是精度下降单精度大概只有7位有效数字对于K-means这种对距离精度不敏感的任务完全够用但对于需要精细后处理的场景比如计算RBF核矩阵后再做数值优化还是建议保留double。进阶操作是X和Y用single参与矩阵乘法最后输出D时再根据需要转回double两头的好处都占一点。5.5 全矩阵存不下试试“在线argmin”模式很多算法其实不需要完整的成对距离矩阵只需要知道每个样本的最近邻。这时候存全矩阵是纯粹浪费。K-means里分配样本到最近的质心就是典型的“只要argmin不要D本身”的场景。下面这个函数展示了如何在线算距离、在线取最小值function labels nearestCentroid(X, C, blockSize) [m, ~] size(X); k size(C, 1); labels zeros(m, 1); C2 sum(C.^2, 2); X2 sum(X.^2, 2); if nargin 3 blockSize 1000; end for i 1:blockSize:m idx i:min(i blockSize - 1, m); D X2(idx) C2. - 2 * (X(idx, :) * C.); [~, labels(idx)] min(D, [], 2); end end这个模式把内存占用从m×k降到了blockSize×k特别适合样本数m很大、质心数k适中几十到几百的聚类任务。它也是整个sqdistance思路最有价值的变体因为它彻底改变了你的内存思维不需要为了一个中间结果去扩容服务器。6. 实战把sqdistance嵌进K-means迭代里6.1 传统写法的距离瓶颈在哪经典K-means每次迭代分两步分配样本到最近质心、根据分配结果更新质心。分配这一步很多教程代码是这么写的for i 1:m dmin inf; for j 1:k dist sum((X(i,:) - C(j,:)).^2); if dist dmin dmin dist; labels(i) j; end end end当m有十万、k有一百、d有50时这种写法每轮迭代要执行上千万次内积计算再加上解释器循环开销一轮下来轻松超过十秒。换成pdist2(X, C, squaredeuclidean)虽然快但会一次性生成100000×100的double矩阵也就是80MB每轮迭代都重新分配一次频繁的内存分配同样让人肉疼。6.2 用sqdistance重写距离分配利用上一节的nearestCentroid函数K-means的主循环可以变得很干净% 初始质心 C for iter 1:maxIter labels nearestCentroid(X, C, 1000); newC zeros(k, d); count zeros(k, 1); for j 1:k idxs (labels j); if ~any(idxs) newC(j, :) C(j, :); else newC(j, :) mean(X(idxs, :), 1); end count(j) sum(idxs); end if norm(newC - C, fro) tol break; end C newC; end这里面的距离分配一次只生成1000×100的临时矩阵只有0.8MB完全不会造成内存压力。同时每次调用nearestCentroid时都在执行一次大的矩阵乘法BLAS的并行能力被完全利用十万样本一轮分配的时间通常在几十毫秒到几百毫秒之间比循环写法快两个数量级都不夸张。6.3 我踩过的坑质心为空簇的处理用sqdistance加速K-means后最需要注意的反而不是性能问题而是空簇。当某个质心在分配后没有任何样本时mean(X(idxs,:),1)会对空集求平均直接产生NaN。我早期跑文本聚类时常被这个坑折磨后来习惯在更新质心时加一个保护空簇保留原质心或者随机初始化为某个远离全局均值的新点。另一个细节是K-means的目标函数本身就是“各点到最近质心平方距离之和”。用nearestCentroid在每次分配时顺便把距离最小值带出来就能轻松计算目标函数值用来监控收敛曲线。这个思路也可以推广到KNN的最近邻搜索里把“算全距离”改成“分块算距离并保留top几”对大规模检索非常有帮助。7. 常见问题速查表与避坑指南我把自己用sqdistance时遇到过的典型问题整理成一张速查表方便收藏现象可能原因处理方式D出现很小的负值浮点舍入误差D(D 0) 0D出现NaNX或Y含NaN调用前用isfinite检查并清洗D出现Inf数据元素过大平方溢出先做标准化或截断极端值内存不足全矩阵m×n太大改用分块版或在线argmin模式计算速度没有提升数据量太小BLAS优势不明显数据较小就别优化保持可读性结果和pdist2对不上误差在1e-10以内属正常用绝对误差阈值比较别用isequalX是稀疏矩阵时崩溃结果矩阵密度失控转稠密子块或按行分块复数计算出错.和用混了复数场景用X * Y并配合abs(X).^2排查问题有个总原则先看输入数据再看中间量。sqdistance本身公式足够简单一旦出现问题八成是数据里有脏值而不是函数写错了。你可以单独检查X2和Y2是否正常再看X*Y.是否满足基本形状就能快速定位。关于BLAS线程还要多说一句如果你发现矩阵乘法速度异常先查一下maxNumCompThreads返回的线程数。默认线程数一般没问题但在共享服务器上线程数过高反而会导致缓存颠簸适当调低有时能收获意外加速。最后再分享一点个人经验这套sqdistance思路我用了好几年后来把所有涉及距离矩阵的算法都重构了一遍。最大的体会不是“快”本身而是当你把内存占用变得可预测之后整个程序的设计空间会大很多。以前写K-means要反复估算“这个矩阵放不放得下内存”改成分块加在线argmin之后脑子里只需要记住一个公式blockSize乘输出维度不要超过5000万。很多看似需要分布式框架的大规模问题在普通单机上用这个思路就能跑起来。如果你想把sqdistance继续深化还有两个方向可以玩一是配合MATLAB的gpuArray把X * Y.放到GPU上算矩阵乘法是GPU最擅长的工作提升还能再上一个台阶二是把分块逻辑参数化做成一个能够自动根据memory函数剩余量调节blockSize的自适应版本。我自己做聚类工具库时已经把这些都封装进去了日常使用几乎不再为距离计算发愁。希望这篇东西也能帮你的MATLAB代码跑得更快一些。