题解:P15721 [JAG 2023 Summer Camp #3] Many-hued Tree

· · 题解

操作就相当于是将两个相邻的颜色差为 1 的点合并为一个点,考虑操作肯定是通过 1n 开始,到最后一次操作,只剩下了一条边,左右两个点分别 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; } ```