集合幂级数与 FWT

· · 算法·理论

集合幂级数与 FWT

集合幂级数其实和多项式很像。具体来说:

f(x) = \sum_{s \subseteq [n]} f_s x^s

没错,“集合幂”就是字面意思。幂上面是个集合。

举个更形象的例子:

f(x) = 3 x^{\varnothing} + 2 x^{\{1\}} - 5 x^{\{2\}} - x^{\{1,2\}}

平时我们可以直接把它当作一个长度为 2^n 的数组,存储每一项的系数。

定义了集合幂级数就可以定义其乘法卷积。比如我们定义 a \cdot b = c:

c_i = \sum_{j \otimes k = i} a_j b_k

其中 \otimes 是按位与、按位或、按位异或之一。

通过 a,b 计算 c 的算法就是 FWT,快速沃尔什变化。

我们的思路是,对集合幂级数 a 构造一个变换函数 fwt(a),满足:

只要能构造,问题就很简单了。先 a \to fwt(a),b \to fwt(b),然后 \mathcal O(n) 做点乘,再 fwt(c) \to c。

考虑神奇的构造。

或

fwt(a)_i = \sum_{j \subseteq i} a_j

即高维前缀和。同样的,逆变换就是高维差分。

::::success[为什么对?]

\begin{aligned} fwt(a)_i \times fwt(b)_i &= \left(\sum_{j \subseteq i} a_j\right) \times \left(\sum_{k \subseteq i} b_k\right) \\ &= \sum_{j \subseteq i,k \subseteq i} a_j b_k \\ &= \sum_{j | k \subseteq i} a_j b_k \\ &= fwt(c)_i \end{aligned}

::::

虽然说咱有经典高维前缀和写法,但为了和后面的兼容,我们这样写:

void fwt(int *a, int o) {   // o=1 正变换,o=-1 逆变换
    for (int i = 1; i < 1 << n; i <<= 1)
        for (int j = 0; j < 1 << n; j += i + i)
            for (int k = 0; k < i; ++ k ) {
                a[i + j + k] += o * a[j + k];
            }
}

具体来说,我们枚举每个长度为 2i 的区间。那么这个区间内的前 i 个数在当前这一维上是 0,后 i 个数是 1。高维前缀和的话,就是后面的每一项加上前面对应的项。这个代码还是很酷的。

与

和或是反着的。高维前缀和变成高维后缀和。代码:

void fwt(int *a, int o) {   // o=1 正变换,o=-1 逆变换
    for (int i = 1; i < 1 << n; i <<= 1)
        for (int j = 0; j < 1 << n; j += i + i)
            for (int k = 0; k < i; ++ k ) {
                a[j + k] += o * a[i + j + k];
            }
}

异或

这才是重头戏。

fwt(a)_i = \sum_{j} (-1)^{|i \cap j|} a_j

::::success[为什么对?]

\begin{aligned} fwt(a)_i \times fwt(b)_i &= \left(\sum_{j}(-1)^{|i \cap j|} a_j\right) \times \left(\sum_{k}(-1)^{|i \cap k|} b_k\right) \\ &= \sum_{j,k} (-1)^{|i \cap j| + |i \cap k|} a_j b_k \\ \end{aligned}

考虑 |i \cap j| + |i \cap k|。拆位考虑,如果 i 的这一位是 0,那么无论 j,k 是什么都没有贡献。否则如果 i 这一位是 1,那么 |i \cap j| + |i \cap k| 的奇偶性发生改变,当且仅当 j,k 的这一位不同。

即,(-1)^{|i \cap j| + |i \cap k|} = (-1)^{|i \cap (j \oplus k)|}。

\begin{aligned} fwt(a)_i \times fwt(b)_i &= \sum_{j,k} (-1)^{|i \cap (j \oplus k)|} a_jb_k \\ &= fwt(c)_i \end{aligned}

:::: 考虑求 fwt(a)。分治,设 a 的前一半和后一半分别是 a_0,a_1。

fwt(a) = merge(fwt(a_0) + fwt(a_1), fwt(a_0) - fwt(a_1))

这是为什么呢?即我考虑求每个 fwt(a)_i,那么我首先应该枚举所有 j,然后把 (-1)^{|i \cap j|} a_j 贡献过去。现在我们已经知道了,所有忽略掉 i,j 的最高位后的答案,即 fwt(a_0), fwt(a_1),那么想让 (-1)^{|i \cap j|} 发生变化,当且仅当 i,j 最高位都是 1。所以最后一项有个负号,别的都是正号。

正变换是这样的:

void fwt(int *a) {
    for (int i = 1; i < 1 << n; i <<= 1)
        for (int j = 0; j < 1 << n; j += i << 1)
            for (int k = 0; k < i; ++ k ) {
                int x = a[j + k], y = a[i + j + k];
                a[j + k] = x + y;
                a[i + j + k] = x - y;
            }
}

逆变换的话,稍微推一推,\frac{(x+y)+(x-y)}{2} = x, \frac{(x+y)-(x-y)}2 = y。因此在上述基础上除以二即可。

void fwt(int *a, int o) {       // o=1 是正变换,o=inv2 是负变换
    for (int i = 1; i < 1 << n; i <<= 1)
        for (int j = 0; j < 1 << n; j += i << 1)
            for (int k = 0; k < i; ++ k ) {
                int x = a[j + k], y = a[i + j + k];
                a[j + k] = x + y;
                a[i + j + k] = x - y;
                a[j + k] *= o;
                a[i + j + k] *= o;
            }
}

例题 1 CF2194F2

给定一棵包含 n 个顶点的树。每个顶点上写有一个非负整数 a_v。同时给定 k 个互不相同的非负整数 b_1 \dots b_k。我们称一组边为美丽的,如果在移除这些边后,树被分成若干连通块,并且每个连通块内所有顶点的数 a_v 的按位异或值属于集合 b。你需要计算该树中美丽的边集的数量,对 10^9+7 取模。

先考虑 Easy Version。

暴力做法是设 f_{u,S} 表示考虑 u 的子树,且 u 当前所在连通块的异或和为 S 的方案数。转移有两种:

复杂度 n \times (10^9)^2。

设 s_u 表示 u 的子树异或和。注意到 f_{u,S} \ne 0 的必要条件是,s_u \oplus S 是 b 的一个子集的异或和。即除了自己所在的连通块,其他已经完成的连通块的异或和的取值只有 2^k 种。于是第二维状态数减少到了 2^k。总复杂度 \mathcal O(n 4^k)。

具体的,设 f_{u,S} 表示考虑 u 的子树且除了 u 所在连通块外,b 中每个元素在这些连通块内出现次数的奇偶性为 S 的方案数(出现两次就能消掉,所以只关心奇偶性)。转移:

:::success[Code]

#include <bits/stdc++.h>

using namespace std;

const int N = 1e5 + 10, P = 1e9 + 7;

int n, a[N], b[4], k, sum[N];
vector<int> g[N];

void add(int &a, int b) {
    a += b;
    if (a >= P) a -= P;
}

int id(int x) {
    for (int i = 0; i < k; ++ i )
        if (b[i] == x) return i;
    return -1;
}

int f[N][1 << 4], h[1 << 4], p[1 << 4];

void dfs(int u, int F) {
    sum[u] = a[u];
    for (int v : g[u])
        if (v != F) {
            dfs(v, u);
            sum[u] ^= sum[v];
        }

    memset(f[u], 0, sizeof(int) * 16);
    f[u][0] = 1;
    for (int v : g[u])
        if (v != F) {
            for (int i = 0; i < 1 << k; ++ i ) h[i] = f[u][i], f[u][i] = 0;

            for (int i = 0; i < 1 << k; ++ i )
                for (int j = 0; j < 1 << k; ++ j ) {
                    add(f[u][i ^ j], 1ll * h[i] * f[v][j] % P);
                    int x = id(sum[v] ^ p[j]);
                    if (~x) add(f[u][i ^ j ^ (1 << x)], 1ll * h[i] * f[v][j] % P);
                }
        }
}

int solve() {
    cin >> n >> k;
    for (int i = 1; i <= n; ++ i ) g[i].clear();
    for (int i = 1; i < n; ++ i ) {
        int a, b;
        cin >> a >> b;
        g[a].push_back(b), g[b].push_back(a);
    }
    for (int i = 1; i <= n; ++ i ) cin >> a[i];
    for (int i = 0; i < k; ++ i ) cin >> b[i];
    for (int s = 0; s < 1 << k; ++ s ) {
        p[s] = 0;
        for (int i = 0; i < k; ++ i ) {
            if (s >> i & 1) p[s] ^= b[i];
        }
    }

    dfs(1, 0);

    int res = 0;
    for (int s = 0; s < 1 << 4; ++ s ) {
        if (~id(sum[1] ^ p[s])) add(res, f[1][s]);
    }
    return res;
}

signed main() {
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    int T;
    cin >> T;
    while (T -- ) cout << solve() << '\n';
    return 0;
}

:::

然后考虑 Hard Version。

这个转移与 FWT 极像。不过在第二种转移里,S,T 一起转移到了 S \oplus T \oplus \{i\}。而这个 i 可以只由 v,T 确定。于是,如果我们先更新 f_{v,T} \to f_{v,T \oplus \{i\}},那么剩下的就是裸的异或卷积了。

不过很可惜,一次 FWT 复杂度是 \mathcal O(2^k k) 的,总复杂度 \mathcal O(n2^k k) 依然无法接受。我们希望优化掉这个 k。

现在我们尝试直接维护 f' 表示 f 做完 FWT 后的结果。那么如果只有第一种转移,直接做 \mathcal O(2^k) 点乘即可。考虑如何处理第二种转移。

注意到第二种转移可以进行(即可以找到 \{i\})当且仅当 s_v 本身就可以被 b 的某个子集表示出来,但这个子集可能有多个,所以先把 b 放到线形基里。这样在做 f_{v,T} \to f_{v,T \oplus \{i\}} 时,T \oplus \{i\} 就是定值了。设这个定值是 A。具体来说,A 就是 s_v 在线形基中的那些位置。

那么现在要做的是对 f_{v,A} 单点加 \sum_{i=1}^{k} f_{v,A \oplus pos_i},其中 pos_i 为 b_i 在线形基中的那些位置。现在有两个问题,一是如何求出 \sum_{i=1}^{k} f_{v,A \oplus pos_i},二是求出这个值后,单点加这个操作会如何更新 f'_v。

第二个问题是容易的。考虑定义 FWT(a)_S = \sum_T (-1)^{|S \cap T|} a_T。那么当 a_T \to a_T + c 后,对每个位置 S 增加 (-1)^{|S \cap T|} c 即可。

那么 \sum_{i=1}^{k} f_{v,A \oplus pos_i} 怎么求呢?考虑一个神奇的构造。设 g 为只有下标 A 为 1 的集合幂级数。做 f_v \gets f_v \ast g(\ast 为异或卷积),现在问题变成求 \sum_{i=1}^{k} f_{v,pos_i}。再设 h 为只有所有下标 pos_i 为 1 的集合幂级数。做 f_v \gets f_v \ast h,现在问题变成求 f_{v,0}。

上面两次异或卷积,可以通过预处理 g,h 的 FWT 后的结果,直接做 \mathcal O(2^k) 的点乘。这样就得到了 F = f'_v \cdot FWT(g) \cdot FWT(h)。通过这个,求单点 f_{v,0} 的值是容易的。具体的 f_{v,0} = \dfrac 1{2^k} \sum_i (-1)^{|\varnothing\cap i|}F_i= \dfrac 1{2^k} \sum_i F_i,\mathcal O(2^k) 暴力求即可。

:::success[Code]

#include <bits/stdc++.h>

using namespace std;

const int N = 1e5 + 10, P = 1e9 + 7, inv2 = P + 1 >> 1;

int n, k, s[N], pos[10], b[10], w[N];
vector<int> adj[N], vec;
int a[32], ppc[N], inv2m, m;

void insert(int x) {
    for (int i = 31; ~i; -- i ) {
        if (x >> i & 1) {
            if (a[i]) x ^= a[i];
            else {
                a[i] = x;
                break;
            }
        }
    }
}

int f[N][1 << 10];
int A[N], g[1 << 10], h[1 << 10];

void FWT(int *a) {
    for (int i = 1; i < 1 << m; i <<= 1)
        for (int j = 0; j < 1 << m; j += i + i)
            for (int u = 0; u < i; ++ u ) {
                const int x = a[j + u], y = a[j + u + i];
                a[j + u] = x + y, a[j + u + i] = x - y;
                if (a[j + u] >= P) a[j + u] -= P;
                if (a[j + u + i] < 0) a[j + u + i] += P;
            }
}

void IFWT(int *a) {
    for (int i = 1; i < 1 << m; i <<= 1)
        for (int j = 0; j < 1 << m; j += i + i)
            for (int u = 0; u < i; ++ u ) {
                const int x = a[j + u], y = a[j + u + i];
                a[j + u] = 1ll * (x + y) * inv2 % P;
                a[j + u + i] = 1ll * (x - y + P) * inv2 % P;
            }
}

void dfs(int u, int F) {
    for (int i = 0; i < 1 << m; ++ i ) {
        f[u][i] = 1;
    }

    s[u] = w[u];
    for (int v : adj[u])
        if (v != F) {
            dfs(v, u);
            if (~A[v]) {
                int sum = 0;
                for (int s = 0; s < 1 << m; ++ s ) {
                    sum = (sum + 1ll * f[v][s] * (ppc[s & A[v]] ? P - h[s] : h[s])) % P;
                }
                sum = 1ll * sum * inv2m % P;
                if (sum)
                for (int s = 0; s < 1 << m; ++ s ) {
                    f[v][s] += ppc[s & A[v]] ? P - sum : sum;
                    if (f[v][s] >= P) f[v][s] -= P;
                }
            }
            for (int s = 0; s < 1 << m; ++ s ) {
                f[u][s] = 1ll * f[u][s] * f[v][s] % P;
            }
            s[u] ^= s[v];
        }

    int x = s[u];
    A[u] = 0;
    for (int i = 0; i < m; ++ i )
        if (x >> vec[i] & 1) {
            A[u] |= 1 << i;
            x ^= a[vec[i]];
        }
    if (x) A[u] = -1;
}

int solve() {
    cin >> n >> k;
    for (int i = 1; i <= n; ++ i ) adj[i].clear();
    for (int i = 1; i < n; ++ i ) {
        int a, b;
        cin >> a >> b;
        adj[a].push_back(b);
        adj[b].push_back(a);
    }

    for (int i = 1; i <= n; ++ i ) {
        cin >> w[i];
    }

    memset(a, 0, sizeof a);
    for (int i = 0; i < k; ++ i ) {
        cin >> b[i];
        insert(b[i]);
    }

    vec.clear();
    for (int j = 31; ~j; -- j ) {
        if (a[j]) vec.push_back(j);
    }
    m = vec.size();

    inv2m = 1;
    for (int i = 0; i < m; ++ i ) inv2m = 1ll * inv2m * inv2 % P;

    memset(h, 0, sizeof h);
    for (int i = 0; i < k; ++ i ) {
        int res = 0;
        for (int j = 0; j < m; ++ j ) {
            if (b[i] >> vec[j] & 1) {
                res |= 1 << j;
                b[i] ^= a[vec[j]];
            }
        }
        pos[i] = res;
        h[pos[i]] = 1;
    }
    FWT(h);

    dfs(1, 0);

    IFWT(f[1]);
    if (A[1] == -1) return 0;
    int res = 0;
    for (int i = 0; i < k; ++ i ) {
        res = (res + f[1][A[1] ^ pos[i]]) % P;
    }
    return res;
}

signed main() {
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    for (int i = 0; i < N; ++ i ) {
        ppc[i] = __builtin_popcount(i) & 1;
    }
    int T;
    cin >> T;
    while (T -- ) cout << solve() << '\n';
    return 0;
}

:::

例题 2 CF1119H

给定 x,y,z。有 n 个数组,第 i 个数组用三元组 (a_i,b_i,c_i) 描述,表示第 i 个数组由 x 个 a_i,y 个 b_i,z 个 c_i 构成。对每个 t=0\dots 2^k-1,求从每个数组中选出一个数,且选出数的异或和为 t 的方案数。

构造 n 个幂级数 F_i,使得 F_i[a_i]=x,F_i[b_i]=y,F_i[c_i]=z。那么它们的异或卷积和就是答案。

如果先把所有 F_i 都变成 FWT(F_i),然后全部点乘起来,最后逆变换回去,复杂度 \mathcal O(nk2^k)。注意到 F_i 中只有三个位置非零,所以 FWT(F_i)[s] = (-1)^{|s \cap a_i|}x+(-1)^{|s \cap b_i|}y + (-1)^{|s \cap c_i|}z。复杂度 \mathcal O(n 2^k)。但还是爆了。

由于 x,y,z 是固定的,因此 FWT(F_i)[s] 的取值事实上只有八种,即每项前面的系数是 1 还是 -1。八个有点多,考虑 a_i \gets 0,b_i \gets b_i \oplus a_i, c_i \gets c_i \oplus a_i,即把 b,c 都异或 a,那么答案会变成原来的异或上所有 a 的异或和,本质是相同的。此时 a_i=0,因此 FWT(F_i)[s] 的取值只有 x+y+z,x+y-z,x-y+z,x-y-z 四种。

考虑一个固定的 s。我们要计算 \prod FWT(F_i)[s],现在问题变成了求 x+y+z,x+y-z,x-y+z,x-y-z 这四种分别的出现次数。最后快速幂乘起来即可。设这四个出现次数分别为 c_0,c_1,c_2,c_3。

考虑构造四个方程把他们求出来。

首先显然有一个 c_0+c_1+c_2+c_3=n。

如果构造 F_i[b_i]=1,其它位置都是 0。那么 FWT(F_i)[s] = (-1)^{|s \cap b_i|}。设 p=\sum_{i=1}^n FWT(F_i)[s]。注意到 p=c_0+c_1-c_2-c_3。

同理,构造 $F_i[c_i]=1$,其它位都是 $0$。得到 $q=c_0-c_1+c_2-c_3$。 然后构造 $F_i[b_i \oplus c_i]=1$,其它位都是 $0$。那么 $FWT(F_i)[s] = (-1)^{|s \cap (b_i \oplus c_i)|}= (-1)^{|s \cap b_i|} (-1)^{|s \cap c_i|}$。设 $r=\sum_{i=1}^n FWT(F_i)[s]$,注意到 $r=c_0-c_1-c_2+c_3$。 解方程就好啦! :::success[Code] ```cpp #include <bits/stdc++.h> using namespace std; const int N = 1e5 + 10, K = 17, P = 998244353, inv2 = P + 1 >> 1; int n, k; long long x, y, z; int a[N], b[N], c[N]; int f[1 << K]; int fpm(long long a, int b) { a %= P; int res = 1; while (b) { if (b & 1) res = 1ll * res * a % P; b >>= 1, a = 1ll * a * a % P; } return res; } void FWT(int *f) { for (int i = 1; i < 1 << k; i <<= 1) for (int j = 0; j < 1 << k; j += i + i) for (int u = 0; u < i; ++ u ) { int x = f[j + u], y = f[i + j + u]; f[j + u] = (x + y) % P; f[i + j + u] = (x - y + P) % P; } } void IFWT(int *f) { for (int i = 1; i < 1 << k; i <<= 1) for (int j = 0; j < 1 << k; j += i + i) for (int u = 0; u < i; ++ u ) { int x = f[j + u], y = f[i + j + u]; f[j + u] = (x + y) % P; f[i + j + u] = (x - y + P) % P; f[j + u] = 1ll * f[j + u] * inv2 % P; f[i + j + u] = 1ll * f[i + j + u] * inv2 % P; } } int p[1 << K], q[1 << K], r[1 << K], res[1 << K]; signed main() { cin >> n >> k >> x >> y >> z; int sum = 0; for (int i = 1; i <= n; ++ i ) { cin >> a[i] >> b[i] >> c[i]; b[i] ^= a[i], c[i] ^= a[i], sum ^= a[i]; } for (int i = 1; i <= n; ++ i ) { p[b[i]] ++ ; } FWT(p); for (int i = 1; i <= n; ++ i ) { q[c[i]] ++ ; } FWT(q); for (int i = 1; i <= n; ++ i ) { r[b[i] ^ c[i]] ++ ; } FWT(r); for (int s = 0; s < 1 << k; ++ s ) { int c0, c1, c2, c3; c0 = (1ll * p[s] + q[s] + r[s] + n) * fpm(4, P - 2) % P; c1 = (1ll * (n + p[s]) * inv2 - c0 + P) % P; c2 = (1ll * (n + q[s]) * inv2 - c0 + P) % P; c3 = ((n - c0 - c1 - c2) % P + P) % P; res[s] = 1ll * fpm(x + y + z, c0) * fpm(x + y - z, c1) % P * fpm(x - y + z, c2) % P * fpm(x - y - z, c3) % P; if (res[s] < 0) res[s] += P; } IFWT(res); for (int s = 0; s < 1 << k; ++ s ) { cout << res[s ^ sum] << ' '; } return 0; } ```