1 条题解

  • 1
    @ 2026-8-5 16:31:30

    更好的阅读体验

    专题:再来一遍一定记住的算法_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

    按位与卷积(Bitwise AND Convolution)

    信息

    ID
    3227
    时间
    500ms
    内存
    1024MiB
    难度
    10
    标签
    递交数
    4
    已通过
    2
    上传者