洛谷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;
}
浙公网安备 33010602011771号