题解:AT_arc222_e [ARC222E] XOR Matching
dyc2022
·
·
题解
更好的阅读体验
到底是谁在出这种题?太魔怔了。
以下不妨设值域 2^m = O(n)。
假设数字 i 的出现次数为 c_i。x = 0 的时候显然答案是 \sum \left \lfloor \frac{c_i}{2}\right \rfloor,先提前特判掉。当 x \neq 0 时,我们考虑一对能够可以配对的数字 i 和 i \oplus x,它们总共能产生 \min(c_i, c_{i \oplus x}) 对配对的数字。所以答案就是
ans_x = \frac{1}{2} \sum_{i \oplus j = x} \min(c_i, c_j)
我们只会带乘法的异或卷积,因此我们希望将 $\min$ 变成一些东西乘起来的形式。我们注意到 $\min(x, y) = \sum \limits_{i = 1} ^ {+ \infty} [i \le x] [i \le y]$,因此
$$
ans_x = \sum_{k = 1} ^{+\infty} \sum_{i \oplus j = x} [c_i \ge k] [c_j \ge k]
$$
这个时候式子的后半部分就是可以 FWT 的形式了。但是如果对于每个 $k$ 都做一遍这个事情的话就会获得 $O(n^2 \log n)$ 的优秀复杂度。
因此我们考虑阈值分治。假设有阈值 $B$,我们对于 $k = 1 \sim B$ 分别去做 FWT。这个时候,所有 $\min(c_i, c_j) \le B$ 的 $(i, j)$ 的贡献我们都求出来了。那么剩下的就是 $c_i > B$ 的数字内部产生的贡献,由于这种数字最多有 $\frac{n}{B}$ 个,因此直接暴力卷积就可以。综合这两种情况,算法的复杂度为 $O \left( Bn \log n + \left( \frac{n}{B}\right)^2\right)$。
取 $B = \left(\frac{n}{\log n}\right)^{\frac{1}{3}}$ 即可做到 $O\left(n^{4/3} \log^{2/3} n\right)$。
```cpp
#include<bits/stdc++.h>
#define endl '\n'
#define N 1048582
#define MOD 998244353
using namespace std;
constexpr int B=20;
inline void add(int &x,int y) {x+=y,x-=x>=MOD?MOD:0;}
inline void dec(int &x,int y) {x+=MOD-y,x-=x>=MOD?MOD:0;}
int n,m,b,cnt[N],ans[N],d[N];
void fwt_xor(int *f,int opt)
{
for(int o=2,k=1;o<=b;o<<=1,k<<=1)
for(int i=0;i<b;i+=o)
for(int j=0;j<k;j++)
{
f[i+j]=(f[i+j]+f[i+j+k])%MOD;
f[i+j+k]=(f[i+j]-2ll*f[i+j+k]%MOD+MOD)%MOD;
f[i+j]=1ll*f[i+j]*opt%MOD;
f[i+j+k]=1ll*f[i+j+k]*opt%MOD;
}
}
main()
{
scanf("%d%d",&n,&m),b=1<<m;
for(int i=1,x;i<=n;i++)scanf("%d",&x),cnt[x]++;
for(int i=1;i<=B;i++)
{
for(int j=0;j<b;j++)d[j]=(cnt[j]>=i);
fwt_xor(d,1);
for(int j=0;j<b;j++)d[j]=1ll*d[j]*d[j]%MOD;
fwt_xor(d,499122177);
for(int j=0;j<b;j++)add(ans[j],d[j]);
}
vector<int> big;
for(int i=0;i<b;i++)
if(cnt[i]>B)big.push_back(i);
for(int i:big)for(int j:big)
add(ans[i^j],min(cnt[i],cnt[j])-B);
ans[0]=0;
for(int i=0;i<b;i++)ans[0]+=cnt[i]/2;
for(int i=1;i<b;i++)
ans[i]=499122177ll*ans[i]%MOD;
int pw10=1,s=0;
for(int i=0;i<b;i++)
add(s,1ll*pw10*ans[i]%MOD),pw10=10ll*pw10%MOD;
printf("%d\n",s);
return 0;
}
```