【问题标题】:Why is my Strassen's Matrix Multiplication slow?为什么我的 Strassen 矩阵乘法很慢?
【发布时间】:2012-11-13 15:20:41
【问题描述】:

我用 C++ 编写了两个矩阵乘法程序:Regular MM (source) 和 Strassen 的 MM (source),它们都对大小为 2^kx 2^k 的方阵进行运算(换句话说,偶数大小的方阵) )。

结果很糟糕。对于 1024 x 1024 矩阵,Regular MM 采用 46.381 sec,而 Strassen 的 MM 采用 1484.303 sec (25 minutes !!!!)。

我试图使代码尽可能简单。在网上找到的其他 Strassen 的 MM 示例与我的代码没有太大区别。 Strassen 的代码的一个问题是显而易见的——我没有切换到常规 MM 的截止点。

我的 Strassen 的 MM 代码还有什么其他问题???

谢谢!

直接链接到来源
http://pastebin.com/HqHtFpq9
http://pastebin.com/USRQ5tuy

编辑1。 拳头,很多很好的建议。感谢您抽出宝贵时间分享知识。

我实施了更改(保留了我的所有代码),添加了截止点。 2048x2048 矩阵的 MM,截止 512 已经给出了很好的结果。 普通MM:191.49s 施特拉森的 MM:112.179s 很显着的提高。 结果是使用 Visual Studio 2012 在配备英特尔迅驰处理器的史前联想 X61 平板电脑上获得的。 我会做更多的检查(以确保我得到正确的结果),并将公布结果。

【问题讨论】:

  • @LuchianGrigore:哦,这很微妙。并且可能也是问题的很大一部分。可能比我实际发现的问题更大。
  • @Mysticial 我假设其中一种算法是缓存无意识的,因为尺寸固定为 2^k。可能是错的。
  • @LuchianGrigore:一种算法非常简单(并且不会忘记缓存)。另一个也不是缓存遗忘,尽管它应该快很多。
  • @LuchianGrigore:不幸的是,我相信施特拉森的算法是一种递归的分治算法。我怀疑它可以用于非 2 的幂。但我也怀疑 OPs 的实现不能。
  • @newprint 那么好吧,原因是here

标签: c++ performance optimization matrix-multiplication strassen


【解决方案1】:

Strassen 代码的一个问题很明显 - 我没有截止点, 切换到常规 MM。

可以公平地说,递归到 1 点是大部分(如果不是全部)问题。试图猜测其他性能瓶颈而不解决这个问题几乎没有实际意义,因为它会带来巨大的性能损失。 (换句话说,您是在将苹果与橙子进行比较。)

正如 cmets 中所讨论的,缓存对齐可能会产生影响,但不会达到这种规模。此外,缓存对齐对常规算法的伤害可能比 Strassen 算法更大,因为后者是缓存无意识的。

void strassen(int **a, int **b, int **c, int tam) {

    // trivial case: when the matrix is 1 X 1:
    if (tam == 1) {
            c[0][0] = a[0][0] * b[0][0];
            return;
    }

这太小了。虽然 Strassen 算法的复杂度较小,但它的 Big-O 常数要大得多。一方面,您的函数调用开销一直到 1 个元素。

这类似于使用合并或快速排序并一直递归到一个元素。为了提高效率,您需要在尺寸变小时停止递归并回退到经典算法。

在快速/合并排序中,您将退回到开销较低的O(n^2) 插入或选择排序。在这里,您将退回到正常的 O(n^3) 矩阵乘法。


您回退经典算法的阈值应该是一个可调阈值,该阈值可能会因硬件和编译器优化代码的能力而异。

对于像 Strassen 乘法,其优势仅在于 O(2.8074) 优于经典的 O(n^3),如果这个阈值非常高,请不要感到惊讶。 (数千个元素?)


在某些应用程序中,可能有许多算法,每个算法的复杂度都在降低,但 Big-O 会增加。结果是多种算法在不同大小下变得最优。

大整数乘法是一个臭名昭著的例子:

*请注意,这些示例阈值是近似值,可能会发生巨大变化 - 通常超过 10 倍。

【讨论】:

  • 这是我喜欢 StackOverflow 的原因之一。通过一个问题,我看到了现实世界的例子,其中可能产生性能问题的微妙影响被放大并以明显的方式展示。然后,当然,这个答案很可能是导致应该更快的算法变慢的原因。
  • 我测试过,对于天真的算法,我发现的问题很重要。但它应该会显着影响两种算法,因此可能不是所谓的更快算法性能不佳的原因。
  • @Omnifarious 是的,我希望二次方对齐惩罚不超过 3 倍。因为allfourmajorexamples所以只有大约 3 倍的性能影响。在这里,OP 有 30 倍的性能差异。
  • 你的分界线实在是太高了。
  • @Mystical - 哦,不,我的问题更简单。令我恼火的是每一行都是单独分配的,这会破坏引用的局部性并引入对每个元素访问的间接性。
【解决方案2】:

因此,可能还有更多问题,但您的第一个问题是您正在使用指向数组的指针数组。由于您使用的是 2 的幂的数组大小,因此与连续分配元素和使用整数除法将长数字数组折叠成行相比,这对性能造成了特别大的影响。

无论如何,这是我对问题的第一个猜测。正如我所说,可能还有更多,当我发现它们时,我会添加到这个答案中。

编辑:这可能只会导致问题的一小部分。问题很可能是Luchian Grigore所指的涉及cache line contention issues with powers of two

我验证了我的担忧对于朴素算法是有效的。如果数组是连续的,那么朴素算法的时间将减少近 50%。这里是the code for this (using a SquareMatrix class that is C++11 dependent) on pastebin

【讨论】:

  • 感谢您的帮助!早上我会看看你的代码。
  • @newprint:我犯了几个小错误,它们对您的代码没有影响,但使SquareMatrix 类对于一般使用不安全。我会修复它们。
  • 这里是简单版本- void madd(int N, int Xpitch, const double X[], int Ypitch, const double Y[], int Spitch, double S[]) { for (int i = 0; i Spitch + j] = X[iXpitch + j] + Y[i*Ypitch + j];
  • @newprint:我不喜欢那个版本,因为你必须记住在每次矩阵访问时都使用乘法。但它非常好用 C 语言,根本不使用 C++ 特性。 :-) 我的版本(带有内联函数)允许编译器做出一些有趣的假设并进行一些非常好的优化,同时允许实际的乘法算法看起来仍然干净利落地使用多维数组访问。
猜你喜欢
  • 2012-07-14
  • 1970-01-01
  • 2021-07-01
  • 2010-12-27
  • 1970-01-01
  • 2012-06-22
  • 1970-01-01
  • 1970-01-01
  • 2023-03-12
相关资源
最近更新 更多