题解:AT_arc222_e [ARC222E] XOR Matching

· · 题解

更好的阅读体验

到底是谁在出这种题?太魔怔了。

以下不妨设值域 2^m = O(n)

假设数字 i 的出现次数为 c_ix = 0 的时候显然答案是 \sum \left \lfloor \frac{c_i}{2}\right \rfloor,先提前特判掉。当 x \neq 0 时,我们考虑一对能够可以配对的数字 ii \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; } ```