1 条题解
-
1
专题:再来一遍一定记住的算法_proMatheus的博客-CSDN博客
🎶当你的天空突然下起了大雨🌧
🍀那是我在为你炸乌云❤
快速沃尔什变换(Fast Walsh-Hadamard Transform, FWT)
是一种用于处理位运算卷积的算法。
它与 FFT 类 似,但 FFT 处理的是加法卷积。
而 FWT 处理的是按位或(OR)、按位与(AND)、按位异或(XOR)等位运算卷积。
洛谷模板:https://www.luogu.com.cn/problem/P4717

void fwt_or(LL a[], int flag) { for (int mid = 1; mid < limit; mid <<= 1) { // 从小到大枚举二进制位 for (int i = 0, R = (mid << 1); i < limit; i += R) { // 根据二进制位划分段 for (int j = 0; j < mid; j ++) { // 枚举段内数 // 此时 i + mid + j 是保证 mid 那一位为 1 // i + j 则保证 mid 位为 0 if (flag == false) { // 正变换 a[i + mid + j] = (a[i + mid + j] + a[i + j]) % P; // 加上自己的子集 } else { // 逆变换 a[i + mid + j] = (a[i + mid + j] - a[i + j] + P) % P; // 减掉自己的子集 } } } } // 因为 mid 是从小到大枚举,所以能保证 i 在用到它之前就把自己的所有子集都加好了 // 而且每次给 a[i + mid + j] 绝对是不同的 mid // 这样就能保证不重不漏 }

void fwt_and(LL a[], int flag) { for (int mid = 1; mid < limit; mid <<= 1) { // 从小到大枚举二进制位 for (int i = 0, R = (mid << 1); i < limit; i += R) { // 根据二进制位划分段 for (int j = 0; j < mid; j ++) { // 枚举段内数 // 此时 i + mid + j 是保证 mid 那一位为 1 // i + j 则保证 mid 位为 0 if (flag == false) { // 正变换 a[i + j] = (a[i + j] + a[i + mid + j]) % P; // 加上自己的超集 } else { // 逆变换 a[i + j] = (a[i + j] - a[i + mid + j] + P) % P; // 减掉自己的超集 } } } } // 因为 mid 是从小到大枚举,所以能保证 i 在用到它之前就把自己的所有超集都加好了 // 而且每次给 a[i + mid + j] 绝对是不同的 mid // 这样就能保证不重不漏 }

void fwt_xor(LL a[], bool flag) { for (int mid = 1; mid < limit; mid <<= 1) { // 从小到大枚举二进制位 for (int i = 0, R = (mid << 1); i < limit; i += R) { // 根据二进制位划分段 for (int j = 0; j < mid; j ++) { // 枚举段内数 // 此时 i + mid + j 是保证 mid 那一位为 1 // i + j 则保证 mid 位为 0 LL x = a[i + j]; LL y = a[i + j + mid]; a[i + j] = (x + y) % P; // 对于 i + j 来说 i + j + mid 在它的范围内 1 是都有的 // 也就是 (i + j + mid) & (i + j) 的 1 的个数和 (i + j) 的 1 的个数相同 // -1 的幂次也一定一样,所以正负号相同 a[i + j + mid] = (x - y + P) % P; // 对于 i + j + mid 来说 i + j 在它的范围内 1 少一个 // 也就是 (i + j) & (i + j + mid) 的 1 的个数比 (i + j + mid) 的 1 的个数少一个 // -1 的幂次也一定不一样,所以正负号不相同 if (flag) { // 逆变换要除以 2 // 你每个数都除以 log_n^2 个 2,那不就是除以 n 吗 a[i + j] = a[i + j] * inv2 % P; a[i + j + mid] = a[i + j + mid] * inv2 % P; } } } } }

#include<bits/stdc++.h> using namespace std; typedef long long LL; const LL P = 998244353; const LL inv2 = 499122177; // 2 的逆元 const int N = 20; LL a[1 << N], b[1 << N], c[1 << N]; LL aa[1 << N], bb[1 << N]; int limit; void fwt_or(LL a[], int flag) { for (int mid = 1; mid < limit; mid <<= 1) { // 从小到大枚举二进制位 for (int i = 0, R = (mid << 1); i < limit; i += R) { // 根据二进制位划分段 for (int j = 0; j < mid; j ++) { // 枚举段内数 // 此时 i + mid + j 是保证 mid 那一位为 1 // i + j 则保证 mid 位为 0 if (flag == false) { // 正变换 a[i + mid + j] = (a[i + mid + j] + a[i + j]) % P; // 加上自己的子集 } else { // 逆变换 a[i + mid + j] = (a[i + mid + j] - a[i + j] + P) % P; // 减掉自己的子集 } } } } // 因为 mid 是从小到大枚举,所以能保证 i 在用到它之前就把自己的所有子集都加好了 // 而且每次给 a[i + mid + j] 绝对是不同的 mid // 这样就能保证不重不漏 } void fwt_and(LL a[], int flag) { for (int mid = 1; mid < limit; mid <<= 1) { // 从小到大枚举二进制位 for (int i = 0, R = (mid << 1); i < limit; i += R) { // 根据二进制位划分段 for (int j = 0; j < mid; j ++) { // 枚举段内数 // 此时 i + mid + j 是保证 mid 那一位为 1 // i + j 则保证 mid 位为 0 if (flag == false) { // 正变换 a[i + j] = (a[i + j] + a[i + mid + j]) % P; // 加上自己的超集 } else { // 逆变换 a[i + j] = (a[i + j] - a[i + mid + j] + P) % P; // 减掉自己的超集 } } } } // 因为 mid 是从小到大枚举,所以能保证 i 在用到它之前就把自己的所有超集都加好了 // 而且每次给 a[i + mid + j] 绝对是不同的 mid // 这样就能保证不重不漏 } // flag = false 正变换,flag = true 逆变换 void fwt_xor(LL a[], bool flag) { for (int mid = 1; mid < limit; mid <<= 1) { // 从小到大枚举二进制位 for (int i = 0, R = (mid << 1); i < limit; i += R) { // 根据二进制位划分段 for (int j = 0; j < mid; j ++) { // 枚举段内数 // 此时 i + mid + j 是保证 mid 那一位为 1 // i + j 则保证 mid 位为 0 LL x = a[i + j]; LL y = a[i + j + mid]; a[i + j] = (x + y) % P; // 对于 i + j 来说 i + j + mid 在它的范围内 1 是都有的 // 也就是 (i + j + mid) & (i + j) 的 1 的个数和 (i + j) 的 1 的个数相同 // -1 的幂次也一定一样,所以正负号相同 a[i + j + mid] = (x - y + P) % P; // 对于 i + j + mid 来说 i + j 在它的范围内 1 少一个 // 也就是 (i + j) & (i + j + mid) 的 1 的个数比 (i + j + mid) 的 1 的个数少一个 // -1 的幂次也一定不一样,所以正负号不相同 if (flag) { // 逆变换要除以 2 // 你每个数都除以 log_n^2 个 2,那不就是除以 n 吗 a[i + j] = a[i + j] * inv2 % P; a[i + j + mid] = a[i + j + mid] * inv2 % P; } } } } } int main () { ios::sync_with_stdio(false); cin.tie(0); int n; cin >> n; limit = (1 << n); for (int i = 0; i < limit; i ++) { cin >> aa[i]; } for (int i = 0; i < limit; i ++) { cin >> bb[i]; } // or for (int i = 0; i < limit; i ++) { a[i] = aa[i]; b[i] = bb[i]; } fwt_or(a, 0); fwt_or(b, 0); for (int i = 0; i < limit; i ++) { c[i] = a[i] * b[i] % P; } fwt_or(c, 1); for (int i = 0; i < limit; i ++) { cout << c[i] << " "; } cout << "\n"; // and for (int i = 0; i < limit; i ++) { a[i] = aa[i]; b[i] = bb[i]; } fwt_and(a, 0); fwt_and(b, 0); for (int i = 0; i < limit; i ++) { c[i] = a[i] * b[i] % P; } fwt_and(c, 1); for (int i = 0; i < limit; i ++) { cout << c[i] << " "; } cout << "\n"; // xor for (int i = 0; i < limit; i ++) { a[i] = aa[i]; b[i] = bb[i]; } fwt_xor(a, 0); fwt_xor(b, 0); for (int i = 0; i < limit; i ++) { c[i] = a[i] * b[i] % P; } fwt_xor(c, 1); for (int i = 0; i < limit; i ++) { cout << c[i] << " "; } cout << "\n"; return 0; }
- 1
信息
- ID
- 3227
- 时间
- 500ms
- 内存
- 1024MiB
- 难度
- 10
- 标签
- 递交数
- 4
- 已通过
- 2
- 上传者