矩阵乘法太慢了怎么办?
masonxiong · · 算法·理论
省流:对于
n=10^3 规模矩阵乘法用时50\space\rm ms ,约为朴素算法用时的\frac 1{55} 。
说在前面
作为线性代数的重要部分,矩阵乘法在 OI 界有非常广泛的应用,例如加速线性递推、表达修改标记、定长路径计数等等。因此,矩阵乘法的效率就显得至关重要。
很多人追求简洁,依照矩阵乘法的定义
为了打破这种瓶颈,追求更高效率,我们将从头开始,逐步引入复杂度更优的 Strassen 算法、大幅度优化常数的 AVX2 指令集等各种技术,逐步实现一个效率极其优秀的矩阵乘法。
记号约定
本文的矩阵乘法是指最常见的模意义下的
-
本文的所有同余都是模
998244353 意义下的。 -
所有测试均在本机
Intel(R) Core(TM) Ultra 9 285H进行。不包含 I/O 时间。 -
编译器为
g++ 15.2.0,编译参数-std=c++23 -Ofast -march=native。
朴素做法
我们从一些相对简单易懂的实现开始,逐步进行优化。
翻译定义
直接按照定义写出代码,这似乎是绝大部分人的选择。
:::success[
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;
:::
缓存友好
注意到朴素做法对
解决方案比较简单:
- 交换
j, k两层循环的顺序。交换后三个矩阵均为逐行访问。 - 将二维数组手动铺平为一维数组。
:::success[
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;
:::
分治思想
朴素做法若想继续优化需要考虑指令集。但是在使用指令集之前,让我们考虑一个问题:矩阵乘法最优只能
答案是否定的。接下来将详细介绍由 Strassen 提出的
八次乘法
考虑
上述写法看似只适用于标量意义下的
其中
也就是说:
证明:任取
\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} 。其余三个子块同理。
因此对于大小为
:::success[
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);
}
};
:::
七次乘法
上述分治做法虽然复杂度未变,但是它启发我们通过减少
下面是 Strassen 构造的
先构造若干中间和差
牛批。这种东西现在 AI 能构造出来吗?
展开即可验证正确。于是每层只需
:::success[
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 + v) % m可优化为x += v; x = std::min(x, x - m)。这可被向量化为如下代码:
X = _mm256_add_epi32(X, V);
X = _mm256_min_epu32(X, _mm256_sub_epi32(X, M));
- 减法:
(x - v) % m可优化为x -= v; x = std::min(x, x + m)。这可被向量化为如下代码:
X = _mm256_sub_epi32(X, V);
X = _mm256_min_epu32(X, _mm256_add_epi32(X, M));
注意向量化要求矩阵大小至少为
另外,尽量保证数组地址按照 load / store 而非 loadu / storeu。
:::success[
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);
}
};
:::
蒙哥马利
上述实现仍有优化空间。观察到递归边界的 -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 约减,这在运算数均为
为了避免这一影响,我们采用 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
这样我们成功避免了
直接的解决方案是:在递归前的矩阵转化时将每个元素转至蒙域。但是有一种更加高明的方法:在递归后的矩阵转化时将每个元素乘上
值得一提的是,这一操作也可以向量化!首先考虑标量形式:
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);
}
然后向量化。第一步是得到 _mm256_mul_epu32 得到两组每组四个 _mm256_blend_epi32 将两组结果按顺序重新混起来。
最后的结果介于 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[
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);
}
};
:::
还能再凹
叶子乘法也能向量化!
朴素
固定下标
这比标量“每个
:::success[
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
}
:::