AT2180
我想用函数求逆过
但是卡不过
求帮忙卡常
#include<bits/stdc++.h>
#define LL long long
#define fr(x) freopen(#x".in","r",stdin);freopen(#x".out","w",stdout);
using namespace std;
const LL mod=998244353,_g=332748118;//逆元
const int g=3,N=5e5+5;
LL fan[N],n,m,a[N],b[N],c[N];
inline void swap(LL &x,LL &y){x^=y^=x^=y;}
inline LL rd()
{
LL 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(LL x)
{
if(x==0) return putchar('0'),putchar(' '),void();
short 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 LL ksm(LL x,LL p)
{
LL s=1;
for(;p;p>>=1) (p&1)and(s=s*x%mod),x=x*x%mod;
return s;
}
inline void ntt(LL *c,const LL zf,const LL mmax)
{
for(LL i=0;i<mmax;i++) if(i<fan[i]) swap(c[i],c[fan[i]]);
for(LL k=(mmax>>1);k>=1;k>>=1)
{
LL nu=mmax/k/2,mul=ksm(zf==1?g:_g,(mod-1)/(nu<<1));
for(LL l=0;l<k;l++)
for(LL i=0,j=1;i<nu;i++,j=j*mul%mod){LL xx=c[l*nu*2+i],yy=j*c[l*nu*2+i+nu]%mod;c[l*nu*2+i]=(xx+yy)%mod,c[l*nu*2+i+nu]=(xx-yy+mod)%mod;}
}
if(zf==1) return;
LL inv=ksm(mmax,mod-2);
for(LL i=0;i<mmax;i++) c[i]=c[i]*inv%mod;
}
inline void Ntt(LL *a,LL *b,LL *c,LL mmax)
{
static LL t1[N],t2[N];
for(int i=0;i<=mmax;i++) t1[i]=a[i],t2[i]=b[i];
ntt(t1,1,mmax);ntt(t2,1,mmax);
for(int i=0;i<=mmax;i++) c[i]=t1[i]*t2[i]%mod;
ntt(c,-1,mmax);
}
inline LL dfs(LL p,LL *a,LL *b,LL mmax)
{
static LL c[N];LL len1=0,len2=n+2;
for(;p;p>>=1)
{
LL i;
memcpy(c,b,sizeof(c));
for(i=1;i<=len2;i+=2) c[i]=mod-c[i];
Ntt(a,c,a,mmax);Ntt(b,c,b,mmax);len1+=len2;len2+=len2;
for(i=p&1;i<=len1;i+=2) a[i/2]=a[i];
for(i/=2;i<=len1;i++) a[i]=0;len1>>=1;len1=min(len1,len2>>1);
for(i=0;i<=len2;i+=2) b[i/2]=b[i];
for(i/=2;i<=len2;i++) b[i]=0;len2>>=1;
}
return a[0]*ksm(b[0],mod-2)%mod;
}
inline void ksm1(LL *a,LL p,LL *b,LL mmax)
{
b[0]=1;
for(;p;p>>=1)
{
if(p&1) Ntt(a,b,b,mmax);
Ntt(a,a,a,mmax);
}
}
signed main()
{
n=rd();m=rd();n--;m--;
LL mmax=1;for(;mmax<=(n<<1)+10;mmax<<=1);
for(LL i=0;i<mmax;i++) fan[i]=(fan[i>>1]>>1)|((i&1)?(mmax>>1):0);
for(LL i=0,j=1;i<=n;i++)
{
LL zf=((n-i)&1)?-1:1;
b[n-i]=(zf*j+mod)%mod;j=j*(n-i)%mod*ksm(i+1,mod-2)%mod;
}
memcpy(c,b,sizeof(c));
for(LL i=0;i<=n;i++)
{
b[i+1]=(b[i+1]-c[i]+mod)%mod;
b[i+2]=(b[i+2]-c[i]+mod)%mod;
}
memset(a,0,sizeof(a));a[0]=1;
printf("%lld\n",dfs(m,a,b,mmax));
return 0;
}