1.核心:
FFT:
正常版本:
#include<bits/stdc++.h>
#define maxn 400005
using namespace std;
const double PI = acos(-1);
struct cplx
{
double r,i;
cplx(double r=0,double i=0):r(r),i(i){}
cplx operator +(const cplx &B)const{ return cplx(r+B.r,i+B.i); }
cplx operator -(const cplx &B)const{ return cplx(r-B.r,i-B.i); }
cplx operator *(const cplx &B)const{ return cplx(r*B.r-i*B.i,i*B.r+r*B.i); }
cplx conj(){ return cplx(r,-i); }
}a[maxn],b[maxn];
int r[maxn]={};
cplx w[maxn] = {1};
inline void FFT(cplx *A,int lgn,int tp)
{
int n = 1<<lgn;
for(int i=1;i<n;i++) r[i] = (r[i>>1]>>1) | ((i&1)<<(lgn-1));
for(int i=1;i<n;i++) if(i < r[i]) swap(A[i] , A[r[i]]);
for(int len=2;len<=n;len<<=1){
int l = len >> 1;cplx wn(cos(PI / l) , sin(PI / l) * tp);
for(int i=1;i<l;i++) w[i] = w[i-1] * wn;
for(int st = 0;st < n;st += len) for(int k=0;k<l;k++)
{
cplx tmp = w[k] * A[st + k + l];
A[st + k + l] = A[st + k] - tmp , A[st + k] = A[st + k] + tmp;
}
}
if(tp==-1) for(int i=0;i<n;i++) A[i].r /= n;
}
int main()
{
int n,m;
scanf("%d%d",&n,&m);
for(int i=0;i<=n;i++) scanf("%lf",&a[i].r);
for(int i=0;i<=m;i++) scanf("%lf",&a[i].i);
n++,m++;
int len = 0;
for(;n+m>(1<<len);len++);
FFT(a,len,1);
for(int i=0,ci,Len = 1<<len;i<Len;i++)
{
ci = (Len - i) & (Len - 1);
cplx A = (a[i] + a[ci].conj())*cplx(0.5,0) , B = (a[i] - a[ci].conj())*cplx(0,-0.5);
b[i] = A * B;
}
FFT(b,len,-1);
for(int i=0;i<n+m-1;i++) printf("%d ",int(b[i].r+0.5));
}
预处理单位元(精度高):
#include<bits/stdc++.h>
#define maxn 300005
using namespace std;
const double Pi = 3.1415926535897932384626433832795;
struct cplx
{
double r,i;
cplx(double r=0,double i=0):r(r),i(i){}
cplx operator +(const cplx &B)const{ return cplx(r+B.r,i+B.i); }
cplx operator -(const cplx &B)const{ return cplx(r-B.r,i-B.i); }
cplx operator *(const cplx &B)const{ return cplx(r*B.r-i*B.i,i*B.r+r*B.i); }
cplx conj()const{ return cplx(r,-i); }
}w[maxn],A[maxn],B[maxn];
int r[maxn];
inline void FFT(cplx A[maxn],int lgn,int tp)
{
int n = 1<<lgn;
for(int i=0;i<n;i++) w[i]=cplx(cos(i*Pi/n),sin(i*Pi/n));
for(int i=0;i<n;i++) r[i] = (r[i>>1]>>1)|((i&1)<<(lgn-1));
for(int i=0;i<n;i++) if(i < r[i]) swap(A[i] , A[r[i]]);
for(int L=2;L<=n;L<<=1)
for(int st=0,l=L>>1;st<n;st+=L)
for(int k=0,lc=0,inc=n/l;k<l;k++,lc+=inc)
{
cplx tmp = (tp==1 ? w[lc] : w[lc].conj()) * A[st+k+l];
A[st+k+l]=A[st+k]-tmp,A[st+k]=A[st+k]+tmp;
}
if(tp==-1) for(int i=0;i<n;i++) A[i].r/=n,A[i].i/=n;
}
int main()
{
int n,m;
scanf("%d%d",&n,&m);
for(int i=0;i<=n;i++) scanf("%lf",&A[i].r);
for(int i=0;i<=m;i++) scanf("%lf",&A[i].i);
int lgn=0;for(;n+m>=(1<<lgn);lgn++);
FFT(A,lgn,1);
for(int i=0,len=1<<lgn;i<len;i++)
{
cplx u=A[i],v=A[(len-1)&(len-i)].conj();
B[i]=(u+v)*(u-v)*cplx(0,-0.25);
}
FFT(B,lgn,-1);
for(int i=0;i<n+m;i++) printf("%d ",int(round(B[i].r)));
printf("%d\n",(int)round(B[n+m].r));
}
U P D : 1.0 K B F F T \mathrm {UPD :1.0 KB\ FFT}UPD:1.0KB FFT
#include<bits/stdc++.h>
#define maxn 300005
#define cp complex<double>
#define Pi 3.1415926535897932384626433832795
#define rep(i,j,k) for(int i=(j);i<=(k);i++)
#define per(i,j,k) for(int i=(j);i>=(k);i--)
using namespace std;
int Wl,lg[maxn],r[maxn];
cp W[maxn];
void init(int n){
for(Wl=1;n>=2*Wl;Wl<<=1);
rep(i,0,Wl<<1) W[i]=exp(cp(0,i*Pi/Wl)),(i>1)&&(lg[i]=lg[i>>1]+1);
}
void FFT(cp *A,int n,int tp){
rep(i,1,n-1) (i<(r[i]=(r[i>>1]>>1)|((i&1)<<(lg[n]-1))))&&(swap(A[i],A[r[i]]),0);cp t;
for(int L=1,B=Wl;L<n;L<<=1,B>>=1) for(int s=0;s<n;s+=L<<1) for(int k=s,x=0;k<s+L;k++,x+=B)
t=(tp==1?W[x]:conj(W[x]))*A[k+L],A[k+L]=A[k]-t,A[k]+=t;
if(tp^1) rep(i,0,n-1) A[i]/=n;
}
int n,m;cp A[maxn],B[maxn];
int main(){
scanf("%d%d",&n,&m);
double x;
rep(i,0,n) scanf("%lf",&x),A[i].real(x);
rep(i,0,m) scanf("%lf",&x),B[i].real(x);
init(n+m);
FFT(A,Wl<<1,1),FFT(B,Wl<<1,1);
rep(i,0,(Wl<<1)-1) A[i]*=B[i];
FFT(A,Wl<<1,-1);
rep(i,0,n+m) printf("%d%c",(int)round(A[i].real())," \n"[i==n+m]);
}
MTT
合并DFT详见myy论文。
合并IDFT其实不需要任何技巧因为:
I D F T ( D F T ( A ( i ) ) + i D F T ( B ( i ) ) ) = A ( i ) + i B ( i ) IDFT(DFT(A(i)) + iDFT(B(i))) = A(i) + iB(i)IDFT(DFT(A(i))+iDFT(B(i)))=A(i)+iB(i)
如果觉得慢的话可以将l o n g d o u b l e \mathrm {long\ double}long double改为d o u b l e \mathrm {double}double
然后预处理单位根照样可以满足1 0 5 10^5105的精度要求
#include<bits/stdc++.h>
#define maxn 300005
#define LL long long
#define M ((1<<15)-1)
#define ld long double
using namespace std;
char cb[1<<15],*cs=cb,*ct=cb;
#define getc() (cs==ct&&(ct=(cs=cb)+fread(cb,1,1<<15,stdin),cs==ct)?0:*cs++)
inline void read(int &res){ char ch;for(;!isdigit(ch=getc()););for(res=ch-'0';isdigit(ch=getc());res=res*10+ch-'0'); }
int p;
const ld Pi = 3.1415926535897932384626433832795;
struct cplx
{
ld r,i;
cplx(ld r=0,ld i=0):r(r),i(i){}
cplx operator +(const cplx &B)const{ return cplx(r+B.r,i+B.i); }
cplx operator -(const cplx &B)const{ return cplx(r-B.r,i-B.i); }
cplx operator *(const cplx &B)const{ return cplx(r*B.r-i*B.i,i*B.r+B.i*r); }
cplx conj(){ return cplx(r,-i); }
}w[maxn]={1};
int a[maxn],b[maxn],c[maxn],r[maxn];
inline void FFT(cplx A[maxn],int lgn,int tp)
{
int n = 1<<lgn;
for(int i=1;i<n;i++) r[i] = (r[i>>1]>>1)|((i&1)<<(lgn-1));
for(int i=1;i<n;i++) if(i<r[i])swap(A[i],A[r[i]]);
for(int L=2;L<=n;L<<=1)
{ int l=L>>1;w[1]=cplx(cos(Pi/l),sin(Pi/l)*tp);
for(int i=2;i<l;i++) w[i] = w[i-1] * w[1];
for(int st=0;st<n;st+=L)
for(int k=0;k<l;k++)
{
cplx tmp = w[k] * A[st+k+l];
A[st+k+l] = A[st+k]-tmp , A[st+k] = A[st+k] + tmp;
}
}
if(tp == -1) for(int i=0;i<n;i++) A[i].r/=n,A[i].i/=n;
}
cplx s[4][maxn];
inline void mul(int a[maxn],int b[maxn],int lgn,int c[maxn])
{
int n = 1<<lgn;
for(int i=0;i<n;i++) s[0][i] = cplx(a[i]>>15,b[i]>>15) , s[1][i] = cplx(a[i]&M,b[i]&M);
FFT(s[0],lgn,1),FFT(s[1],lgn,1);
for(int i=0;i<n;i++)
{
cplx a[4] = {s[0][i] , s[0][(n-i)&(n-1)].conj() , s[1][i] , s[1][(n-1)&(n-i)].conj()};
cplx b[4] = {(a[0]+a[1])*cplx(0.5,0),(a[0]-a[1])*cplx(0,-0.5),
(a[2]+a[3])*cplx(0.5,0),(a[2]-a[3])*cplx(0,-0.5)};
s[2][i] = b[0]*b[1]+cplx(0,1)*(b[2]*b[3]) , s[3][i] = b[0]*b[3]+cplx(0,1)*b[1]*b[2];//IDFT(DFT(A(x))+DFT(iB(x))) = A(x) + iB(x)
}
FFT(s[2],lgn,-1),FFT(s[3],lgn,-1);
for(int i=0;i<n;i++)
{
LL a[4] = {llround(s[2][i].r)%p,llround(s[2][i].i)%p,llround(s[3][i].r)%p,llround(s[3][i].i)%p};
c[i] = (a[1] + (((a[2]+a[3])%p)<<15) + (a[0]<<30)) % p;
}
}
int main()
{
int n,m;p=1000000007;
read(n),read(m);
for(int i=0;i<=n;i++) read(a[i]);
for(int i=0;i<=m;i++) read(b[i]);
int lgn = 0;
for(;n+m>=(1<<lgn);lgn++);
mul(a,b,lgn,c);
for(int i=0;i<n+m;i++) printf("%d ",(c[i]+p)%p);
printf("%d\n",(c[n+m]+p)%p);
}
2.多项式求逆:
MTT版本:
#include<bits/stdc++.h>
#define maxn 300005
#define mod 1000000007
#define LL long long
#define M ((1<<15)-1)
#define ld long double
using namespace std;
const ld Pi = 3.1415926535897932384626433832795;
int wlen=0;
struct cplx
{
ld r,i;
cplx(ld r=0,ld i=0):r(r),i(i){}
cplx operator +(const cplx &B)const{return cplx(r+B.r,i+B.i);}
cplx operator -(const cplx &B)const{return cplx(r-B.r,i-B.i);}
cplx operator *(const cplx &B)const{return cplx(r*B.r-i*B.i,i*B.r+r*B.i);}
cplx conj(){return cplx(r,-i);}
}w[maxn];
inline int Pow(int base,int k)
{
int ret = 1;
for(;k;k>>=1,base=1ll*base*base%mod) if(k&1) ret=1ll*ret*base%mod;
return ret;
}
int r[maxn];
inline void FFT(cplx A[maxn],int lgn,int tp)
{
int n = 1<<lgn;
for(int i=0;i<n;i++) r[i] = (r[i>>1]>>1)|((i&1)<<(lgn-1));
for(int i=0;i<n;i++) if(i < r[i]) swap(A[i] , A[r[i]]);
for(int L=2;L<=n;L<<=1)
for(int st=0,l=L>>1,inc=wlen/l;st<n;st+=L)
for(int k=0,lc=0;k<l;k++,lc+=inc)
{
cplx tmp = (tp==1 ? w[lc] : w[lc].conj()) * A[st+k+l];
A[st+k+l]=A[st+k]-tmp,A[st+k]=A[st+k]+tmp;
}
if(tp==-1) for(int i=0;i<n;i++) A[i].r/=n,A[i].i/=n;
}
cplx s[4][maxn];
inline void mul(int a[maxn],int b[maxn],int lgn,int c[maxn])
{
int n = 1<<lgn;
for(int i=0;i<n;i++) s[0][i]=cplx(a[i]>>15,b[i]>>15),s[1][i]=cplx(a[i]&M,b[i]&M);
FFT(s[0],lgn,1),FFT(s[1],lgn,1);
for(int i=0;i<n;i++)
{
cplx a[4]={s[0][i],s[0][(n-i)&(n-1)].conj(),s[1][i],s[1][(n-1)&(n-i)].conj()};
cplx b[4]={(a[0]+a[1])*cplx(0.5,0),(a[0]-a[1])*cplx(0,-0.5),
(a[2]+a[3])*cplx(0.5,0),(a[2]-a[3])*cplx(0,-0.5)};
s[2][i] = b[0]*b[1]+cplx(0,1)*b[2]*b[3],s[3][i]=b[0]*b[3]+cplx(0,1)*b[1]*b[2];
}
FFT(s[2],lgn,-1),FFT(s[3],lgn,-1);
for(int i=0;i<n;i++)
{
LL a[4]={llround(s[2][i].r)%mod,llround(s[2][i].i)%mod,llround(s[3][i].r)%mod,llround(s[3][i].i)%mod};
c[i] = ((a[0]<<30)+(a[1])+((a[2]+a[3])<<15)) % mod;
}
}
void Inv(int a[maxn],int lgn,int b[maxn])
{
if(lgn==0){ b[0]=Pow(a[0],mod-2);return; }
Inv(a,lgn-1,b);
int n = (1<<(lgn+1));
static int tmp[3][maxn];
for(int i=0;i<n;i++)
{
if(i<(n>>1)) tmp[0][i] = a[i]; else tmp[0][i] = 0;
if(i<(n>>2)) tmp[1][i] = b[i]; else tmp[1][i] = 0;
}
mul(tmp[0],tmp[1],lgn+1,tmp[2]);
for(int i=0;i<n;i++) tmp[2][i] = (i==0) * 2 - tmp[2][i];
mul(tmp[2],tmp[1],lgn+1,b);
for(int i=(n>>1);i<n;i++) b[i] = 0;
}
int n,a[maxn],b[maxn];
int main()
{
int n;
scanf("%d",&n);
for(int i=0;i<n;i++) scanf("%d",&a[i]);
int lgn = 0;
for(;n>(1<<lgn);lgn++);
wlen = 1<<(lgn+1);
for(int i=0;i<wlen;i++) w[i] = cplx(cos(i*Pi/wlen),sin(i*Pi/wlen));
Inv(a,lgn,b);
printf("%d",(b[0]+mod)%mod);
for(int i=1;i<n;i++) printf(" %d",(b[i]+mod)%mod);
}
多项式除法:
求A ( x ) ≡ B ( x ) ∗ C ( x ) + R ( x ) ( m o d x n ) A(x)\equiv B(x)*C(x)+R(x)\pmod{x^n}A(x)≡B(x)∗C(x)+R(x)(modxn)
其中A ( x ) , B ( x ) A(x),B(x)A(x),B(x)已知。∣ R ( x ) ∣ < ∣ B ( x ) ∣ |R(x)|<|B(x)|∣R(x)∣<∣B(x)∣
不知道怎么理解。
多项式除法
还是手推一下式子比较好:
A ( x ) ≡ B ( x ) C ( x ) + R ( x ) A ( 1 x ) ≡ B ( 1 x ) C ( 1 x ) + R ( 1 x ) x n A ( 1 x ) ≡ x ∣ B ( x ) ∣ B ( x ) ∗ x ∣ C ( x ) ∣ C ( x ) + x ∣ B ( x ) ∣ − 1 R ( x ) ∗ x n − ∣ B ( x ) ∣ + 1 将 R ( x ) 看 做 ∣ B ( x ) ∣ − 1 次 多 项 式 , A R ( x ) = x ∣ A ( x ) ∣ A ( 1 x ) A R ( x ) ≡ B R ( x ) C R ( x ) + x n − ∣ B ( x ) ∣ + 1 R R ( x ) 在 ( m o d x n − ∣ B ( x ) ∣ + 1 ) 意 义 下 做 多 项 式 逆 元 \begin{aligned} &A(x)\equiv B(x)C(x)+R(x)\\ &A(\frac 1x)\equiv B(\frac 1x)C(\frac 1x)+R(\frac 1x)\\ &x^nA(\frac 1x) \equiv x^{|B(x)|}B(x) * x^{|C(x)|}C(x)+x^{|B(x)|-1}R(x)*x^{n-|B(x)|+1}\\ &将R(x)看做|B(x)|-1次多项式,A^R(x) = x^{|A(x)|}A(\frac 1x)\\ &A^R(x)\equiv B^R(x)C^R(x)+x^{n-|B(x)|+1}R^R(x) 在\pmod{x^{n-|B(x)|+1}}意义下做多项式逆元 \end{aligned}A(x)≡B(x)C(x)+R(x)A(x1)≡B(x1)C(x1)+R(x1)xnA(x1)≡x∣B(x)∣B(x)∗x∣C(x)∣C(x)+x∣B(x)∣−1R(x)∗xn−∣B(x)∣+1将R(x)看做∣B(x)∣−1次多项式,AR(x)=x∣A(x)∣A(x1)AR(x)≡BR(x)CR(x)+xn−∣B(x)∣+1RR(x)在(modxn−∣B(x)∣+1)意义下做多项式逆元
Code:
l u o g u 4512 luogu 4512luogu4512
注意NTT最大范围在2 n 2n2n
10行左右??
注意多项式求逆带n-m+1,可以快很多。。。。。。
#include<bits/stdc++.h>
#define maxn 600005
#define mod 998244353 // 119 * 2^23 + 1
using namespace std;
int n,m,lg[maxn];
int w[maxn<<1]={1},wlen,r[maxn];
int Pow(int base,int k)
{ int ret=1;
for(;k;k>>=1,base=1ll*base*base%mod) if(k&1) ret=1ll*ret*base%mod;
return ret;
}
void Init()
{
w[2*wlen]=w[0]=1;
for(int i=1,inc=(mod-1)/(2*wlen),pw=Pow(3,inc);i<2*wlen;i++)
w[i] = 1ll * w[i-1] * pw % mod;
}
inline void NTT(int A[maxn],int lgn,int tp)
{
int n = 1<<lgn;
for(int i=1;i<n;i++) r[i]=(r[i>>1]>>1)|((i&1)<<(lgn-1));
for(int i=1;i<n;i++) if(i<r[i]) swap(A[i],A[r[i]]);
for(int L=2;L<=n;L<<=1)
for(int st=0,l=L>>1,inc=wlen/l;st<n;st+=L)
for(int k=0,lc=0;k<l;k++,lc+=inc)
{
int tmp = 1ll * (tp==1?w[lc]:w[2*wlen-lc]) * A[st+k+l] % mod;
A[st+k+l]=(A[st+k]-tmp)%mod,A[st+k]=(A[st+k]+tmp)%mod;
}
if(tp==-1)for(int i=0,inv=Pow(n,mod-2);i<n;i++) A[i]=1ll*A[i]*inv%mod;
}
void Inv(int a[maxn],int lgn,int b[maxn])
{
static int tmp[maxn];
if(lgn == 0){b[0]=Pow(a[0],mod-2);return;}
Inv(a,lgn-1,b);
int n = 1<<(lgn+1);
for(int i=0;i<n;i++) if(i<(n>>1)) tmp[i] = a[i];else tmp[i] = 0;
NTT(tmp,lgn+1,1),NTT(b,lgn+1,1);
for(int i=0;i<n;i++) b[i] = 1ll * b[i] * (2ll - 1ll * b[i] * tmp[i] % mod) % mod;
NTT(b,lgn+1,-1);
for(int i=n>>1;i<n;i++) b[i] = 0;
}
int sta[2][maxn];
inline void mul(int a[maxn],int b[maxn],int lgn,int c[maxn],int n=0,int m=0)
{
int len = 1<<lgn;n?0:n=len>>1,m?0:m=len>>1;
for(int i=0;i<len;i++) sta[0][i]=i<n?a[i]:0,sta[1][i]=i<m?b[i]:0;
NTT(sta[0],lgn,1),NTT(sta[1],lgn,1);
for(int i=0;i<len;i++) c[i] = 1ll * sta[0][i] * sta[1][i] % mod;
NTT(c,lgn,-1);
}
int ra[maxn],rb[maxn],dc[maxn];
inline void Div(int a[maxn],int b[maxn],int n,int m,int c[maxn],int d[maxn])
{
int lgn = lg[2*n]+1 , len = 1<<lgn;
for(int i=0;i<len;i++){
if(i<=n) ra[i] = a[n-i]; else ra[i] = 0;
if(i<=m) rb[i] = b[m-i]; else rb[i] = 0;}
Inv(rb,lg[n-m+1]+1,c);
mul(ra,c,lgn,c);
for(int i=0;i<len;i++) if(i<n-m-i) swap(c[i],c[n-m-i]); else if(i>n-m) c[i] = 0;
for(int i=0;i<(m-i);i++) swap(rb[i],rb[m-i]);
for(int i=0;i<len;i++) dc[i] = c[i];
mul(dc,rb,lgn,d);
for(int i=0;i<len;i++) d[i] = (a[i] - d[i]) % mod;
}
int a[maxn],b[maxn],c[maxn],d[maxn];
int main()
{
scanf("%d%d",&n,&m);
for(int i=0;i<=n;i++) scanf("%d",&a[i]);
for(int i=0;i<=m;i++) scanf("%d",&b[i]);
for(int i=2;i<maxn;i++) lg[i] = lg[i>>1] + 1;
wlen = 1<<(lg[2*n]+1);
Init();
Div(a,b,n,m,c,d);
for(int i=0;i<n-m;i++) printf("%d ",(c[i]+mod)%mod);
printf("%d\n",(c[n-m]+mod)%mod);
for(int i=0;i<m-1;i++) printf("%d ",(d[i]+mod)%mod);
if(m)printf("%d\n",(d[m-1]+mod)%mod);
}
PS : 试了一发(小于某数时暴力,发现效果并不显著)
upd:个人钻的超短代码(58行)
#include<bits/stdc++.h>
#define maxn 2100005
#define LL long long
#define mod 998244353
using namespace std;
int lg[maxn],r[maxn],w[maxn]={1},wlen;
int Pow(int base,int k)
{ int ret=1;
for(;k;k>>=1,base=1ll*base*base%mod) if(k&1) ret=1ll*ret*base%mod;
return ret;}
inline void NTT(int A[maxn],int n,int tp)//n=2^k
{ int lgn = lg[n];
for(int i=1;i<n;i++) r[i] = (r[i>>1]>>1)|((i&1)<<(lgn-1));
for(int i=1;i<n;i++) if(i<r[i]) swap(A[i],A[r[i]]);
for(int L=2;L<=n;L<<=1)
for(int st=0,l=L>>1,inc=wlen/l;st<n;st+=L)
for(int k=st,x=0;k<st+l;k++,x+=inc)
{ int tmp = 1ll *(tp==1?w[x]:w[2*wlen-x])*A[k+l] % mod;
A[k+l]=(A[k]-tmp)%mod,A[k]=(A[k]+tmp)%mod;}
if(tp==-1) for(int i=0,inv=Pow(n,mod-2);i<n;i++) A[i]=1ll*A[i]*inv%mod;}
void mul(int A[maxn],int B[maxn],int C[maxn],int n,int m)
{ int lgn = lg[n+m]+1,len = 1<<lgn;
static int sta[maxn];
for(int i=0;i<len;i++) sta[i]=i<=n?A[i]:0,C[i]=i<=m?B[i]:0;
NTT(sta,len,1),NTT(C,len,1);
for(int i=0;i<len;i++) C[i] = 1ll * C[i] * sta[i] % mod;
NTT(C,len,-1);}
void Inv(int A[maxn],int B[maxn],int n)
{ B[0] = Pow(A[0],mod-2);
static int tmp[maxn];
for(int i=2;i<(n<<1);i<<=1)// mod x^i
{ for(int j=0;j<i;j++) tmp[j] = j<n?A[j]:0;
NTT(tmp,i<<1,1),NTT(B,i<<1,1);
for(int j=0;j<(i<<1);j++) B[j]=1ll*B[j]*(2-1ll*B[j]*tmp[j]%mod)%mod;
NTT(B,i<<1,-1); for(int j=min(n,i);j<(i<<1);j++) B[j] = 0;}}
void Div(int A[maxn],int B[maxn],int C[maxn],int R[maxn],int n,int m)
{ static int ra[maxn],rb[maxn];
for(int i=0,lim=max(n,m);i<=lim;i++) ra[i]=i<=n?A[n-i]:0,rb[i]=i<=m?B[m-i]:0;
Inv(rb,C,n-m+1),mul(ra,C,C,n,n-m);
for(int i=0;i<=2*n-m;i++) if(i<n-m-i)swap(C[i],C[n-m-i]);else if(i>n-m)C[i]=0;
mul(C,B,R,n-m,m);for(int i=0;i<=n;i++) R[i]=i<m?(A[i]-R[i])%mod:0;}
void Init(int n)
{ for(wlen=1;n>=(wlen<<1);wlen<<=1);
for(int i=1,pw=Pow(3,(mod-1)/(2*wlen));i<=2*wlen;i++) w[i]=1ll*w[i-1]*pw%mod;
for(int i=2;i<=2*wlen;i++) lg[i] = lg[i>>1] + 1;}
int a[maxn],b[maxn],c[maxn],R[maxn];
int main()
{
int n,m;
scanf("%d%d",&n,&m);
for(int i=0;i<=n;i++) scanf("%d",&a[i]);
for(int i=0;i<=m;i++) scanf("%d",&b[i]);
Init(n<<1);
Div(a,b,c,R,n,m);
for(int i=0;i<=n-m;i++) printf("%d%c",(c[i]+mod)%mod,i==n-m?'\n':' ');
for(int i=0;i<m;i++) printf("%d%c",(R[i]+mod)%mod,i==m-1?'\n':' ');
}
常系数线性递推
这篇看起来又线代又通俗
这篇告诉你如果避开高深的定理直接背结论
Code:
#include<bits/stdc++.h>
#define maxn 140005
#define LL long long
#define mod 998244353// 119 * 2 ^ 23 + 1
using namespace std;
int w[maxn]={1} , lg[maxn] , wlen = 1<<16 , f[maxn] , a[maxn] , rea[maxn] , ans[maxn] = {1} , base[maxn] = {0,1} , r[maxn];
int n,k;
int Pow(int base,int k)
{ int ret=1;
for(;k;k>>=1,base=1ll*base*base%mod) if(k&1) ret=1ll*ret*base%mod;
return ret;}
inline void NTT(int A[maxn],int lgn,int tp)
{ int n = 1<<lgn;
for(int i=1;i<n;i++) r[i] = (r[i>>1]>>1)|((i&1)<<(lgn-1));
for(int i=1;i<n;i++) if(i<r[i]) swap(A[i],A[r[i]]);
for(int L=2;L<=n;L<<=1)
for(int st=0,l=L>>1,inc=wlen/l;st<n;st+=L)
for(int k=0,lc=0,tmp;k<l;k++,lc+=inc)
tmp = 1ll * (tp==1?w[lc]:w[2*wlen-lc]) * A[st+k+l] % mod,
A[st+k+l] = (A[st+k] - tmp)%mod , A[st+k] = (A[st+k] + tmp)%mod;
if(tp==-1) for(int i=0,inv=Pow(n,mod-2);i<n;i++) A[i]=1ll*A[i]*inv%mod;}
void mul(int A[maxn],int B[maxn],int lgn,int C[maxn],int n=0,int m=0)
{
static int s[2][maxn];
int len = 1<<lgn;n?0:n=len>>1,m?0:m=len>>1;
for(int i=0;i<len;i++) s[0][i]=i<n?A[i]:0,s[1][i]=i<m?B[i]:0;
NTT(s[0],lgn,1),NTT(s[1],lgn,1);
for(int i=0;i<len;i++) C[i] = s[0][i] * 1ll * s[1][i] % mod;
NTT(C,lgn,-1);
}
void Inv(int A[maxn],int lgn,int B[maxn])
{ if(lgn==0){ B[0]=Pow(A[0],mod-2);return; }
Inv(A,lgn-1,B);
int n=1<<(lgn+1);
static int tmp[maxn];
for(int i=0;i<n;i++) tmp[i] = i<(n>>1)?A[i]:0;
NTT(tmp,lgn+1,1),NTT(B,lgn+1,1);
for(int i=0;i<n;i++) B[i] = 1ll * B[i] * (2 - 1ll * tmp[i] * B[i] % mod) % mod;
NTT(B,lgn+1,-1);
for(int i=n>>1;i<n;i++) B[i]=0;
}
void Mod(int A[maxn])
{
int n=2*k;for(;!A[n];n--);if(n<k) return;
int lgn = lg[2*n]+1, len = 1<<lgn;
static int ra[maxn],res[maxn];
for(int i=0;i<len;i++) ra[i]=i<=n?A[n-i]:0;
mul(ra,rea,lgn,res);
for(int i=0;i<len;i++) if(i<n-k-i)swap(res[i],res[n-k-i]);else if(i>n-k) res[i]=0;
mul(res,a,lgn,ra);
for(int i=0;i<len;i++) A[i]=(A[i]-ra[i])%mod;
}
int main()
{
scanf("%d%d",&n,&k);
for(int i=2;i<maxn;i++) lg[i] = lg[i>>1] + 1;
for(int i=1,inc=(mod-1)/(2*wlen),pw=Pow(3,inc);i<=2*wlen;i++) w[i] = 1ll * w[i-1] * pw % mod;
for(int i=k-1;i>=0;i--) scanf("%d",&a[i]),a[i]=-a[i];
a[k] = 1;
for(int i=1;i<=k;i++) scanf("%d",&f[i]);
for(int i=0;i<(k-i);i++) swap(a[i],a[k-i]);
int lgn = 0;for(;2*k>=(1<<lgn);lgn++);
Inv(a,lgn,rea);
for(int i=0;i<(k-i);i++) swap(a[i],a[k-i]);
Mod(base);
for(;n;n>>=1,mul(base,base,lgn,base),Mod(base))
if(n&1) mul(base,ans,lgn,ans),Mod(ans);
int ret = 0;
for(int i=1;i<=k;i++)
ret = (ret + 1ll * f[i] * ans[i-1]) % mod;
printf("%d",(ret+mod)%mod);
}
多项式求对数:
求B ( x ) = ln A ( x ) B(x)=\ln A(x)B(x)=lnA(x)
求导B ( x ) ′ = A ′ ( x ) A ( x ) B(x)' = \frac {A'(x)}{A(x)}B(x)′=A(x)A′(x)
B ( x ) = ∫ A ′ ( x ) A ( x ) B(x)=\int \frac {A'(x)}{A(x)}B(x)=∫A(x)A′(x)
Code:
#include<bits/stdc++.h>
#define maxn 300005
#define LL long long
#define mod 998244353
using namespace std;
int lg[maxn],r[maxn],w[maxn]={1},wlen;
int Pow(int base,int k)
{ int ret = 1;
for(;k;k>>=1,base=1ll*base*base%mod) if(k&1) ret=1ll*ret*base%mod;
return ret;}
void Init(int n)
{ for(wlen=1;n>=(2*wlen);wlen<<=1);
for(int i=1,pw=Pow(3,(mod-1)/(2*wlen));i<=2*wlen;i++) w[i] = 1ll * w[i-1] * pw % mod;
for(int i=2;i<=2*wlen;i++) lg[i] = lg[i>>1]+1; }
inline void NTT(int A[maxn],int n,int tp)
{ int lgn = lg[n];
for(int i=1;i<n;i++) r[i]=(r[i>>1]>>1)|((i&1)<<(lgn-1));
for(int i=1;i<n;i++) if(i<r[i]) swap(A[i],A[r[i]]);
for(int L=2;L<=n;L<<=1)
for(int st=0,l=L>>1,inc=wlen/l;st<n;st+=L)
for(int k=st,x=0;k<st+l;k++,x+=inc)
{ int tmp = 1ll * (tp==1?w[x]:w[2*wlen-x]) * A[k+l] % mod;
A[k+l]=(A[k]-tmp)%mod,A[k]=(A[k]+tmp)%mod;}
if(tp==-1) for(int i=0,inv=Pow(n,mod-2);i<n;i++) A[i]=1ll*A[i]*inv%mod;}
void Inv(int A[maxn],int B[maxn],int n)
{ B[0] = Pow(A[0],mod-2);
static int tmp[maxn];
for(int k=2;k<(n<<1);k<<=1)
{ for(int i=0;i<k;i++) tmp[i]=i<n?A[i]:0;
NTT(tmp,k<<1,1),NTT(B,k<<1,1);
for(int i=0;i<(k<<1);i++) B[i]=1ll*B[i]*(2-1ll*B[i]*tmp[i]%mod)%mod;
NTT(B,k<<1,-1);for(int i=min(n,k);i<(k<<1);i++) B[i]=0;}}
void cLn(int A[maxn],int B[maxn],int n)
{ Inv(A,B,n);
static int tmp[maxn];
int lgn = lg[n-1] + 2 , len = 1<<lgn;
for(int i=0;i<n-1;i++) tmp[i] = 1ll * A[i+1] * (i+1) % mod;
NTT(tmp,len,1),NTT(B,len,1);
for(int i=0;i<len;i++) tmp[i]=1ll*tmp[i]*B[i]%mod;
NTT(tmp,len,-1);B[0]=0;
for(int i=1;i<n;i++) B[i] = 1ll * tmp[i-1] * Pow(i,mod-2) % mod;}
int a[maxn],b[maxn];
int main()
{
int n;
scanf("%d",&n);
for(int i=0;i<n;i++) scanf("%d",&a[i]);
Init(n<<1);
cLn(a,b,n);
for(int i=0;i<n;i++) printf("%d%c",(b[i]+mod)%mod,i==n-1?'\n':' ');
}
多项式的exp用下面的方法。
多项式牛顿迭代
B ( x ) = A e ( x ) ( x ) ln B ( x ) = A ( x ) F ( B ( x ) ) = ln B ( x ) − A ( x ) = 0 F ′ ( B ( x ) ) = 1 B ( x ) B ( x ) = B 0 ( x ) − B 0 ( x ) ∗ ( ln B 0 ( x ) − A ( x ) ) = B 0 ( x ) ( 1 − ln B 0 ( x ) + A ( x ) ) \begin{aligned} &B(x) = A^{e(x)}(x)\\ &\ln B(x) = A(x)\\ &F(B(x)) = \ln B(x) - A(x) = 0\\ &F'(B(x)) = \frac 1{B(x)} \\ &B(x) = B_0(x) - B_0(x) * (\ln B_0(x) - A(x)) = B_0(x)(1 - \ln B_0(x) + A(x)) \end{aligned}B(x)=Ae(x)(x)lnB(x)=A(x)F(B(x))=lnB(x)−A(x)=0F′(B(x))=B(x)1B(x)=B0(x)−B0(x)∗(lnB0(x)−A(x))=B0(x)(1−lnB0(x)+A(x))
AC Code:(PS : 一旦到了需要多次调用函数的时候很多细节(清零)问题就出来了,以上的代码好像都没有考虑的说。。。。。。求逆,求L n LnLn,求e x p expexp都以本代码为准)
63行还是有点短的说。
#include<bits/stdc++.h>
#define maxn 300005
#define mod 998244353
using namespace std;
int lg[maxn],r[maxn],w[maxn]={1},wlen,inv[maxn]={1,1};
int Pow(int base,int k)
{ int ret = 1;
for(;k;k>>=1,base=1ll*base*base%mod) if(k&1) ret=1ll*ret*base%mod;
return ret; }
void Init(int n)
{ for(wlen=1;n>=(wlen*2);wlen<<=1);
for(int i=1,pw=Pow(3,(mod-1)/(2*wlen));i<=2*wlen;i++) w[i] = 1ll * w[i-1] * pw % mod;
for(int i=2;i<=2*wlen;i++) lg[i] = lg[i>>1] + 1;
for(int i=2;i<=2*wlen;i++) inv[i] = 1ll * (mod - mod / i) * inv[mod % i] % mod; }
inline void NTT(int *A,int n,int tp)
{ int lgn = lg[n];
for(int i=1;i<n;i++) r[i] = (r[i>>1]>>1)|((i&1)<<(lgn-1));
for(int i=1;i<n;i++) if(i<r[i]) swap(A[i],A[r[i]]);
for(int L=2;L<=n;L<<=1)
for(int st=0,l=L>>1,inc=wlen/l;st<n;st+=L)
for(int k=st,x=0;k<st+l;k++,x+=inc)
{ int tmp = 1ll * (tp==1?w[x]:w[2*wlen-x]) * A[k+l] % mod;
A[k+l]=(A[k]-tmp)%mod,A[k]=(A[k]+tmp)%mod;}
if(tp==-1) for(int i=0,inv=Pow(n,mod-2);i<n;i++) A[i]=1ll*A[i]*inv%mod;}
void Inv(int *A,int *B,int n)
{ B[0] = Pow(A[0] , mod-2);
static int tmp[maxn];
for(int k=2;k<(n<<1);k<<=1)
{ for(int i=0,lim=min(k,n);i<(k<<1);i++) tmp[i]=i<lim?A[i]:0,B[i]=i<lim?B[i]:0;
NTT(tmp,k<<1,1),NTT(B,k<<1,1);
for(int i=0,lim=k<<1;i<lim;i++) B[i]=1ll*B[i]*(2-1ll*B[i]*tmp[i]%mod)%mod;
NTT(B,k<<1,-1);for(int i=min(k,n);i<(k<<1);i++) B[i] = 0;}
}
void cLn(int *A,int *B,int n)
{ Inv(A,B,n);
static int tmp[maxn];
int lgn = lg[n-1]+2 , len = 1<<lgn;
for(int i=0;i<len;i++) tmp[i]=i<n-1?1ll*A[i+1]*(i+1)%mod:0;
NTT(tmp,len,1),NTT(B,len,1);
for(int i=0;i<len;i++) tmp[i]=1ll*B[i]*tmp[i]%mod;
NTT(tmp,len,-1);B[0]=0;
for(int i=1;i<n;i++) B[i] = 1ll * tmp[i-1] * inv[i] % mod;
for(int i=n;i<len;i++) B[i] = 0;}
void eXp(int *A,int *B,int n)
{ B[0] = 1;
static int tmp[maxn];
for(int k=2;k<(n<<1);k<<=1)
{ cLn(B,tmp,k);
for(int i=0,lim=min(n,k);i<(k<<1);i++) tmp[i]=i<lim?((i==0)-tmp[i]+A[i])%mod:0,B[i]=i<lim?B[i]:0;
NTT(B,k<<1,1),NTT(tmp,k<<1,1);
for(int i=0;i<(k<<1);i++) B[i]=1ll*B[i]*tmp[i]%mod;
NTT(B,k<<1,-1);for(int i=min(n,k);i<(k<<1);i++) B[i]=0;}}
int a[maxn],b[maxn];
int main()
{
int n;
scanf("%d",&n);
for(int i=0;i<n;i++) scanf("%d",&a[i]);
Init(n<<1);
eXp(a,b,n);
for(int i=0;i<n;i++) printf("%d%c",(b[i]+mod)%mod,i==n-1?'\n':' ');
}
一个奇怪的分支:多项式开根
#include<bits/stdc++.h>
#define maxn 300005
#define mod 998244353
using namespace std;
int w[maxn]={1},lg[maxn],r[maxn],wlen;
int Pow(int base,int k)
{ int ret = 1;
for(;k;k>>=1,base=1ll*base*base%mod) if(k&1) ret=1ll*ret*base%mod;
return ret;}
void Init(int n)
{ for(wlen=1;n>=2*wlen;wlen<<=1);
for(int i=1,pw=Pow(3,(mod-1)/(2*wlen));i<=2*wlen;i++) w[i]=1ll*w[i-1]*pw%mod;
for(int i=2;i<=2*wlen;i++) lg[i]=lg[i>>1]+1;}
inline void NTT(int *A,int n,int tp)
{ int lgn = lg[n];
for(int i=1;i<n;i++) r[i]=(r[i>>1]>>1)|((i&1)<<(lgn-1));
for(int i=1;i<n;i++) if(i<r[i]) swap(A[i],A[r[i]]);
for(int L=2;L<=n;L<<=1)
for(int st=0,l=L>>1,inc=wlen/l;st<n;st+=L)
for(int k=st,x=0,tmp;k<st+l;k++,x+=inc)
tmp=1ll*(tp==1?w[x]:w[2*wlen-x])*A[k+l]%mod,
A[k+l]=(A[k]-tmp)%mod,A[k]=(A[k]+tmp)%mod;
if(tp==-1) for(int i=0,inv=Pow(n,mod-2);i<n;i++) A[i]=1ll*A[i]*inv%mod;}
void Inv(int *A,int *B,int n)
{ B[0]=Pow(A[0],mod-2),B[1]=B[2]=B[3]=0;
static int tmp[maxn];
for(int k=2;k<(n<<1);k<<=1)
{ for(int i=0;i<(k<<1);i++) tmp[i]=i<k?A[i]:0,B[i]=i<k?B[i]:0;
NTT(tmp,k<<1,1),NTT(B,k<<1,1);
for(int i=0;i<(k<<1);i++) B[i]=1ll*B[i]*(2-1ll*B[i]*tmp[i]%mod)%mod;
NTT(B,k<<1,-1);for(int i=min(k,n);i<(k<<1);i++) B[i]=0;}
}
void Sqt(int *A,int *B,int n)
{ B[0]=1,B[1]=B[2]=B[3]=0;
static int tmp[2][maxn];
for(int k=2;k<(n<<1);k<<=1)
{ for(int i=0;i<(k<<1);i++) tmp[0][i]=i<k?A[i]:0,B[i]=i<k?B[i]:0;
Inv(B,tmp[1],k),NTT(tmp[0],k<<1,1),NTT(tmp[1],k<<1,1),NTT(B,k<<1,1);
for(int i=0;i<(k<<1);i++) B[i] = 1ll * (mod+1)/2 * tmp[1][i] % mod * (tmp[0][i]+1ll*B[i]*B[i]%mod) % mod;
NTT(B,k<<1,-1);for(int i=k;i<(k<<1);i++) B[i]=0;}}
int n,m;
int a[maxn],b[maxn],c[maxn],d[maxn];
int main()
{
scanf("%d%d",&n,&m);m++;
for(int i=0,x;i<n;i++) scanf("%d",&x),a[x]++;
Init(m<<1);
for(int i=0;i<=m;i++) b[i] = (i==0) - 4ll * a[i];
Sqt(b,c,m);
c[0]++;
Inv(c,d,m);
for(int i=1;i<m;i++)
printf("%d\n",(2ll*d[i]%mod+mod)%mod);
}
还是清零要注意
多点求值和快速插值
51nod 1387