1 条题解

  • 0
    @ 2026-5-17 14:24:58
    #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
    上传者