【问题标题】:How to compute the sum of squares of outer products of two matrices minus a common matrix in Matlab?如何计算两个矩阵的外积平方和减去Matlab中的公共矩阵?
【发布时间】:2018-11-12 11:30:57
【问题描述】:

假设有三个 n * n 矩阵XYS。如何快速计算以下标量 b

for i = 1:n
  b = b  + sum(sum((X(i,:)' * Y(i,:) - S).^2));
end

计算成本为 O(n^3)。存在一个fast way to compute the outer product of two matrices。具体来说,矩阵 C

for i = 1:n
  C = C + X(i,:)' * Y(i,:);
end

可以在没有 for 循环 C = A.'*B 的情况下计算,它只有 O(n^2)。是否有更快的方法来计算b

【问题讨论】:

  • 您的意思是X(i,:)'*Y(i,:) 吗?否则没有提到Y
  • 是的,感谢您指出错误。我已经更正了。

标签: matlab matrix-multiplication


【解决方案1】:

你可以使用:

X2 = X.^2;
Y2 = Y.^2;
S2 = S.^2;
b = sum(sum(X2.' * Y2 - 2 * (X.' * Y ) .* S + n * S2));

举个例子

b=0;
for i = 1:n
   b = b  + sum(sum((X(i,:).' * Y(i,:) - S).^2));
end

我们可以先把求和带出循环:

b=0;
for i = 1:n
  b = b  + (X(i,:).' * Y(i,:) - S).^2;
end
b=sum(b(:))

知道我们可以把(a - b)^2写成a^2 - 2*a*b + b^2

b=0;
for i = 1:n
  b = b  + (X(i,:).' * Y(i,:)).^2 - 2.* (X(i,:).' * Y(i,:)) .*S + S.^2;
end
b=sum(b(:))

而且我们知道(a * b) ^ 2a^2 * b^2 相同:

X2 = X.^2;
Y2 = Y.^2;
S2 = S.^2;
b=0;
for i = 1:n
  b = b  + (X2(i,:).' * Y2(i,:)) - 2.* (X(i,:).' * Y(i,:)) .*S + S2;
end
b=sum(b(:))

现在我们可以分别计算每一项:

 b = sum(sum(X2.' * Y2 - 2 * (X.' * Y ) .* S + n * S2));

这是 Octave 中的测试结果,该测试比较了我的方法和@AndrasDeak 提供的其他两种方法,以及针对大小为500*500 的输入的原始基于循环的解决方案:

===rahnema1 (B)===
Elapsed time is 0.0984299 seconds.

===Andras Deak (B2)===
Elapsed time is 7.86407 seconds.

===Andras Deak (B3)===
Elapsed time is 2.99158 seconds.

===Loop solution===
Elapsed time is 2.20357 seconds


n=500;
X= rand(n);
Y= rand(n);
S= rand(n);

disp('===rahnema1 (B)===')
tic
    X2 = X.^2;
    Y2 = Y.^2;
    S2 = S.^2;
    b=sum(sum(X2.' * Y2 - 2 * (X.' * Y ) .* S + n * S2));
toc
disp('===Andras Deak (B2)===')
tic
    b2 = sum(reshape((permute(reshape(X, [n, 1, n]).*Y, [3,2,1]) - S).^2, 1, []));
toc
disp('===Andras Deak (B3)===')
tic
    b3 = sum(reshape((reshape(X, [n, 1, n]).*Y - reshape(S.', [1, n, n])).^2, 1, []));
toc
tic
    b=0;
    for i = 1:n
      b = b  + sum(sum((X(i,:)' * Y(i,:) - S).^2));
    end
toc

【讨论】:

  • 非常好! -- 我猜sum(reshape(...))sum(sum(...)) 贵?这是因为 reshape 函数的开销吗? -- 为避免双重和,在 Octave 中您可以执行 sum((...)(:)),在 MATLAB R2018b 中您现在可以执行 sum(...,'all'),两者都比双重和更优雅,应该 更快。​​
  • 谢谢。我不确切知道 sum-reshape 或 sum-sum 哪个表现更好。这里我使用 sum-sum 来反映原始代码,可能更具可读性。 sum((...)(:)) 可能会使可读性复杂化,sum(...,'all') 使用丑陋的字符串参数传递。比较它们,我认为当前问题的速度差异可以忽略不计。我们是否需要一个更优雅的函数/运算符?
  • double sum 需要一个中间数组,所以原则上它会做更多的工作。但是,是的,我确信在这种情况下这不是一个可测量的差异。我同意您的可读性 cmets。一个新的运营商会很棒。像sumall=@(x)sum(x(:)) 这样的东西。但这确实意味着有 20 个左右的新运营商(meanallstdallmaxall 等),你不能只做一个,而不能做所有其他的。 :)
【解决方案2】:

您可能无法节省时间复杂度,但您可以利用向量化来摆脱循环并尽可能利用低级代码和缓存。它实际上是否更快取决于您的尺寸,因此您需要进行一些时序测试,看看这是否值得:

% dummy data
n = 3;
X = rand(n);
Y = rand(n);
S = rand(n);

% vectorize
b2 = sum(reshape((permute(reshape(X, [n, 1, n]).*Y, [3,2,1]) - S).^2, 1, []));

% check
b - b2 % close to machine epsilon i.e. zero

发生的情况是,我们在其中一个数组中插入了一个新的单一维度,最终得到一个大小为 [n, 1, n] 的数组,而另一个大小为 [n, n],后者隐含地与 [n, n, 1] 相同。重叠的第一个索引对应于循环中的i,其余两个索引对应于每个i 的二元乘积的矩阵索引。然后我们对索引进行置换,以便将“i”索引放在最后,这样我们就可以再次以(隐式)大小为[n, n, 1]S 广播结果。然后我们得到的是一个大小为[n, n, n] 的矩阵,其中前两个索引是原始矩阵中的索引,最后一个对应于i。然后我们只需要取平方并对每个项求和(而不是求和两次,我将数组重新整形为一行并求和一次)。

上述转置 S 的轻微变化而不是可能更快的 3d 数组(同样,您应该计时):

b3 = sum(reshape((reshape(X, [n, 1, n]).*Y - reshape(S.', [1, n, n])).^2, 1, []));

在性能方面,reshape 是免费的(它只会重新解释数据,不会复制),但permute/transpose 在复制数据时通常会导致性能下降。

【讨论】:

    猜你喜欢
    • 2019-07-12
    • 2014-11-10
    • 2012-02-07
    • 1970-01-01
    • 1970-01-01
    • 2020-07-09
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多