题解:P17340 【MX-X30-T6】布谷鸟钟
lailai0916
·
·
题解
题意简述
给定一棵有根树。选择点 u 时,若 c_u 不是 d_u 的倍数,就把 u 到根路径上的数全部加 1。求所有可达数组的数量。
解题思路
设点 u 被选择了 x_u 次,s_u 表示 u 子树内的操作总数。最终的 c_u 会增加 s_u,故最终数组与所有 s_u 一一对应,并且:
x_u=s_u-\sum_{v\in C_u}s_v
其中 C_u 是 u 的儿子集合。记:
a_u=(-c_u)\bmod d_u
观察点 $u$ 经历的 $s_u$ 次加一。第 $t$ 次加一前,点 $u$ 的值为 $c_u+t$。当 $t=a_u+k\cdot d_u$ 时,当前操作不能选择 $u$,只能来自真子树。因此前 $s_u$ 个位置中的每个此类位置都要分配给 $y_u$ 次真子树操作。该条件等价于:
$$
y_u\le s_u\le a_u+d_u\cdot y_u
$$
这个条件也足够。先任取各个儿子子树内的合法操作顺序,再将它们合并,得到长度为 $y_u$ 的序列。将序列中的操作依次放到所有禁止选择 $u$ 的位置。其余操作放到任意空位,剩余位置选择 $u$。插入对 $u$ 的操作不会改变后代的值,所以各个子树内的合法性保持不变。
设 $F_u(z)$ 中 $z^s$ 的系数表示满足 $s_u=s$ 的方案数。各个儿子相互独立,定义:
$$
H_u(z)=\prod_{v\in C_u}F_v(z)
$$
根据 $y_u\le s_u\le a_u+d_u\cdot y_u$ 求和,可以得到:
$$
\begin{aligned}
F_u(z) & =\sum_{y\ge0}[z^y]H_u(z)\sum_{s=y}^{a_u+d_u\cdot y}z^s \\
& =\frac{H_u(z)-z^{a_u+1}H_u(z^{d_u})}{1-z}
\end{aligned}
$$
多项式的次数可能极大,但最终只需求 $F_1(1)$。令 $A_u(t)=F_u(e^t)$、$B_u(t)=H_u(e^t)$,上式变为:
$$
A_u(t)=\frac{e^{(a_u+1)t}B_u(d_u\cdot t)-B_u(t)}{e^t-1}
$$
若 $F_u(z)=\sum_s w_sz^s$,则:
$$
[t^i]A_u(t)=\frac{1}{i!}\sum_s w_ss^i
$$
因此可以直接在 $t=0$ 附近维护形式幂级数。预处理:
$$
Q(t)=\frac{e^t-1}{t}=\sum_{i\ge0}\frac{t^i}{(i+1)!}
$$
以及 $Q(t)^{-1}$。计算 $A_u$ 时,先求 $e^{(a_u+1)t}B_u(d_u\cdot t)-B_u(t)$,删去常数项并整体除以 $t$,再乘 $Q(t)^{-1}$。
由于 $Q(0)=1$,且所需阶数小于模数,相关逆元均存在。此时 $B_u(t)=\prod_{v\in C_u}A_v(t)$,可直接合并儿子的幂级数。
若 $u$ 的深度为 $r$,仅维护 $A_u$ 的 $0\sim r$ 次项。上述转移因为除以 $t$,需要 $B_u$ 的 $0\sim r+1$ 次项。每个儿子的深度都是 $r+1$,已有的信息恰好足够。根的深度为 $0$,最后保留的常数项就是 $A_1(0)=F_1(1)$。
代码用数论变换(Number Theoretic Transform,NTT)完成多项式乘法。所有卷积的长度之和为 $O(n^2)$。因此时间复杂度为 $O(n^2\log n)$,空间复杂度为 $O(n^2)$。
## 参考代码
```cpp
#include <bits/stdc++.h>
using namespace std;
using ll=long long;
const int N=2005;
const int M=4101;
const int mod=998244353;
const int g=3;
int c[N],d[N],dep[N],fa[N],fac[N],ifac[N],iq[N],ord[N];
int f[N][N],h[N],p[N],e[N],tmp[N],ta[M],tb[M],rev[M];
int rt[M]={0,1};
int rn=2;
vector<int> G[N];
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 ntt(int a[],int n,int op)
{
for(int i=0;i<n;i++)rev[i]=(rev[i>>1]>>1)|((i&1)?n>>1:0);
for(int i=0;i<n;i++)if(i<rev[i])swap(a[i],a[rev[i]]);
while(rn<n)
{
int w=int(Pow(g,(mod-1)/(rn<<1)));
for(int i=rn>>1;i<rn;i++)
{
rt[i<<1]=rt[i];
rt[i<<1|1]=int((ll)rt[i]*w%mod);
}
rn<<=1;
}
for(int k=1;k<n;k<<=1)
{
for(int i=0;i<n;i+=k<<1)
{
for(int j=0;j<k;j++)
{
int u=a[i+j],v=int((ll)a[i+j+k]*rt[k+j]%mod);
a[i+j]=u+v<mod?u+v:u+v-mod;
a[i+j+k]=u-v>=0?u-v:u-v+mod;
}
}
}
if(op==-1)
{
reverse(a+1,a+n);
int iv=int(Pow(n,mod-2));
for(int i=0;i<n;i++)a[i]=int((ll)a[i]*iv%mod);
}
}
void mul(const int a[],int na,const int b[],int nb,int res[],int k)
{
int m=min(k,na+nb-1);
if((ll)na*nb<=4096)
{
fill(tmp,tmp+m,0);
for(int i=0;i<na&&i<m;i++)
{
for(int j=0;j<nb&&i+j<m;j++)tmp[i+j]=int((tmp[i+j]+(ll)a[i]*b[j])%mod);
}
copy(tmp,tmp+m,res);
return;
}
int n=1;
while(n<na+nb-1)n<<=1;
fill(ta,ta+n,0);
fill(tb,tb+n,0);
copy(a,a+na,ta);
copy(b,b+nb,tb);
ntt(ta,n,1);
ntt(tb,n,1);
for(int i=0;i<n;i++)ta[i]=int((ll)ta[i]*tb[i]%mod);
ntt(ta,n,-1);
copy(ta,ta+m,res);
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int n;
cin>>n;
for(int i=1;i<=n;i++)cin>>c[i]>>d[i];
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;
ord[0]=1;
for(int i=0;i<cnt;i++)
{
int u=ord[i];
for(auto v:G[u])if(v!=fa[u])
{
fa[v]=u;
dep[v]=dep[u]+1;
ord[cnt++]=v;
}
}
fac[0]=1;
for(int i=1;i<=n+1;i++)fac[i]=int((ll)fac[i-1]*i%mod);
ifac[n+1]=int(Pow(fac[n+1],mod-2));
for(int i=n+1;i;i--)ifac[i-1]=int((ll)ifac[i]*i%mod);
iq[0]=1;
for(int i=1;i<=n;i++)
{
ll sum=0;
for(int j=1;j<=i;j++)sum=(sum+(ll)ifac[j+1]*iq[i-j])%mod;
iq[i]=int((mod-sum)%mod);
}
for(int i=n-1;i>=0;i--)
{
int u=ord[i],k=dep[u]+2;
int len=1;
h[0]=1;
for(auto v:G[u])if(fa[v]==u)
{
if(len==1&&h[0]==1)copy(f[v],f[v]+k,h);
else mul(h,len,f[v],k,h,k);
len=k;
}
fill(h+len,h+k,0);
int a=(d[u]-c[u]%d[u])%d[u],b=a+1;
ll pw=1;
for(int j=0;j<k;j++)
{
p[j]=int(h[j]*pw%mod);
pw=pw*(d[u]%mod)%mod;
}
pw=1;
for(int j=0;j<k;j++)
{
e[j]=int(pw*ifac[j]%mod);
pw=pw*(b%mod)%mod;
}
mul(p,k,e,k,p,k);
for(int j=0;j<k-1;j++)p[j]=(p[j+1]-h[j+1]+mod)%mod;
mul(p,k-1,iq,k-1,f[u],k-1);
}
cout<<f[1][0]<<'\n';
return 0;
}
```