1 条题解
-
0
#include<bits/stdc++.h> #define fr(x) freopen(#x".in","r",stdin);freopen(#x".out","w",stdout); using namespace std; const int mod=998244353,N=8e5+5; int n,m,a[N],b[N],c[N],I[N],w[N],mmax,ans; inline int rd() { int x=0,zf=1; char ch=getchar(); while(ch<'0'||ch>'9') (ch=='-')and(zf=-1),ch=getchar(); while(ch>='0'&&ch<='9') x=x*10+ch-'0',ch=getchar(); return x*zf; } inline void wr(int x) { if(x==0) return putchar('0'),putchar(' '),void(); int num[35],len=0; while(x) num[++len]=x%10,x/=10; for(int i=len;i>=1;i--) putchar(num[i]+'0'); putchar(' '); } inline int bger(int x){return x|=x>>1,x|=x>>2,x|=x>>4,x|=x>>8,x|=x>>16,x+1;} inline int md(int x){return x>=mod?x-mod:x;} inline int ksm(int x,int p){int s=1;for(;p;(p&1)&&(s=1ll*s*x%mod),x=1ll*x*x%mod,p>>=1);return s;} inline void dao(int *a,int n){for(int i=1;i<n;i++) a[i-1]=1ll*i*a[i]%mod;a[n-1]=0;} inline void ji(int *a,int n){for(int i=n-1;i>=1;i--) a[i]=1ll*ksm(i,mod-2)*a[i-1]%mod;a[0]=0;} inline void init(int mmax) { for(int i=1,j,k;i<mmax;i<<=1) for(w[j=i]=1,k=ksm(3,(mod-1)/(i<<1)),j++;j<(i<<1);j++) w[j]=1ll*w[j-1]*k%mod; } inline void DNT(int *a,int mmax) { for(int i,j,k=mmax>>1,L,*W,*x,*y,z;k;k>>=1) for(L=k<<1,i=0;i<mmax;i+=L) for(j=0,W=w+k,x=a+i,y=x+k;j<k;j++,W++,x++,y++) *y=1ll*(*x+mod-(z=*y))* *W%mod,*x=md(*x+z); } inline void IDNT(int *a,int mmax) { for(int i,j,k=1,L,*W,*x,*y,z;k<mmax;k<<=1) for(L=k<<1,i=0;i<mmax;i+=L) for(j=0,W=w+k,x=a+i,y=x+k;j<k;j++,W++,x++,y++) z=1ll* *W* *y%mod,*y=md(*x+mod-z),*x=md(*x+z); reverse(a+1,a+mmax); for(int inv=ksm(mmax,mod-2),i=0;i<mmax;i++) a[i]=1ll*a[i]*inv%mod; } inline void NTT(int *a,int *b,int n,int m) { mmax=bger(n+m);init(mmax); DNT(a,mmax);DNT(b,mmax); for(int i=0;i<mmax;i++) a[i]=1ll*a[i]*b[i]%mod; IDNT(a,mmax); } void INV(int num,int *a,int *b) { if(num==1) return b[0]=ksm(a[0],mod-2),void(); INV((num+1)>>1,a,b); int mmax=bger(num<<1);init(mmax); static int c[N]; for(int i=0;i<num;i++) c[i]=a[i];for(int i=num;i<mmax;i++) c[i]=0; DNT(c,mmax);DNT(b,mmax); for(int i=0;i<mmax;i++) b[i]=1ll*(2-1ll*c[i]*b[i]%mod+mod)%mod*b[i]%mod; IDNT(b,mmax); for(int i=num;i<mmax;i++) b[i]=0; } inline void Ln(int *a,int n){static int b[N];for(int i=0;i<bger(n<<1);i++) b[i]=0;INV(n,a,b);dao(a,n);NTT(a,b,n,n);ji(a,n);for(int i=n;i<bger(n<<1);i++) a[i]=0;} inline void Exp(int *a,int *b,int n) { if(n==1) return b[0]=1,void(); Exp(a,b,(n+1)>>1);static int c[N];for(int i=0;i<bger(n<<1);i++) c[i]=0; for(int i=0;i<n;i++) c[i]=b[i];Ln(c,n); for(int i=0;i<n;i++) c[i]=md(mod-c[i]+a[i]);c[0]=md(c[0]+1); NTT(b,c,n,n);for(int i=n;i<bger(n<<1);i++) b[i]=0; } int main() { n=rd(),m=rd();for(int i=1;i<=m;i++) c[rd()]++; I[1]=1;for(int i=2;i<=n;i++) I[i]=mod-1ll*I[mod%i]*(mod/i)%mod; for(int i=1;i<=n;i++) for(int j=1;j<=n/i;j++) a[i*j]=(a[i*j]+1ll*c[i]*I[j])%mod; for(int i=1;i<=n;i++) a[i]=md(mod-a[i]);Exp(a,b,n+1); for(int i=1;i<=n;i++) ans=(ans+1ll*b[i]*I[i])%mod; return wr(1ll*(mod-n)*ans%mod),0; }
- 1
信息
- ID
- 8296
- 时间
- 2000ms
- 内存
- 1024MiB
- 难度
- 9
- 标签
- 递交数
- 11
- 已通过
- 3
- 上传者