题解:P17149 [ICPC 2017 Xi'an R] Island
lailai0916
·
·
题解
题意简述
给定一棵以 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;
}
```