Matlab计算点到质心笛卡尔距离如何避免冗余reshape操作?
优化实现方案
方案1:纯基础Matlab(无额外工具箱依赖)
Matlab R2016b及以上版本支持隐式广播,同时可以用索引中插入[]的语法糖直接添加单例维度,完全不需要手动写reshape函数,优化后代码如下:
assert(size(X,2) == size(centroids,2), "Dimensions do not match.") distance = sqrt(sum((X(:, [], :) - centroids([], :, :)).^2, 3));
说明:
X(:, [], :)会自动把形状为(M,N)的X插入第二维单例,变成(M,1,N),等价于原来的reshape操作,语法更简洁centroids([], :, :)会自动把形状为(K,N)的centroids插入第一维单例,变成(1,K,N)- 两个矩阵直接做减法会触发隐式广播,得到
(M,K,N)的差值矩阵,后续求和、开根号逻辑和原实现一致,整体代码量大幅减少
如果是更早不支持隐式广播的版本,也可以用bsxfun实现:
assert(size(X,2) == size(centroids,2), "Dimensions do not match.") distance = sqrt(sum(bsxfun(@minus, X(:, [], :), centroids([], :, :)).^2, 3));
方案2:有统计工具箱的极简写法
如果你安装了Statistics and Machine Learning Toolbox,可以直接调用内置的pdist2函数,一行代码就能得到结果,且内部做了高度优化,运行效率比手动实现更高:
distance = pdist2(X, centroids);
补充:大矩阵内存友好型实现
如果M、K取值很大,三维广播会占用过多内存,可以用平方差展开公式实现纯二维运算,内存占用更低、运算速度更快:
assert(size(X,2) == size(centroids,2), "Dimensions do not match.") X_sq = sum(X.^2, 2); cent_sq = sum(centroids.^2, 2)'; distance = sqrt(X_sq + cent_sq - 2 * X * centroids');
内容的提问来源于stack exchange,提问作者pete
相关产品推荐
相关产品推荐

