1 条题解

  • 0
    @ 2026-10-2 2:11:16

    先找一个判定条件。

    设 sus_u 表示 uu 的子树中总共进行了多少次操作。对于节点 uu,记 S=∑v∈son(u)svS=\sum_{v\in son(u)}s_v,即所有儿子子树一共进行了 SS 次操作。

    定义 ru=(du−cu) mod dur_u=(d_u-c_u)\bmod d_u。

    当 SS 确定后,节点 uu 自己可以额外操作 xux_u 次,其中 0≤xu≤ru+(du−1)S0\le x_u\le r_u+(d_u-1)S。

    因此 su=S+xus_u=S+x_u,所以有 su∈[S,duS+ru]s_u\in[S,d_uS+r_u]。

    于是可以做树形 DP。设 fu[x]f_u[x] 表示 su=xs_u=x 时的方案数。对于节点 uu,先把所有儿子的 DP 数组卷积起来,得到所有儿子的操作次数之和为 SS 的方案数,然后将 SS 转移到整个区间 [S,duS+ru][S,d_uS+r_u]。

    考虑生成函数。定义 Fu(x)=∑ifu[i]xiF_u(x)=\sum_i f_u[i]x^i,并设所有儿子的 GF 乘积为 F(x)=∏v∈son(u)Fv(x)F(x)=\prod_{v\in son(u)}F_v(x)。

    若儿子的操作次数之和为 SS,那么它对 FuF_u 的贡献为 xS+xS+1+⋯+xduS+rux^S+x^{S+1}+\cdots+x^{d_uS+r_u},即 xS−xduS+ru+11−x\frac{x^S-x^{d_uS+r_u+1}}{1-x}。

    因此可以得到转移:

    Fu(x)=F(x)−xru+1F(xdu)1−xF_u(x)=\frac{F(x)-x^{r_u+1}F(x^{d_u})}{1-x}

    最终答案为 F1(1)F_1(1)。但这个式子在 x=1x=1 处不方便计算,因此做代换 x→exx\to e^x。

    定义 Gu(x)=Fu(ex)G_u(x)=F_u(e^x),同时令 G(x)=F(ex)G(x)=F(e^x),则有 F(edux)=G(dux)F(e^{d_ux})=G(d_ux),于是转移变为

    Gu(x)=G(x)−G(dux)e(ru+1)x1−exG_u(x)=\frac{G(x)-G(d_ux)e^{(r_u+1)x}}{1-e^x}

    原来的答案 F1(1)F_1(1) 也变成了 G1(0)G_1(0),所以最终只需要求根节点生成函数的常数项。

    又因为 1−ex=−x+O(x2)1-e^x=-x+O(x^2),每进行一次除以 1−ex1-e^x 的操作,只会使我们需要的项数增加 11。因此,对于 uu,就只需要保留 Gu(x)G_u(x) 的 dep 项。

    使用 NTT 优化,时间复杂度 O(n2log⁡n)O(n^2\log n)。

    #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
    上传者