【问题标题】:matlab/octave - Generalized matrix multiplicationmatlab/octave - 广义矩阵乘法
【发布时间】:2014-08-06 08:50:55
【问题描述】:

我想做一个泛化矩阵乘法的函数。基本上,它应该能够进行标准的矩阵乘法,但它应该允许通过任何其他函数来更改两个二元运算符的乘积/和。

目标是在 CPU 和内存方面尽可能高效。当然,它总是比 A*B 效率低,但操作员的灵活性是这里的重点。

以下是我在阅读variousinterestingthreads后可以提出的一些命令:

A = randi(10, 2, 3);
B = randi(10, 3, 4);

% 1st method
C = sum(bsxfun(@mtimes, permute(A,[1 3 2]),permute(B,[3 2 1])), 3)
% Alternative: C = bsxfun(@(a,b) mtimes(a',b), A', permute(B, [1 3 2]))

% 2nd method
C = sum(bsxfun(@(a,b) a*b, permute(A,[1 3 2]),permute(B,[3 2 1])), 3)

% 3rd method (Octave-only)
C = sum(permute(A, [1 3 2]) .* permute(B, [3 2 1]), 3)

% 4th method (Octave-only): multiply nxm A with nx1xd B to create a nxmxd array
C = bsxfun(@(a, b) sum(times(a,b)), A', permute(B, [1 3 2]));
C = C2 = squeeze(C(1,:,:)); % sum and turn into mxd

方法 1-3 的问题是它们会在使用 sum() 折叠它们之前生成 n 个矩阵。 4 更好,因为它在 bsxfun 中执行 sum(),但 bsxfun 仍然生成 n 个矩阵(除了它们大部分是空的,仅包含一个非零值向量作为总和,其余填充为 0 以匹配尺寸要求)。

我想要的是第 4 种方法,但没有无用的 0 来节省内存。

有什么想法吗?

【问题讨论】:

  • 为什么不尝试使用稀疏矩阵来节省内存分配?你也许可以让它发挥作用。 spfun 类似于 bsxfun,但是对于稀疏矩阵,所以我假设它在后台也保持内存使用率非常低。
  • 已经完成了,确实第 4 种方法应该能够从稀疏中获利,但不幸的是它不适用于 Octave,因为它的 bsxfun 运算符对稀疏不友好,所以所有内容都将存储在内存中.
  • 您的第三个和第四个示例不起作用。您的第一个示例不适用于 MATLAB R2010b 及更早版本。
  • 我的另一个问题是,您正在处理的矩阵有多大,以至于您如此关心内存。
  • @RodyOldenhuis:感谢您的反馈,是的,确实第三和第四只在 Octave 上工作。寻找替代方案的另一个原因,因为第 4 种方法的问题正是我要解决的问题:输出维度不正确。

标签: matlab matrix octave matrix-multiplication


【解决方案1】:

在检查了 bsxfun 等几个处理函数后,似乎无法使用这些函数进行直接矩阵乘法(我的意思是直接的意思是临时乘积不存储在内存中,而是尽快求和,然后其他sum-products 被处理),因为它们具有固定大小的输出(或者与输入相同,或者使用 bsxfun 单例扩展两个输入维度的笛卡尔积)。然而,可以稍微欺骗 Octave(这不适用于检查输出尺寸的 MatLab):

C = bsxfun(@(a,b) sum(bsxfun(@times, a, B))', A', sparse(1, size(A,1)))
C = bsxfun(@(a,b) sum(bsxfun(@times, a, B))', A', zeros(1, size(A,1), 2))(:,:,2)

但是不要使用它们,因为输出的值不可靠(Octave 可以破坏甚至删除它们并返回 0!)。

所以现在我只是实现一个半矢量化版本,这是我的功能:

function C = genmtimes(A, B, outop, inop)
% C = genmtimes(A, B, inop, outop)
% Generalized matrix multiplication between A and B. By default, standard sum-of-products matrix multiplication is operated, but you can change the two operators (inop being the element-wise product and outop the sum).
% Speed note: about 100-200x slower than A*A' and about 3x slower when A is sparse, so use this function only if you want to use a different set of inop/outop than the standard matrix multiplication.

if ~exist('inop', 'var')
    inop = @times;
end

if ~exist('outop', 'var')
    outop = @sum;
end

[n, m] = size(A);
[m2, o] = size(B);

if m2 ~= m
    error('nonconformant arguments (op1 is %ix%i, op2 is %ix%i)\n', n, m, m2, o);
end


C = [];
if issparse(A) || issparse(B)
    C = sparse(o,n);
else
    C = zeros(o,n);
end

A = A';
for i=1:n
    C(:,i) = outop(bsxfun(inop, A(:,i), B))';
end
C = C';

end

用稀疏矩阵和正常矩阵进行测试:稀疏矩阵(慢 3 倍)的性能差距比正常矩阵(慢约 100 倍)要小得多。

我认为这比 bsxfun 实现要慢,但至少不会溢出内存:

A = randi(10, 1000);
C = genmtimes(A, A');

如果有人能提供更好的,我仍在寻找更好的选择!

【讨论】:

    【解决方案2】:

    无需深入了解细节,mtimesxMMX 等工具是快速通用矩阵和标量运算例程。您可以查看他们的代码并根据您的需要进行调整。 它很可能比 matlab 的 bsxfun 更快。

    【讨论】:

    • +1 本地化 C/C++ 代码绝对是这里的路
    • 我同意这将是速度方面和内存方面的最佳解决方案,但我不确定您是否可以将函数作为 GMM 的参数传递,尽管我想可以定义最常见的运算符在 mex 文件中并将字符串作为参数传递给选择。还有一个 LAPACK 直接调用 FEX 可能适合这里的账单,甚至无需重写任何 MEX 代码:mathworks.com/matlabcentral/fileexchange/16777-lapack/content/…
    • 我尝试了 LAPACK FEX 库,但它不适用于 Octave,我不知道最新的 MatLab 版本。我发现了另一个有趣的项目:Mc2For 项目,旨在从 MatLab 函数源代码生成 Fortran 95 代码:sable.mcgill.ca/mclab/matlab_fortran.html
    【解决方案3】:

    为什么不直接利用bsxfun 接受任意函数的能力?

    C = shiftdim(feval(f, (bsxfun(g, A.', permute(B,[1 3 2])))), 1);
    

    这里

    • f外部函数(对应于矩阵乘法中的 sum)。它应该接受任意大小的 3D 数组 mxnxp 并沿其列操作以返回 1xmxp 数组。
    • g内部函数(对应于矩阵乘法中的 product)。根据bsxfun,它应该接受两个相同大小的列向量,或者一个列向量和一个标量作为输入,并返回一个与输入大小相同的列向量作为输出。

    这在 Matlab 中有效。我没有在 Octave 中测试过。


    示例 1:矩阵乘法:

    >> f = @sum;   %// outer function: sum
    >> g = @times; %// inner function: product
    >> A = [1 2 3; 4 5 6];
    >> B = [10 11; -12 -13; 14 15];
    >> C = shiftdim(feval(f, (bsxfun(g, A.', permute(B,[1 3 2])))), 1)
    C =
        28    30
        64    69
    

    检查:

    >> A*B
    ans =
        28    30
        64    69
    

    示例 2:考虑上述两个矩阵与

    >> f = @(x,y) sum(abs(x));     %// outer function: sum of absolute values
    >> g = @(x,y) max(x./y, y./x); %// inner function: "symmetric" ratio
    >> C = shiftdim(feval(f, (bsxfun(g, A.', permute(B,[1 3 2])))), 1)
    C =
       14.8333   16.1538
        5.2500    5.6346
    

    检查:手动计算C(1,2)

    >> sum(abs( max( (A(1,:))./(B(:,2)).', (B(:,2)).'./(A(1,:)) ) ))
    ans =
       16.1538
    

    【讨论】:

    • 感谢您的详细回答,但这与 Divakar 上面提出的非常相似,问题是内存爆炸,因为 bsxfun 无法对每个单例展开求和积(而不是产生整个展开然后在第三维求和)。例如,尝试在 Octave 上使用这个矩阵: A = randi(2, 1000, 500)-1;这将产生索引溢出。此外,这不适用于稀疏矩阵,因为如果稀疏则无法置换到第 3 维。
    【解决方案4】:

    这是您发布的解决方案的稍微完善的版本,并进行了一些小的改进。

    我们检查行数是否多于列数或相反,然后通过选择行与矩阵相乘或矩阵与列相乘来进行相应的乘法运算(从而执行最少的循环迭代)。

    注意:这可能并不总是最好的策略(按行而不是按列),即使行数少于列数; MATLAB 数组存储在内存中的 column-major order 中这一事实使得按列切片更有效,因为元素是连续存储的。而访问行涉及通过strides 遍历元素(这对缓存不友好——想想spatial locality)。

    除此之外,代码应处理双/单、实数/复数、完整/稀疏(以及不可能组合的错误)。它还尊重空矩阵和零维。

    function C = my_mtimes(A, B, outFcn, inFcn)
        % default arguments
        if nargin < 4, inFcn = @times; end
        if nargin < 3, outFcn = @sum; end
    
        % check valid input
        assert(ismatrix(A) && ismatrix(B), 'Inputs must be 2D matrices.');
        assert(isequal(size(A,2),size(B,1)),'Inner matrix dimensions must agree.');
        assert(isa(inFcn,'function_handle') && isa(outFcn,'function_handle'), ...
            'Expecting function handles.')
    
        % preallocate output matrix
        M = size(A,1);
        N = size(B,2);
        if issparse(A)
            args = {'like',A};
        elseif issparse(B)
            args = {'like',B};
        else
            args = {superiorfloat(A,B)};
        end
        C = zeros(M,N, args{:});
    
        % compute matrix multiplication
        % http://en.wikipedia.org/wiki/Matrix_multiplication#Inner_product
        if M < N
            % concatenation of products of row vectors with matrices
            % A*B = [a_1*B ; a_2*B ; ... ; a_m*B]
            for m=1:M
                %C(m,:) = A(m,:) * B;
                %C(m,:) = sum(bsxfun(@times, A(m,:)', B), 1);
                C(m,:) = outFcn(bsxfun(inFcn, A(m,:)', B), 1);
            end
        else
            % concatenation of products of matrices with column vectors
            % A*B = [A*b_1 , A*b_2 , ... , A*b_n]
            for n=1:N
                %C(:,n) = A * B(:,n);
                %C(:,n) = sum(bsxfun(@times, A, B(:,n)'), 2);
                C(:,n) = outFcn(bsxfun(inFcn, A, B(:,n)'), 2);
            end
        end
    end
    

    比较

    这个函数无疑是整个过程都比较慢,但是对于更大的尺寸,它比内置的矩阵乘法差几个数量级:

            (tic/toc times in seconds)
          (tested in R2014a on Windows 8)
    
        size      mtimes       my_mtimes 
        ____    __________     _________
         400     0.0026398       0.20282
         600      0.012039       0.68471
         800      0.014571        1.6922
        1000      0.026645        3.5107
        2000       0.20204         28.76
        4000        1.5578        221.51
    

    这里是测试代码:

    sz = [10:10:100 200:200:1000 2000 4000];
    t = zeros(numel(sz),2);
    for i=1:numel(sz)
        n = sz(i); disp(n)
        A = rand(n,n);
        B = rand(n,n);
    
        tic
        C = A*B;
        t(i,1) = toc;
        tic
        D = my_mtimes(A,B);
        t(i,2) = toc;
    
        assert(norm(C-D) < 1e-6)
        clear A B C D
    end
    
    semilogy(sz, t*1000, '.-')
    legend({'mtimes','my_mtimes'}, 'Interpreter','none', 'Location','NorthWest')
    xlabel('Size N'), ylabel('Time [msec]'), title('Matrix Multiplication')
    axis tight
    

    额外

    为了完整起见,下面是实现广义矩阵乘法的两种更简单的方法(如果要比较性能,请将my_mtimes 函数的最后一部分替换为其中任何一种)。我什至不会费心发布他们经过的时间:)

    C = zeros(M,N, args{:});
    for m=1:M
        for n=1:N
            %C(m,n) = A(m,:) * B(:,n);
            %C(m,n) = sum(bsxfun(@times, A(m,:)', B(:,n)));
            C(m,n) = outFcn(bsxfun(inFcn, A(m,:)', B(:,n)));
        end
    end
    

    另一种方式(使用三重循环):

    C = zeros(M,N, args{:});
    P = size(A,2); % = size(B,1);
    for m=1:M
        for n=1:N
            for p=1:P
                %C(m,n) = C(m,n) + A(m,p)*B(p,n);
                %C(m,n) = plus(C(m,n), times(A(m,p),B(p,n)));
                C(m,n) = outFcn([C(m,n) inFcn(A(m,p),B(p,n))]);
            end
        end
    end
    

    接下来要尝试什么?

    如果您想获得更多性能,您将不得不迁移到 C/C++ MEX 文件以减少解释 MATLAB 代码的开销。您仍然可以通过从 MEX 文件中调用优化的 BLAS/LAPACK 例程来利用它们(例如,请参阅 the second part of this post)。 MATLAB 附带 Intel MKL 库,坦率地说,在 Intel 处理器上进行线性代数计算时,您无法击败它。

    其他人已经在 File Exchange 上提到了一些将通用矩阵例程实现为 MEX 文件的提交(请参阅@natan 的答案)。如果您将它们与优化的 BLAS 库链接起来,它们会特别有效。

    【讨论】:

    • M 开关的妙招,但是不应该是 M N(因为我们想减少循环的数量?),不应该是 C(m, :) = outFcn(bsxfun(inFcn, A(m,:), B), 1);是 C(m,:) = outFcn(bsxfun(inFcn, A(:,m), B), 1);而是完全使用列主顺序(对于稀疏矩阵更重要)?还有一点需要注意的是,您的代码使用了一些不适用于 Octave 3.8.1 的函数(args{:},superiorfloat)。但是,我确实使用您的代码获得了更快的速度,因此您现在有了最接近的答案。
    • @user1121352: 1) 是的,我的错,应该与M&lt;N 正好相反。我会修复它。 2) 不,它是正确的(我认为你错过了A(m,:)' 中的转置)。这个想法是水平连接A的行“乘以”矩阵B。 3)我没有用Octave测试它,但它应该很容易适应代码。语法 zeros(.., 'like',X) 还没有进入 Octave,你可以用类似的 zeros(.., class(X)) 替换它(虽然不会选择稀疏属性)。
    • ... 至于superiorfloat 调用,您可以将其替换为手动检查以在single/double 之间选择适当的类型(这是因为the way data types propagate in MATLAB 组合时不同的数字类)
    • 在之前的评论中,我应该说“垂直连接”而不是“水平连接”。我还翻转了代码 cmets 中的描述。现在已修复:/ 我还添加了关于按行与按列遍历的注释...
    • 感谢您的精确度,但您真的确定访问 A'(:,m) 而不是 A(m,:)' 会更好,这样 A 会被按列切片大步的?同样,C 可以转置,然后可以使用 C(:,m) 而不是 C(m,:) 填充?无论如何,我会奖励你赏金,但我会尝试的最后一件事是使用 LAPACK 。如果这不起作用,我会接受你的回答。
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2020-04-09
    • 2015-01-07
    • 2012-09-25
    • 1970-01-01
    相关资源
    最近更新 更多