容斥原理+哈夫曼式多项式乘法NTT:ABC462G

本文搬运自本人高中时期CSDN博客,若图片加载不出来,可到原文查看:https://blog.csdn.net/zhangtingxiqwq/article/details/162099398

https://atcoder.jp/contests/abc462/tasks/abc462_g

首先根据容斥原理,我们相当于求:

在这里插入图片描述

我们对颜色进行分类,对于颜色 kk ,我们假设有 XkX_k 个球, YkY_k 个盒子。

我们现在枚举它有 DkD_k 个球放在相应颜色的盒子里,方案有:

在这里插入图片描述

因为颜色间不相互影响,所以这里的总方案数位:

在这里插入图片描述

而剩余的球的位置可以随便放,共:

在这里插入图片描述

因此,在固定 AA 下,答案为:

在这里插入图片描述


此时我们采用多项式处理,定义:

在这里插入图片描述

所以有:

在这里插入图片描述

我们现在相当于求:

在这里插入图片描述


多个多项式之间合并的方法是采用小根堆,每次取出两个最小的多项式NTT一下即可。

在这里插入图片描述


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
#include<bits/stdc++.h>
#include<atcoder/all>
using namespace std;
using namespace atcoder;
using mint = modint998244353;
#ifdef LOCAL
#define debug(...) fprintf(stdout, ##__VA_ARGS__)
#else
#define debug(...) void(0)
#endif
//#define int long long
inline int read(){int x=0,f=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')f=-1;
ch=getchar();}while(ch>='0'&&ch<='9'){x=(x<<1)+
(x<<3)+(ch^48);ch=getchar();}return x*f;}
#define Z(x) (x)*(x)
#define pb push_back
#define fi first
#define se second
//#define M
//#define mo
#define N 200010
int n, m, i, j, k, T;
int X[N], Y[N], lim, pk;

signed main()
{
#ifdef LOCAL
freopen("in.txt", "r", stdin);
freopen("out.txt", "w", stdout);
#endif
// srand(time(NULL));
// T = read();
// while(T--) {
//
// }
n = read();
for(i = 1; i <= n; ++i) k = read(), X[k]++;
for(i = 1; i <= n; ++i) k = read(), Y[k]++;
for(i = 1; i <= n; ++i) debug("%lld ", X[i]); debug("\n");
for(i = 1; i <= n; ++i) debug("%lld ", Y[i]); debug("\n");
vector<mint> fac(N), ifac(N);
for(i = 1, fac[0] = 1; i <= n; ++i) fac[i] = fac[i - 1] * i;
ifac[n] = fac[n].inv();
for(i = n; i >= 1; --i) ifac[i - 1] = ifac[i] * i;
auto C = [&] (int n, int r) -> mint {
if(r < 0 || r > n) return 0;
return fac[n] * ifac[r] * ifac[n - r];
};
vector<vector<mint> > polys;
for(k = 1; k <= n; ++k) {
if((lim = min(X[k], Y[k])) == 0) continue;
vector<mint>poly(lim + 1);
for(i = 0, pk = 1; i <= lim; ++i, pk = - pk) {
mint cnt = pk * C(X[k], i) * C(Y[k], i) * fac[i];
poly[i] = cnt;
}
polys.pb(poly);
}
auto cmp = [&] (const vector<mint>& a, const vector<mint>& b) {
return a.size() > b.size();
};
priority_queue<vector<mint>, vector<vector<mint> >, decltype(cmp)> q(cmp);
for(auto& p : polys) q.push(p);
vector<mint> f;
if(q.empty()) f = {1};
else {
while(q.size() > 1) {
auto a = q.top(); q.pop();
auto b = q.top(); q.pop();
auto c= convolution(a, b);
q.push(c);
}
f = q.top();
}
mint ans = 0;
for(i = 0; i < (int)f.size(); ++i)
ans += f[i] * fac[n - i];
ans = ans * ifac[n];
cout << ans.val();
return 0;
}