是否可以加速这个MATLAB脚本?

min*_*ing 4 performance matlab matrix

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

脚本是

    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
Run Code Online (Sandbox Code Playgroud)

Div*_*kar 6

讨论和解决方案代码

bsxfun这里可以提出很少的方法.另外,请继续阅读以了解如何30x+在这样的问题上获得加速!

方法#1(天真矢量化方法)

为了适应行之间的两次减法操作,X然后在它们之间进行随后的逐元素乘法,基于朴素bsxfun的方法将导致对应于的4D中间阵列((X(i,:)-X(j,:))' * (X(i,:)-X(j,:))).在那之后,需要乘以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,[]);
Run Code Online (Sandbox Code Playgroud)

方法#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,[])
Run Code Online (Sandbox Code Playgroud)

方法#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
Run Code Online (Sandbox Code Playgroud)

标杆

比较原始方法与方法#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
Run Code Online (Sandbox Code Playgroud)

运行时结果

------------------------------ With Original Approach
Elapsed time is 29.728571 seconds.
------------------------------ With Proposed Approach #3
Elapsed time is 0.839726 seconds.
Run Code Online (Sandbox Code Playgroud)

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

  • 嗯,你可能会做Divakar所建议的,但是以"块状"的方式.这样你在外面仍然会有2个for`循环,但实际的计算将使用更高效的`bsxfun`来完成.另外,如果你不需要'double`的精度,也许可以尝试将矩阵转换为更小的数据类型(例如`single`) - 这也有助于解决内存问题...... (2认同)
  • @mining真棒!实际上我从这个特定的问题中学到了很多东西,所以非常感谢你把它带到Stackoverflow! (2认同)