浅谈FWT的一些进阶模板
- 本文需要一点 FWT 和 FFT 的基础。
学过 FWT 的大佬,也建议看一看本蒟蒻写的 FWT 介绍,后文可能会用其中的一些符号和定义。
卷积形式
我们有两个序列
求序列
万能卷积
为了方便,设
延续 FFT 的思路,我们把多项式转成点值表达。
我们做一次线性变换,并设出变换系数。
即
那么若
可以推出
这是
能发现此性质只限制了同一个
为了方便我们定义删最高位操作。
那么新的
逆变换
考虑把上述矩阵求逆。
即
那么一个合法
合法矩阵
如何构造?
考场上如果忘记矩阵,那么可以考虑解方程。
这个方程存在多个解,于是可以考虑把他们作为矩阵的不同的行。
多尝试几次一般就能成功。
按位或矩阵
其逆矩阵:
按位与矩阵
其逆矩阵:
按位异或矩阵
其逆矩阵:
模版一:普通 FWT
P4717 【模板】快速莫比乌斯 / 沃尔什变换 (FMT / FWT)
使用 for 而不是递归。
#include<bits/stdc++.h>
using namespace std;
typedef long long ll;
const ll mod=998244353;
const int N=1e6+10;
int n,len;
ll a[N],b[N],f[N],g[N];
void OR(ll *f,int len,ll v)
{
for(int h=2;h<=len;h<<=1)
{
int mid=h/2;
for(int i=0;i<len;i+=h)
{
for(int j=i;j<i+mid;j++)
{
f[j+mid]=(f[j+mid]+f[j]*v)%mod;
}
}
}
}
void AND(ll *f,int len,ll v)
{
for(int h=2;h<=len;h<<=1)
{
int mid=h/2;
for(int i=0;i<len;i+=h)
{
for(int j=i;j<i+mid;j++)
{
f[j]=(f[j]+f[j+mid]*v)%mod;
}
}
}
}
void XOR(ll *f,int len,ll v)
{
for(int h=2;h<=len;h<<=1)
{
int mid=h/2;
for(int i=0;i<len;i+=h)
{
for(int j=i;j<i+mid;j++)
{
ll a=f[j],b=f[j+mid];
f[j]=(a+b)*v%mod;f[j+mid]=(a-b)*v%mod;
}
}
}
}
int main()
{
scanf("%d",&n);len=(1<<n);
for(int i=0;i<len;i++) scanf("%lld",&a[i]);
for(int i=0;i<len;i++) scanf("%lld",&b[i]);
for(int i=0;i<len;i++) f[i]=a[i],g[i]=b[i];
OR(f,len,1);OR(g,len,1);
for(int i=0;i<len;i++) f[i]=f[i]*g[i]%mod;
OR(f,len,-1);
for(int i=0;i<len;i++) printf("%lld ",(f[i]%mod+mod)%mod);printf("\n");
for(int i=0;i<len;i++) f[i]=a[i],g[i]=b[i];
AND(f,len,1);AND(g,len,1);
for(int i=0;i<len;i++) f[i]=f[i]*g[i]%mod;
AND(f,len,-1);
for(int i=0;i<len;i++) printf("%lld ",(f[i]%mod+mod)%mod);printf("\n");
for(int i=0;i<len;i++) f[i]=a[i],g[i]=b[i];
XOR(f,len,1);XOR(g,len,1);
for(int i=0;i<len;i++) f[i]=f[i]*g[i]%mod;
XOR(f,len,(mod+1)/2);
for(int i=0;i<len;i++) printf("%lld ",(f[i]%mod+mod)%mod);
return 0;
}
模板二:子集卷积
问题描述:
有两个长度为
如果没有
能发现若
这提示着我们把集合大小加入状态。
即定义
有
那么答案
为了减少时间复杂度,我们一开始把所有序列 FWT,那么中间的所有卷积都变成直积了。
时间复杂度:
参考代码:
#include<bits/stdc++.h>
using namespace std;
typedef long long ll;
const ll mod=1e9+9;
const int N=(1<<20)+10;
void FWT(ll *a,int len,ll v)
{
for(int h=2;h<=len;h<<=1)
{
int mid=h/2;
for(int i=0;i<len;i+=h)
{
for(int j=i;j<i+mid;j++)
{
ll x=a[j],y=a[j+mid];
a[j]=(x+y)*v%mod;a[j+mid]=(x-y+mod)*v%mod;
}
}
}
}
ll f[22][N],g[22][N],h[22][N];
int n,len=0,c[N];
int main()
{
scanf("%d",&n);len=(1<<n);
for(int i=0;i<len;i++) c[i]=c[i>>1]+(i&1);
for(int i=0;i<len;i++)
{
int x;scanf("%d",&x);
f[c[i]][i]=x;
}
for(int i=0;i<len;i++)
{
int x;scanf("%d",&x);
g[c[i]][i]=x;
}
for(int i=0;i<=n;i++) FWT(f[i],len,1),FWT(g[i],len,1);
for(int i=0;i<=n;i++)
{
for(int j=0;j<=i;j++)
{
for(int k=0;k<len;k++) h[i][k]=(h[i][k]+f[j][k]*g[i-j][k])%mod;
}
}
for(int i=0;i<=n;i++) FWT(h[i],len,(mod+1)/2);
for(int i=0;i<len;i++) printf("%lld ",h[c[i]][i]);
return 0;
}
核的引入
定义,对于一个序列
:::info[引理一]
一个序列若存在核,则他一定最多有
我们来分类讨论
i,j 的取值范围。
i<2^{\beta}$ 或 $j<2^{\beta}$,不妨设 $i<2^{\beta}\leq j<2^{\beta+1}$,此时任意 $i$ 所生成的 $j$ 的值都等于 $B_k$,因为 $\text{highbit}(j)=\beta$,不难发现贡献为:$B_k\sum\limits_{i<2^{\beta}} A_i
因此证毕。 :::info[引理三] 若 A 的核为x ,B 的核为y ,那么C=A\oplus B (异或卷积)的核为x\oplus y 。::: :::success[证明] 构造 A'_{i}=A_{i\oplus x},B'_{j}=B_{j\oplus y} ,那么A'\oplus B'=C' ,由引理一可知C' 的核为0 ,又因C_i=C'_{i\oplus x\oplus y} ,所以C 的核为x\oplus y 。因此 Q.E.D. ::: # 模板三:有核序列的异或卷积。
根据上述证明,我们可以得到如何求出两个核为
若核不为
原题连接: #3073. 「2019 集训队互测 Day 2」序列
参考代码:
struct st
{
ll val[25],vis[25],sp;
void init(ll x,ll v)
{
for(ll i=19;i>=0;i--)
{
if(x>>i&1)
{
vis[i]=1;ll tmp=x-(x&((1ll<<i)-1));
val[i]=(tmp*(tmp^v)%mod);
}
else vis[i]=0;
}
}
friend st operator*(st a,st b)
{
st c;
ll sum1=a.sp,sum2=b.sp,s=0;
for(ll i=0;i<=19;i++) sum1=(sum1+a.val[i]*(1ll<<i))%mod,sum2=(sum2+b.val[i]*(1ll<<i))%mod;
for(ll i=19;i>=0;i--)
{
c.vis[i]=a.vis[i]^b.vis[i];
sum1=(sum1-a.val[i]*(1ll<<i))%mod;sum2=(sum2-b.val[i]*(1ll<<i))%mod;
c.val[i]=(s+sum1*b.val[i]+sum2*a.val[i])%mod;
s+=a.val[i]*b.val[i]%mod*(1ll<<i)%mod;s%=mod;
}
c.sp=a.sp*b.sp+s;
return c;
}
ll query(ll d)
{
for(int i=19;i>=0;i--) if(((d>>i&1)^vis[i])==1) return (val[i]%mod+mod)%mod;
return (sp%mod+mod)%mod;
}
}f[N];
int main()
{
scanf("%d%d%d",&n,&m,&q);
for(int i=1;i<=n;i++)
{
ll x,v;scanf("%lld%lld",&x,&v);
f[i].init(x,v);
}
for(int i=2;i<=n;i++) f[i]=f[i-1]*f[i];
while(q--)
{
int c,d;scanf("%d%d",&c,&d);
printf("%lld\n",f[c].query(d));
}
return 0;
}
模板四:稀疏序列的异或卷积(系数可重集相同)
以一道题目为例。
:::info[题目描述]
有
n 个非负整数三元组(a_i, b_i, c_i) 和四个非负整数x,y,z,k 。 你利用这n 个三元组填充了n 个数组,其中第i 个数组中有x 个a_i ,y 个b_i ,z 个c_i (所以第i 个数组长度为(x+y+z) )。 对于i=0,1,\dots,2^k−1 ,回答以下询问:
- 从每个数组中选择恰好一个数,使得这些数的
\mathrm{xor} 和为i ,方案数是多少?
你只需要输出方案数对 998,244,353 取模后得到的结果。定义 f_i 表示所有序列 FWT 后直积的结果。
能发现稀疏序列的异或卷积有一个前提:所有序列的
这提示着我们有些序列之间可以划分到一个类中。
先分析
显然:
为了简化计算,我们让每个序列的值都异或上
解得:
原题链接:CF1119H Triple
参考代码:
#include<bits/stdc++.h>
using namespace std;
typedef long long ll;
const int N=1e6+10;
const ll mod=998244353;
int n,k,len;
ll x,y,z;
void FWT(ll *a,int len,ll v)
{
for(int h=2;h<=len;h<<=1)
{
int mid=h/2;
for(int i=0;i<len;i+=h)
{
for(int j=i;j<i+mid;j++)
{
ll x=a[j],y=a[j+mid];
a[j]=(x+y)*v%mod;a[j+mid]=(x-y)*v%mod;
}
}
}
}
ll fb[N],fc[N],fbc[N],ta=0,f[N];
ll ku(ll a,ll b)
{
ll ans=1;
while(b)
{
if(b&1) ans=ans*a%mod;
a=a*a%mod;b>>=1;
}
return ans;
}
int main()
{
scanf("%d%d",&n,&k);len=(1<<k);
scanf("%lld%lld%lld",&x,&y,&z);
for(int i=1;i<=n;i++)
{
ll a,b,c;scanf("%lld%lld%lld",&a,&b,&c);
ta^=a;b^=a;c^=a;
fb[b]++;fc[c]++;fbc[b^c]++;
}
FWT(fb,len,1);FWT(fc,len,1);FWT(fbc,len,1);
for(int i=0;i<len;i++)
{
ll p1=fb[i],p2=fc[i],p3=fbc[i];
ll c1=(n+p1+p2+p3)/4,c2=(n+p1-2*c1)/2,c3=(n+p2-2*c1)/2,c4=(n+p3-2*c1)/2;
f[i]=ku(x+y+z,c1)*ku(x+y-z,c2)%mod*ku(x-y+z,c3)%mod*ku(x-y-z,c4)%mod;
}
FWT(f,len,(mod+1)/2);
for(int i=0;i<len;i++) printf("%lld ",(f[i^ta]%mod+mod)%mod);
return 0;
}
能发现
假设每个序列有
定义
定义
我们要先计算出
代码就咕掉了(小编太懒了...)。
优化
我们发现可以这个方程刚好是
即把高斯消元替换为一次 IFWT。
时间复杂度:
参考代码(用的是EA的代码(依旧小编太懒了...)):
#include <bits/stdc++.h>
using ll = long long;
const int mod = 998244353;
const int inv2 = (mod + 1) / 2;
const int maxn = (1 << 20) + 12;
int n,m,k,w[maxn],z[maxn][12],F[maxn],G[maxn],W[maxn],sum[maxn],lg[maxn];
int qpow(int a,ll b) {
if (b == 0) return 1;
ll d = qpow(a,b>>1); d = d * d % mod;
if (b & 1) d = d * a % mod;
return d;
}
void DFT(int *a,int lim,int flag){
for (int i = 1; i < lim; i <<= 1)
for (int j = 0; j < lim; j += (i << 1))
for (int k = 0; k < i; ++ k) {
ll A0 = a[j+k], A1 = a[j+k+i];
a[j+k] = (A0 + A1) % mod; a[j+k+i] = (A0 - A1 + mod) % mod;
if (flag == -1) {
a[j+k] = (ll) a[j+k] * inv2 % mod;
a[j+k+i] = (ll) a[j+k+i] * inv2 % mod;
}
}
}
int main() {
scanf("%d%d%d",&n,&m,&k);
int H[(1<<k)+10][(1<<m)+10];
std::memset(H,0,sizeof(H));
for (int i = 0; i < k; ++ i) scanf("%d",&w[i]);
lg[0] = -1;
for (int i = 0; i < (1 << k); ++ i) {
if (i) lg[i] = lg[i/2] + 1;
for (int j = 0; j < k; ++ j) {
if (i & (1 << j)) W[i] = (W[i] - w[j] + mod) % mod;
else W[i] = (W[i] + w[j]) % mod;
}
}
for (int i = 1; i <= n; ++ i) {
for (int j = 0; j < k; ++ j)
scanf("%d",&z[i][j]);
H[0][0] += 1;
for (int mask = 1; mask < (1 << k); ++ mask) {
int lowbit = mask & -mask;
int b = lg[lowbit];
sum[mask] = sum[mask - lowbit] ^ z[i][b];
H[mask][sum[mask]] += 1;
}
}
for (int mask = 0; mask < (1 << k); ++ mask)
DFT(H[mask],(1<<m),1);
for (int i = 0; i < (1 << m); ++ i) {
for (int j = 0; j < (1 << k); ++ j)
G[j] = H[j][i];
DFT(G,(1<<k),-1);
F[i] = 1;
for (int j = 0; j < (1 << k); ++ j)
F[i] = (ll) F[i] * qpow(W[j],G[j]) % mod;
} DFT(F,(1<<m),-1);
for (int i = 0; i < (1 << m); ++ i)
printf("%d ",F[i]);
return 0;
}
模板五:稀疏序列的异或卷积(只有两个非 0 位,)
问题描述:求
为了方便描述,我们使用形式幂级数:
我们一开始显然可以提出一个
可以变为:
单独处理前面的
现在问题转化为计算
考虑采用分治乘法,分治的每一段区间形如
于是将
更进一步,我们可以仅维护这两个集合幂级数在
参考代码:
#include<bits/stdc++.h>
using namespace std;
typedef long long ll;
const int N=(1<<22)+10;
const ll mod=998244353;
int n,q,len=0;
ll f[N],g[N];
void FWT(int len)
{
for(int h=2;h<=len;h<<=1)
{
int mid=h/2;
for(int i=0;i<len;i+=h)
{
for(int j=i;j<i+mid;j++)
{
ll a=f[j],b=f[j+mid],c=g[j],d=g[j+mid];
f[j]=(a*b+c*d)%mod;f[j+mid]=(a*b-c*d)%mod;
g[j]=(b*c+a*d)%mod;g[j+mid]=(b*c-a*d)%mod;
}
}
}
}
void IFWT(int len,ll v)
{
for(int h=2;h<=len;h<<=1)
{
int mid=h/2;
for(int i=0;i<len;i+=h)
{
for(int j=i;j<i+mid;j++)
{
ll a=f[j],b=f[j+mid];
f[j]=(a+b)*v%mod;f[j+mid]=(a-b+mod)*v%mod;
}
}
}
}
int main()
{
scanf("%d%d",&n,&q);len=(1<<n);
for(int i=0;i<len;i++) f[i]=1,g[i]=0;
for(int i=1;i<=q;i++)
{
ll a,b,c;scanf("%lld%lld%lld",&a,&b,&c);
ll x=f[c],y=g[c];
f[c]=(x*a+y*b)%mod;
g[c]=(y*a+x*b)%mod;
}
FWT(len);
for(int i=0;i<len;i++) f[i]=(f[i]+g[i])%mod;
IFWT(len,(mod+1)/2);
ll ans=0;
for(ll i=0,sum=1;i<len;i++,sum=sum*2ll%mod) ans^=((f[i]*sum%mod+mod)%mod);
printf("%lld",ans);
return 0;
}
模板六:k 维异或卷积
众所周知异或是不进位的加法,那如果我们也就可以把模
为了方便,我们先只探讨
先分析贡献系数:
那么我们的
可以构造为单位根。
因为维度提升了,所以我们的分治每次就要分成
即我们需要
我们可以考虑范德蒙德矩阵。
下面记
即
其逆矩阵为:
即
这里我们也能看出范德蒙德矩阵的内在性质,他的每一行都是
但在一些题目中,精度会发生问题,我们需要一个比较通用的一个做法。
我们考虑用一个代数结构来刻画单位根,我们可以让
性质一:
那么我们不如直接多项式模
这样任何的
但这真的对吗?
为何不能直接模 x^k-1
我们先思考什么结构能让新的范德蒙德矩阵存在逆矩阵。
我们让范德蒙德矩阵乘他的期望逆矩阵。
对于
乘的结果为:
这满足条件。
那么
乘的结果为
这个结果无法整除
因为要想整除
这等价于代入
但能发现这可能做不到。
比如
代入
所以我们该怎么半?
结论:模分圆多项式是对的
为何可以模分圆多项式
先要证
根据分圆多项式的性质:
所以
再证范德蒙德矩阵乘他的期望逆矩阵得到单位矩阵。
和上文一样,对于
那么
乘的结果为
然后这相当于每次走
这次我们只需要代入
能发现对于这样的
考虑反证法,若存在
因为
所以
因此
回到矩阵乘法结果:
根据上文结论:
则原式
下面来分类讨论:
- 若
\gcd(h,k)=1 ,那么qh\bmod k 两两不同,所以结果等于0 。 - 否则,定义
g=\gcd(h,k) ,原式=\frac{1}{10}\sum\limits_{q=0}^9 \omega^{qgp'}=\frac{1}{10}\sum\limits_{q=0}^9 (\omega_k^g)^{qh'}=\frac{1}{10}\sum\limits_{q=0}^9 (\omega_{\frac{k}{g}})^{q\frac{h}{g}} ,因为\gcd(\frac{h}{g},k)=1 ,所以q\frac{h}{g}\bmod \frac{k}{g} 两两不同,所以结果等于0 。
因此我们就证明了结果等于
自此我们就证明了结果可以模分圆多项式。
最终实现
但是模分圆多项式复杂度太大了,因为
原题链接:CF1103E Radix sum
参考代码:
#include<bits/stdc++.h>
using namespace std;
typedef unsigned long long ll;
const int N=1e5+10;
int n;
ll a[N],c[11][11];
struct poly
{
ll a[10];
void clear(){for(int i=0;i<10;i++) a[i]=0;}
poly operator*(poly b)
{
poly c;
for(int i=0;i<10;i++) c.a[i]=0;
for(int i=0;i<10;i++)
for(int j=0;j<10;j++)
c.a[(i+j)%10]+=a[i]*b.a[j];
return c;
}
poly operator*(ll b)
{
poly c;
for(int i=0;i<10;i++) c.a[i]=0;
for(int i=0;i<10;i++) c.a[(i+b)%10]+=a[i];
return c;
}
poly operator+(poly b)
{
poly c;
for(int i=0;i<10;i++) c.a[i]=a[i]+b.a[i];
return c;
}
poly operator-(poly b)
{
poly c;
for(int i=0;i<10;i++) c.a[i]=a[i]-b.a[i];
return c;
}
}f[N],tmp[N];
ll len=100000;
ll ku(ll a,ll b)
{
ll ans=1;
while(b)
{
if(b&1) ans=ans*a;
a=a*a;b>>=1;
}
return ans;
}
ll b[N];
void divide(ll *a)
{
int n=10,m=5;
ll inv=b[m-1];
for(int i=n-1;i>=m-1;i--)
{
ll d=a[i]*inv;
for(int j=i;j>=i-m+1;j--) a[j]=a[j]-d*b[m-(i-j)-1];
}
}
void FWT(poly *f,int op)
{
for(int h=10;h<=100000;h*=10)
{
int mid=h/10;
for(int i=0;i<100000;i+=h)
{
for(int j=i;j<i+mid;j++)
{
for(int k=0;k<10;k++) tmp[k]=f[j+k*mid],f[j+k*mid].clear();
for(int k=0;k<10;k++)
{
for(int p=0;p<10;p++)
{
f[j+k*mid]=(f[j+k*mid]+(tmp[p]*((op*c[k][p]+10)%10)));
}
}
}
}
}
}
poly ksm(poly a,int b)
{
poly ans;ans.clear();ans.a[0]=1;
while(b)
{
if(b&1) ans=ans*a;
a=a*a;b>>=1;
}
return ans;
}
int main()
{
scanf("%d",&n);
for(int i=1;i<=n;i++) scanf("%llu",&a[i]),f[a[i]].a[0]++;
for(int i=0;i<10;i++)
{
for(int j=0;j<10;j++)
{
c[i][j]=(i*j%10+10)%10;
}
}
FWT(f,1);
for(int i=0;i<100000;i++) f[i]=ksm(f[i],n);
FWT(f,-1);
b[0]=1;b[1]=-1;b[2]=1;b[3]=-1;b[4]=1;
for(int i=0;i<n;i++)
{
divide(f[i].a);
ll mod=(1ll<<58);
ll x=(f[i].a[0]*6723469279985657373)<<1>>6;
printf("%llu\n",x%mod);
}
return 0;
}
致谢和参考资料
感谢 CJX,CXM老师,ZZH与我激情讨论。
参考资料
- EternalAlexander:P6097 【模板】子集卷积
- 武林的2026年集训队论文:集合幂级数中的稀疏多项式乘法
- Fucious_Yin的FWT 小记