讨论和解决方案代码
这里可以建议使用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 倍以上 的加速!?