【问题标题】:Is it possible to speed up this MATLAB script?是否可以加快这个 MATLAB 脚本的速度?
【发布时间】:2015-03-07 09:00:47
【问题描述】:

我遇到了一些性能问题,因此我想加快那些运行缓慢的脚本。但是我对如何加快它们没有更多的想法。因为我发现我经常被索引阻塞。我发现抽象思维对我来说非常困难。

脚本是

    tic,
    n = 1000;
    d = 500;
    X = rand(n, d);
    R = rand(n, n);
    F = zeros(d, d);
    for i=1:n
        for j=1:n
           F = F + R(i,j)* ((X(i,:)-X(j,:))' * (X(i,:)-X(j,:)));
        end
    end
    toc

【问题讨论】:

  • 嗯。这看起来好熟悉。你到底在执行什么?也许我们可以使用一些数学来加快速度。
  • 嗨,@knedlsepp,我正在尝试实现这篇论文 [papers.nips.cc/paper/….

标签: performance matlab matrix


【解决方案1】:

讨论和解决方案代码

这里可以建议使用bsxfun 的几种方法。另外,请继续阅读以了解如何在此类问题上获得 30x+ 加速!

方法 #1(朴素矢量化方法)

为了适应X 行之间的减法运算以及它们之间的后续逐元素乘法运算,基于bsxfun 的简单方法将导致对应于((X(i,:)-X(j,:))' * (X(i,:)-X(j,:))) 的4D 中间数组。之后,需要将R 相乘以得到最终输出F。这实现如下所示 -

v1 = bsxfun(@minus,X,permute(X,[3 2 1]));
v2 = bsxfun(@times,permute(v1,[1 3 2]),permute(v1,[1 3 4 2]));
F = reshape(R(:).'*reshape(v2,[],d^2),d,[]);

方法 #2(不那么幼稚的矢量化方法)

前面提到的方法适用于 4D,这可能会减慢速度。因此,您可以通过重塑将中间数据保留到 3D。下面列出了 -

sub1 = bsxfun(@minus,X,permute(X,[3 2 1]));
sub1_2d = reshape(permute(sub1,[1 3 2]),n^2,[])
mult1 = bsxfun(@times,sub1_2d,permute(sub1_2d,[1 3 2]))
F = reshape(R(:).'*reshape(mult1,[],d^2),d,[])

方法 #3(混合方法)

现在,您可以基于方法 #2 (vectorized subtractions + loopy multiplications) 制作混合方法。这种方法的好处是它使用fast matrix multiplication 来执行乘法并将复杂性从早期的 O(n^2) 降低到 O(n),这应该会提高效率。感谢@Dev-iL,提出这个想法!这是代码-

sub1 = bsxfun(@minus,X,permute(X,[3 2 1]));
sub1 = bsxfun(@times,sub1,permute(sqrt(R),[1 3 2]));

F = zeros(d);
for k = 1:size(sub1,3)
    blk = sub1(:,:,k);    
    F = F + blk.'*blk;
end

基准测试

比较原始方法与方法#3

的基准代码
%// Parameters
n = 500;
d = 250;
X = rand(n, d);
R = rand(n, n);

%// Warm up tic/toc.
for k = 1:100000
    tic(); elapsed = toc();
end

disp('------------------------------ With Original Approach')
tic
F1 = zeros(d, d);
for i=1:n
    for j=1:n
        F1 = F1 + R(i,j)*((X(i,:)-X(j,:))' * (X(i,:)-X(j,:)));
    end
end
toc, clear F1 i j

disp('------------------------------ With Proposed Approach #3')
tic
sub1 = bsxfun(@minus,X,permute(X,[3 2 1]));
sub1 = bsxfun(@times,sub1,permute(sqrt(R),[1 3 2]));

F = zeros(d);
for k = 1:size(sub1,3)
    blk = sub1(:,:,k);    
    F = F + blk.'*blk;
end
toc

运行时结果

------------------------------ With Original Approach
Elapsed time is 29.728571 seconds.
------------------------------ With Proposed Approach #3
Elapsed time is 0.839726 seconds.

那么,谁准备好迎接 30 倍以上 的加速!?

【讨论】:

  • 好吧,你可能可以按照 Divakar 的建议去做,但要以“块式”的方式进行。这样你仍然会有 2 个for 循环在外部,但实际计算将使用更高效的bsxfun 完成。此外,如果您不需要 double 的精度,可以尝试将矩阵转换为更小的数据类型(例如 single) - 这也有助于解决内存问题...
  • @mining 查看刚刚添加的方法#3?此外,如果一些精度下降不会打扰您,则可以考虑转换为single
  • @Dev-iL 谢谢!实现了这一点,似乎它提供了巨大的加速!
  • 不客气 - 我很高兴知道我的想法不仅对我有意义 :)。更重要的是,可以从这个答案中提出的bsxfun 深层魔法中学到很多东西。
  • @mining 太棒了!实际上,我自己从这个特定问题中学到了很多东西,所以也感谢你把它带到 Stackoverflow 上!
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 2020-05-06
  • 1970-01-01
  • 2022-01-16
  • 1970-01-01
  • 1970-01-01
  • 2012-08-16
  • 1970-01-01
相关资源
最近更新 更多