矩阵乘法:Strassen vs. Standard

mul*_*lle 4 c++ performance matrix matrix-multiplication strassen

我尝试使用C++ 实现Strassen算法进行矩阵乘法,但结果不是我所期望的.正如您所看到的,strassen总是花费更多时间,然后标准实现,并且只有2的幂的维度与标准实现一样快.什么地方出了错? 替代文字

matrix mult_strassen(matrix a, matrix b) {
if (a.dim() <= cut)
    return mult_std(a, b);

matrix a11 = get_part(0, 0, a);
matrix a12 = get_part(0, 1, a);
matrix a21 = get_part(1, 0, a);
matrix a22 = get_part(1, 1, a);

matrix b11 = get_part(0, 0, b);
matrix b12 = get_part(0, 1, b);
matrix b21 = get_part(1, 0, b);
matrix b22 = get_part(1, 1, b);

matrix m1 = mult_strassen(a11 + a22, b11 + b22); 
matrix m2 = mult_strassen(a21 + a22, b11);
matrix m3 = mult_strassen(a11, b12 - b22);
matrix m4 = mult_strassen(a22, b21 - b11);
matrix m5 = mult_strassen(a11 + a12, b22);
matrix m6 = mult_strassen(a21 - a11, b11 + b12);
matrix m7 = mult_strassen(a12 - a22, b21 + b22);

matrix c(a.dim(), false, true);
set_part(0, 0, &c, m1 + m4 - m5 + m7);
set_part(0, 1, &c, m3 + m5);
set_part(1, 0, &c, m2 + m4);
set_part(1, 1, &c, m1 - m2 + m3 + m6);

return c; 
}
Run Code Online (Sandbox Code Playgroud)


PROGRAM
matrix.h http://pastebin.com/TYFYCTY7
matrix.cpp http://pastebin.com/wYADLJ8Y
main.cpp http://pastebin.com/48BSqGJr

g++ main.cpp matrix.cpp -o matrix -O3.

ues*_*esp 8

一些想法:

  • 您是否优化过它以考虑用零填充非功率的两个大小的矩阵?我认为该算法假设您不打扰这些术语的倍增.这就是为什么你得到的运行时间在2 ^ n和2 ^(n + 1)-1之间的平坦区域.通过不将您知道的术语乘以零,您应该能够改进这些区域.或许Strassen只能用于2 ^ n大小的矩阵.
  • 考虑到"大"矩阵是任意的,并且该算法仅略微优于天真的情况,O(N ^ 3)对O(N ^ 2.8).在尝试更大的矩阵之前,您可能看不到可衡量的收益.例如,我做了一些有限元建模,其中10,000x10,000矩阵被认为是"小".很难从你的图表中看出来,但看起来在Stassen案例中511案例可能会更快.
  • 尝试使用各种优化级别进行测试,包括根本不进行优化.
  • 该算法似乎假设乘法比加法要昂贵得多.这在40年前首次开发时确实如此,但我相信更现代的处理器,加法和乘法之间的差异变小了.这可能会降低算法的有效性,这似乎会减少乘法,但会增加相加.
  • 你有没有看过其他一些Strassen实现的想法?尝试对已知良好的实现进行基准测试,以确切了解您可以获得多快的速度.

  • +1.600x600矩阵实际上非常小. (4认同)