3 条题解

  • 3
    @ 2026-2-25 14:26:49
    #include<bits/stdc++.h>
    using namespace std;
    #define int long long
    #define fu(i,j,k) for(int i=j;i<=k;i++)
    #define fd(i,j,k) for(int i=j;i>=k;i--)
    const int N=8e5+10,P=998244353;
    int qpow(int a,int b){int ans=1;for(;b;b>>=1,a=a*a%P)if(b&1)ans=ans*a%P;return ans;}
    int a[N],b[N],d[N];
    void ntt(int s[],int n,int x)
    {
    	if(n==1)return;
    	int s1[n/2],s2[n/2];
    	fu(i,0,n/2-1)s1[i]=s[i*2],s2[i]=s[i*2+1];
    	ntt(s1,n/2,x*x%P);ntt(s2,n/2,x*x%P);
    	for(int i=0,xi=1;i<n/2;i++,xi=xi*x%P)
    	{
    		s[i]=(s1[i]+s2[i]*xi)%P;
    		s[i+n/2]=((s1[i]-s2[i]*xi)%P+P)%P;
    	} 
    }
    int merge(int s1[],int len1,int s2[],int len2,int s[])
    {
    	int D=1;while(D<len1+len2-1)D<<=1;
    	fu(i,len1,D-1)s1[i]=0;
    	fu(i,len2,D-1)s2[i]=0;
    	int inv=qpow(D,P-2),x=qpow(3,(P-1)/D);
    	ntt(s1,D,x);ntt(s2,D,x);
    	fu(i,0,D-1)s[i]=s1[i]*s2[i]%P;
    	int inv1=qpow(x,P-2);
    	ntt(s,D,inv1);
    	fu(i,0,len1+len2-2)s[i]=s[i]*inv%P;
    	return len1+len2-1;
    }
    int divide(int s[],int l,int r)
    {
    	if(l==r){s[0]=1;s[1]=d[l];return 2;}
    	int tmp1[N],tmp2[N];
    	int mid=(l+r)>>1;
    	int len1=divide(tmp1,l,mid),len2=divide(tmp2,mid+1,r);
    	return merge(tmp1,len1,tmp2,len2,s);
    }
    signed main()
    {
    	int n;cin>>n;
    	fu(i,1,n)cin>>d[i];
    	int s=0;fu(i,1,n)s+=d[i];
    	divide(a,1,n);
    	int ans=0,sum=1;
    	fu(i,0,n)
    	{
    		ans=(ans+sum*a[i])%P;
    		sum=sum*(i+1)%P*qpow(s-i,P-2)%P;
    	}
    	cout<<ans;
    	return 0;
    }
    • 1
      @ 2026-2-25 15:43:15
      #include<bits/stdc++.h>
      using namespace std;
      typedef long long ll;
      const ll mod=998244353;
      ll qpow(ll a,ll b){
      	ll ans=1;
      	for(;b;b>>=1,a=a*a%mod)if(b&1)ans=ans*a%mod;
      	return ans;
      }
      int r[2000010];
      void NTT(vector<ll> &a,ll n,ll x){
      	for(int i=0;i<n;i++)if(i<r[i])swap(a[i],a[r[i]]);
      	for(int m=2;m<=n;m<<=1){
      		ll g1=qpow(x,n/m);
      		for(int i=0;i<n;i+=m){
      			ll gk=1;
      			for(int j=0;j<m/2;j++){
      				ll x=a[i+j],y=a[i+j+m/2]*gk%mod;
      				a[i+j]=(x+y)%mod;a[i+j+m/2]=(x-y+mod)%mod;
      				gk=gk*g1%mod;
      			}
      		}
      	}
      }
      vector<ll> solve(vector<ll> a,vector<ll> b){
      	if(a.empty()||b.empty())return {0};
      	int n=1;
      	while(n<a.size()+b.size()-1)n<<=1;
      	a.resize(n);b.resize(n);
      	ll x=qpow(3,(mod-1)/n);
      	for(int i=0;i<n;i++)r[i]=r[i/2]/2+(i&1)*n/2;
      	NTT(a,n,x);NTT(b,n,x);
      	for(int i=0;i<n;i++)a[i]=a[i]*b[i]%mod;
      	NTT(a,n,qpow(x,mod-2));
      	ll inv=qpow(n,mod-2);
      	for(int i=0;i<n;i++)a[i]=a[i]*inv%mod;
      	return a; 
      }
      vector<ll> f(vector<ll> &a,int l,int r){
      	if(l==r)return {1,a[l]};
      	int mid=(l+r)>>1;
      	return solve(f(a,l,mid),f(a,mid+1,r));
      }
      int main(){
      	ios::sync_with_stdio(0);
      	cin.tie(0);
      	int n;
      	cin>>n;
      	vector<ll> a(n);
      	ll cnt=0;
      	for(int i=0;i<n;i++){
      		cin>>a[i];
      		cnt+=a[i];
      	}
      	vector<ll> b=f(a,0,n-1);
      	ll ans=1,s1=1,s2=1;
      	for(int i=1;i<=n&&i<b.size();i++){
      		s1=s1*i%mod;
      		s2=s2*(cnt-i+1)%mod;
      		ans=(ans+b[i]*s1%mod*qpow(s2,mod-2)%mod)%mod;
      	}
      	cout<<ans;
      	return 0;
      }
      
      • 0
        @ 2026-2-25 14:47:39

        改了一下码风

        #include<bits/stdc++.h>
        using namespace std;
        typedef long long ll;
        const ll M=8e5+10,P=998244353;
        inline ll qpow(ll a,ll b){
        	ll res=1;
        	for(;b;b>>=1,a=a*a%P){
        		if(b&1){
        			res=res*a%P;
        		}
        	}
        	return res;
        }
        inline void NTT(ll a[],ll n,ll x){
        	if(n==1)return;
        	ll a1[n/2],a2[n/2];
        	for(ll i=0;i<n/2;i++){
        		a1[i]=a[2*i];
        		a2[i]=a[2*i+1];
        	}
        	NTT(a1,n/2,x*x%P);
        	NTT(a2,n/2,x*x%P);
        	ll xi=1;
        	for(ll i=0;i<n/2;i++){
        		a[i]=(a1[i]+a2[i]*xi)%P;
        		a[i+n/2]=((a1[i]-a2[i]*xi)%P+P)%P;
        		xi=xi*x%P;
        	}
        }
        ll merge(ll A[],ll l1,ll B[],ll l2,ll C[]){
            ll N=1ll<<(ll)log2(l1+l2-1)+1;
            fill(A+l1,A+N,0);
        	fill(B+l2,B+N,0);
            ll inv_N=qpow(N,P-2),x=qpow(3,(P-1)/N),inv_x=qpow(x,P-2);
            NTT(A,N,x);
        	NTT(B,N,x);
            for(ll i=0;i<N;i++)
        		C[i]=1ll*A[i]*B[i]%P;
            NTT(C,N,inv_x);
            for(ll i=0;i<l1+l2-1;i++)C[i]=1ll*C[i]*inv_N%P;
            return l1+l2-1;
        }
        ll n,d[M];
        ll divide(ll C[],ll l,ll r){
            if(l==r){
        		C[0]=1;
        		C[1]=d[l];
        		return 2;
        	}
            ll tmp1[M],tmp2[M],mid=l+r>>1;
            return merge(tmp1,divide(tmp1,l,mid),tmp2,divide(tmp2,mid+1,r),C);
        }
        ll a[M];
        int main(){
            scanf("%lld",&n);
            for(ll i=1;i<=n;i++)scanf("%lld",&d[i]);
            ll S=0;
        	for(ll i=1;i<=n;i++)S+=d[i];
        	divide(a,1,n);
            ll ans=0,sum=1;
            for(ll i=0;i<=n;i++){
                ans=(ans+1ll*sum*a[i])%P;
                sum=1ll*sum*(i+1)%P*qpow(S-i,P-2)%P;
            }
            printf("%lld",ans);
            return 0;
        }
        
        • 1

        信息

        ID
        1612
        时间
        3000ms
        内存
        1024MiB
        难度
        8
        标签
        递交数
        79
        已通过
        11
        上传者