数学相关
容斥原理相关
一般容斥
定义
「LibreOJ NOI Round #2」不等关系
题意
给定一个长度为 \(n\),仅包含 \(<,>\) 的字符串 \(S\),求有多少个长度为 \(n+1\) 的排列 \(p\) 满足对于任意 \(1\le i\le n\),若 \(S_i\) 为 \(<\) 则 \(p_i < p_{i+1}\),否则 \(p_i > p_{i+1}\)。
\(1\le n\le 10^5\)
solution
很经典的容斥。
考虑不管 \(>\) 的限制,只考虑 \(<\) 怎么做。
只考虑 \(<\) 的限制时简单的,只需要简单组合计数一下即可算出答案,加上 \(>\) 号会很麻烦,所以我们考虑容斥掉 \(>\) 的限制。
钦定一个集合 \(S\),满足 \(S\) 内的 \(>\) 号一定不能满足,即对于 \(S\) 内的 \(>\) 号,若其在第 \(i\) 个位置,那么在实际的排列上有 \(p_i<p_{i+1}\)。
令 \(f(S)\) 表示让 \(S\) 内的 \(>\) 号一定不满足,其他 \(>\) 号没有限制的方案数,根据容斥,答案即为 \(\sum (-1)^{\left \lvert S \right \rvert}f(S)\)。
考虑用 dp 去维护这个容斥,我们关心的的其实是使用不在 \(S\) 内的 \(>\) 号进行分割后得到的若干部分的长度(分割后每个部分形成一个递增的序列,即段内全部满足 \(<\) 的限制,这里的 \(<\) 包括 \(S\) 内的 \(>\)),令 \(dp_i\) 表示考虑第 \(i\) 个符号前的容斥和,那么我们枚举上一个不出现在 \(S\) 内的 \(>\) 号,令位置为 \(j\),则 \((j,i)\) 内所有 \(>\) 号都在 \(S\) 内,并且排列在 \([j,i)\) 上单调递增,所以我们可以列出转移方程如下:
其中,\(cnt_i\) 表示 \([1,i]\) 中 \(s_i='>'\) 的 \(i\) 的个数,\((-1)^{cnt_{i-1}-cnt_j}\) 表示了将 \((j,i)\) 内所有 \(>\) 号选入集合 \(S\) 所带来的容斥系数,\(\binom{i}{i-j}\) 表示了从 \(i\) 个数中选出 \(i-j\) 个数构成 \([j,i)\) 上的递增序列。
于是我们得到了一个 \(O(n^2)\) 的 dp 算法。
后面就是优化了,把转移方程两组合数拆成阶乘形式,并在两边同时除以 \(i!\),即可得到:
这就是一个喜闻乐见的卷积形式了,用分治 NTT 即可做到 \(O(n\log^2 n)\)。
Code
#include<cstdio>
#include<algorithm>
#include<cstring>
#include<vector>
#include<cmath>
using namespace std;
#define ll long long
#define qwq Ff472130
#define f(i,l,r) for (int i=l;i<=r;i++)
#define F(i,l,r) for (int i=l;i>=r;i--)
constexpr int N=4e5+10;
constexpr int inf=1e9+10;
namespace Poly {
constexpr int N=2e6+10;
constexpr double pi=acos(-1.0);
constexpr int mod=998244353,gen=3;
inline ll qpow(ll a,int b) {
ll res=1;
while (b) {
if (b&1) res=res*a%mod;
a=a*a%mod;b>>=1;
}
return res;
}
static int len,rev[N];
inline void initrev() {
const int lt=log(len)/log(2)-1;
f(i,1,len-1) rev[i]=((rev[i>>1]>>1)|((i&1)<<lt));
}
static int g[30],flag_init_gen;
inline void init_gen() {f(i,1,22) g[i]=qpow(gen,(mod-1)>>i);}
inline int ad(int x,int y) {return (x+y>=mod)?(x+y-mod):(x+y);}
inline void add(int &x,int y) {x=ad(x,y);}
inline void NTT(int *a,int type) {
if (!flag_init_gen) init_gen(),flag_init_gen=1;
f(i,1,len-1) if (rev[i]<i) swap(a[i],a[rev[i]]);
for (int mid=1,t=1;mid<len;mid<<=1,t++) {
int Wn=g[t];
for (int i=mid<<1,j=0;j<len;j+=i) {
int w=1;
for (int k=0;k<mid;k++,w=(ll)(w)*Wn%mod) {
int x=a[j+k],y=(ll)(w)*a[j+k+mid]%mod;
a[j+k]=ad(x,y);a[j+k+mid]=ad(x,mod-y);
}
}
}
if (type==-1) {
reverse(a+1,a+len);
const ll invn=qpow(len,mod-2);
f(i,0,len-1) a[i]=invn*a[i]%mod;
}
}
inline void poly_px(int *f,int *g,int len_mx) {for(int i=0;i<len_mx;i++) f[i]=(ll)(f[i])*g[i]%mod;}
inline void poly_cpy(int *f,int *g,int len_mx) {for(int i=0;i<len_mx;i++) f[i]=g[i];}
inline void poly_clr(int *f,int l,int r) {f(i,l,r) f[i]=0;}
inline void NTT_mul(int *f,int *g,int *ret,int lenf,int leng) {
static int X0[N],Y0[N];
for (len=1;len<lenf+leng;len<<=1);
initrev();
poly_cpy(X0,f,lenf);
poly_cpy(Y0,g,leng);
NTT(X0,1);NTT(Y0,1);poly_px(X0,Y0,len);NTT(X0,-1);
f(i,0,lenf+leng) ret[i]=X0[i];
f(i,0,len) X0[i]=Y0[i]=0;
}
}
constexpr int mod=998244353;
inline int ad(int x,int y) {return ((x+y>=mod)?(x+y-mod):(x+y));}
inline void add(int &x,int y) {x=ad(x,y);}
inline ll qpow(ll a,int b) {
ll res=1;
while (b) {
if (b&1) res=res*a%mod;
a=a*a%mod;
b>>=1;
}
return res;
}
inline void read(int &x) {
x=0;
char ch=getchar();
while (ch<48) ch=getchar();
while (ch>=48) x=(x<<3)+(x<<1)+(ch^48),ch=getchar();
}
ll fac[N],inv[N];
inline void init(int n) {
fac[0]=1;
f(i,1,n) fac[i]=fac[i-1]*i%mod;
inv[n]=qpow(fac[n],mod-2);
F(i,n,1) inv[i-1]=inv[i]*i%mod;
}
inline ll C(int n,int m) {
if (n<0||m<0||n<m) return 0;
return fac[n]*inv[m]%mod*inv[n-m]%mod;
}
inline vector<int> NTT(vector<int> &x,vector<int> &y) {
int lenx=x.size(),leny=y.size();
static int X[N],Y[N],H[N];
f(i,0,lenx-1) X[i]=x[i];
f(i,0,leny-1) Y[i]=y[i];
Poly::NTT_mul(X,Y,H,lenx,leny);
int lenh=lenx+leny-1;
vector<int> ret;
f(i,0,lenh-1) ret.push_back(H[i]);
return ret;
}
int n;
int c[N],dp[N];
char s[N];
inline void solve(int l,int r) {
if (l==r) return;
int mid=l+r>>1;
solve(l,mid);
vector<int> vx,vy;
f(i,l,mid) vx.push_back((s[i]=='<')?0:((c[i]&1)?(mod-dp[i]):dp[i]));
f(i,0,r-l+1) vy.push_back(inv[i]);
vector<int> ret=NTT(vx,vy);
f(i,mid+1,r) {
if (c[i-1]&1) add(dp[i],mod-ret[i-l]);
else add(dp[i],ret[i-l]);
}
solve(mid+1,r);
}
int main() {
init(N-1);
scanf("%s",s+1);
n=strlen(s+1);
f(i,1,n+1) c[i]=c[i-1]+(s[i]=='>');
dp[0]=1;solve(0,n+1);
printf("%d\n",(int)(dp[n+1]*fac[n+1]%mod));
return 0;
}
Tree Coloring
题意
给定一棵有根树,你需要将第 \(i\) 个节点染成颜色 \(c_i\),求有多少种排列 \(c\) 满足对于每个点 \(u\) 的颜色都不为它的父亲 \(fa_u\) 的颜色减去 \(1\),即 \(c_u\ne c_{fa_u}-1\)。
\(1\le n\le 2.5\times 10^5\)。
solution
容斥经典题。
考虑钦定一个点集 \(S\),使点集内的所有点 \(u\),都不满足条件,即 \(c_u= c_{fa_u}-1\),对于不在点集 \(S\) 的点,没有限制,那么答案就是 \(\sum\limits_{S 合法} (-1)^{\lvert S\rvert }(n- \lvert S\rvert)!\),其中乘上 \((n-\lvert S\rvert )!\) 的意义是,对于所有有限制的点,将其与父亲缩成一个点,则最终会形成 \(n-\lvert S\rvert\) 个连通块,那么每个连通块的权值可以任意排列,由于连通块内的大小已经完全确定(即一根由 \(S\) 内的点组成的链的顶端一定是最大值),所以系数应该是 \((n-\lvert S\rvert)!\),而非 \(A_{n}^{n-\lvert S\rvert}\)。
因为一个节点只能由一个子节点在 \(S\) 内,而系数又只与 \(S\) 的大小相关,所以对于每个节点考虑它是否选择一个子节点加入 \(S\),令 \(s_u\) 为 \(u\) 的子节点数量,那么计算出 \(\prod\limits_{i=1}^{n} (s_i x+1)\) 得到的多项式中 \(x^k\) 的系数即为 \(\lvert S\rvert =k\) 的合法 \(S\) 个数,直接上分治 NTT 即可做到 \(O(n\log ^2 n)\)。
计算 \(\prod\limits_{i=1}^{n} (s_i x+1)\) 这个式子还有更优秀的做法。
由于 \(\sum s_i=n-1\),令满足 \(s_i=d\) 的 \(i\) 共有 \(cnt_d\) 个,将子节点数相同的节点合并为 \((d x+1)^{cnt_d}\),用二项式展开,再将所有式子卷起来,一侧长度小就暴力乘,否则用 NTT,这个做法的时间复杂度神秘地达到了 \(O(n\log n)\),证明似乎可以用哈夫曼树。
Code
#include<cstdio>
#include<algorithm>
#include<vector>
#include<cmath>
using namespace std;
#define ll long long
#define qwq Ff472130
#define f(i,l,r) for (int i=l;i<=r;i++)
#define F(i,l,r) for (int i=l;i>=r;i--)
inline void read(int &x) {
x=0;
char ch=getchar();
while (ch<48) ch=getchar();
while (ch>=48) x=(x<<3)+(x<<1)+(ch^48),ch=getchar();
}
namespace Poly {
constexpr int N=2e6+10;
constexpr double pi=acos(-1.0);
constexpr int mod=998244353,gen=3;
inline ll qpow(ll a,int b) {
ll res=1;
while (b) {
if (b&1) res=res*a%mod;
a=a*a%mod;b>>=1;
}
return res;
}
static int len,rev[N];
inline void initrev() {
const int lt=log(len)/log(2)-1;
f(i,1,len-1) rev[i]=((rev[i>>1]>>1)|((i&1)<<lt));
}
static int g[30],flag_init_gen;
inline void init_gen() {f(i,1,22) g[i]=qpow(gen,(mod-1)>>i);}
inline int ad(int x,int y) {return (x+y>=mod)?(x+y-mod):(x+y);}
inline void add(int &x,int y) {x=ad(x,y);}
inline void NTT(int *a,int type) {
if (!flag_init_gen) init_gen(),flag_init_gen=1;
f(i,1,len-1) if (rev[i]<i) swap(a[i],a[rev[i]]);
for (int mid=1,t=1;mid<len;mid<<=1,t++) {
int Wn=g[t];
for (int i=mid<<1,j=0;j<len;j+=i) {
int w=1;
for (int k=0;k<mid;k++,w=(ll)(w)*Wn%mod) {
int x=a[j+k],y=(ll)(w)*a[j+k+mid]%mod;
a[j+k]=ad(x,y);a[j+k+mid]=ad(x,mod-y);
}
}
}
if (type==-1) {
reverse(a+1,a+len);
const ll invn=qpow(len,mod-2);
f(i,0,len-1) a[i]=invn*a[i]%mod;
}
}
inline void poly_px(int *f,int *g,int len_mx) {for(int i=0;i<len_mx;i++) f[i]=(ll)(f[i])*g[i]%mod;}
inline void poly_cpy(int *f,int *g,int len_mx) {for(int i=0;i<len_mx;i++) f[i]=g[i];}
inline void poly_clr(int *f,int l,int r) {f(i,l,r) f[i]=0;}
inline void NTT_mul(int *f,int *g,int *ret,int lenf,int leng) {
static int X0[N],Y0[N];
for (len=1;len<lenf+leng;len<<=1);
initrev();
poly_cpy(X0,f,lenf);
poly_cpy(Y0,g,leng);
NTT(X0,1);NTT(Y0,1);poly_px(X0,Y0,len);NTT(X0,-1);
f(i,0,lenf+leng) ret[i]=X0[i];
f(i,0,len) X0[i]=Y0[i]=0;
}
}
constexpr int N=3e5+10;
constexpr int mod=998244353;
inline int ad(int x,int y) {return (x+y>=mod)?(x+y-mod):(x+y);}
inline void add(int &x,int y) {x=ad(x,y);}
int n;
int val[N],fac[N];
inline vector<int> mul(vector<int> X,vector<int> Y) {
static int F[N],G[N],H[N];
int lenf=X.size(),leng=Y.size(),lenh=lenf+leng-1;
f(i,0,lenf-1) F[i]=X[i];
f(i,0,leng-1) G[i]=Y[i];
vector<int> ret;
Poly::NTT_mul(F,G,H,lenf,leng);
f(i,0,lenh-1) ret.push_back(H[i]);
return ret;
}
inline vector<int> solve(int l,int r) {
if (l==r) {
vector<int> ret;
ret.push_back(1);
ret.push_back(val[l]);
return ret;
}
int mid=l+r>>1;
return mul(solve(l,mid),solve(mid+1,r));
}
int main() {
read(n);fac[0]=1;
f(i,2,n) {
int x,y;read(x);read(y);
val[x]++;val[y]++;
}
f(i,2,n) val[i]--;
vector<int> ret=solve(1,n);
f(i,1,n) fac[i]=1ll*fac[i-1]*i%mod;
int ans=0;
f(i,0,n-1) {
int val=1ll*((i&1)?-1:1)*fac[n-i]*ret[i]%mod;
if (val<0) val+=mod;
add(ans,val);
}
printf("%d\n",ans);
return 0;
}
Awkward
题意
给定一棵 \(n\) 个节点的树,你需要求有多少个排列满足对于相邻两项都不为树上的一条边,答案对 \(10^9+7\) 取模。
\(2\le n\le 5\times 10^3\)。
solution
只会链的情况喵,不会排列计数喵。
考虑容斥,钦定一个边集 \(S \subseteq E\) 一定满足这些边在排列中都相邻,其他的边随意,令方案数为 \(F(S)\),则答案为 \(\sum (-1)^{\lvert S \rvert }F(S)\)。
注意到对于一个边集 \(S\) 一定不会使得一个点的度数大于 \(2\),因为若一个点度数大于 \(2\) 则会使得它需要相邻的点超过两个,不符合排列的需求。
这样的话一个边集就可以看作若干条链,而每条长度不为 \(1\) 的链上的方案数只有两个(正着放和反着放),将每个连通块缩成一个点,这些点之间可以任意排列,所以对于一个边集 \(S\),令 \(c(S)\) 为其形成长度不为 \(1\) 的链的数量,\(\lvert S \rvert =k\),则有 \(F(S)=2^{c(S)}(n-k)!\)。
那么令 \(A_k=\sum\limits_{\lvert S\rvert =k}2^{c(S)}\),答案就为 \(\sum\limits_{k=0}^{n-1} (-1)^{k}A_k(n-k)!\)。
考虑用个树形 dp 求出 \(A_k\),令 \(f_{i,j,k}\) 表示在第 \(i\) 个节点子树及其连向父亲的边中共选了 \(j\) 条,连向父亲的边是否被选上的每种方案权值之和。
我们将 \(2\) 倍的贡献统一在链顶加上。
我们先将儿子合并起来,令 \(g_{i,t}\) 表示考虑完若干子节点,其中共选了 \(i\) 条边,有 \(t\) 个子节点连向当前节点的所有方案权值和,这里的 \(i\) 条边不包括当前节点与父亲的连边,同时 \(0\le t\le 2\),因为这个点连边数量不超过 \(2\)。
对每个子节点 \(v\),\(g\) 的转移是:
处理完所有子节点后,\(g\) 对 \(f\) 的转移是:
这样就做完了,由树形背包的时间复杂度,两个节点仅在 \(\operatorname{lca}\) 处贡献一次,所以时间复杂度是 \(O(n^2)\) 的。
Code
#include<cstdio>
#include<algorithm>
#include<vector>
using namespace std;
#define ll long long
#define qwq Ff472130
#define f(i,l,r) for (int i=l;i<=r;i++)
#define F(i,l,r) for (int i=l;i>=r;i--)
constexpr int N=5e3+10;
constexpr int inf=1e9+10;
constexpr int mod=1e9+7;
inline int ad(int x,int y) {return ((x+y>=mod)?(x+y-mod):(x+y));}
inline void add(int &x,int y) {x=ad(x,y);}
inline void read(int &x) {
x=0;
char ch=getchar();
while (ch<48) ch=getchar();
while (ch>=48) x=(x<<3)+(x<<1)+(ch^48),ch=getchar();
}
int n;
int siz[N],f[N][N][2],g[N][N][3],tmp[N][3];
ll fac[N];
vector<int> e[N];
inline void dfs(int now) {
siz[now]=1;
g[now][0][0]=1;
for (int v:e[now]) {
dfs(v);
f(i,0,siz[now]) f(k,0,2) tmp[i][k]=g[now][i][k],g[now][i][k]=0;
f(i,0,siz[now]) f(j,0,siz[v]) {
add(g[now][i+j][0],1ll*tmp[i][0]*f[v][j][0]%mod);
add(g[now][i+j][1],(1ll*tmp[i][0]*f[v][j][1]+1ll*tmp[i][1]*f[v][j][0])%mod);
add(g[now][i+j][2],(1ll*tmp[i][1]*f[v][j][1]+1ll*tmp[i][2]*f[v][j][0])%mod);
}
siz[now]+=siz[v];
}
f(i,0,siz[now]) {
f[now][i][0]=(g[now][i][0]+2ll*g[now][i][1]+2ll*g[now][i][2])%mod;
if (i) f[now][i][1]=ad(g[now][i-1][0],g[now][i-1][1]);
}
}
int main() {
freopen("permutation.in","r",stdin);
freopen("permutation.out","w",stdout);
read(n);fac[0]=1;
f(i,1,n) fac[i]=fac[i-1]*i%mod;
f(i,2,n) {
int fa;read(fa);
e[fa].push_back(i);
}
dfs(1);
ll ans=0;
f(i,0,n-1) {
int op=((i&1)?-1:1);
ans+=op*f[1][i][0]*fac[n-i]%mod;
}
printf("%lld\n",(ans%mod+mod)%mod);
return 0;
}
Ban Permutation
题意
给定 \(n,X\),求有多少个长度为 \(n\) 的排列 \(P\) 满足:
- \(\forall 1\leq i\leq n\),\(|P_i-i|\geq X(X\leq 5)\)。
\(1\le n\le 100,1\le X\le 5\)。
solution
这个条件看起来很难做,但是把它反过来变成 \(|P_i-i|< X\) 就好做了。
于是考虑容斥,钦定一个位置集合 \(S\),\(S\) 内所有位置不满足条件,令这样的方案数为 \(f(S)\),根据容斥可以得到答案即为 \(\sum (-1)^{\left\lvert S\right\rvert} \times f(S)\)。
这样就好做了,令 \(dp_{i,j,s}\) 表示选到第 \(i\) 个位置,有 \(j\) 个位置被选入集合 \(S\),并且选入集合 \(S\) 中的位置的数确定,并占用 \([i-X+1,i+X-1]\) 内的数的状态为 \(s\),只考虑选入集合 \(S\) 内的位置的方案数。
转移是简单的,枚举第 \(i\) 个位置选或不选,填入哪个值即可。
最后答案即为 \(\sum (-1)^{j} \times dp_{n,j,s}\times (n-j)!\),其中 \((n-j)!\) 表示不选入 \(S\) 位置的方案数。
时间复杂度 \(O(n^2 4^X X)\)。
Code
#include<cstdio>
#include<algorithm>
using namespace std;
#define ll long long
#define qwq Ff472130
#define f(i,l,r) for (int i=l;i<=r;i++)
#define F(i,l,r) for (int i=l;i>=r;i--)
constexpr int N=100+10;
constexpr int V=(1<<9)+10;
constexpr int inf=1e9+10;
constexpr int mod=998244353;
inline int ad(int x,int y) {return ((x+y>=mod)?(x+y-mod):(x+y));}
inline void add(int &x,int y) {x=ad(x,y);}
inline void read(int &x) {
x=0;
char ch=getchar();
while (ch<48) ch=getchar();
while (ch>=48) x=(x<<3)+(x<<1)+(ch^48),ch=getchar();
}
int n,X;
int dp[N][N][V];
ll fac[N];
int main() {
read(n);read(X);X--;
if (!X) {
fac[1]=0;fac[2]=1;
f(i,3,n) fac[i]=(fac[i-1]+fac[i-2])*(i-1)%mod;
printf("%d\n",(int)(fac[n]));
return 0;
}
dp[0][0][0]=1;
int num=X*2+1,S=(1<<num)-1;
f(i,1,n) f(j,0,i) f(s,0,S) {
int rs=(s>>1);
add(dp[i][j][rs],dp[i-1][j][s]);
if (!j) continue;
f(k,0,num-1) {
int val=(i-X+k);
if (val<1||val>n||((rs>>k)&1)) continue;
add(dp[i][j][rs|(1<<k)],dp[i-1][j-1][s]);
}
}
fac[0]=1;
f(i,1,n) fac[i]=fac[i-1]*i%mod;
ll ans=0;
f(i,0,n) f(s,0,S) {
if (i&1) ans-=dp[n][i][s]*fac[n-i]%mod;
else ans+=dp[n][i][s]*fac[n-i]%mod;
}
printf("%d\n",(int)((ans%mod+mod)%mod));
return 0;
}
[SNOI2024] 公交线路
题意
给定一棵 \(n\) 个节点的树,你需要在一些点对间的简单路径上建立公交线路,使得每两个点间至多只需换乘一次公交线路就能抵达,求方案数。
形式化地说,考虑树上的所有 \(\frac{n (n - 1)}{2}\) 条两个端点不同的简单路径。对于这些路径的一个子集 \(S\),称它是好的当且仅当:
- 考虑一张新的图 \(G\),对于一对点 \(u, v\),当且仅当存在 \(S\) 中的一条路径 \(P\),满足 \(u\) 和 \(v\) 都在 \(P\) 上,我们会在 \(u, v\) 之间连上边权为 \(1\) 的无向边。
- 要求 \(G\) 中任意两点之间的距离都不超过 \(2\)。
你需要求出有多少个子集 \(S\) 是好的。由于答案可能很大,输出对 \(998244353\) 取模的结果。
\(1\le n\le 3000\)。
solution
神秘的容斥,做完发现欸我怎么在天上飞。
首先 \(n\le 2\) 时答案为 \(1\),以下只考虑 \(n\ge 3\) 的情况,首先找到一个非叶子节点作根。
每两个点间至多换乘一次的限制只需要考虑叶子节点是否满足即可,而叶子节点满足条件当且仅当考虑叶子节点直接经过的路径并集,所有叶子覆盖路径的交为一个连通块,这是显然的。
对连通块相关计数,由树的性质结合平面图欧拉定理可得 \(V-E=1\),所以可以变成对点减边计数,即令 \(F_x\) 为点 \(x\) 被所有叶子节点所选的路径共同覆盖(即包含在交集中)的方案数,\(G_e\) 为边 \(e\) 被所有叶子路径覆盖方案数,那么答案即为 \(\sum F_x -\sum G_e\)。
接下来需要计算 \(F_x\) 和 \(G_e\)。
\(F_x\) 似乎不是很好直接计算,于是考虑容斥,钦定一个叶子节点的集合 \(S\),使得 \(S\) 内的所有叶子节点直接经过的路径都不经过 \(x\),其他叶子没有限制,令这样的方案数为 \(f(S)\),那么 \(F_x=\sum (-1)^{\left \lvert S\right \rvert} f(S)\)。
考虑进行一个 dp,令 \(dp_i\) 表示 \(x\) 子树内有 \(i\) 个点没有限制(非叶子节点或叶子节点且不在 \(S\) 内)的方案数,这里的方案数只考虑了跨过 \(x\) 的路径,其他的路径不会计入 dp 内,加入一棵子树 \(v\),令 \(v\) 有 \(siz_v\) 个节点(包括叶子),\(lef_v\) 个叶子节点,那么枚举 \(v\) 中有 \(j\) 个叶子节点加入 \(S\) 内,有以下转移:
其中 \((-1)^j\) 是容斥系数,\(\binom{lef_v}{j}\) 表示在 \(v\) 子树内选出 \(j\) 个叶子,\(2^{i\times (siz_v-j)}\) 表示 \(v\) 子树内无限制节点对于其他子树的无限制节点之间的路径可以选或不选,然后无限制点数变成 \(i+siz_v-j\)。
求出子树内的 dp 后,还需要考虑父亲方向的部分带来的贡献,父亲方向就不用容斥了,只需要要求其所有叶子节点直接经过路径经过了 \(x\) 即可。
那么令父亲方向一整块有 \(siz_p\) 个节点,\(lef_p\) 个叶子,那么有以下式子:
其中,\(2^{\binom{siz_p}{2} + \sum\limits_v \binom{siz_v}{2}}\) 表示,不跨过 \(x\) 的路径可以任选,\(2^{i\times (siz_p-lef_p)}\) 表示父亲方向的无限制节点与 \(x\) 子树内无限制节点之间的路径可以任选,\((2^i-1)^{lef_p}\) 表示父亲方向的每个叶子节点与 \(x\) 子树内的无限制节点路径中至少要选一条。
那么 \(F_x\) 做完了,那么考虑 \(G_e\) 怎么求。
实际上和 \(F_x\) 的求法很类似,对于一条边 \(e\),令其连接了 \(u,v\) 两个节点,这两个方向的节点数量和叶子数量分别为 \(siz,lef\),那么枚举 \(u\) 这一侧的限制节点数(在 \(S\) 内的节点数)即可得到:
各项系数和上面求点的系数几乎一样,不再解释了。
预处理出各种幂和组合数,树上背包复杂度 \(O(n^2)\),所以计算点贡献 \(O(n^2)\),每条边 \(O(n)\) 计算贡献,所以计算边贡献 \(O(n^2)\)。
总时间复杂度 \(O(n^2)\),空间复杂度 \(O(n^2)\),空间瓶颈在于预处理出的幂数组。
Code
#include<cstdio>
#include<algorithm>
#include<vector>
using namespace std;
#define ll long long
#define qwq Ff472130
#define f(i,l,r) for (int i=l;i<=r;i++)
#define F(i,l,r) for (int i=l;i>=r;i--)
constexpr int N=3000+10;
constexpr int inf=1e9+10;
constexpr int mod=998244353;
inline int ad(int x,int y) {return ((x+y>=mod)?(x+y-mod):(x+y));}
inline void add(int &x,int y) {x=ad(x,y);}
inline ll qpow(ll a,int b) {
ll res=1;
while (b) {
if (b&1) res=res*a%mod;
a=a*a%mod;
b>>=1;
}
return res;
}
ll fac[N],inv[N],pw2[N*N],pw0[N][N];
inline void init(int n,int nn) {
fac[0]=pw2[0]=1;
f(i,1,n) pw2[i]=ad(pw2[i-1],pw2[i-1]);
n=nn;
f(i,1,n) fac[i]=fac[i-1]*i%mod;
inv[n]=qpow(fac[n],mod-2);
F(i,n,1) inv[i-1]=inv[i]*i%mod;
f(i,1,n) {
ll mul=ad(pw2[i],mod-1);
pw0[i][0]=1;
f(j,1,n) pw0[i][j]=pw0[i][j-1]*mul%mod;
}
}
inline ll C(int n,int m) {
if (n<0||m<0||n<m) return 0;
return fac[n]*inv[m]%mod*inv[n-m]%mod;
}
inline void read(int &x) {
x=0;
char ch=getchar();
while (ch<48) ch=getchar();
while (ch>=48) x=(x<<3)+(x<<1)+(ch^48),ch=getchar();
}
int n,sum_lef;
vector<int> e[N];
ll ans;
int siz[N],lef[N];
int f[N],g[N];
inline void dfs(int now,int fa) {
siz[now]=1;
lef[now]=(e[now].size()==1);
int Cs=0;
for (int v:e[now]) if (v^fa) dfs(v,now);
f[1]=1;
for (int v:e[now]) if (v^fa) {
f(i,1,siz[now]) g[i]=f[i],f[i]=0;
f(i,1,siz[now]) f(j,0,lef[v]) {
if (j&1) add(f[i+siz[v]-j],mod-C(lef[v],j)*pw2[i*(siz[v]-j)]%mod*g[i]%mod);
else add(f[i+siz[v]-j],C(lef[v],j)*pw2[i*(siz[v]-j)]%mod*g[i]%mod);
}
siz[now]+=siz[v];
lef[now]+=lef[v];
Cs+=siz[v]*(siz[v]-1)/2;
}
int sF=n-siz[now],lF=sum_lef-lef[now];
Cs+=sF*(sF-1)/2;
ll sum=0;
f(i,1,siz[now]) {
sum+=f[i]*pw0[i][lF]%mod*pw2[i*(sF-lF)]%mod;
f[i]=0;
}
ans+=(sum%mod)*pw2[Cs]%mod;
if (!fa) return;
sum=0;
f(i,0,lef[now]) {
int k=siz[now]-i;
sum+=((i&1)?-1:1)*pw2[k*(sF-lF)]*pw0[k][lF]%mod*C(lef[now],i)%mod;
}
ans-=sum%mod*pw2[sF*(sF-1)/2+siz[now]*(siz[now]-1)/2]%mod;
}
int main() {
read(n);init(n*n,n);
if (n<=2) return puts("1"),0;
f(i,2,n) {
int u,v;read(u);read(v);
e[u].push_back(v);
e[v].push_back(u);
}
f(i,1,n) sum_lef+=(e[i].size()==1);
f(i,1,n) if (e[i].size()>=2) {dfs(i,0);break;}
printf("%d\n",(int)((ans%mod+mod)%mod));
return 0;
}
一些习题
DAG 容斥
数 DAG 定向
给定无向图,求给边定向使得其是 DAG 的方案数。
对于一张 DAG,一个明显的性质就是存在拓扑排序,于是从这里入手,我们每次删去无出度(或无入度)结点,就能递归到子问题。
但是有一个明显的缺点就是:如果只考虑一个无出度点,由于拓扑排序个数不唯一,那么删点的顺序也是不唯一的。所以我们考虑每次删去所有无出度点,转化成一个子问题解决。
回到题目,令 \(dp_S\) 表示导出子图为 \(S\) 时,上述问题的答案。
我们每次枚举无出度点集,需要保证它是独立集,可以得到转移方程:\(dp_S \leftarrow dp_{S \setminus T}\),\(T \subseteq S\) 且 \(T\) 为独立集。
这个转移会出错,因为它在转移时并没有保证除去 \(T\) 后剩余的点集都有出度,导致记重。具体的,若一个集合 \(T_0\) 内的点没有出度,那么对于任意 \(T \subseteq T_0 \land T \ne \emptyset\) 的 \(T\) 都会造成一次贡献,即造成 \(2^{\left\lvert T_0\right\rvert}-1\) 次贡献,但是实际上应该只有一次贡献。
考虑容斥,对于 \(dp_{S \setminus T}\) 配上 \((-1)^{\left\lvert T\right\rvert+1}\) 的容斥系数,那么对于每一个无环定向,令其无出度点集为 \(S\),则会被计入 \(\sum_{T\subseteq S \land T \ne \emptyset} (-1)^{\left\lvert T\right\rvert+1}=1\)(由组合数学可以推得),这样就解决了数重的问题。
转移方程 \(dp_S \leftarrow (-1)^{\left\lvert T\right\rvert+1}dp_{S \setminus T}\)。
直接 dp 是 \(O(3^n)\) 的,考虑多项式优化。
令 \(g_S=\sum_{T \subseteq S \land T \ne \emptyset}(-1)^{\left\lvert T\right\rvert+1}[T是独立集]\),那么 \(f = f \ * g+1\),其中乘法为子集卷积,得到 \(f=\frac{1}{1-g}\),用集合幂级数求逆可以做到 \(O(n^22^n)\)。
在对 DAG 的容斥中,我们可以考虑如何配系数使得合法方案正好被计入 \(1\) 次,非法方案正好被计入 \(0\) 次,即需要做到不重不漏不记错,这种技巧被称为“DAG 容斥”。
相关题目:
有标号 DAG 计数
题意
对 \(n\) 个点的有标号 DAG 计数,要求:弱连通图。
\(1\le n\le 10^5\)。
solution
先考虑有标号不要求弱连通 DAG 怎么计数。
令 \(f_i\) 表示 \(i\) 个点有标号不要求弱连通 DAG 数量,那么枚举删去无出度点后剩余点个数 \(j\),根据上面提到的 DAG 容斥技巧,可以得到转移:
其中 \((-1)^{i-j+1}\) 是容斥系数,\(\binom{i}{i-j}\) 是选出无出度点的方案,\(2^{j\times(i-j)}\) 表示无出度点可以向非无出度点随意连边。
把组合数拆开,经典的两边除以 \(i!\) 即可处理掉 \(\binom{i}{i-j}\)。
处理 \(2^{j\times(i-j)}\) 有个经典的 trick:\(j\times(i-j)=\binom{i}{2}-\binom{j}{2}-\binom{i-j}{2}\),两边除以 \(\binom{i}{2}\) 即可解决。
最后得到式子:
设出两个生成函数:
得到 \(F(x)=F(x)G(x)+1\),即 \(F(x)=\frac{1}{1-G(x)}\)。
这个式子可以用分治 NTT 做到 \(O(n\log^2 n)\) 或多项式求逆做到 \(O(n\log n)\)。
然后发现不要求弱连通实际上就是若干个弱连通 DAG 方案乘积之和,这个组合意义就是 \(\exp\) 的意义,所以先转成 \(EGF\) 再求 \(\ln\) 即可。
Code
#include<cstdio>
#include<algorithm>
#include<vector>
#include<cmath>
using namespace std;
#define ll long long
#define qwq Ff472130
#define f(i,l,r) for (int i=l;i<=r;i++)
#define F(i,l,r) for (int i=l;i>=r;i--)
constexpr int N=2e5+10;
constexpr int inf=1e9+10;
namespace Poly {
constexpr int N=4e6+10;
constexpr double pi=acos(-1.0);
constexpr int mod=998244353,gen=3;
inline ll qpow(ll a,int b) {
ll res=1;
while (b) {
if (b&1) res=res*a%mod;
a=a*a%mod;b>>=1;
}
return res;
}
static int len,rev[N];
inline void initrev() {
const int lt=log(len)/log(2)-1;
f(i,1,len-1) rev[i]=((rev[i>>1]>>1)|((i&1)<<lt));
}
static int g[30],flag_init_gen;
inline void init_gen() {f(i,1,22) g[i]=qpow(gen,(mod-1)>>i);}
inline int ad(int x,int y) {return (x+y>=mod)?(x+y-mod):(x+y);}
inline void add(int &x,int y) {x=ad(x,y);}
inline void NTT(int *a,int type) {
if (!flag_init_gen) init_gen(),flag_init_gen=1;
f(i,1,len-1) if (rev[i]<i) swap(a[i],a[rev[i]]);
for (int mid=1,t=1;mid<len;mid<<=1,t++) {
int Wn=g[t];
for (int i=mid<<1,j=0;j<len;j+=i) {
int w=1;
for (int k=0;k<mid;k++,w=(ll)(w)*Wn%mod) {
int x=a[j+k],y=(ll)(w)*a[j+k+mid]%mod;
a[j+k]=ad(x,y);a[j+k+mid]=ad(x,mod-y);
}
}
}
if (type==-1) {
reverse(a+1,a+len);
const ll invn=qpow(len,mod-2);
f(i,0,len-1) a[i]=invn*a[i]%mod;
}
}
inline void poly_px(int *f,int *g,int len_mx) {for(int i=0;i<len_mx;i++) f[i]=(ll)(f[i])*g[i]%mod;}
inline void poly_cpy(int *f,int *g,int len_mx) {for(int i=0;i<len_mx;i++) f[i]=g[i];}
inline void poly_clr(int *f,int l,int r) {f(i,l,r) f[i]=0;}
inline void NTT_mul(int *f,int *g,int *ret,int lenf,int leng) {
static int X0[N],Y0[N];
for (len=1;len<lenf+leng;len<<=1);
initrev();
poly_cpy(X0,f,lenf);
poly_cpy(Y0,g,leng);
NTT(X0,1);NTT(Y0,1);poly_px(X0,Y0,len);NTT(X0,-1);
f(i,0,lenf+leng) ret[i]=X0[i];
f(i,0,len) X0[i]=Y0[i]=0;
}
inline void poly_inv(int *F,int *G,int lenf) {
static int X1[N],Y1[N];
int Rlen=1;
for (;Rlen<lenf;Rlen<<=1);
f(i,0,Rlen) G[i]=0;
G[0]=qpow(F[0],mod-2);
for (int l=2;l<=Rlen;l<<=1) {
poly_cpy(X1,F,l);poly_cpy(Y1,G,l);
len=(l<<1);initrev();
NTT(Y1,1);poly_px(Y1,Y1,len);
NTT(X1,1);poly_px(Y1,X1,len);
NTT(Y1,-1);poly_clr(Y1,l,len);
for (int i=0;i<l;i++) add(G[i],ad(G[i],mod-Y1[i]));
}
poly_clr(G,lenf,Rlen);
Rlen<<=1;
f(i,0,Rlen) X1[i]=Y1[i]=0;
}
static int inv[N],flag_init_inv;
inline void init_inv() {inv[1]=1;for(int i=2;i<N;i++)inv[i]=1ll*inv[mod%i]*(mod-mod/i)%mod;}
inline void poly_dao(int *f,int *g,int len_mx) {
g[len_mx-1]=0;
for (int i=1;i<len_mx;i++) g[i-1]=1ll*f[i]*i%mod;
}
inline void poly_jif(int *f,int *g,int len_mx) {
if (!flag_init_inv) init_inv(),flag_init_inv=1;g[0]=0;
for (int i=1;i<len_mx;i++) g[i]=1ll*f[i-1]*inv[i]%mod;
}
inline void poly_ln(int *F,int *G,int lenf) {
static int X3[N],Y3[N];
poly_dao(F,X3,lenf);
poly_inv(F,G,lenf);
NTT_mul(X3,G,Y3,lenf,lenf);
poly_jif(Y3,G,lenf);
lenf<<=2;
f(i,0,lenf) X3[i]=Y3[i]=0;
}
}
constexpr int mod=998244353;
inline int ad(int x,int y) {return ((x+y>=mod)?(x+y-mod):(x+y));}
inline void add(int &x,int y) {x=ad(x,y);}
inline ll qpow(ll a,int b) {
ll res=1;
while (b) {
if (b&1) res=res*a%mod;
a=a*a%mod;
b>>=1;
}
return res;
}
inline void read(int &x) {
x=0;
char ch=getchar();
while (ch<48) ch=getchar();
while (ch>=48) x=(x<<3)+(x<<1)+(ch^48),ch=getchar();
}
ll fac[N],inv[N];
inline void init(int n) {
fac[0]=1;
f(i,1,n) fac[i]=fac[i-1]*i%mod;
inv[n]=qpow(fac[n],mod-2);
F(i,n,1) inv[i-1]=inv[i]*i%mod;
}
inline vector<int> NTT(vector<int> &x,vector<int> &y) {
int lenx=x.size(),leny=y.size();
static int X[N<<1],Y[N<<1],H[N<<1];
f(i,0,lenx-1) X[i]=x[i];
f(i,0,leny-1) Y[i]=y[i];
Poly::NTT_mul(X,Y,H,lenx,leny);
int lenh=lenx+leny-1;
vector<int> ret;
f(i,0,lenh-1) ret.push_back(H[i]);
return ret;
}
int dp[N],ans[N];
inline void solve(int l,int r) {
if (l==r) return;
int mid=l+r>>1;
solve(l,mid);
vector<int> f,g;
f(i,l,mid) f.push_back(dp[i]);
f(i,0,r-l) {
if (i&1) g.push_back(inv[i]*qpow(2,(mod-1)-(1ll*i*(i-1)/2%(mod-1)))%mod);
else g.push_back(mod-inv[i]*qpow(2,(mod-1)-(1ll*i*(i-1)/2%(mod-1)))%mod);
}
vector<int> h=NTT(f,g);
f(i,mid+1,r) add(dp[i],h[i-l]);
solve(mid+1,r);
}
int main() {
int n;read(n);init(n);
dp[0]=1;solve(0,n);
f(i,0,n) dp[i]=dp[i]%mod*qpow(2,(1ll*i*(i-1)/2)%(mod-1))%mod;
Poly::poly_ln(dp,ans,n+1);
f(i,1,n) printf("%d\n",(int)(ans[i]*fac[i]%mod));
return 0;
}
反演相关
二项式反演
两种形式
高维形式
一层层用一维的形式展开即可得到。
组合意义
令 \(f_i\) 表示钦定选 \(i\) 个元素的方案数,\(g_j\) 表示恰好选 \(j\) 个元素的方案数,那么对于 \(i\le j\),\(g_j\) 在 \(f_i\) 中计入了 \(\binom{j}{i}\) 次,所以有 \(f_i=\sum\limits_{j=i}^{m} \binom{j}{i} g_j\),其中 \(m\) 为上界。
二项式反演多用于钦定和恰好的转化,令 \(f_i\) 表示钦定了 \(i\) 个元素的方案,\(g_j\) 表示恰好选择了 \(j\) 个元素的方案,那么我们可以求出 \(f\) 后利用二项式反演求出 \(g\)。
二项式反演是容斥中仅关心集合大小时的特殊形式。
[JSOI2015] 染色问题
题意
有一个 \(n\times m\) 的棋盘和 \(c\) 种颜色,你需要在棋盘每个格子上决定不填颜色或填上 \(c\) 种颜色中的一种,你需要求出满足以下条件填入颜色的方案数:
- 每行和每列至少有一个位置填入颜色;
- 每种颜色至少在棋盘上出现一次。
\(1\le n,m,c\le 400\)。
solution
令 \(f_{n,m,c}\) 表示 \(n\) 行 \(m\) 列 \(c\) 中颜色没有行列和颜色至少出现限制的总方案数,那么显然有 \(f_{n,m,c}=(c+1)^{nm}\)。
令 \(g_{n,m,c}\) 表示有 \(n\) 行 \(m\) 列染上颜色,用过 \(c\) 种颜色的方案数,那么 \(f\) 和 \(g\) 有以下关系:
那么根据二项式反演可以得到:
答案即为 \(g_{n,m,c}\)。
这个式子可以用二项式定理优化一下:
\(n,m,k\) 同阶则时间复杂度 \(O(n^2)\),代码偷了个懒没有预处理幂数组多带一只 \(\log\)。
Code
#include<cstdio>
#include<algorithm>
using namespace std;
#define ll long long
#define qwq Ff472130
#define f(i,l,r) for (int i=l;i<=r;i++)
#define F(i,l,r) for (int i=l;i>=r;i--)
constexpr int N=500+10;
constexpr int inf=1e9+10;
constexpr int mod=1e9+7;
inline ll qpow(ll a,int b) {
ll res=1;
while (b) {
if (b&1) res=res*a%mod;
a=a*a%mod;
b>>=1;
}
return res;
}
ll fac[N],inv[N];
inline void init(int n) {
fac[0]=1;
f(i,1,n) fac[i]=fac[i-1]*i%mod;
inv[n]=qpow(fac[n],mod-2);
F(i,n,1) inv[i-1]=inv[i]*i%mod;
}
inline ll C(int n,int m) {
if (n<0||m<0||n<m) return 0;
return fac[n]*inv[m]%mod*inv[n-m]%mod;
}
int main() {
int n,m,c;init(N-1);
scanf("%d%d%d",&n,&m,&c);
ll ans=0;
f(i,0,n) f(k,0,c) ans+=(((n-i+c-k)&1)?-1:1)*C(n,i)*C(c,k)%mod*qpow(qpow(k+1,i)-1,m)%mod;
printf("%d\n",(int)((ans%mod+mod)%mod));
return 0;
}
已经没有什么好害怕的了
题意
给定两个长度为 \(n\) 的数组 \(a,b\),你需要给 \(a\) 中的每个元素与 \(b\) 中的一个元素配对,每个元素只能参与一次配对,求配对后每对中 \(a\) 的值比 \(b\) 的值大的对数比 \(a\) 的值比 \(b\) 的值小的对数多 \(k\) 的配对方案数,保证 \(a,b\) 中给出的元素均不相同。
\(1\le n\le 2000,0\le k\le n\)。
solution
等价于配对使恰好 \(\frac{n+k}{2}\) 对中 \(a\) 的值大于 \(b\) 的值。
套路地将恰好转化为钦定,令 \(f_i\) 表示钦定了 \(i\) 个配对 \(a\) 值大于 \(b\) 值的方案数,用 dp 可以 \(O(n^2)\) 求出。
具体地,将 \(a\) 排序,令 \(dp_{i,j}\) 表示考虑到第 \(i\) 个元素,已经钦定了 \(j\) 个元素的方案数,令 \(cnt\) 表示 \(b\) 中小于 \(a_i\) 的元素数量,那么有以下转移:
分别表示第 \(i\) 个元素是否钦定的方案,最后算入不钦定的元素随意匹配的方案数,即 \(f_i=dp_{n,i}\times (n-i)!\)。
那么根据二项式反演可以得到:
时间复杂度 \(O(n^2)\),空间复杂度 \(O(n)\)。
Code
#include<cstdio>
#include<algorithm>
using namespace std;
#define ll long long
#define qwq Ff472130
#define f(i,l,r) for (int i=l;i<=r;i++)
#define F(i,l,r) for (int i=l;i>=r;i--)
constexpr int N=2000+10;
constexpr int inf=1e9+10;
constexpr int mod=1e9+9;
inline int ad(int x,int y) {return ((x+y>=mod)?(x+y-mod):(x+y));}
inline void add(int &x,int y) {x=ad(x,y);}
inline ll qpow(ll a,int b) {
ll res=1;
while (b) {
if (b&1) res=res*a%mod;
a=a*a%mod;b>>=1;
}
return res;
}
inline void read(int &x) {
x=0;
char ch=getchar();
while (ch<48) ch=getchar();
while (ch>=48) x=(x<<3)+(x<<1)+(ch^48),ch=getchar();
}
int n,k;
int a[N],b[N],dp[N];
ll fac[N],inv[N];
inline void init(int n) {
fac[0]=1;
f(i,1,n) fac[i]=fac[i-1]*i%mod;
inv[n]=qpow(fac[n],mod-2);
F(i,n,1) inv[i-1]=inv[i]*i%mod;
}
inline ll C(int n,int m) {
if (n<0||m<0||n<m) return 0;
return fac[n]*inv[m]%mod*inv[n-m]%mod;
}
int main() {
read(n);read(k);init(n);
f(i,1,n) read(a[i]);
f(i,1,n) read(b[i]);
if ((n+k)&1) return puts("0"),0;
sort(a+1,a+1+n);
sort(b+1,b+1+n);
dp[0]=1;k=(n+k)/2;
int now=0;
f(i,1,n) {
while (now<n&&b[now+1]<a[i]) now++;
F(j,i,1) {
ll val=now-j+1;
if (val>0) add(dp[j],dp[j-1]*val%mod);
}
}
ll ans=0;
f(j,k,n) ans+=(((j-k)&1)?-1:1)*C(j,k)*dp[j]%mod*fac[n-j]%mod;
printf("%d\n",(int)((ans%mod+mod)%mod));
return 0;
}

浙公网安备 33010602011771号