1 条题解

  • 0
    @ 2026-7-4 22:55:27

    #include <cstdio>
    #include <iostream>
    using namespace std;
    const int M = 100005;
    #define int long long
    const int MOD = 998244353;
    int read()
    {
    	int x=0,f=1;char c;
    	while((c=getchar())<'0' || c>'9') {if(c=='-') f=-1;}
    	while(c>='0' && c<='9') {x=(x<<3)+(x<<1)+(c^48);c=getchar();}
    	return x*f;
    }
    int n,m,w,fac[M],inv[M],f[2][M],rev[M],A[M],B[M],C[M];
    void init(int n)
    {
    	fac[0]=inv[0]=inv[1]=1;
    	for(int i=1;i<=n;i++) fac[i]=fac[i-1]*i%MOD;
    	for(int i=2;i<=n;i++) inv[i]=inv[MOD%i]*(MOD-MOD/i)%MOD;
    	for(int i=2;i<=n;i++) inv[i]=inv[i-1]*inv[i]%MOD;
    }
    int comb(int n,int m)
    {
    	if(n<m || m<0) return 0;
    	return fac[n]*inv[m]%MOD*inv[n-m]%MOD;
    }
    int qkpow(int a,int b)
    {
    	int r=1;
    	while(b>0)
    	{
    		if(b&1) r=r*a%MOD;
    		a=a*a%MOD;
    		b>>=1;
    	}
    	return r;
    }
    void NTT(int *a,int len,int op)
    {
    	for(int i=0;i<len;i++)
    	{
    		rev[i]=(rev[i>>1]>>1)|((len/2)*(i&1));
    		if(i<rev[i]) swap(a[i],a[rev[i]]);
    	}
    	for(int s=2;s<=len;s<<=1)
    	{
    		int t=s/2,w=(op==1)?qkpow(3,(MOD-1)/s):
    		qkpow(3,MOD-1-(MOD-1)/s);
    		for(int i=0;i<len;i+=s)
    			for(int j=0,x=1;j<t;j++,x=x*w%MOD)
    			{
    				int fe=a[i+j],fo=a[i+j+t];
    				a[i+j]=(fe+x*fo)%MOD;
    				a[i+j+t]=(fe-x*fo%MOD+MOD)%MOD;
    			}
    	}
    	if(op==1) return ;
    	int inv=qkpow(len,MOD-2);
    	for(int i=0;i<len;i++) a[i]=a[i]*inv%MOD;
    }
    signed main()
    {
    	n=read();m=read();init(1e5);
    	f[0][0]=f[1][0]=1;
    	while(m--)
    	{
    		int len=1;while(len<2*n) len<<=1;
    		for(int i=0;i<len;i++) A[i]=B[i]=C[i]=0;
    		w^=1;
    		for(int i=0;i<n;i++)
    			A[i]=f[w^1][i]*inv[i]%MOD,B[i]=inv[i+3];
    		NTT(A,len,1);NTT(B,len,1);
    		for(int i=0;i<len;i++) C[i]=A[i]*B[i]%MOD;
    		NTT(C,len,-1);
    		for(int i=1;i<=n;i++)
    			f[w][i]=((i*i+i+2)/2*f[w^1][i]+fac[i+2]*C[i-1])%MOD;
    	}
    	int ans=0;
    	for(int i=0;i<=n;i++)
    		ans=(ans+comb(n,i)*f[w][i])%MOD;
    	printf("%lld\n",ans);
    }
    
    
    • 1