1 条题解
-
0
先找一个判定条件。
设 表示 的子树中总共进行了多少次操作。对于节点 ,记 ,即所有儿子子树一共进行了 次操作。
定义 。
当 确定后,节点 自己可以额外操作 次,其中 。
因此 ,所以有 。
于是可以做树形 DP。设 表示 时的方案数。对于节点 ,先把所有儿子的 DP 数组卷积起来,得到所有儿子的操作次数之和为 的方案数,然后将 转移到整个区间 。
考虑生成函数。定义 ,并设所有儿子的 GF 乘积为 。
若儿子的操作次数之和为 ,那么它对 的贡献为 ,即 。
因此可以得到转移:
最终答案为 。但这个式子在 处不方便计算,因此做代换 。
定义 ,同时令 ,则有 ,于是转移变为
原来的答案 也变成了 ,所以最终只需要求根节点生成函数的常数项。
又因为 ,每进行一次除以 的操作,只会使我们需要的项数增加 。因此,对于 ,就只需要保留 的 dep 项。
使用 NTT 优化,时间复杂度 。
#include<bits/stdc++.h> #define For(i,a,b) for(int i=(a);i<=(b);++i) #define Rep(i,a,b) for(int i=(a);i>=(b);--i) #define ll long long #define ull unsigned long long #define SZ(x) ((int)((x).size())) #define ALL(x) (x).begin(),(x).end() using namespace std; inline int read(){ char c=getchar();int x=0;bool f=0; for(;!isdigit(c);c=getchar())f^=!(c^45); for(;isdigit(c);c=getchar())x=(x<<1)+(x<<3)+(c^48); return f?-x:x; } #define mod 998244353 struct modint{ unsigned int x; modint(int o=0){x=o;} modint &operator = (int o){return x=o,*this;} modint &operator +=(modint o){return x=x+o.x>=mod?x+o.x-mod:x+o.x,*this;} modint &operator -=(modint o){return x=x<o.x?x-o.x+mod:x-o.x,*this;} modint &operator *=(modint o){return x=1ull*x*o.x%mod,*this;} modint &operator ^=(int b){ modint a=*this,c=1; for(;b;b>>=1,a*=a)if(b&1)c*=a; return x=c.x,*this; } modint &operator /=(modint o){return *this *=o^=mod-2;} friend modint operator +(modint a,modint b){return a+=b;} friend modint operator -(modint a,modint b){return a-=b;} friend modint operator *(modint a,modint b){return a*=b;} friend modint operator /(modint a,modint b){return a/=b;} friend modint operator ^(modint a,int b){return a^=b;} friend bool operator ==(modint a,modint b){return a.x==b.x;} friend bool operator !=(modint a,modint b){return a.x!=b.x;} bool operator ! () {return !x;} modint operator - () {return x?mod-x:0;} bool operator <(const modint&b)const{return x<b.x;} }; inline modint qpow(modint x,int y){return x^y;} vector<modint> fac,ifac,iv; inline void initC(int n) { if(iv.empty())fac=ifac=iv=vector<modint>(2,1); int m=iv.size(); ++n; if(m>=n)return; iv.resize(n),fac.resize(n),ifac.resize(n); For(i,m,n-1){ iv[i]=iv[mod%i]*(mod-mod/i); fac[i]=fac[i-1]*i,ifac[i]=ifac[i-1]*iv[i]; } } inline modint C(int n,int m){ if(m<0||n<m)return 0; return initC(n),fac[n]*ifac[m]*ifac[n-m]; } inline modint sign(int n){return (n&1)?(mod-1):(1);} #define fi first #define se second #define pb push_back #define mkp make_pair typedef pair<int,int>pii; typedef vector<int>vi; #define poly vector<modint> const modint G=3,Ginv=modint(1)/3; inline poly one(){poly a;a.push_back(1);return a;} vector<int>rev; int rts[2100000]; inline int ext(int n){ int k=0; while((1<<k)<n)++k;return k; } inline void init(int k){ int n=1<<k; rts[0]=1,rts[1<<k]=qpow(31,1<<(21-k)).x; Rep(i,k,1)rts[1<<(i-1)]=1ull*rts[1<<i]*rts[1<<i]%mod; For(i,1,n-1)rts[i]=1ull*rts[i&(i-1)]*rts[i&-i]%mod; } void ntt(poly&a,int k,int typ){ int n=1<<k; static ull tmp[2100000]; for(int i=0;i<n;++i)tmp[i]=a[i].x; if(typ==1){ for(int l=n>>1;l>=1;l>>=1){ ull*k=tmp; for(int*g=rts;k<tmp+n;k+=(l<<1),++g){ for(ull*x=k;x<k+l;++x){ int o=x[l]%mod*(*g)%mod; x[l]=*x+mod-o,*x+=o; } } } for(int i=0;i<n;++i)a[i].x=tmp[i]%mod; }else{ for(int l=1;l<n;l<<=1){ ull*k=tmp; for(int*g=rts;k<tmp+n;k+=(l<<1),++g){ for(ull*x=k;x<k+l;++x){ int o=x[l]%mod; x[l]=(*x+mod-o)*(*g)%mod,*x+=o; } } } int iv=qpow(n,mod-2).x; for(int i=0;i<n;++i)a[i].x=tmp[i]%mod*iv%mod; reverse(a.begin()+1,a.end()); } } poly operator +(poly a,poly b){ int n=max(a.size(),b.size());a.resize(n),b.resize(n); For(i,0,n-1)a[i]+=b[i];return a; } poly operator -(poly a,poly b){ int n=max(a.size(),b.size());a.resize(n),b.resize(n); For(i,0,n-1)a[i]-=b[i];return a; } poly operator *(poly a,modint b){ int n=a.size(); For(i,0,n-1)a[i]*=b;return a; } poly operator *(poly a,poly b){ if(!a.size()||!b.size())return {}; if((int)a.size()<=32 || (int)b.size()<=32){ poly c(a.size()+b.size()-1,0); for(int i=0;i<a.size();++i)for(int j=0;j<b.size();++j)c[i+j]+=a[i]*b[j]; return c; } int n=(int)a.size()+(int)b.size()-1,k=ext(n); a.resize(1<<k),b.resize(1<<k); ntt(a,k,1),ntt(b,k,1); For(i,0,(1<<k)-1)a[i]*=b[i]; ntt(a,k,-1),a.resize(n);return a; } poly Tmp; poly pmul(poly a,poly b,int n,bool ok=0) { int k=ext(n); a.resize(1<<k),ntt(a,k,1); if(!ok) b.resize(1<<k),ntt(b,k,1),Tmp=b; For(i,0,(1<<k)-1)a[i]*=Tmp[i]; ntt(a,k,-1),a.resize(n); return a; } poly inv(poly a,int n) { a.resize(n); if(n==1){ poly f(1,1/a[0]); return f; } poly f0=inv(a,(n+1)>>1),f=f0; poly now=pmul(a,f0,n,0); for(int i=0;i<f0.size();++i)now[i]=0; now=pmul(now,poly(0),n,1); f.resize(n); for(int i=f0.size();i<n;++i)f[i]=-now[i]; return f; } poly inv(poly a){return inv(a,a.size());} poly deriv(poly a){ int n=(int)a.size()-1; For(i,0,n-1)a[i]=a[i+1]*(i+1); a.resize(n);return a; } poly inter(poly a){ int n=a.size()+1;a.resize(n); initC(n); Rep(i,n-1,1)a[i]=a[i-1]*iv[i]; a[0]=0;return a; } poly ln(poly a){ int n=a.size(); a=deriv(a)*inv(a),a.resize(n-1);return inter(a); } poly exp(poly a,int k){ int n=1<<k;a.resize(n); if(n==1)return one(); poly f0=exp(a,k-1);f0.resize(n); return f0*(one()+a-ln(f0)); } poly exp(poly a){ int n=a.size(); a=exp(a,ext(n));a.resize(n);return a; } poly div(poly a,poly b){ int n=a.size(),m=b.size(),k=ext(n-m+1); reverse(a.begin(),a.end()),reverse(b.begin(),b.end()); a.resize(n-m+1),b.resize(n-m+1); a=a*inv(b),a.resize(n-m+1),reverse(a.begin(),a.end()); return a; } poly modulo(poly a,poly b){ if(b.size()>a.size())return a; int n=b.size()-1; a=a-div(a,b)*b;a.resize(n);return a; } #define maxn 200005 #define inf 0x3f3f3f3f int n,c[maxn],d[maxn],fa[maxn],dep[maxn]; vi e[maxn],ord; poly f[maxn],ber; inline poly mul_lim(poly a,poly b,int lim) { if(a.empty()||b.empty()||lim<=0)return {}; if(SZ(a)>lim)a.resize(lim); if(SZ(b)>lim)b.resize(lim); if(SZ(a)<=32||SZ(b)<=32){ poly res(min(lim,SZ(a)+SZ(b)-1),0); For(i,0,SZ(a)-1){ int R=min(SZ(b)-1,lim-1-i); For(j,0,R)res[i+j]+=a[i]*b[j]; } return res; } poly res=a*b; if(SZ(res)>lim)res.resize(lim); return res; } signed main() { n=read(); For(i,1,n)c[i]=read(),d[i]=read(); For(i,1,n-1){ int u=read(),v=read(); e[u].pb(v),e[v].pb(u); } ord.pb(1),fa[1]=0,dep[1]=0; for(int p=0;p<SZ(ord);++p){ int u=ord[p]; for(int v:e[u])if(v!=fa[u]){ fa[v]=u; dep[v]=dep[u]+1; ord.pb(v); } } int D=0; For(i,1,n)D=max(D,dep[i]); initC(D+2); init(ext(2*D+5)); ber.resize(D+1); ber[0]=1; For(k,1,D){ modint s=0; For(i,1,k)s+=ifac[i+1]*ber[k-i]; ber[k]=-s; } Rep(p,SZ(ord)-1,0){ int u=ord[p]; int r=dep[u],lim=r+2; poly h(1,1); for(int v:e[u])if(fa[v]==u){ h=mul_lim(h,f[v],lim); } h.resize(lim); poly hd(lim); modint pw=1; For(i,0,lim-1){ hd[i]=h[i]*pw; pw*=d[u]; } ll a=(d[u]-1ll*(c[u]%d[u]))%d[u]; modint A=a+1; poly ex(lim); ex[0]=1; For(i,1,lim-1)ex[i]=ex[i-1]*A*iv[i]; poly q=mul_lim(ex,hd,lim); q.resize(lim); For(i,0,lim-1)q[i]-=h[i]; poly g(r+1); For(i,0,r)g[i]=q[i+1]; poly b(r+1); For(i,0,r)b[i]=ber[i]; f[u]=mul_lim(g,b,r+1); f[u].resize(r+1); } printf("%u\n",f[1][0].x); return 0; }
- 1
信息
- ID
- 12706
- 时间
- 1500ms
- 内存
- 350MiB
- 难度
- 10
- 标签
- 递交数
- 4
- 已通过
- 1
- 上传者