洛谷P4717 快速沃尔什变换(FWT)模板题

题目链接:https://www.luogu.com.cn/problem/P4717

解题思路完全来自 oi.wiki

补充理解:

或(or)运算 中,

\[FWT[A]_i \cdot FWT[B]_i \]

\[= \left( \sum_{i \cup j=i} A_j \right) \left( \sum_{i \cup k=i} B_k \right) \]

\[= \sum_{i \cup j=i} \sum_{i \cup k=i} A_j B_k \]

\[= \sum_{i \cup (j \cup k)=i} A_jB_k \]

\[= FWT[C]_i \]

所以此时的 \(FWT[C]_i\) 表示的其实是所有满足 \(i \cup (j \cup k) = i\)\(A_j B_k\) 之和。

但是 \(j \cup k\) 不一定等于 \(i\),而是 \(i\) 的子集。即:

\[j \cup k \subseteq i \]

所以,逆变换要做的事情是,从 \(i\) 里面删去所有 \(i\) 的真子集。

所以:

  • 顺变换其实是:集合本身 \(\Rightarrow\) 集合的子集。
  • 逆变换其实是:集合的子集 \(\Rightarrow\) 集合本身。

示例程序:

#include <bits/stdc++.h>
using namespace std;
const long long mod = 998244353;
const int maxn = (1<<17) + 5;

long long fpow(long long a, long long b) {
    long long res = 1;
    for (long long t = a % mod; b; b >>= 1, t = t * t % mod)
        if (b & 1ll)
            res = res * t % mod;
    return res;
}
const long long inv2 = fpow(2, mod-2);

int n;
long long A[maxn], B[maxn], a[maxn], b[maxn];

void add(long long &a, long long b) {
	a = (a + b % mod + mod) % mod;
}

void Or(long long a[], int tp) {
	for (int x = 2; x <= n; x <<= 1) {
		int k = x >> 1;
		for (int i = 0; i < n; i += x) {
			for (int j = 0; j < k; j++) {
				add(a[i+j+k], a[i+j]*tp);
			}
		}
	}
}

void And(long long a[], int tp) {
	for (int x = 2; x <= n; x <<= 1) {
		int k = x >> 1;
		for (int i = 0; i < n; i += x) {
			for (int j = 0; j < k; j++) {
				add(a[i+j], a[i+j+k]*tp);
			}
		}
	}
}

void Xor(long long a[], int tp) {
    for (int x = 2; x <= n; x <<= 1) {
		int k = x >> 1;
		for (int i = 0; i < n; i += x) {
			for (int j = 0; j < k; j++) {
                long long x = a[i+j], y = a[i+j+k];
				a[i+j] = (x + y) % mod;
				a[i+j+k] = (x - y + mod) % mod;
                (a[i+j] *= tp) %= mod;
                (a[i+j+k] *= tp) %= mod;
			}
		}
	}
}

void cal(void (*f)(long long[], int), int tp2 = -1) {
	copy(A, A+n, a);
	copy(B, B+n, b);
	f(a, 1);
	f(b, 1);
	for (int i = 0; i < n; i++)
		(a[i] *= b[i]) %= mod;
	f(a, tp2); // -1 or 1/2
	for (int i = 0; i < n; i++) {
		if (i) putchar(' ');
		printf("%lld", a[i]);
	}
	putchar('\n');
}

int main() {
	scanf("%d", &n);
	n = 1<<n;
	for (int i = 0; i < n; i++) scanf("%lld", A+i);
	for (int i = 0; i < n; i++) scanf("%lld", B+i);
	cal(Or);
	cal(And);
	cal(Xor, inv2);
	return 0;
}
posted @ 2026-04-30 15:30  quanjun  阅读(9)  评论(0)    收藏  举报