矩阵乘法太慢了怎么办?

· · 算法·理论

省流:对于 n=10^3 规模矩阵乘法用时 50\space\rm ms,约为朴素算法用时的 \frac 1{55}。

说在前面

作为线性代数的重要部分,矩阵乘法在 OI 界有非常广泛的应用,例如加速线性递推、表达修改标记、定长路径计数等等。因此,矩阵乘法的效率就显得至关重要。

很多人追求简洁,依照矩阵乘法的定义 \mathrm C_{i,j}=\sum\mathrm A_{i,k}\cdot\mathrm B_{k,j} 直接实现,然而这种做法无论时间复杂度还是常数都不够优秀,经常成为程序效率的瓶颈。

为了打破这种瓶颈,追求更高效率,我们将从头开始,逐步引入复杂度更优的 Strassen 算法、大幅度优化常数的 AVX2 指令集等各种技术,逐步实现一个效率极其优秀的矩阵乘法。

记号约定

本文的矩阵乘法是指最常见的模意义下的 (+,\times) 矩阵乘法。也就是说,对于矩阵 \mathrm A_{u,v} 和 \mathrm B_{v,w} 定义乘积矩阵 \mathrm C_{u,w} 为:

\mathrm C_{i,j}=\left(\sum_{k=1}^v\mathrm A_{i,k}\cdot\mathrm B_{k,j}\right)\bmod 998244353

朴素做法

我们从一些相对简单易懂的实现开始,逐步进行优化。

翻译定义

直接按照定义写出代码,这似乎是绝大部分人的选择。

:::success[u=v=w=10^3 用时 \approx 28000\space\rm ms]

for (auto i = 0U; i != u; ++i)
    for (auto j = 0U; j != w; ++j)
        for (auto k = 0U; k != v; ++k)
            C[i][j] = (C[i][j] + 1ULL * A[i][k] * B[k][j]) % Modulus;

:::

缓存友好

注意到朴素做法对 \rm B 矩阵的访问是逐列进行的,地址并不连续,缓存命中率低。同时一般的动态二维数组写法两行之间地址不连续,这也会对效率产生影响。

解决方案比较简单:

:::success[u=v=w=10^3 用时 \approx 550\space\rm ms]

for (auto i = 0U; i != u; ++i)
    for (auto k = 0U; k != v; ++k)
        for (auto j = 0U; j != w; ++j)
            C[i * u + j] = (C[i * u + j] + 1ULL * A[i * u + k] * B[k * v + j]) % Modulus;

:::

分治思想

朴素做法若想继续优化需要考虑指令集。但是在使用指令集之前,让我们考虑一个问题:矩阵乘法最优只能 \mathcal O(n^3) 吗?

答案是否定的。接下来将详细介绍由 Strassen 提出的 \mathcal O(n^{\log_27})\approx\mathcal O(n^{2.80735}) 基于分治思想的矩阵乘法算法,这是首个时间复杂度低于 \mathcal O(n^3) 的矩阵乘法算法。目前最优算法可以做到 \mathcal O(n^{2.37134}),但是 Strassen 算法在常数上有绝对优势。

八次乘法

考虑 2\times 2 矩阵乘法 \rm C=\rm A\times\rm B:

\begin{bmatrix} \mathrm C_{0,0}&\mathrm C_{0,1}\\ \mathrm C_{1,0}&\mathrm C_{1,1} \end{bmatrix}= \begin{bmatrix} \mathrm A_{0,0}&\mathrm A_{0,1}\\ \mathrm A_{1,0}&\mathrm A_{1,1} \end{bmatrix}\times \begin{bmatrix} \mathrm B_{0,0}&\mathrm B_{0,1}\\ \mathrm B_{1,0}&\mathrm B_{1,1} \end{bmatrix}\\ \begin{aligned}\\ \mathrm C_{0,0}&=\mathrm A_{0,0}\cdot\mathrm B_{0,0}+\mathrm A_{0,1}\cdot\mathrm B_{1,0}\\ \mathrm C_{0,1}&=\mathrm A_{0,0}\cdot\mathrm B_{0,1}+\mathrm A_{0,1}\cdot\mathrm B_{1,1}\\ \mathrm C_{1,0}&=\mathrm A_{1,0}\cdot\mathrm B_{0,0}+\mathrm A_{1,1}\cdot\mathrm B_{1,0}\\ \mathrm C_{1,1}&=\mathrm A_{1,0}\cdot\mathrm B_{0,1}+\mathrm A_{1,1}\cdot\mathrm B_{1,1}\\ \end{aligned}

上述写法看似只适用于标量意义下的 2\times 2 矩阵,但实际上可以推广到任意偶数大小方阵的分块乘法。设 n 为偶数,m=\frac n2。将 n\times n 矩阵 \mathrm A,\mathrm B,\mathrm C 按如下方式分成四个 m\times m 子块:

\mathrm A= \begin{bmatrix} \mathrm A_{00}&\mathrm A_{01}\\ \mathrm A_{10}&\mathrm A_{11} \end{bmatrix},\quad \mathrm B= \begin{bmatrix} \mathrm B_{00}&\mathrm B_{01}\\ \mathrm B_{10}&\mathrm B_{11} \end{bmatrix},\quad \mathrm C= \begin{bmatrix} \mathrm C_{00}&\mathrm C_{01}\\ \mathrm C_{10}&\mathrm C_{11} \end{bmatrix}

其中 \mathrm A_{ij},\mathrm B_{ij},\mathrm C_{ij} 均为 m\times m 矩阵。若 \mathrm C=\mathrm A\times\mathrm B 则有

\begin{aligned} \mathrm C_{00}&=\mathrm A_{00}\mathrm B_{00}+\mathrm A_{01}\mathrm B_{10}\\ \mathrm C_{01}&=\mathrm A_{00}\mathrm B_{01}+\mathrm A_{01}\mathrm B_{11}\\ \mathrm C_{10}&=\mathrm A_{10}\mathrm B_{00}+\mathrm A_{11}\mathrm B_{10}\\ \mathrm C_{11}&=\mathrm A_{10}\mathrm B_{01}+\mathrm A_{11}\mathrm B_{11} \end{aligned}

也就是说:2\times 2 情形中的“标量乘加”可原样替换为“子矩阵乘加”,公式形式不变。

证明:任取 \mathrm C 中位于左上子块的下标 (i,j)\space(i,j\in[0,m))。根据矩阵乘法定义:

\mathrm C_{i,j} =\sum_{k=0}^{n-1}\mathrm A_{i,k}\mathrm B_{k,j} =\sum_{k=0}^{m-1}\mathrm A_{i,k}\mathrm B_{k,j} +\sum_{k=m}^{n-1}\mathrm A_{i,k}\mathrm B_{k,j}

第一段求和恰为 (\mathrm A_{00}\mathrm B_{00})_{i,j}:行在 \mathrm A_{00} 内、列在 \mathrm B_{00} 内。第二段中令 k'=k-m,则 0\le k'<m,且 \mathrm A_{i,k}=\mathrm A_{i,\,m+k'} 属于 \mathrm A_{01},\mathrm B_{k,j}=\mathrm B_{m+k',\,j} 属于 \mathrm B_{10},故该段等于 (\mathrm A_{01}\mathrm B_{10})_{i,j}。

于是 \mathrm C_{i,j}=(\mathrm A_{00}\mathrm B_{00}+\mathrm A_{01}\mathrm B_{10})_{i,j}=\mathrm A_{00}\mathrm B_{00}+\mathrm A_{01}\mathrm B_{10}。其余三个子块同理。

因此对于大小为 n 的矩阵,可递归将问题四分,每层做 8 次规模减半的矩阵乘法和 4 次矩阵加法。时间复杂度 \mathcal T(n)=8\mathcal T(\frac n2)+\mathcal O(n^2)=\mathcal O(n^3) 与朴素做法相同。

:::success[u=v=w=10^3 用时 \approx 105480\space\rm ms]

auto multiplies = [](auto& self, auto A, auto B, auto C, auto n) -> void {
    if (n == 1) {
        C[0] = static_cast<unsigned>(1ULL * A[0] * B[0] % Modulus);
        return;
    }

    auto m = n / 2;
    std::vector<unsigned> P(m * m);
    std::vector<unsigned> Q(m * m);

    std::vector<unsigned> A00(m * m), A01(m * m), A10(m * m), A11(m * m);
    std::vector<unsigned> B00(m * m), B01(m * m), B10(m * m), B11(m * m);
    std::vector<unsigned> C00(m * m), C01(m * m), C10(m * m), C11(m * m);

    for (auto i = 0ULL; i != m; ++i) {
        std::copy_n(A + i * n + 0, m, A00.data() + i * m);
        std::copy_n(A + i * n + m, m, A01.data() + i * m);
        std::copy_n(A + (i + m) * n + 0, m, A10.data() + i * m);
        std::copy_n(A + (i + m) * n + m, m, A11.data() + i * m);
    }
    for (auto i = 0ULL; i != m; ++i) {
        std::copy_n(B + i * n + 0, m, B00.data() + i * m);
        std::copy_n(B + i * n + m, m, B01.data() + i * m);
        std::copy_n(B + (i + m) * n + 0, m, B10.data() + i * m);
        std::copy_n(B + (i + m) * n + m, m, B11.data() + i * m);
    }

    self(self, A00.data(), B00.data(), P.data(), m);
    self(self, A01.data(), B10.data(), Q.data(), m);
    for (auto i = 0ULL; i != m * m; ++i) C00[i] = (P[i] + Q[i]) % Modulus;

    self(self, A00.data(), B01.data(), P.data(), m);
    self(self, A01.data(), B11.data(), Q.data(), m);
    for (auto i = 0ULL; i != m * m; ++i) C01[i] = (P[i] + Q[i]) % Modulus;

    self(self, A10.data(), B00.data(), P.data(), m);
    self(self, A11.data(), B10.data(), Q.data(), m);
    for (auto i = 0ULL; i != m * m; ++i) C10[i] = (P[i] + Q[i]) % Modulus;

    self(self, A10.data(), B01.data(), P.data(), m);
    self(self, A11.data(), B11.data(), Q.data(), m);
    for (auto i = 0ULL; i != m * m; ++i) C11[i] = (P[i] + Q[i]) % Modulus;

    for (auto i = 0ULL; i != m; ++i) {
        std::copy_n(C00.data() + i * m, m, C + i * n + 0);
        std::copy_n(C01.data() + i * m, m, C + i * n + m);
        std::copy_n(C10.data() + i * m, m, C + (i + m) * n + 0);
        std::copy_n(C11.data() + i * m, m, C + (i + m) * n + m);
    }
};

:::

七次乘法

上述分治做法虽然复杂度未变,但是它启发我们通过减少 2\times 2 矩阵乘法所需乘法次数来优化矩阵乘法的时间复杂度。基于这一启示,Strassen 构造了一系列匪夷所思的算式,将 2\times 2 矩阵乘法所需乘法次数降低到了 7 次,进而将矩阵乘法时间复杂度降低到了 \mathcal O(n^{\log_27})。

下面是 Strassen 构造的 7 次乘法实现 2\times 2 矩阵乘法 \rm C=\rm A\times\rm B:

\begin{bmatrix} \mathrm C_{0,0}&\mathrm C_{0,1}\\ \mathrm C_{1,0}&\mathrm C_{1,1} \end{bmatrix}= \begin{bmatrix} \mathrm A_{0,0}&\mathrm A_{0,1}\\ \mathrm A_{1,0}&\mathrm A_{1,1} \end{bmatrix}\times \begin{bmatrix} \mathrm B_{0,0}&\mathrm B_{0,1}\\ \mathrm B_{1,0}&\mathrm B_{1,1} \end{bmatrix}

先构造若干中间和差 \mathrm S_i,再做 7 次乘法得到 \mathrm P_i:

\begin{aligned} \mathrm S_0&=\mathrm B_{0,1}-\mathrm B_{1,1},& \mathrm P_0&=\mathrm A_{0,0}\mathrm S_0,\\ \mathrm S_1&=\mathrm A_{0,0}+\mathrm A_{0,1},& \mathrm P_1&=\mathrm S_1\mathrm B_{1,1},\\ \mathrm S_2&=\mathrm A_{1,0}+\mathrm A_{1,1},& \mathrm P_2&=\mathrm S_2\mathrm B_{0,0},\\ \mathrm S_3&=\mathrm B_{1,0}-\mathrm B_{0,0},& \mathrm P_3&=\mathrm A_{1,1}\mathrm S_3,\\ \mathrm S_4&=\mathrm A_{0,0}+\mathrm A_{1,1},& \mathrm S_7&=\mathrm B_{0,0}+\mathrm B_{1,1},& \mathrm P_4&=\mathrm S_4\mathrm S_7,\\ \mathrm S_5&=\mathrm A_{1,0}-\mathrm A_{0,0},& \mathrm S_8&=\mathrm B_{0,0}+\mathrm B_{0,1},& \mathrm P_5&=\mathrm S_5\mathrm S_8,\\ \mathrm S_6&=\mathrm A_{0,1}-\mathrm A_{1,1},& \mathrm S_9&=\mathrm B_{1,0}+\mathrm B_{1,1},& \mathrm P_6&=\mathrm S_6\mathrm S_9. \end{aligned}\\ \begin{aligned} \mathrm C_{0,0}&=\mathrm P_4+\mathrm P_3-\mathrm P_1+\mathrm P_6,\\ \mathrm C_{0,1}&=\mathrm P_0+\mathrm P_1,\\ \mathrm C_{1,0}&=\mathrm P_2+\mathrm P_3,\\ \mathrm C_{1,1}&=\mathrm P_4-\mathrm P_2+\mathrm P_0+\mathrm P_5. \end{aligned}

牛批。这种东西现在 AI 能构造出来吗?

展开即可验证正确。于是每层只需 7 次规模减半的矩阵乘法,时间复杂度 \mathcal T(n)=7\mathcal T(\frac n2)+\mathcal O(n^2)=\mathcal O(n^{\log_27})。

:::success[u=v=w=10^3 用时 \approx 68840\space\rm ms]

auto multiplies = [](auto& self, auto A, auto B, auto C, auto n) -> void {
    if (n == 1) {
        C[0] = static_cast<unsigned>(1ULL * A[0] * B[0] % Modulus);
        return;
    }

    auto m = n / 2;
    std::vector<unsigned> P0(m * m);
    std::vector<unsigned> P1(m * m);
    std::vector<unsigned> P2(m * m);
    std::vector<unsigned> P3(m * m);
    std::vector<unsigned> P4(m * m);
    std::vector<unsigned> P5(m * m);
    std::vector<unsigned> P6(m * m);
    std::vector<unsigned> S0(m * m), S1(m * m), S2(m * m), S3(m * m), S4(m * m);
    std::vector<unsigned> S5(m * m), S6(m * m), S7(m * m), S8(m * m), S9(m * m);
    std::vector<unsigned> A00(m * m), A01(m * m), A10(m * m), A11(m * m);
    std::vector<unsigned> B00(m * m), B01(m * m), B10(m * m), B11(m * m);
    std::vector<unsigned> C00(m * m), C01(m * m), C10(m * m), C11(m * m);

    for (auto i = 0ULL; i != m; ++i) {
        std::copy_n(A + i * n + 0, m, A00.data() + i * m);
        std::copy_n(A + i * n + m, m, A01.data() + i * m);
        std::copy_n(A + (i + m) * n + 0, m, A10.data() + i * m);
        std::copy_n(A + (i + m) * n + m, m, A11.data() + i * m);
    }
    for (auto i = 0ULL; i != m; ++i) {
        std::copy_n(B + i * n + 0, m, B00.data() + i * m);
        std::copy_n(B + i * n + m, m, B01.data() + i * m);
        std::copy_n(B + (i + m) * n + 0, m, B10.data() + i * m);
        std::copy_n(B + (i + m) * n + m, m, B11.data() + i * m);
    }

    for (auto i = 0ULL; i != m * m; ++i) {
        S0[i] = (B01[i] - B11[i] + Modulus) % Modulus;
        S1[i] = (A00[i] + A01[i] + Modulus) % Modulus;
        S2[i] = (A10[i] + A11[i] + Modulus) % Modulus;
        S3[i] = (B10[i] - B00[i] + Modulus) % Modulus;
        S4[i] = (A00[i] + A11[i] + Modulus) % Modulus;
        S5[i] = (B00[i] + B11[i] + Modulus) % Modulus;
        S6[i] = (A01[i] - A11[i] + Modulus) % Modulus;
        S7[i] = (B10[i] + B11[i] + Modulus) % Modulus;
        S8[i] = (A00[i] - A10[i] + Modulus) % Modulus;
        S9[i] = (B00[i] + B01[i] + Modulus) % Modulus;
    }

    self(self, A00.data(), S0.data(), P0.data(), m);
    self(self, S1.data(), B11.data(), P1.data(), m);
    self(self, S2.data(), B00.data(), P2.data(), m);
    self(self, A11.data(), S3.data(), P3.data(), m);
    self(self, S4.data(), S5.data(), P4.data(), m);
    self(self, S6.data(), S7.data(), P5.data(), m);
    self(self, S8.data(), S9.data(), P6.data(), m);

    for (auto i = 0ULL; i != m * m; ++i) {
        C00[i] = (P4[i] + P3[i] - P1[i] + P5[i] + Modulus) % Modulus;
        C01[i] = (P0[i] + P1[i]) % Modulus;
        C10[i] = (P2[i] + P3[i]) % Modulus;
        C11[i] = (P4[i] + P0[i] - P2[i] - P6[i] + Modulus + Modulus) % Modulus;
    }

    for (auto i = 0ULL; i != m; ++i) {
        std::copy_n(C00.data() + i * m, m, C + i * n + 0);
        std::copy_n(C01.data() + i * m, m, C + i * n + m);
        std::copy_n(C10.data() + i * m, m, C + (i + m) * n + 0);
        std::copy_n(C11.data() + i * m, m, C + (i + m) * n + m);
    }
};

:::

力大砖飞

上述实现效率低下主要几点原因以及解决方案:

然后优化加法 & 减法取模部分:

X = _mm256_add_epi32(X, V);
X = _mm256_min_epu32(X, _mm256_sub_epi32(X, M));
X = _mm256_sub_epi32(X, V);
X = _mm256_min_epu32(X, _mm256_add_epi32(X, M));

注意向量化要求矩阵大小至少为 8\times 8。当矩阵大小为 8 时,直接调用朴素做法。特别地,由于此时矩阵规模小,可以统一求和后取模,并且不用在意求和顺序(8\times 8 大小的矩阵肯定全进缓存)。

另外,尽量保证数组地址按照 32 字节对齐,这样可以使用更高效的 load / store 而非 loadu / storeu。

:::success[u=v=w=10^3 用时 \approx 130\space\rm ms]

using Scalar = uint32_t;
using Vector = __m256i;

static const Scalar sModulus = 998244353U;
static const Scalar sInverse = 998244351U;
static const Scalar sRsquare = 932051910U;

static const Vector vModulus = _mm256_set1_epi32(sModulus);
static const Vector vInverse = _mm256_set1_epi32(sInverse);
static const Vector vRsquare = _mm256_set1_epi32(sRsquare);

static const unsigned sSize = sizeof(Scalar);
static const unsigned vSize = sizeof(Vector) / sSize;

static auto shrink = [](auto V) -> Vector {
    return _mm256_min_epu32(V, _mm256_sub_epi32(V, vModulus));
};

static auto dilate = [](auto V) -> Vector {
    return _mm256_min_epu32(V, _mm256_add_epi32(V, vModulus));
};

auto place = [](auto& self, auto u, auto v, auto n, auto N, auto o, auto O) -> void {
    if (n == vSize) {
        for (auto i = 0U; i != vSize; ++i) {
            auto value = _mm256_load_si256(reinterpret_cast<Vector*>(o + (v + i) * N + u));
            _mm256_store_si256(reinterpret_cast<Vector*>(O + i * vSize), value);
        }
    } else {
        auto m = n / 2;
        auto M = m * m;
        self(self, u + 0 * m, v + 0 * m, m, N, o, O + 0 * M);
        self(self, u + 1 * m, v + 0 * m, m, N, o, O + 1 * M);
        self(self, u + 0 * m, v + 1 * m, m, N, o, O + 2 * M);
        self(self, u + 1 * m, v + 1 * m, m, N, o, O + 3 * M);
    }
};

auto antiplace = [](auto& self, auto u, auto v, auto n, auto N, auto o, auto O) -> void {
    if (n == vSize) {
        for (auto i = 0U; i != vSize; ++i) {
            auto value = _mm256_load_si256(reinterpret_cast<Vector*>(O + i * vSize));
            _mm256_store_si256(reinterpret_cast<Vector*>(o + (v + i) * N + u), value);
        }
    } else {
        auto m = n / 2;
        auto M = m * m;
        self(self, u + 0 * m, v + 0 * m, m, N, o, O + 0 * M);
        self(self, u + 1 * m, v + 0 * m, m, N, o, O + 1 * M);
        self(self, u + 0 * m, v + 1 * m, m, N, o, O + 2 * M);
        self(self, u + 1 * m, v + 1 * m, m, N, o, O + 3 * M);
    }
};

auto multiplies = [](auto& self, auto A, auto B, auto C, auto n, auto P, auto Q, auto R) -> void {
    if (n == vSize) {
        for (auto i = 0U; i != vSize; ++i) {
            for (auto j = 0U; j != vSize; ++j) {
                auto sum = 0ULL;
                for (auto k = 0U; k != vSize; ++k)
                    sum += 1ULL * A[i * vSize + k] * B[k * vSize + j];
                C[i * vSize + j] = Scalar(sum % sModulus);
            }
        }
    } else {
        auto m = n / 2;
        auto M = m * m;

        auto plus = [M](auto U, auto V, auto W) -> void {
            for (auto i = 0U; i != M; i += vSize) {
                auto Ui = _mm256_load_si256(reinterpret_cast<Vector*>(U + i));
                auto Vi = _mm256_load_si256(reinterpret_cast<Vector*>(V + i));
                auto Wi = shrink(_mm256_add_epi32(Ui, Vi));
                _mm256_store_si256(reinterpret_cast<Vector*>(W + i), Wi);
            }
        };

        auto minus = [M](auto U, auto V, auto W) -> void {
            for (auto i = 0U; i != M; i += vSize) {
                auto Ui = _mm256_load_si256(reinterpret_cast<Vector*>(U + i));
                auto Vi = _mm256_load_si256(reinterpret_cast<Vector*>(V + i));
                auto Wi = dilate(_mm256_sub_epi32(Ui, Vi));
                _mm256_store_si256(reinterpret_cast<Vector*>(W + i), Wi);
            }
        };

        auto i00 = M * 0, i01 = M * 1;
        auto i10 = M * 2, i11 = M * 3;

        minus(A + i01, A + i11, P);
        plus(B + i10, B + i11, Q);
        self(self, P, Q, C + i00, m, P + M, Q + M, R);

        plus(A + i00, A + i01, P);
        self(self, P, B + i11, C + i01, m, P + M, Q + M, R);
        minus(C + i00, C + i01, C + i00);

        minus(B + i10, B + i00, Q);
        self(self, A + i11, Q, C + i10, m, P + M, Q + M, R);
        plus(C + i00, C + i10, C + i00);

        minus(A + i10, A + i00, P);
        plus(B + i00, B + i01, Q);
        self(self, P, Q, C + i11, m, P + M, Q + M, R);

        plus(A + i10, A + i11, P);
        self(self, P, B + i00, R, m, P + M, Q + M, R + M);
        plus(C + i10, R, C + i10);
        minus(C + i11, R, C + i11);

        plus(A + i00, A + i11, P);
        plus(B + i00, B + i11, Q);
        self(self, P, Q, R, m, P + M, Q + M, R + M);
        plus(C + i00, R, C + i00);
        plus(C + i11, R, C + i11);

        minus(B + i01, B + i11, Q);
        self(self, A + i00, Q, R, m, P + M, Q + M, R + M);
        plus(C + i01, R, C + i01);
        plus(C + i11, R, C + i11);
    }
};

:::

蒙哥马利

上述实现仍有优化空间。观察到递归边界的 8\times 8 矩阵乘法(称其为叶子部分)会调用 64\space\rm bits 常量取模,在 Compiler Explorer 进行测试(-std=c++23 -Ofast -march=native):

auto M64S(uint64_t sum) -> uint32_t {
    static constexpr uint32_t Modulus = 998244353U;
    return uint32_t(sum % Modulus);
}
"M64S(unsigned long)":
        movabs  rax, -8525806094425994177
        mul     rdi  ; RAX * RDI -> 128 位结果在 RDX:RAX
        mov     eax, edi
        shr     rdx, 29
        imul    rdx, rdx, 998244353
        sub     eax, edx
        ret

编译器底层对模运算的优化一般使用 Barrett 约减,这在运算数均为 64 位整数时必须调用 128\space\rm bits 乘法,延迟大同时依赖链长,对效率产生极大影响。

为了避免这一影响,我们采用 Montgomery Reduction 手动进行优化。

auto M64M(uint64_t sum) -> uint32_t {
    static constexpr uint32_t Modulus = 998244353U;
    static constexpr uint32_t Inverse = 998244351U;
    auto reduced = uint32_t((sum + uint64_t(uint32_t(sum) * Inverse) * Modulus) >> 32);
    reduced = std::min(reduced, reduced - Modulus * 2);
    reduced = std::min(reduced, reduced - Modulus * 1);
    return reduced;
}
"M64M(unsigned long)":
        imul    eax, edi, 998244351
        imul    rax, rax, 998244353
        add     rax, rdi
        shr     rax, 32
        mov     edx, eax
        add     edx, -1996488706
        cmovnc  edx, eax
        mov     eax, edx
        add     eax, -998244353
        cmovnc  eax, edx
        ret

这样我们成功避免了 128\space\rm bits 乘法带来的效率影响,但是同时带来了另一个问题:Montgomery Reduction 要求运算数均在蒙域(即:表示为 x\rm R^{-1} 的形式,其中 \rm R=2^{32}),因为会造成 \rm R^{-1} 的缩放。

直接的解决方案是:在递归前的矩阵转化时将每个元素转至蒙域。但是有一种更加高明的方法:在递归后的矩阵转化时将每个元素乘上 \rm R^2 消除约减带来的影响。这样做的优势在于,只用对一个矩阵进行操作。

值得一提的是,这一操作也可以向量化!首先考虑标量形式:

auto MR2(uint32_t value) -> uint32_t {
    static constexpr uint32_t Modulus = 998244353U;
    static constexpr uint32_t Rsquare = 932051910U;
    uint64_t product = 1ULL * value * Rsquare;
    auto result = uint32_t((product + uint64_t(uint32_t(product) * Inverse) * Modulus) >> 32);
    return std::min(result, result - Modulus);
}

然后向量化。第一步是得到 64\space\rm bits 乘积,需要拆奇偶位然后使用 _mm256_mul_epu32 得到两组每组四个 64\space\rm bits 乘积。接下来每组各自约减,完成后将偶数位组结果右移 32\space\rm bits 到应有的位置,使用 _mm256_blend_epi32 将两组结果按顺序重新混起来。

最后的结果介于 [0,2\rm M) 之间,再调用一次 shrink 将其规约即可。

auto value = _mm256_load_si256(reinterpret_cast<Vector*>(O + i * vSize));
auto P = _mm256_mul_epu32(value, vRsquare);
auto Q = _mm256_mul_epu32(_mm256_srli_epi64(value, 32), vRsquare);
auto U = _mm256_mul_epu32(_mm256_mul_epu32(P, vInverse), vModulus);
auto V = _mm256_mul_epu32(_mm256_mul_epu32(Q, vInverse), vModulus);
auto X = _mm256_srli_epi64(_mm256_add_epi64(U, P), 32);
auto Y = _mm256_add_epi64(V, Q);
auto T = shrink(_mm256_blend_epi32(X, Y, 0xAA));

:::success[u=v=w=10^3 用时 \approx 80\space\rm ms]

using Scalar = uint32_t;
using Vector = __m256i;

static const Scalar sModulus = 998244353U;
static const Scalar sInverse = 998244351U;
static const Scalar sRsquare = 932051910U;

[[maybe_unused]] static const Vector vModulus = _mm256_set1_epi32(sModulus);
[[maybe_unused]] static const Vector vInverse = _mm256_set1_epi32(sInverse);
[[maybe_unused]] static const Vector vRsquare = _mm256_set1_epi32(sRsquare);

static const unsigned sSize = sizeof(Scalar);
static const unsigned vSize = sizeof(Vector) / sSize;

static auto shrink = [](auto V) -> Vector {
    return _mm256_min_epu32(V, _mm256_sub_epi32(V, vModulus));
};

static auto dilate = [](auto V) -> Vector {
    return _mm256_min_epu32(V, _mm256_add_epi32(V, vModulus));
};

auto place = [](auto& self, auto u, auto v, auto n, auto N, auto o, auto O) -> void {
    if (n == vSize) {
        for (auto i = 0U; i != vSize; ++i) {
            auto value = _mm256_load_si256(reinterpret_cast<Vector*>(o + (v + i) * N + u));
            _mm256_store_si256(reinterpret_cast<Vector*>(O + i * vSize), value);
        }
    } else {
        auto m = n / 2;
        auto M = m * m;
        self(self, u + 0 * m, v + 0 * m, m, N, o, O + 0 * M);
        self(self, u + 1 * m, v + 0 * m, m, N, o, O + 1 * M);
        self(self, u + 0 * m, v + 1 * m, m, N, o, O + 2 * M);
        self(self, u + 1 * m, v + 1 * m, m, N, o, O + 3 * M);
    }
};

auto antiplace = [](auto& self, auto u, auto v, auto n, auto N, auto o, auto O) -> void {
    if (n == vSize) {
        for (auto i = 0U; i != vSize; ++i) {
            auto value = _mm256_load_si256(reinterpret_cast<Vector*>(O + i * vSize));
            auto P = _mm256_mul_epu32(value, vRsquare);
            auto Q = _mm256_mul_epu32(_mm256_srli_epi64(value, 32), vRsquare);
            auto U = _mm256_mul_epu32(_mm256_mul_epu32(P, vInverse), vModulus);
            auto V = _mm256_mul_epu32(_mm256_mul_epu32(Q, vInverse), vModulus);
            auto X = _mm256_srli_epi64(_mm256_add_epi64(U, P), 32);
            auto Y = _mm256_add_epi64(V, Q);
            auto T = shrink(_mm256_blend_epi32(X, Y, 0xAA));
            _mm256_store_si256(reinterpret_cast<Vector*>(o + (v + i) * N + u), T);
        }
    } else {
        auto m = n / 2;
        auto M = m * m;
        self(self, u + 0 * m, v + 0 * m, m, N, o, O + 0 * M);
        self(self, u + 1 * m, v + 0 * m, m, N, o, O + 1 * M);
        self(self, u + 0 * m, v + 1 * m, m, N, o, O + 2 * M);
        self(self, u + 1 * m, v + 1 * m, m, N, o, O + 3 * M);
    }
};

auto multiplies = [](auto& self, auto A, auto B, auto C, auto n, auto P, auto Q, auto R) -> void {
    if (n == vSize) {
        for (auto i = 0U; i != vSize; ++i) {
            for (auto j = 0U; j != vSize; ++j) {
                auto sum = 0ULL;
                for (auto o = 0U; o != vSize; ++o)
                    sum += 1ULL * A[i * vSize + o] * B[o * vSize + j];
                auto reduced = Scalar((sum + std::uint64_t(Scalar(sum) * sInverse) * sModulus) >> 32);
                reduced = std::min(reduced, reduced - sModulus * 2);
                reduced = std::min(reduced, reduced - sModulus * 1);
                C[i * vSize + j] = reduced;
            }
        }
    } else {
        auto m = n / 2;
        auto M = m * m;

        auto plus = [M](auto U, auto V, auto W) -> void {
            for (auto i = 0U; i != M; i += vSize) {
                auto Ui = _mm256_load_si256(reinterpret_cast<Vector*>(U + i));
                auto Vi = _mm256_load_si256(reinterpret_cast<Vector*>(V + i));
                auto Wi = shrink(_mm256_add_epi32(Ui, Vi));
                _mm256_store_si256(reinterpret_cast<Vector*>(W + i), Wi);
            }
        };

        auto minus = [M](auto U, auto V, auto W) -> void {
            for (auto i = 0U; i != M; i += vSize) {
                auto Ui = _mm256_load_si256(reinterpret_cast<Vector*>(U + i));
                auto Vi = _mm256_load_si256(reinterpret_cast<Vector*>(V + i));
                auto Wi = dilate(_mm256_sub_epi32(Ui, Vi));
                _mm256_store_si256(reinterpret_cast<Vector*>(W + i), Wi);
            }
        };

        auto i00 = M * 0, i01 = M * 1;
        auto i10 = M * 2, i11 = M * 3;

        minus(A + i01, A + i11, P);
        plus(B + i10, B + i11, Q);
        self(self, P, Q, C + i00, m, P + M, Q + M, R);

        plus(A + i00, A + i01, P);
        self(self, P, B + i11, C + i01, m, P + M, Q + M, R);
        minus(C + i00, C + i01, C + i00);

        minus(B + i10, B + i00, Q);
        self(self, A + i11, Q, C + i10, m, P + M, Q + M, R);
        plus(C + i00, C + i10, C + i00);

        minus(A + i10, A + i00, P);
        plus(B + i00, B + i01, Q);
        self(self, P, Q, C + i11, m, P + M, Q + M, R);

        plus(A + i10, A + i11, P);
        self(self, P, B + i00, R, m, P + M, Q + M, R + M);
        plus(C + i10, R, C + i10);
        minus(C + i11, R, C + i11);

        plus(A + i00, A + i11, P);
        plus(B + i00, B + i11, Q);
        self(self, P, Q, R, m, P + M, Q + M, R + M);
        plus(C + i00, R, C + i00);
        plus(C + i11, R, C + i11);

        minus(B + i01, B + i11, Q);
        self(self, A + i00, Q, R, m, P + M, Q + M, R + M);
        plus(C + i01, R, C + i01);
        plus(C + i11, R, C + i11);
    }
};

:::

还能再凹

叶子乘法也能向量化!

朴素 8\times 8 叶子需要三重循环:对每个 (i,j) 累加 8 次乘积,再调用一次 Montgomery Reduction。问题在于,叶子会被 Strassen 递归调用极多次,每次输出都要约减仍然较慢;且标量循环无法复用 \mathrm B 的行数据。

固定下标 x 可得 \mathrm C\xleftarrow+\mathrm A_{:,x}\mathrm B_{x,:} 就是 \rm A 的 x 列 8 个数与 \rm B 的 x 行 8 个数相乘。此处 \mathrm B_{x,:} 地址连续可以直接读入,\mathrm A_{:,x} 对于每个输出行只需要广播一次即可同时更新该行 8 个输出。

这比标量“每个 (i,j) 枚举 k”效率更高,因为 \rm B 的一行被 8 个输出行复用。另外注意此处仍需得到 64\space\rm bits 乘积,需要拆奇偶位。所有累加结束之后使用前面所提的向量化约减即可。

:::success[u=v=w=10^3 用时 \approx 50\space\rm ms]

using Scalar = uint32_t;
using Vector = __m256i;

static const Scalar sModulus = 998244353U;
static const Scalar sInverse = 998244351U;
static const Scalar sRsquare = 932051910U;

[[maybe_unused]] static const Vector vModulus = _mm256_set1_epi32(sModulus);
[[maybe_unused]] static const Vector vInverse = _mm256_set1_epi32(sInverse);
[[maybe_unused]] static const Vector vRsquare = _mm256_set1_epi32(sRsquare);

static const unsigned sSize = sizeof(Scalar);
static const unsigned vSize = sizeof(Vector) / sSize;

static auto shrink = [](auto V) -> Vector {
    return _mm256_min_epu32(V, _mm256_sub_epi32(V, vModulus));
};

static auto dilate = [](auto V) -> Vector {
    return _mm256_min_epu32(V, _mm256_add_epi32(V, vModulus));
};

auto place = [](auto& self, auto u, auto v, auto n, auto N, auto o, auto O) -> void {
    if (n == vSize) {
        for (auto i = 0U; i != vSize; ++i) {
            auto value = _mm256_load_si256(reinterpret_cast<Vector*>(o + (v + i) * N + u));
            _mm256_store_si256(reinterpret_cast<Vector*>(O + i * vSize), value);
        }
    } else {
        auto m = n / 2;
        auto M = m * m;
        self(self, u + 0 * m, v + 0 * m, m, N, o, O + 0 * M);
        self(self, u + 1 * m, v + 0 * m, m, N, o, O + 1 * M);
        self(self, u + 0 * m, v + 1 * m, m, N, o, O + 2 * M);
        self(self, u + 1 * m, v + 1 * m, m, N, o, O + 3 * M);
    }
};

auto antiplace = [](auto& self, auto u, auto v, auto n, auto N, auto o, auto O) -> void {
    if (n == vSize) {
        for (auto i = 0U; i != vSize; ++i) {
            auto value = _mm256_load_si256(reinterpret_cast<Vector*>(O + i * vSize));
            auto P = _mm256_mul_epu32(value, vRsquare);
            auto Q = _mm256_mul_epu32(_mm256_srli_epi64(value, 32), vRsquare);
            auto U = _mm256_mul_epu32(_mm256_mul_epu32(P, vInverse), vModulus);
            auto V = _mm256_mul_epu32(_mm256_mul_epu32(Q, vInverse), vModulus);
            auto X = _mm256_srli_epi64(_mm256_add_epi64(U, P), 32);
            auto Y = _mm256_add_epi64(V, Q);
            auto T = shrink(_mm256_blend_epi32(X, Y, 0xAA));
            _mm256_store_si256(reinterpret_cast<Vector*>(o + (v + i) * N + u), T);
        }
    } else {
        auto m = n / 2;
        auto M = m * m;
        self(self, u + 0 * m, v + 0 * m, m, N, o, O + 0 * M);
        self(self, u + 1 * m, v + 0 * m, m, N, o, O + 1 * M);
        self(self, u + 0 * m, v + 1 * m, m, N, o, O + 2 * M);
        self(self, u + 1 * m, v + 1 * m, m, N, o, O + 3 * M);
    }
};

auto multiplies = [](auto& self, auto A, auto B, auto C, auto n, auto P, auto Q, auto R) -> void {
    if (n == vSize) {
        Vector s[vSize]{};
        Vector S[vSize]{};
        for (auto w = 0U; w != vSize; ++w) {
            auto o = _mm256_load_si256(reinterpret_cast<Vector*>(B + w * vSize));
            auto O = _mm256_srli_epi64(o, 32);
            for (auto i = 0U; i != vSize; ++i) {
                s[i] = _mm256_add_epi64(s[i], _mm256_mul_epu32(_mm256_set1_epi32(int(A[i * vSize + w])), o));
                S[i] = _mm256_add_epi64(S[i], _mm256_mul_epu32(_mm256_set1_epi32(int(A[i * vSize + w])), O));
            }
        }
        for (auto i = 0U; i != vSize; ++i) {
            auto U = _mm256_mul_epu32(_mm256_mul_epu32(s[i], vInverse), vModulus);
            auto V = _mm256_mul_epu32(_mm256_mul_epu32(S[i], vInverse), vModulus);
            auto X = _mm256_srli_epi64(_mm256_add_epi64(U, s[i]), 32);
            auto Y = _mm256_add_epi64(V, S[i]);
            auto T = shrink(_mm256_blend_epi32(X, Y, 0xAA));
            T = shrink(_mm256_min_epu32(T, _mm256_sub_epi32(T, _mm256_set1_epi32(sModulus * 2))));
            _mm256_store_si256(reinterpret_cast<Vector*>(C + i * vSize), T);
        }
    } else {
        auto m = n / 2;
        auto M = m * m;

        auto plus = [M](auto U, auto V, auto W) -> void {
            for (auto i = 0U; i != M; i += vSize) {
                auto Ui = _mm256_load_si256(reinterpret_cast<Vector*>(U + i));
                auto Vi = _mm256_load_si256(reinterpret_cast<Vector*>(V + i));
                auto Wi = shrink(_mm256_add_epi32(Ui, Vi));
                _mm256_store_si256(reinterpret_cast<Vector*>(W + i), Wi);
            }
        };

        auto minus = [M](auto U, auto V, auto W) -> void {
            for (auto i = 0U; i != M; i += vSize) {
                auto Ui = _mm256_load_si256(reinterpret_cast<Vector*>(U + i));
                auto Vi = _mm256_load_si256(reinterpret_cast<Vector*>(V + i));
                auto Wi = dilate(_mm256_sub_epi32(Ui, Vi));
                _mm256_store_si256(reinterpret_cast<Vector*>(W + i), Wi);
            }
        };

        auto i00 = M * 0, i01 = M * 1;
        auto i10 = M * 2, i11 = M * 3;

        minus(A + i01, A + i11, P);
        plus(B + i10, B + i11, Q);
        self(self, P, Q, C + i00, m, P + M, Q + M, R);

        plus(A + i00, A + i01, P);
        self(self, P, B + i11, C + i01, m, P + M, Q + M, R);
        minus(C + i00, C + i01, C + i00);

        minus(B + i10, B + i00, Q);
        self(self, A + i11, Q, C + i10, m, P + M, Q + M, R);
        plus(C + i00, C + i10, C + i00);

        minus(A + i10, A + i00, P);
        plus(B + i00, B + i01, Q);
        self(self, P, Q, C + i11, m, P + M, Q + M, R);

        plus(A + i10, A + i11, P);
        self(self, P, B + i00, R, m, P + M, Q + M, R + M);
        plus(C + i10, R, C + i10);
        minus(C + i11, R, C + i11);

        plus(A + i00, A + i11, P);
        plus(B + i00, B + i11, Q);
        self(self, P, Q, R, m, P + M, Q + M, R + M);
        plus(C + i00, R, C + i00);
        plus(C + i11, R, C + i11);

        minus(B + i01, B + i11, Q);
        self(self, A + i00, Q, R, m, P + M, Q + M, R + M);
        plus(C + i01, R, C + i01);
        plus(C + i11, R, C + i11);
    }
};

:::

说在后面

本来在 Library Checker 上面是最优解的,但是在本文完成之前被人抢了。等我抢回去会再讲讲怎么接着卡(

:::success[完整代码]

#include <bits/extc++.h>
#include <immintrin.h>

auto main() -> int {
    std::cin.tie(nullptr)->sync_with_stdio(false);

    using Scalar = uint32_t;
    using Vector = __m256i;

    static const Scalar sModulus = 998244353U;
    static const Scalar sInverse = 998244351U;
    static const Scalar sRsquare = 932051910U;

    [[maybe_unused]] static const Vector vModulus = _mm256_set1_epi32(sModulus);
    [[maybe_unused]] static const Vector vInverse = _mm256_set1_epi32(sInverse);
    [[maybe_unused]] static const Vector vRsquare = _mm256_set1_epi32(sRsquare);

    static const unsigned sSize = sizeof(Scalar);
    static const unsigned vSize = sizeof(Vector) / sSize;

    static auto shrink = [](auto V) -> Vector {
        return _mm256_min_epu32(V, _mm256_sub_epi32(V, vModulus));
    };

    static auto dilate = [](auto V) -> Vector {
        return _mm256_min_epu32(V, _mm256_add_epi32(V, vModulus));
    };

    auto place = [](auto& self, auto u, auto v, auto n, auto N, auto o, auto O) -> void {
        if (n == vSize) {
            for (auto i = 0U; i != vSize; ++i) {
                auto value = _mm256_load_si256(reinterpret_cast<Vector*>(o + (v + i) * N + u));
                _mm256_store_si256(reinterpret_cast<Vector*>(O + i * vSize), value);
            }
        } else {
            auto m = n / 2;
            auto M = m * m;
            self(self, u + 0 * m, v + 0 * m, m, N, o, O + 0 * M);
            self(self, u + 1 * m, v + 0 * m, m, N, o, O + 1 * M);
            self(self, u + 0 * m, v + 1 * m, m, N, o, O + 2 * M);
            self(self, u + 1 * m, v + 1 * m, m, N, o, O + 3 * M);
        }
    };

    auto antiplace = [](auto& self, auto u, auto v, auto n, auto N, auto o, auto O) -> void {
        if (n == vSize) {
            for (auto i = 0U; i != vSize; ++i) {
                auto value = _mm256_load_si256(reinterpret_cast<Vector*>(O + i * vSize));
                auto P = _mm256_mul_epu32(value, vRsquare);
                auto Q = _mm256_mul_epu32(_mm256_srli_epi64(value, 32), vRsquare);
                auto U = _mm256_mul_epu32(_mm256_mul_epu32(P, vInverse), vModulus);
                auto V = _mm256_mul_epu32(_mm256_mul_epu32(Q, vInverse), vModulus);
                auto X = _mm256_srli_epi64(_mm256_add_epi64(U, P), 32);
                auto Y = _mm256_add_epi64(V, Q);
                auto T = shrink(_mm256_blend_epi32(X, Y, 0xAA));
                _mm256_store_si256(reinterpret_cast<Vector*>(o + (v + i) * N + u), T);
            }
        } else {
            auto m = n / 2;
            auto M = m * m;
            self(self, u + 0 * m, v + 0 * m, m, N, o, O + 0 * M);
            self(self, u + 1 * m, v + 0 * m, m, N, o, O + 1 * M);
            self(self, u + 0 * m, v + 1 * m, m, N, o, O + 2 * M);
            self(self, u + 1 * m, v + 1 * m, m, N, o, O + 3 * M);
        }
    };

    auto multiplies = [](auto& self, auto A, auto B, auto C, auto n, auto P, auto Q, auto R) -> void {
        if (n == vSize) {
            Vector s[vSize]{};
            Vector S[vSize]{};
            for (auto w = 0U; w != vSize; ++w) {
                auto o = _mm256_load_si256(reinterpret_cast<Vector*>(B + w * vSize));
                auto O = _mm256_srli_epi64(o, 32);
                for (auto i = 0U; i != vSize; ++i) {
                    s[i] = _mm256_add_epi64(s[i], _mm256_mul_epu32(_mm256_set1_epi32(int(A[i * vSize + w])), o));
                    S[i] = _mm256_add_epi64(S[i], _mm256_mul_epu32(_mm256_set1_epi32(int(A[i * vSize + w])), O));
                }
            }
            for (auto i = 0U; i != vSize; ++i) {
                auto U = _mm256_mul_epu32(_mm256_mul_epu32(s[i], vInverse), vModulus);
                auto V = _mm256_mul_epu32(_mm256_mul_epu32(S[i], vInverse), vModulus);
                auto X = _mm256_srli_epi64(_mm256_add_epi64(U, s[i]), 32);
                auto Y = _mm256_add_epi64(V, S[i]);
                auto T = shrink(_mm256_blend_epi32(X, Y, 0xAA));
                T = shrink(_mm256_min_epu32(T, _mm256_sub_epi32(T, _mm256_set1_epi32(sModulus * 2))));
                _mm256_store_si256(reinterpret_cast<Vector*>(C + i * vSize), T);
            }
        } else {
            auto m = n / 2;
            auto M = m * m;

            auto plus = [M](auto U, auto V, auto W) -> void {
                for (auto i = 0U; i != M; i += vSize) {
                    auto Ui = _mm256_load_si256(reinterpret_cast<Vector*>(U + i));
                    auto Vi = _mm256_load_si256(reinterpret_cast<Vector*>(V + i));
                    auto Wi = shrink(_mm256_add_epi32(Ui, Vi));
                    _mm256_store_si256(reinterpret_cast<Vector*>(W + i), Wi);
                }
            };

            auto minus = [M](auto U, auto V, auto W) -> void {
                for (auto i = 0U; i != M; i += vSize) {
                    auto Ui = _mm256_load_si256(reinterpret_cast<Vector*>(U + i));
                    auto Vi = _mm256_load_si256(reinterpret_cast<Vector*>(V + i));
                    auto Wi = dilate(_mm256_sub_epi32(Ui, Vi));
                    _mm256_store_si256(reinterpret_cast<Vector*>(W + i), Wi);
                }
            };

            auto i00 = M * 0, i01 = M * 1;
            auto i10 = M * 2, i11 = M * 3;

            minus(A + i01, A + i11, P);
            plus(B + i10, B + i11, Q);
            self(self, P, Q, C + i00, m, P + M, Q + M, R);

            plus(A + i00, A + i01, P);
            self(self, P, B + i11, C + i01, m, P + M, Q + M, R);
            minus(C + i00, C + i01, C + i00);

            minus(B + i10, B + i00, Q);
            self(self, A + i11, Q, C + i10, m, P + M, Q + M, R);
            plus(C + i00, C + i10, C + i00);

            minus(A + i10, A + i00, P);
            plus(B + i00, B + i01, Q);
            self(self, P, Q, C + i11, m, P + M, Q + M, R);

            plus(A + i10, A + i11, P);
            self(self, P, B + i00, R, m, P + M, Q + M, R + M);
            plus(C + i10, R, C + i10);
            minus(C + i11, R, C + i11);

            plus(A + i00, A + i11, P);
            plus(B + i00, B + i11, Q);
            self(self, P, Q, R, m, P + M, Q + M, R + M);
            plus(C + i00, R, C + i00);
            plus(C + i11, R, C + i11);

            minus(B + i01, B + i11, Q);
            self(self, A + i00, Q, R, m, P + M, Q + M, R + M);
            plus(C + i01, R, C + i01);
            plus(C + i11, R, C + i11);
        }
    };

    unsigned u, v, w;
    u = v = w = 1000;

    std::mt19937 rng;
    auto n = std::max(vSize, std::bit_ceil(std::max({u, v, w})));
    auto m = std::align_val_t(alignof(Vector));

    auto A = static_cast<Scalar*>(::operator new(n * n * sSize, m));
    auto B = static_cast<Scalar*>(::operator new(n * n * sSize, m));
    auto C = static_cast<Scalar*>(::operator new(n * n * sSize, m));
    std::memset(A, 0, n * n * sSize);
    std::memset(B, 0, n * n * sSize);
    std::memset(C, 0, n * n * sSize);

    for (auto i = 0U; i != u; ++i)
        for (auto j = 0U; j != v; ++j)
            A[i * n + j] = rng() % sModulus;
    for (auto i = 0U; i != v; ++i)
        for (auto j = 0U; j != w; ++j)
            B[i * n + j] = rng() % sModulus;

    auto oA = static_cast<Scalar*>(::operator new(n * n * sSize, m));
    auto oB = static_cast<Scalar*>(::operator new(n * n * sSize, m));
    auto oC = static_cast<Scalar*>(::operator new(n * n * sSize, m));
    auto P = static_cast<Scalar*>(::operator new(n * n * sSize, m));
    auto Q = static_cast<Scalar*>(::operator new(n * n * sSize, m));
    auto R = static_cast<Scalar*>(::operator new(n * n * sSize, m));
    auto begin = std::clock();
    place(place, 0U, 0U, n, n, A, oA);
    place(place, 0U, 0U, n, n, B, oB);
    multiplies(multiplies, oA, oB, oC, n, P, Q, R);
    antiplace(antiplace, 0U, 0U, n, n, C, oC);
    auto end = std::clock();

    auto xorsum = 0U;
    for (auto i = 0U; i != u; ++i)
        for (auto j = 0U; j != w; ++j)
            xorsum ^= C[i * n + j];
    std::cout << "Duration = " << end - begin << " clocks\n"; // 50
    std::cout << "XorSum = " << xorsum << '\n'; // 219566061
}

:::