3 条题解
-
3
#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
#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
改了一下码风
#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
- 上传者