题解:P17149 [ICPC 2017 Xi'an R] Island

· · 题解

题意简述

给定一棵以 1 为起点的树。 每条边都可以独立保留或删除。 对每个 k,求使节点 1 所在连通块恰有 k 个节点的删边方案数。

解题思路

把原树以 1 为根。 记 s_u 为节点 u 的子树大小,

$u$ 所在连通块恰有 $k$ 个节点的方案数。 使用生成函数: $$ F_u(x)=\sum_{k=1}^{s_u}f_{u,k}x^k $$ 考虑 $u$ 的一个子节点 $v$。 若保留边 $(u,v)$,节点 $v$ 所在连通块会接到 $u$,贡献为 $F_v(x)$。 若删除这条边,$v$ 子树剩余的 $s_v-1$ 条边不再影响根连通块, 可以任意保留或删除,贡献为 $2^{s_v-1}$。 不同儿子子树内的边相互独立,再乘上节点 $u$ 自身,得到: $$ F_u(x)=x\prod_{v\in C_u}\left(F_v(x)+2^{s_v-1}\right) $$ 其中 $C_u$ 是 $u$ 的子节点集合。 每个因子分别决定父子边保留还是删除, 删除时又完整枚举该子树内部的边, 所以这个转移与所有删边方案一一对应。 直接逐个节点合并会在链形或梳形树上产生平方级复杂度。 下面同时处理多儿子乘积和长链转移。 对树按子树大小选择重子节点。 记 $h_u$ 为 $u$ 的重子节点,$L_u$ 为轻子节点集合,并定义: $$ H_u(x)=\prod_{v\in L_u}\left(F_v(x)+2^{s_v-1}\right) $$ 空积取 $1$。 再令: $$ P_u(x)=xH_u(x) $$ 若 $u$ 有重子节点 $h_u$,原转移变成: $$ F_u(x)=P_u(x)F_{h_u}(x)+P_u(x)2^{s_{h_u}-1} $$ 若 $u$ 没有重子节点,则 $F_u(x)=P_u(x)$。 把一个多项式 $Y$ 看成变量。 对于重链上的非末节点,定义仿射变换: $$ T_u(Y)=P_uY+P_u2^{s_{h_u}-1} $$ 链尾的变换是常值函数 $T_u(Y)=P_u$。 从链尾向上代入,整条链链首的 $F$ 就是这些变换依次复合的结果。 用二元组 $(a,b)$ 表示 $Y\mapsto aY+b$。 上方区间为 $(a,b)$,下方区间为 $(c,d)$ 时,复合结果为: $$ (a,b)\circ(c,d)=(ac,ad+b) $$ 链尾常值函数表示为 $(0,P_u)$。 因此,分治复合整条重链后,二元组的第二项就是链首的 $F$。 代码只为每条重链的链首保存最终 $F$, 其父亲把它作为一个完整轻子树使用。 现在处理 $H_u$。 每个轻子节点 $v$ 已经得到 $F_v$, 只需给常数项加上 $2^{s_v-1}$,再把所有多项式相乘。 将它们按长度放入小根堆, 每次取出最短的两个相乘,能够避免大多项式被反复合并。 多项式乘法较小时直接朴素计算,较大时使用 NTT。 模数满足: $$ 1811939329=27\times2^{26}+1 $$ $13$ 是该模数的原根。 本题多项式长度不超过 $10^5+1$, 所以最大变换长度 $2^{17}$ 可以直接预处理全部单位根。 树的父子关系、子树大小与重子节点都用迭代遍历求出, 避免链形树导致递归栈溢出。 随后按逆序处理节点, 此时所有轻子树的链首多项式都已经计算完成。 ## 正确性证明 首先证明子树转移正确。 对每条父子边,保留时子节点连通块大小由 $F_v$ 记录; 删除时该子树不再影响 $u$ 的连通块,内部每条边可任意选择, 共有 $2^{s_v-1}$ 种。 各子树边集互不相交,乘法会独立组合它们。 再乘 $x$ 计入节点 $u$,故转移恰好统计所有方案。 接着证明重链复合正确。 轻子节点的全部贡献已经包含在 $H_u$ 中。 对有重子节点的 $u$, 变换 $T_u$ 正是把 $F_{h_u}$ 代入原转移后的式子; 链尾没有重子节点,其结果就是 $P_u$。 由链尾向上归纳,仿射变换复合得到的第二项等于链首的 $F$。 所有轻子树先独立求出 $F_v$, 再通过 $F_v+2^{s_v-1}$ 接入父节点。 每个节点恰属于一条重链, 每条父子边也恰在轻子树乘积或重链转移中出现一次。 因此,算法最终得到的 $F_1$ 与原始子树转移完全一致, 其第 $k$ 次系数就是题目所求答案。 一条根到叶路径至多经过 $O(\log n)$ 条轻边。 因此,所有轻子树多项式参与合并的总长度为 $O(n\log n)$。 小根堆乘积和重链分治各增加至多一个对数层数, 而一次 NTT 乘法还需要一个对数因子。 单组时间复杂度为 $O(n\log^3 n)$,空间复杂度为 $O(n\log n)$。 ## 参考代码 ```cpp #include <bits/stdc++.h> using namespace std; using ll=long long; using pii=pair<int,int>; using poly=vector<int>; const int N=100005; const int K=1<<17; const int mod=1811939329; const int g=13; vector<int> G[N]; poly f[N],h[N]; int fa[N],siz[N],son[N],seq[N],pw[N],rt[K],irt[K]; struct trans { poly a,b; }; ll Pow(ll x,ll y) { x%=mod; ll res=1; while(y) { if(y&1)res=res*x%mod; x=x*x%mod; y>>=1; } return res; } void init() { pw[0]=1; rt[0]=1; irt[0]=1; for(int i=1;i<N;i++)pw[i]=(ll)pw[i-1]*2%mod; int x=Pow(g,(mod-1)/K); int y=Pow(x,mod-2); for(int i=1;i<K;i++) { rt[i]=(ll)rt[i-1]*x%mod; irt[i]=(ll)irt[i-1]*y%mod; } } void ntt(poly &a,bool op) { int n=a.size(); for(int i=1,j=0;i<n;i++) { int k=n>>1; while(j>=k) { j-=k; k>>=1; } j+=k; if(i<j)swap(a[i],a[j]); } for(int i=2;i<=n;i<<=1) { int st=K/i; for(int j=0;j<n;j+=i) { for(int k=0;k<i/2;k++) { int x=a[j+k]; int y=(ll)a[j+k+i/2]*(op?rt[k*st]:irt[k*st])%mod; a[j+k]=x-(mod-y); if(a[j+k]<0)a[j+k]+=mod; a[j+k+i/2]=x-y; if(a[j+k+i/2]<0)a[j+k+i/2]+=mod; } } } if(!op) { int inv=Pow(n,mod-2); for(int i=0;i<n;i++)a[i]=(ll)a[i]*inv%mod; } } poly mul(poly a,poly b) { int n=a.size(); int m=b.size(); int len=n+m-1; if((ll)n*m<=4096) { poly c(len); for(int i=0;i<n;i++) for(int j=0;j<m;j++)c[i+j]=(c[i+j]+(ll)a[i]*b[j])%mod; return c; } int k=1; while(k<len)k<<=1; a.resize(k); b.resize(k); ntt(a,1); ntt(b,1); for(int i=0;i<k;i++)a[i]=(ll)a[i]*b[i]%mod; ntt(a,0); a.resize(len); return a; } poly add(poly a,const poly &b) { int n=b.size(); int m=a.size(); if(m<n)a.resize(n); for(int i=0;i<n;i++) { a[i]-=mod-b[i]; if(a[i]<0)a[i]+=mod; } return a; } poly mul(poly a,int b) { int n=a.size(); for(int i=0;i<n;i++)a[i]=(ll)a[i]*b%mod; return a; } trans calc(vector<int> &c,int l,int r) { if(l==r) { int u=c[l]; h[u].insert(h[u].begin(),0); if(!son[u])return {{},move(h[u])}; poly b=mul(h[u],pw[siz[son[u]]-1]); return {move(h[u]),move(b)}; } int m=(l+r)>>1; trans x=calc(c,l,m); trans y=calc(c,m+1,r); poly b=add(mul(x.a,move(y.b)),x.b); poly a; if(!y.a.empty())a=mul(move(x.a),move(y.a)); return {move(a),move(b)}; } void solve() { int n; cin>>n; for(int i=1;i<n;i++) { int u,v; cin>>u>>v; G[u].push_back(v); G[v].push_back(u); } int cnt=1; seq[1]=1; for(int i=1;i<=cnt;i++) { int u=seq[i]; for(int v:G[u]) if(v!=fa[u]) { fa[v]=u; seq[++cnt]=v; } } for(int i=1;i<=n;i++)siz[i]=1; for(int i=n;i;i--) { int u=seq[i]; if(!fa[u])continue; siz[fa[u]]+=siz[u]; if(siz[u]>siz[son[fa[u]]])son[fa[u]]=u; } for(int i=n;i;i--) { int u=seq[i]; vector<poly> p; priority_queue<pii,vector<pii>,greater<pii>> q; for(int v:G[u]) if(fa[v]==u&&v!=son[u]) { poly x=move(f[v]); x[0]-=mod-pw[siz[v]-1]; if(x[0]<0)x[0]+=mod; int k=p.size(); p.push_back(move(x)); q.push({p[k].size(),k}); } if(q.empty())h[u]={1}; else { while(q.size()>1) { int x=q.top().second; q.pop(); int y=q.top().second; q.pop(); int k=p.size(); p.push_back(mul(move(p[x]),move(p[y]))); q.push({p[k].size(),k}); } h[u]=move(p[q.top().second]); } if(fa[u]&&son[fa[u]]==u)continue; vector<int> c; int v=u; while(v) { c.push_back(v); v=son[v]; } int len=c.size(); trans x=calc(c,0,len-1); f[u]=move(x.b); } for(int i=1;i<=n;i++)cout<<f[1][i]<<(i==n?'\n':' '); for(int i=1;i<=n;i++) { G[i].clear(); f[i].clear(); h[i].clear(); fa[i]=0; siz[i]=0; son[i]=0; seq[i]=0; } } int main() { ios::sync_with_stdio(false); cin.tie(nullptr); init(); int T; cin>>T; while(T--)solve(); return 0; } ```