题解:P15721 [JAG 2023 Summer Camp #3] Many-hued Tree
_zaa_
·
·
题解
操作就相当于是将两个相邻的颜色差为 1 的点合并为一个点,考虑操作肯定是通过 1 和 n 开始,到最后一次操作,只剩下了一条边,左右两个点分别 1\sim i 中的所有点和 i+1\sim n 合并而成的,于是你考虑枚举这一条边,计算左右两边的方案数之积,这样是 O(n^2) 的。
至于如何计算 1\sim i 可以合并的方案数?在 1\sim i 那边单独看作是一棵树,然后设 x 位置上的颜色为 1,那么可以合并成一个点的染色方案就是以 x 为根的树的拓扑序个数,考虑这个怎么求:
设
总方案数也就是说为 $\frac{(siz_x-1)!}{\prod_{v\not=x} siz_v}=\frac{siz_x!}{\prod_v siz_v}$。
$i+1\sim n$ 的方案数同理。
但是如果只算计算每条边左右两边的方案数之积,这样子其实是不对的!因为我们会发现如果有 $x-a-y$ 这种情况($x,y$ 代表一个子树,$a$ 代表一个点,这种情况也就是说有一个度数为 $2$ 的点),我如果最后两步只剩下 $x-a,a-y$ 这两条边,实际上这应该怎么操作都是对的,但是我们却会在这两个边都算上这种情况,这样子就会算重。
考虑这个会算重的情况肯定是最后长 $x-1-2-3-\dots-p-y$ 这样子,也就是说 $1,2,3\dots,p$ 都是二度点,而这种情况正好多算了 $p$ 次,于是考虑对于每个二度点,减去这些情况即可(就是左右都可以合并成一个点)。
## Code
```cpp
#include<bits/stdc++.h>
//#pragma GCC optimize("Ofast,no-stack-protector,unroll-loops,fast-math")
//#pragma GCC optimize(2)
//#pragma GCC optimize(3)
//#pragma GCC optimize("Ofast,unroll-loops")
//#pragma GCC target("sse,sse2,sse3,ssse3,sse4.1,sse4.2,avx,avx2,popcnt,tune=native")
//#include <immintrin.h>
//#include <emmintrin.h>
#define int long long
#define ls(x) ((x)*2)
#define rs(x) ((x)*2+1)
#define pii pair<int,int>
#define fi first
#define se second
#define Debug(...) fprintf(stderr, __VA_ARGS__)
#define For(i,a,b) for(int i=a,i##end=b;i<=i##end;i++)
#define Rof(i,a,b) for(int i=a,i##end=b;i>=i##end;i--)
#define rep(i, b) for(int i=1,i##end=b;i<=i##end;i++)
using namespace std;
const int N=200000+5;
const int M=N*2;
const int Mod=998244353;
inline void chmx(int &x,int y){(x<y)&&(x=y);}
inline void chmn(int &x,int y){(x>y)&&(x=y);}
inline void Add(int &x,int y){(x=x+y+Mod)%=Mod;}
inline int read(){
int f=0,x=0;
char ch=getchar();
while(!isdigit(ch)){f|=(ch=='-');ch=getchar();}
while(isdigit(ch)){x=(x<<3)+(x<<1)+(ch^48);ch=getchar();}
return f?-x:x;
}
void print(int n){
if(n<0){
putchar('-');
n*=-1;
}
if(n>9) print(n/10);
putchar(n%10+'0');
}
int n;
vector<int>q[N];
pii e[N];
int fac[N],inv[N];
int vis[N],tim;
int siz[N],f[N];
inline int ksm(int a,int b){
int res=1;
while(b){
if(b&1) res=res*a%Mod;
a=a*a%Mod;
b>>=1;
}
return res;
}
inline int C(int n,int m){
if(n<m||m<0) return 0;
return fac[n]*inv[m]%Mod*inv[n-m]%Mod;
}
int fa[N];
int sum=0;
int gg[N],tot;
void dfs1(int x,int faz){
siz[x]=1;gg[++tot]=x;
for(auto v:q[x]){
if(v==faz) continue;
dfs1(v,x);
siz[x]+=siz[v];
}
}
void dfs2(int x,int faz){
for(auto v:q[x]){
if(v==faz) continue;
f[v]=f[x]*siz[v]%Mod*inv[tot-siz[v]]%Mod;
dfs2(v,x);
}
sum+=f[x];sum%=Mod;
}
inline int calc(int x,int faz){
tot=0;
dfs1(x,faz);
f[x]=fac[tot];
For(i,1,tot)f[x]=f[x]*inv[siz[gg[i]]]%Mod;
sum=0;
dfs2(x,faz);
return sum;
}
signed main(){
// freopen("perm.in","r",stdin);
// freopen("perm.out","w",stdout);
// ios::sync_with_stdio(false);
// cin.tie(0); cout.tie(0);
fac[0]=1;
For(i,1,N-5) fac[i]=fac[i-1]*i%Mod;
inv[1]=1;
For(i,2,N-5) inv[i]=Mod-Mod/i*inv[Mod%i]%Mod;
n=read();
For(i,1,n-1){
int u=read(),v=read();
e[i]={u,v};
q[u].push_back(v);
q[v].push_back(u);
}
int ans=0;
For(i,1,n-1){
int u=e[i].fi,v=e[i].se;
ans=(ans+2*calc(u,v)%Mod*calc(v,u))%Mod;
}
For(x,1,n){
if(q[x].size()==2){
int a=0,b=0;
for(auto v:q[x]){
if(!a) a=v;
else b=v;
}
ans=(ans-2*calc(a,x)%Mod*calc(b,x)%Mod+Mod)%Mod;
}
}
printf("%lld\n",ans);
#ifdef LOCAL
Debug("\nMy Time: %.3lfms\n",(double)clock()/CLOCKS_PER_SEC);
#endif
return 0;
}
```